From 5b9d3abced730a4259aa76ee4fa7c9908fe336c9 Mon Sep 17 00:00:00 2001 From: kbcoulter Date: Wed, 23 Sep 2026 10:59:10 -0700 Subject: [PATCH 1/2] Added patch_legacy_linear(), which checks for the final linear layer expected when loading legacy models and patches if necessary. Changed some loading around this function to ensure that loading still fails, as expected, with non-legacy models. --- src/embkit/factory/core.py | 31 ++++++++++++++++++++++++++++++- 1 file changed, 30 insertions(+), 1 deletion(-) diff --git a/src/embkit/factory/core.py b/src/embkit/factory/core.py index dcfc729..408039b 100644 --- a/src/embkit/factory/core.py +++ b/src/embkit/factory/core.py @@ -47,6 +47,21 @@ def save(model, path): def load(path, device=None, dtype=None): """Load a serialized model from ``path`` and optionally move its tensors.""" + + def patch_legacy_linear(): + '''Check for final linear layer expected in legacy loading.''' + encoder = getattr(model, "encoder", None) + if encoder is None: + return False + idx = len(encoder.net) + weight_key, bias_key = f"encoder.net.{idx}.weight", f"encoder.net.{idx}.bias" + if weight_key not in result.unexpected_keys or bias_key not in result.unexpected_keys: + return False + from .mapping import Linear + out_features, in_features = state_dict[weight_key].shape + encoder.net.append(Linear(in_features, out_features)) + return True + state_dict = torch.load(path, map_location=device, weights_only=False) desc = state_dict.pop("__model__", None) if desc is None: @@ -55,7 +70,21 @@ def load(path, device=None, dtype=None): "The file does not contain a model description and cannot be loaded." ) model = build(desc) - model.load_state_dict(state_dict) + result = model.load_state_dict(state_dict, strict = False) # Load non-strict + + if result.unexpected_keys and patch_legacy_linear(): + model.load_state_dict(state_dict, strict = True) # Legacy matches, Strict loading + elif result.unexpected_keys: + raise RuntimeError( + f"Error(s) in loading state_dict for {model.__class__.__name__}: " + f"Unexpected key(s) in state_dict: {result.unexpected_keys}." + ) + elif result.missing_keys: + raise RuntimeError( + f"Error(s) in loading state_dict for {model.__class__.__name__}: " + f"Missing key(s) in state_dict: {result.missing_keys}." + ) + if device is not None or dtype is not None: model.to(device=device, dtype=dtype) return model From 3846772617f285e4adf722de0ca6b47e778a76b2 Mon Sep 17 00:00:00 2001 From: kbcoulter Date: Wed, 23 Sep 2026 11:17:41 -0700 Subject: [PATCH 2/2] Minor Doc Fix --- src/embkit/factory/core.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/embkit/factory/core.py b/src/embkit/factory/core.py index 408039b..fe86833 100644 --- a/src/embkit/factory/core.py +++ b/src/embkit/factory/core.py @@ -49,7 +49,7 @@ def load(path, device=None, dtype=None): """Load a serialized model from ``path`` and optionally move its tensors.""" def patch_legacy_linear(): - '''Check for final linear layer expected in legacy loading.''' + '''Check for final linear layer expected in legacy loading and patch if present.''' encoder = getattr(model, "encoder", None) if encoder is None: return False