From 9cf25bd3dbb849198e7f6114b71e8fb1b91fabe2 Mon Sep 17 00:00:00 2001 From: Jack <32371937+jackzhxng@users.noreply.github.com> Date: Fri, 28 Feb 2025 12:26:54 -0800 Subject: [PATCH 1/2] Allow none tensor checkpoint values --- examples/models/checkpoint.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/examples/models/checkpoint.py b/examples/models/checkpoint.py index ee3fb560429..d11ae238565 100644 --- a/examples/models/checkpoint.py +++ b/examples/models/checkpoint.py @@ -64,7 +64,7 @@ def get_checkpoint_dtype(checkpoint: Dict[str, Any]) -> Optional[str]: mismatched_dtypes = [ (key, value.dtype) for key, value in checkpoint.items() - if value.dtype != dtype + if hasattr(value, 'dtype') and value.dtype != dtype ] if len(mismatched_dtypes) > 0: print( From eb76f77eafb3587886df72ad625fbbd296be0586 Mon Sep 17 00:00:00 2001 From: Jack <32371937+jackzhxng@users.noreply.github.com> Date: Fri, 28 Feb 2025 12:46:19 -0800 Subject: [PATCH 2/2] Lint --- examples/models/checkpoint.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/examples/models/checkpoint.py b/examples/models/checkpoint.py index d11ae238565..c84a689b951 100644 --- a/examples/models/checkpoint.py +++ b/examples/models/checkpoint.py @@ -64,7 +64,7 @@ def get_checkpoint_dtype(checkpoint: Dict[str, Any]) -> Optional[str]: mismatched_dtypes = [ (key, value.dtype) for key, value in checkpoint.items() - if hasattr(value, 'dtype') and value.dtype != dtype + if hasattr(value, "dtype") and value.dtype != dtype ] if len(mismatched_dtypes) > 0: print(