From 3a4855b92f9ea1f0300b6fa4f73320b4253c6bc2 Mon Sep 17 00:00:00 2001 From: RecML authors Date: Tue, 29 Sep 2026 19:37:24 -0700 Subject: [PATCH] Internal change PiperOrigin-RevId: 990677193 --- recml/core/utils/keras_utils.py | 346 ++++++++- recml/core/utils/keras_utils_test.py | 1029 ++++++++++++++++++++++++++ 2 files changed, 1360 insertions(+), 15 deletions(-) diff --git a/recml/core/utils/keras_utils.py b/recml/core/utils/keras_utils.py index 1a66e84..ab6af89 100644 --- a/recml/core/utils/keras_utils.py +++ b/recml/core/utils/keras_utils.py @@ -13,7 +13,7 @@ # limitations under the License. """Utilities for training Keras models on Jax backend.""" -from collections.abc import Mapping, Sequence +from collections.abc import Container, Mapping, Sequence import dataclasses import datetime import enum @@ -37,6 +37,7 @@ FORMAT_VERSION_KEY = "format_version" NON_TRAINABLE_PATHS_KEY = "non_trainable_paths" OPTIMIZER_PATHS_KEY = "optimizer_paths" +SHARED_VARIABLE_MAP_KEY = "shared_variable_map" ORBAX_CHECKPOINT_DEFAULT_KEY = "default" @@ -87,6 +88,86 @@ def _variables_to_path_dict( return var_dict +def extract_shared_variable_map(model: keras.Model) -> dict[str, str]: + """Extracts mapping of shared variables across the model layer tree. + + When multiple layers in a model reference the same `keras.Variable` instance, + Keras's `model.trainable_variables` assigns the variable a single canonical + `.path` corresponding to the first layer that tracked it. + + This function traverses the full layer hierarchy to discover alternative + logical paths referencing those shared variables and maps them to their + canonical path. + + An alias path is built as `{sharing_layer_path}/{var.name}`. Layer names in + the path include Keras uniquifying suffixes (e.g. `dense_1`); variable names + are never uniquified by Keras, so `var.name` is the name the variable was + created with. + + Supported scope: + Only layer-level sharing is supported: a layer instance is reused in + several places. The variable is created inside the shared layer, so its + name is the same under every parent, and the alias matches the path an + unshared model would use, e.g. + `{"model/surface_b/text_emb/embeddings": + "model/surface_a/text_emb/embeddings"}`. + + Caveat: + Variable-level sharing is NOT supported. This is when a layer stores + another layer's variable as an attribute under its own name, e.g. + `self.token_table = other_layer.embeddings`. The variable is still + recorded here, but its alias uses the variable's own name + (`model/tower/embeddings`), not the attribute name (`token_table`). A + target model that owns the weight as `model/tower/token_table` will not + match the alias, and restore fails with a missing-path error. `prefix_map` + does not help, since it keeps variable names unchanged; such models need an + explicit per-variable transform. + + Args: + model: The Keras model instance. + + Returns: + A dictionary mapping shared_path -> canonical_path (e.g. + `{"model/search_surface/text_emb/embeddings": + "model/chrome_surface/text_emb/embeddings"}`). + + Raises: + ValueError: If two different shared variables produce the same alias path, + e.g. a layer that stores the `kernel` of two other layers. + """ + shared_var_map = {} + + def _traverse(layer: keras.layers.Layer, prefix: str): + # Walk tracked child layers + for child in getattr(layer, "_layers", []): + # child.name includes any uniquifying suffix (e.g. 'dense_1') assigned + # during layer construction. + child_prefix = f"{prefix}/{child.name}" if prefix else child.name + _traverse(child, child_prefix) + + # Check variables directly attached to this layer. + own_vars = getattr(layer, "_trainable_variables", []) + getattr( + layer, "_non_trainable_variables", [] + ) + for var in own_vars: + # var.name is identical across both paths since it is the same Variable. + logical_path = f"{prefix}/{var.name}" if prefix else var.name + if logical_path == var.path: + continue + existing = shared_var_map.get(logical_path) + if existing is not None and existing != var.path: + raise ValueError( + f"Shared variables {existing} and {var.path} both map to alias " + f"path {logical_path}. This happens when a layer stores several " + "shared variables with the same name as attributes; share whole " + "layers instead." + ) + shared_var_map[logical_path] = var.path + + _traverse(model, model.name or "") + return shared_var_map + + def _to_shape_dtype_struct(x: keras.Variable) -> jax.ShapeDtypeStruct: if not isinstance(x, keras.Variable): raise ValueError(f"Expected a `keras.Variable`, got {type(x)}.") @@ -253,6 +334,7 @@ def __init__( FORMAT_VERSION_KEY, NON_TRAINABLE_PATHS_KEY, OPTIMIZER_PATHS_KEY, + SHARED_VARIABLE_MAP_KEY, ), options=ocp.CheckpointManagerOptions( save_interval_steps=save_interval_epochs, @@ -305,8 +387,11 @@ def save_model_variables( "paths": [v.path for v in model.non_trainable_variables] } optimizer_paths = {"paths": [v.path for v in model.optimizer.variables]} + shared_var_map = extract_shared_variable_map(model) + shared_variable_map = {SHARED_VARIABLE_MAP_KEY: shared_var_map} logging.info("SAVED non_trainable_paths: %s", non_trainable_paths) logging.info("SAVED optimizer_paths: %s", optimizer_paths) + logging.info("SAVED shared_variable_map: %s", shared_variable_map) logging.info("Saving checkpoint for epoch %s...", epoch) self.save( @@ -317,6 +402,7 @@ def save_model_variables( FORMAT_VERSION_KEY: ocp.args.JsonSave({"version": 3}), NON_TRAINABLE_PATHS_KEY: ocp.args.JsonSave(non_trainable_paths), OPTIMIZER_PATHS_KEY: ocp.args.JsonSave(optimizer_paths), + SHARED_VARIABLE_MAP_KEY: ocp.args.JsonSave(shared_variable_map), }), metrics=logs, ) @@ -480,8 +566,15 @@ def _detect_checkpoint_version(checkpoint_path: str) -> CheckpointVersion: Detection order and discriminator criteria: - V1 checkpoints use the legacy item directory layout ('default/'). - - V3 checkpoints contain an explicit 'format_version' item with {"version": 3}. + - V3 checkpoints contain an explicit 'format_version' item with + {"version": 3}. - V2 checkpoints contain 'state/' without the V3 'format_version' marker. + + Args: + checkpoint_path: Path to the checkpoint directory. + + Returns: + The detected CheckpointVersion enum value. """ # V1 uses the legacy 'default' directory layout. if _is_v1_checkpoint_path(checkpoint_path): @@ -489,7 +582,8 @@ def _detect_checkpoint_version(checkpoint_path: str) -> CheckpointVersion: # V3 checkpoints are explicitly tagged with format_version. if _is_v3_checkpoint_path(checkpoint_path): return CheckpointVersion.V3 - # V2 checkpoints have a 'state' directory without the V3 format_version marker. + # V2 checkpoints have a 'state' directory without the V3 format_version + # marker. if gfile.Exists(os.path.join(checkpoint_path, STATE_CHECKPOINT_KEY)): return CheckpointVersion.V2 raise ValueError(f"Unknown checkpoint format at {checkpoint_path}") @@ -507,10 +601,17 @@ def _validate_v3_checkpoint( ) -> Any: """Validates that the V3 checkpoint is healthy and returns its state metadata. - A healthy V3 checkpoint must contain the state directory and companion path - metadata files on disk, and its state metadata must contain all three keys + A healthy V3 checkpoint must contain the state directory and its companion + metadata items on disk (non_trainable_paths, optimizer_paths, and + shared_variable_map), and its state metadata must contain all three keys (trainable_variables, non_trainable_variables, and optimizer_variables). + `shared_variable_map` is required even for models without shared variables, + where it is an empty map. Without it, a restore could not resolve shared + variables referenced by an alias path, so a checkpoint that lacks it is + treated as incomplete rather than silently restored without alias + resolution. + Args: checkpoint_path: Path to the checkpoint directory. saved_state_metadata: Optional pre-loaded state metadata. If None, it will @@ -527,6 +628,7 @@ def _validate_v3_checkpoint( STATE_CHECKPOINT_KEY, NON_TRAINABLE_PATHS_KEY, OPTIMIZER_PATHS_KEY, + SHARED_VARIABLE_MAP_KEY, ) missing_items = [f for f in required_items if not (root / f).exists()] if missing_items: @@ -556,12 +658,131 @@ def _validate_v3_checkpoint( return saved_state_metadata +def _find_stored_path( + path: str, + stored_paths: Container[str], + shared_var_map: Mapping[str, str], +) -> str | None: + """Returns the path under which `path`'s value is stored, or None. + + A variable is stored under its own path, unless it is a shared variable + reached through an alias, in which case it is stored under the canonical path + recorded in `shared_var_map`. A direct match takes precedence, so a real + variable is never shadowed by an alias that happens to share its path. + + Args: + path: A variable path, possibly an alias of a shared variable. + stored_paths: The paths stored in the checkpoint. + shared_var_map: Map from alias path to canonical path. + + Returns: + The stored path holding the value, or None if there is none. + """ + if path in stored_paths: + return path + canonical = shared_var_map.get(path) + if canonical is not None and canonical in stored_paths: + return canonical + return None + + +def _resolve_transform_source( + transform: Any, + key: str, + stored_paths: Container[str], + shared_var_map: Mapping[str, str], +) -> Any: + """Rewrites a transform's source key to its stored path if it names an alias. + + In V3 checkpoints, shared weights are deduplicated so only the canonical path + is physically stored on disk (e.g. `surface_b`), while aliases are + recorded in `shared_var_map` (e.g. `surface_a -> surface_b`). + + If a caller supplies a `Transform(original_key="surface_a")` to restore a + renamed variable (e.g. `surface_a_1`), Orbax cannot find `surface_a` on disk. + This function resolves `original_key` to its physical storage path + (`surface_b`) so Orbax can load the tensor. + + Example: + If `surface_a` was deduplicated to `surface_b` in the checkpoint: + `shared_var_map`: `{"surface_a/kernel": "surface_b/kernel"}` + `stored_paths`: `{"surface_b/kernel"}` + + # User maps renamed variable surface_a_1 -> surface_a: + resolved = _resolve_transform_source( + transform=Transform( + original_key="trainable_variables/surface_a/kernel" + ), + key="trainable_variables", + stored_paths=stored_paths, + shared_var_map=shared_var_map, + ) + # rewritten to: original_key="trainable_variables/surface_b/kernel" + + Args: + transform: A user-supplied transform for one target variable. + key: The state key, e.g. `trainable_variables`. `original_key` may or may + not carry it as a `{key}/` prefix. + stored_paths: The paths stored in the checkpoint under `key`. + shared_var_map: Map from alias path to canonical path. + + Returns: + The transform, with `original_key` rewritten if it named an alias. + """ + if not isinstance(transform, ocp.transform_utils.Transform) or not isinstance( + transform.original_key, str + ): + return transform + prefix = f"{key}/" + has_prefix = transform.original_key.startswith(prefix) + source = transform.original_key.removeprefix(prefix) + stored = _find_stored_path(source, stored_paths, shared_var_map) + if stored is None or stored == source: + # Not an alias. Leave it for Orbax to resolve, or to reject if missing. + return transform + logging.info( + "Resolved transform source %s to stored path %s via shared variable map", + source, + stored, + ) + return dataclasses.replace( + transform, original_key=f"{prefix}{stored}" if has_prefix else stored + ) + + +def _map_prefix( + path: str, prefix_map: Mapping[str, str] +) -> tuple[str, str] | None: + """Rewrites `path` using the longest matching prefix in `prefix_map`. + + Prefixes match whole `/`-separated segments only, so `model/tower` matches + `model/tower/dense/kernel` but not `model/tower_2/dense/kernel`. + + Args: + path: A target variable path. + prefix_map: Map from target path prefix to source path prefix. + + Returns: + A tuple of (matched target prefix, source path), or None if no prefix + matches. + """ + best = None + for target_prefix in prefix_map: + if path == target_prefix or path.startswith(f"{target_prefix}/"): + if best is None or len(target_prefix) > len(best): + best = target_prefix + if best is None: + return None + return best, prefix_map[best] + path[len(best) :] + + def _prepare_v3_restore( checkpoint_path: str, abstract_state: Mapping[str, Any], model: keras.Model | None = None, restore_optimizer_vars: bool = False, transforms: Mapping[str, Any] | None = None, + prefix_map: Mapping[str, str] | None = None, ) -> tuple[Mapping[str, Any], Mapping[str, Any]]: """Prepares the abstract state and constructs transforms for V3 restore. @@ -581,6 +802,8 @@ def _prepare_v3_restore( variables mapping. restore_optimizer_vars: Whether to prepare optimizer variables. transforms: An optional mapping of custom transforms for variable mapping. + prefix_map: An optional map from target path prefix to source path prefix + for trainable variables. See `restore_partial_checkpoint`. Returns: A tuple of (filtered_abstract_state, state_transforms) to be passed to @@ -589,7 +812,18 @@ def _prepare_v3_restore( # Validate that underlying checkpoint is healthy and retrieve state metadata. saved_state_metadata = _validate_v3_checkpoint(checkpoint_path) - # Load index path metadata for non-trainable and optimizer variables if needed. + # The shared variable map is a required V3 item, checked above. + map_path = epath.Path(checkpoint_path) / SHARED_VARIABLE_MAP_KEY + checkpointer = ocp.Checkpointer(ocp.handlers.JsonCheckpointHandler()) + try: + shared_var_map = checkpointer.restore( + os.fspath(map_path), args=ocp.args.JsonRestore() + )[SHARED_VARIABLE_MAP_KEY] + finally: + checkpointer.close() + + # Load index path metadata for non-trainable and optimizer variables if + # needed. non_trainable_paths = None if model is not None and abstract_state.get(NON_TRAINABLE_VARIABLES_KEY): nt_path = epath.Path(checkpoint_path) / NON_TRAINABLE_PATHS_KEY @@ -673,7 +907,10 @@ def _prepare_v3_restore( i, ) else: - # Strict matching for trainable variables unless explicit transforms provided + # Trainable variables are matched by path. Each target is restored from + # the source a user transform names, else from its prefix-mapped path, + # else from its own path; any of these may be an alias of a shared + # variable, resolved by _find_stored_path. missing_paths = [] key_transforms = {} if transforms: @@ -682,15 +919,65 @@ def _prepare_v3_restore( if key in transforms and isinstance(transforms[key], Mapping) else transforms ) + stored_paths = saved_state_metadata[key] + normalized_prefix_map = { + t.rstrip("/"): s.rstrip("/") for t, s in (prefix_map or {}).items() + } + if "" in normalized_prefix_map: + raise ValueError("prefix_map keys must be non-empty path prefixes.") + unused_prefixes = set(normalized_prefix_map) for target_path, struct in abstract_state[key].items(): - if target_path in saved_state_metadata[key]: - filtered_abstract_state[key][target_path] = struct - elif target_path in key_transforms: - filtered_abstract_state[key][target_path] = struct - state_transforms[key][target_path] = key_transforms[target_path] + mapped_source = None + prefix_match = _map_prefix(target_path, normalized_prefix_map) + if prefix_match is not None: + matched_prefix, mapped_source = prefix_match + unused_prefixes.discard(matched_prefix) + # Branch 1: Explicit user transform (remapping/surgery). + # Target and source paths must match exact model variable paths + # (including suffixes like `dense_1`). If `original_key` is an alias, + # it is resolved to its physical storage path. + if target_path in key_transforms: + transform = _resolve_transform_source( + key_transforms[target_path], key, stored_paths, shared_var_map + ) + # Branch 2: Prefix mapping. The target is restored from the source path + # obtained by swapping its longest matching prefix; the source may be + # an alias of a shared variable. + elif mapped_source is not None: + stored = _find_stored_path( + mapped_source, stored_paths, shared_var_map + ) + if stored is None: + missing_paths.append(f"{target_path} (from {mapped_source})") + continue + transform = ( + None + if stored == target_path + else ocp.transform_utils.Transform(original_key=f"{key}/{stored}") + ) + # Branch 3: Default 1-to-1 match. + # Resolves target_path if it is an alias, synthesizing a Transform to + # its canonical storage path. else: - missing_paths.append(target_path) - + stored = _find_stored_path(target_path, stored_paths, shared_var_map) + if stored is None: + missing_paths.append(target_path) + continue + transform = ( + None + if stored == target_path + else ocp.transform_utils.Transform(original_key=f"{key}/{stored}") + ) + filtered_abstract_state[key][target_path] = struct + if transform is not None: + state_transforms[key][target_path] = transform + + if unused_prefixes: + logging.warning( + "prefix_map entries matched no target variables for key %s: %s", + key, + sorted(unused_prefixes), + ) if missing_paths: raise ValueError( f"Failed to restore variables for key {key}. " @@ -965,9 +1252,21 @@ def restore_partial_checkpoint( partial_variables: Mapping[str, Any], epoch: int | None = None, transforms: Mapping[str, Any] | None = None, + prefix_map: Mapping[str, str] | None = None, ) -> Mapping[str, Any]: """Restores partial variables from an Orbax checkpoint. + Each target variable is restored from, in order of precedence: + 1. The source named by its entry in `transforms`, if any. + 2. The source obtained by rewriting its path with `prefix_map`, if a prefix + matches. + 3. Its own path (default 1-to-1 match). + User-specified remappings (1 and 2) take precedence even if the target path + already exists in the checkpoint (e.g. initializing one surface from another + within the same model). Sources in all three cases may be aliases of shared + variables; they are resolved to their stored paths through the checkpoint's + shared variable map. + Args: checkpoint_dir: The directory containing the Orbax checkpoint(s). partial_variables: A dictionary mapping keys (e.g. @@ -976,7 +1275,23 @@ def restore_partial_checkpoint( epoch: The epoch to restore. If None, latest is used. transforms: An optional mapping of custom transforms (e.g. mapping target variable paths to `ocp.transform_utils.Transform` objects) for explicit - variable re-mapping. + variable re-mapping. Keys and `original_key` values must be exact + variable paths, including Keras layer name suffixes (e.g. `dense_1`). + Prefer `prefix_map` or the default resolution unless you need + per-variable control. + prefix_map: An optional map from target path prefix to source path prefix + for subtree remapping, e.g. `{"model/target_tower": "model/src_tower"}`. + Paths omit the state key (e.g. `trainable_variables/`). + Rules: + - Depth matching: A prefix matches all variables at any depth beneath + it on whole path segments (`layer_a` matches `layer_a/var_a` and + `layer_a/layer_b/var_b`, but not `layer_a_2`). + - Overlapping prefixes: The longest (most specific) prefix wins, + allowing subtree overrides (`model/new/tower` overrides `model/new`). + - Subtree structure: The path remainder below the prefix is preserved, + so source and target subtrees must share identical structure. + Logs a warning if a prefix matches no target variables; raises if a + mapped source is missing; shape mismatches fail at restore. Returns: The restored state dictionary (containing Jax Arrays). @@ -1019,6 +1334,7 @@ def restore_partial_checkpoint( model=None, restore_optimizer_vars=False, transforms=transforms, + prefix_map=prefix_map, ) restored_state = _restore_state_pytree( diff --git a/recml/core/utils/keras_utils_test.py b/recml/core/utils/keras_utils_test.py index ce49612..d130f55 100644 --- a/recml/core/utils/keras_utils_test.py +++ b/recml/core/utils/keras_utils_test.py @@ -1357,6 +1357,7 @@ def test_validate_v3_checkpoint_missing_state_key_fails(self): for meta_key in ( keras_utils.NON_TRAINABLE_PATHS_KEY, keras_utils.OPTIMIZER_PATHS_KEY, + keras_utils.SHARED_VARIABLE_MAP_KEY, ): with open(os.path.join(step_dir, meta_key), "w") as f: f.write('{"paths": []}') @@ -1371,6 +1372,1034 @@ def test_validate_v3_checkpoint_missing_state_key_fails(self): ): keras_utils._validate_v3_checkpoint(step_dir, incomplete_metadata) + def test_extract_shared_variable_map(self): + shared_dense = keras.layers.Dense(2, name="shared_dense") + + class SurfaceLayer(keras.layers.Layer): + + def __init__(self, inner, **kwargs): + super().__init__(**kwargs) + self.inner = inner + + def call(self, x): + return self.inner(x) + + model = keras.Sequential( + [ + SurfaceLayer(shared_dense, name="chrome_surface"), + SurfaceLayer(shared_dense, name="search_surface"), + ], + name="my_model", + ) + model.build((1, 2)) + + shared_var_map = keras_utils.extract_shared_variable_map(model) + expected_map = { + "my_model/search_surface/shared_dense/kernel": ( + "my_model/chrome_surface/shared_dense/kernel" + ), + "my_model/search_surface/shared_dense/bias": ( + "my_model/chrome_surface/shared_dense/bias" + ), + } + self.assertEqual(shared_var_map, expected_map) + + def test_extract_shared_variable_map_no_shared_variables(self): + model = keras.Sequential( + [ + keras.layers.Dense(4, name="dense_1"), + keras.layers.Dense(2, name="dense_2"), + ], + name="sequential_model", + ) + model.build((1, 8)) + shared_var_map = keras_utils.extract_shared_variable_map(model) + self.assertEmpty(shared_var_map) + + def test_extract_shared_variable_map_multi_surface(self): + shared_dense = keras.layers.Dense(2, name="shared_dense") + + class SurfaceLayer(keras.layers.Layer): + + def __init__(self, inner, **kwargs): + super().__init__(**kwargs) + self.inner = inner + + def call(self, x): + return self.inner(x) + + model = keras.Sequential( + [ + SurfaceLayer(shared_dense, name="surface_a"), + SurfaceLayer(shared_dense, name="surface_b"), + SurfaceLayer(shared_dense, name="surface_c"), + ], + name="my_model", + ) + model.build((1, 2)) + + shared_var_map = keras_utils.extract_shared_variable_map(model) + expected_map = { + "my_model/surface_b/shared_dense/kernel": ( + "my_model/surface_a/shared_dense/kernel" + ), + "my_model/surface_b/shared_dense/bias": ( + "my_model/surface_a/shared_dense/bias" + ), + "my_model/surface_c/shared_dense/kernel": ( + "my_model/surface_a/shared_dense/kernel" + ), + "my_model/surface_c/shared_dense/bias": ( + "my_model/surface_a/shared_dense/bias" + ), + } + self.assertEqual(shared_var_map, expected_map) + + def test_extract_shared_variable_map_deeply_nested(self): + shared_dense = keras.layers.Dense(2, name="shared_dense") + + class TowerLayer(keras.layers.Layer): + + def __init__(self, inner, **kwargs): + super().__init__(**kwargs) + self.inner = inner + + def call(self, x): + return self.inner(x) + + class SurfaceLayer(keras.layers.Layer): + + def __init__(self, tower, **kwargs): + super().__init__(**kwargs) + self.tower = tower + + def call(self, x): + return self.tower(x) + + model = keras.Sequential( + [ + SurfaceLayer( + TowerLayer(shared_dense, name="tower"), name="surface_a" + ), + SurfaceLayer( + TowerLayer(shared_dense, name="tower"), name="surface_b" + ), + ], + name="nested_model", + ) + model.build((1, 2)) + + shared_var_map = keras_utils.extract_shared_variable_map(model) + expected_map = { + "nested_model/surface_b/tower/shared_dense/kernel": ( + "nested_model/surface_a/tower/shared_dense/kernel" + ), + "nested_model/surface_b/tower/shared_dense/bias": ( + "nested_model/surface_a/tower/shared_dense/bias" + ), + } + self.assertEqual(shared_var_map, expected_map) + + def test_extract_shared_variable_map_with_suffixed_layers(self): + shared_dense = keras.layers.Dense(2) + + class SurfaceLayer(keras.layers.Layer): + + def __init__(self, inner, **kwargs): + super().__init__(**kwargs) + self.inner = inner + + def call(self, x): + return self.inner(x) + + model = keras.Sequential( + [ + SurfaceLayer(shared_dense), + SurfaceLayer(shared_dense), + ], + name="my_model", + ) + model.build((1, 2)) + + # Verify that child layer names with uniquifying suffixes (e.g. dense_1) + # map correctly to the first instance. + shared_var_map = keras_utils.extract_shared_variable_map(model) + self.assertNotEmpty(shared_var_map) + for logical_path, canonical_path in shared_var_map.items(): + self.assertNotEqual(logical_path, canonical_path) + self.assertTrue( + logical_path.endswith("/kernel") or logical_path.endswith("/bias") + ) + + def test_restore_partial_checkpoint_with_shared_variable_map(self): + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + + class MultiSurfaceModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.shared_emb = keras.layers.Embedding(10, 4, name="text_emb") + self.chrome_surface = keras.Sequential( + [self.shared_emb], name="chrome_surface" + ) + self.search_surface = keras.Sequential( + [self.shared_emb], name="search_surface" + ) + + def build(self, input_shape): + self.shared_emb.build(input_shape) + self.chrome_surface.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.chrome_surface(x) + self.search_surface(x) + + model_source = MultiSurfaceModel(name="model") + model_source.compile(optimizer="adam") + model_source.build((1, 1)) + model_source.optimizer.build(model_source.trainable_variables) + test_weights = np.ones((10, 4), dtype=np.float32) * 7.5 + model_source.shared_emb.embeddings.assign(test_weights) + + manager.save_model_variables(model_source, epoch=1) + manager.wait_until_finished() + + # Target model only has search_surface with its own independent embedding + # layer. + class SearchOnlyModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.target_emb = keras.layers.Embedding(10, 4, name="text_emb") + self.search_surface = keras.Sequential( + [self.target_emb], name="search_surface" + ) + + def build(self, input_shape): + self.target_emb.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.search_surface(x) + + model_target = SearchOnlyModel(name="model") + model_target.build((1, 1)) + model_target.target_emb.embeddings.assign( + np.zeros((10, 4), dtype=np.float32) + ) + + # Partial restore targeting only search_surface's embeddings + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + model_target.target_emb.embeddings.path: ( + model_target.target_emb.embeddings + ) + } + } + + # Restores successfully using the shared variable map automatically + restored_state = keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, partial_vars, epoch=1 + ) + + for key, var_dict in partial_vars.items(): + for path, var in var_dict.items(): + var.assign(restored_state[key][path]) + + np.testing.assert_allclose( + model_target.target_emb.embeddings.value, test_weights + ) + manager.close() + + def test_restore_partial_checkpoint_explicit_transforms_precedence(self): + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + + class MultiSurfaceModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.shared_emb = keras.layers.Embedding(10, 4, name="text_emb") + self.chrome_surface = keras.Sequential( + [self.shared_emb], name="chrome_surface" + ) + self.search_surface = keras.Sequential( + [self.shared_emb], name="search_surface" + ) + self.custom_emb = keras.layers.Embedding(10, 4, name="custom_emb") + + def build(self, input_shape): + self.shared_emb.build(input_shape) + self.chrome_surface.build(input_shape) + self.search_surface.build(input_shape) + self.custom_emb.build(input_shape) + super().build(input_shape) + + def call(self, x): + return ( + self.chrome_surface(x) + + self.search_surface(x) + + self.custom_emb(x) + ) + + model_source = MultiSurfaceModel(name="model") + model_source.compile(optimizer="adam") + model_source.build((1, 1)) + model_source.optimizer.build(model_source.trainable_variables) + shared_weights = np.ones((10, 4), dtype=np.float32) * 7.5 + custom_weights = np.ones((10, 4), dtype=np.float32) * 3.0 + model_source.shared_emb.embeddings.assign(shared_weights) + model_source.custom_emb.embeddings.assign(custom_weights) + + manager.save_model_variables(model_source, epoch=1) + manager.wait_until_finished() + + class SearchOnlyModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.target_emb = keras.layers.Embedding(10, 4, name="text_emb") + self.search_surface = keras.Sequential( + [self.target_emb], name="search_surface" + ) + + def build(self, input_shape): + self.target_emb.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.search_surface(x) + + model_target = SearchOnlyModel(name="model") + model_target.build((1, 1)) + model_target.target_emb.embeddings.assign( + np.zeros((10, 4), dtype=np.float32) + ) + + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + model_target.target_emb.embeddings.path: ( + model_target.target_emb.embeddings + ) + } + } + + # Explicit transform mapping target to custom_emb instead of + # shared_variable_map default (chrome_surface). + explicit_transforms = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + model_target.target_emb.embeddings.path: ( + ocp.transform_utils.Transform( + original_key=( + f"{keras_utils.TRAINABLE_VARIABLES_KEY}" + "/model/custom_emb/embeddings" + ) + ) + ) + } + } + + restored_state = keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, + partial_vars, + epoch=1, + transforms=explicit_transforms, + ) + + for key, var_dict in partial_vars.items(): + for path, var in var_dict.items(): + var.assign(restored_state[key][path]) + + np.testing.assert_allclose( + model_target.target_emb.embeddings.value, custom_weights + ) + manager.close() + + def test_restore_partial_checkpoint_suffixed_target_path_fails_without_transform( + self, + ): + """Verifies suffixed paths fail without transform and resolve via shared_var_map.""" + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + + class MultiSurfaceModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.shared_dense = keras.layers.Dense(2, name="dense") + self.surface_b = keras.Sequential( + [self.shared_dense], name="surface_b" + ) + self.surface_a = keras.Sequential( + [self.shared_dense], name="surface_a" + ) + + def build(self, input_shape): + # Building surface_b first ensures shared_dense gets its canonical + # path under surface_b ("model/surface_b/dense/kernel"). + self.surface_b.build(input_shape) + self.surface_a.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.surface_b(x) + self.surface_a(x) + + model_source = MultiSurfaceModel(name="model") + model_source.compile(optimizer="adam") + model_source.build((1, 2)) + model_source.optimizer.build(model_source.trainable_variables) + test_weights = np.ones((2, 2), dtype=np.float32) * 4.2 + model_source.shared_dense.kernel.assign(test_weights) + + manager.save_model_variables(model_source, epoch=1) + manager.wait_until_finished() + + # Target model has a suffixed surface name (surface_a_1 due to + # auto-naming or naming drift) instead of surface_a. + class SuffixedModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.surface_a = keras.Sequential( + [keras.layers.Dense(2, name="dense")], + name="surface_a_1", + ) + + def build(self, input_shape): + self.surface_a.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.surface_a(x) + + model_target = SuffixedModel(name="model") + model_target.build((1, 2)) + target_dense = model_target.surface_a.layers[0] + target_dense.kernel.assign(np.zeros((2, 2), dtype=np.float32)) + + suffixed_target_path = target_dense.kernel.path + self.assertEqual( + suffixed_target_path, "model/surface_a_1/dense/kernel" + ) + + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + suffixed_target_path: target_dense.kernel + } + } + + # 1. Proves that V3 strict matching fails fast when target path has suffix + # drift (surface_a_1) because checkpoint only contains surface_b and alias + # surface_a. + with self.assertRaisesRegex( + ValueError, + ( + r"(?i)missing paths in checkpoint:" + r" \['model/surface_a_1/dense/kernel'\]" + ), + ): + keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, partial_vars, epoch=1 + ) + + # 2. Proves that providing an explicit transform (surface_a_1 -> surface_a) + # succeeds, resolving surface_a through shared_variable_map to surface_b. + model_target_2 = SuffixedModel(name="model") + model_target_2.build((1, 2)) + target_dense_2 = model_target_2.surface_a.layers[0] + target_dense_2.kernel.assign(np.zeros((2, 2), dtype=np.float32)) + + partial_vars_2 = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + suffixed_target_path: target_dense_2.kernel + } + } + explicit_transform = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + suffixed_target_path: ocp.transform_utils.Transform( + original_key=( + f"{keras_utils.TRAINABLE_VARIABLES_KEY}" + "/model/surface_a/dense/kernel" + ) + ) + } + } + restored_state = keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, + partial_vars_2, + epoch=1, + transforms=explicit_transform, + ) + for key, var_dict in partial_vars_2.items(): + for path, var in var_dict.items(): + var.assign(restored_state[key][path]) + + np.testing.assert_allclose( + target_dense_2.kernel.value, test_weights + ) + manager.close() + + def test_restore_keras_checkpoint_v3_with_shared_variable_map(self): + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + + class MultiSurfaceModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.shared_emb = keras.layers.Embedding(10, 4, name="text_emb") + self.chrome_surface = keras.Sequential( + [self.shared_emb], name="chrome_surface" + ) + self.search_surface = keras.Sequential( + [self.shared_emb], name="search_surface" + ) + + def build(self, input_shape): + self.shared_emb.build(input_shape) + self.chrome_surface.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.chrome_surface(x) + self.search_surface(x) + + model_source = MultiSurfaceModel(name="model") + model_source.compile(optimizer="adam") + model_source.build((1, 1)) + model_source.optimizer.build(model_source.trainable_variables) + test_weights = np.ones((10, 4), dtype=np.float32) * 6.0 + model_source.shared_emb.embeddings.assign(test_weights) + + manager.save_model_variables(model_source, epoch=1) + manager.wait_until_finished() + + class SearchOnlyModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.target_emb = keras.layers.Embedding(10, 4, name="text_emb") + self.search_surface = keras.Sequential( + [self.target_emb], name="search_surface" + ) + + def build(self, input_shape): + self.target_emb.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.search_surface(x) + + model_target = SearchOnlyModel(name="model") + model_target.build((1, 1)) + model_target.target_emb.embeddings.assign( + np.zeros((10, 4), dtype=np.float32) + ) + + keras_utils.restore_keras_checkpoint( + self.checkpoint_dir, + model=model_target, + epoch=1, + restore_optimizer_vars=False, + ) + + np.testing.assert_allclose( + model_target.target_emb.embeddings.value, test_weights + ) + manager.close() + + def test_find_stored_path_resolves_alias(self): + stored = {"model/a/dense/kernel"} + shared = {"model/b/dense/kernel": "model/a/dense/kernel"} + self.assertEqual( + keras_utils._find_stored_path("model/b/dense/kernel", stored, shared), + "model/a/dense/kernel", + ) + self.assertEqual( + keras_utils._find_stored_path("model/a/dense/kernel", stored, shared), + "model/a/dense/kernel", + ) + self.assertIsNone( + keras_utils._find_stored_path("model/c/dense/kernel", stored, shared) + ) + + def test_find_stored_path_prefers_direct_match_over_alias(self): + # "model/b/kernel" is both a stored variable and an alias key; the stored + # variable must win. + stored = {"model/a/kernel", "model/b/kernel"} + shared = {"model/b/kernel": "model/a/kernel"} + self.assertEqual( + keras_utils._find_stored_path("model/b/kernel", stored, shared), + "model/b/kernel", + ) + + def test_restore_v3_checkpoint_missing_shared_variable_map_fails(self): + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + model = keras.Sequential( + [keras.layers.Dense(1, name="dense")], name="model" + ) + model.compile(optimizer="sgd") + model.build((1, 1)) + model.optimizer.build(model.trainable_variables) + + manager.save_model_variables(model, epoch=1) + manager.wait_until_finished() + + # The map is written even for models without shared variables... + map_path = os.path.join( + self.checkpoint_dir, "1", keras_utils.SHARED_VARIABLE_MAP_KEY + ) + self.assertTrue(os.path.exists(map_path)) + # ...and a V3 checkpoint without it is rejected as incomplete. + shutil.rmtree(map_path) + + with self.assertRaisesRegex( + ValueError, + r"missing required checkpoint item\(s\): \['shared_variable_map'\]", + ): + keras_utils.restore_keras_checkpoint( + self.checkpoint_dir, + model=model, + epoch=1, + restore_optimizer_vars=False, + ) + manager.close() + + def test_restore_partial_checkpoint_missing_variable_fails(self): + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + model = keras.Sequential([keras.layers.Dense(1, name="dense")]) + model.compile(optimizer="sgd") + model.build((1, 1)) + model.optimizer.build(model.trainable_variables) + + manager.save_model_variables(model, epoch=1) + manager.wait_until_finished() + + dummy_var = keras.Variable( + np.zeros((1, 1), dtype=np.float32), name="missing" + ) + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + "non_existent_path/kernel": dummy_var + } + } + + with self.assertRaisesRegex( + ValueError, "Failed to restore variables for key trainable_variables" + ): + keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, partial_vars, epoch=1 + ) + manager.close() + + def test_restore_partial_checkpoint_multi_variable_layer(self): + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + + class MultiSurfaceDenseModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.shared_dense = keras.layers.Dense(2, name="shared_dense") + self.chrome_surface = keras.Sequential( + [self.shared_dense], name="chrome_surface" + ) + self.search_surface = keras.Sequential( + [self.shared_dense], name="search_surface" + ) + + def build(self, input_shape): + self.shared_dense.build(input_shape) + self.chrome_surface.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.chrome_surface(x) + self.search_surface(x) + + model_source = MultiSurfaceDenseModel(name="model") + model_source.compile(optimizer="adam") + model_source.build((1, 2)) + model_source.optimizer.build(model_source.trainable_variables) + test_kernel = np.ones((2, 2), dtype=np.float32) * 4.2 + test_bias = np.ones((2,), dtype=np.float32) * 1.8 + model_source.shared_dense.kernel.assign(test_kernel) + model_source.shared_dense.bias.assign(test_bias) + + manager.save_model_variables(model_source, epoch=1) + manager.wait_until_finished() + + class SearchOnlyDenseModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.target_dense = keras.layers.Dense(2, name="shared_dense") + self.search_surface = keras.Sequential( + [self.target_dense], name="search_surface" + ) + + def build(self, input_shape): + self.target_dense.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.search_surface(x) + + model_target = SearchOnlyDenseModel(name="model") + model_target.build((1, 2)) + model_target.target_dense.kernel.assign( + np.zeros((2, 2), dtype=np.float32) + ) + model_target.target_dense.bias.assign(np.zeros((2,), dtype=np.float32)) + + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + model_target.target_dense.kernel.path: ( + model_target.target_dense.kernel + ), + model_target.target_dense.bias.path: ( + model_target.target_dense.bias + ), + } + } + + restored_state = keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, partial_vars, epoch=1 + ) + + for key, var_dict in partial_vars.items(): + for path, var in var_dict.items(): + var.assign(restored_state[key][path]) + + np.testing.assert_allclose( + model_target.target_dense.kernel.value, test_kernel + ) + np.testing.assert_allclose( + model_target.target_dense.bias.value, test_bias + ) + manager.close() + + def test_extract_shared_variable_map_variable_level_sharing(self): + source_emb = keras.layers.Embedding(10, 4, name="text_emb") + source_emb.build((1,)) + + class Tower(keras.layers.Layer): + + def __init__(self, table, **kwargs): + super().__init__(**kwargs) + # Shares a single variable, not a layer, under a different attribute. + self.token_table = table + + def call(self, x): + return x + + model = keras.Sequential( + [source_emb, Tower(source_emb.embeddings, name="tower")], + name="my_model", + ) + model.build((1, 1)) + + shared_var_map = keras_utils.extract_shared_variable_map(model) + # The alias uses the variable's own name, not the attribute name. + self.assertEqual( + shared_var_map["my_model/tower/embeddings"], source_emb.embeddings.path + ) + self.assertNotIn("my_model/tower/token_table", shared_var_map) + + def test_extract_shared_variable_map_alias_collision_raises(self): + dense_a = keras.layers.Dense(2, name="dense_a") + dense_b = keras.layers.Dense(2, name="dense_b") + dense_a.build((1, 2)) + dense_b.build((1, 2)) + + class Tower(keras.layers.Layer): + + def __init__(self, kernel_a, kernel_b, **kwargs): + super().__init__(**kwargs) + self.kernel_a = kernel_a + self.kernel_b = kernel_b + + def call(self, x): + return x + + model = keras.Sequential( + [dense_a, dense_b, Tower(dense_a.kernel, dense_b.kernel, name="tower")], + name="my_model", + ) + model.build((1, 2)) + + with self.assertRaisesRegex(ValueError, "both map to alias path"): + keras_utils.extract_shared_variable_map(model) + + def test_map_prefix(self): + prefix_map = { + "model/new": "model/old", + "model/new/tower": "model/other/tower", + } + # Longest prefix wins. + self.assertEqual( + keras_utils._map_prefix("model/new/tower/dense/kernel", prefix_map), + ("model/new/tower", "model/other/tower/dense/kernel"), + ) + self.assertEqual( + keras_utils._map_prefix("model/new/head/kernel", prefix_map), + ("model/new", "model/old/head/kernel"), + ) + # Matches whole path segments only. + self.assertIsNone( + keras_utils._map_prefix("model/new_2/head/kernel", prefix_map) + ) + + def _save_multi_surface_dense_model(self): + """Saves a model whose surfaces share one Dense layer.""" + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + + class MultiSurfaceDenseModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.shared_dense = keras.layers.Dense(2, name="shared_dense") + self.chrome_surface = keras.Sequential( + [self.shared_dense], name="chrome_surface" + ) + self.search_surface = keras.Sequential( + [self.shared_dense], name="search_surface" + ) + + def build(self, input_shape): + self.shared_dense.build(input_shape) + self.chrome_surface.build(input_shape) + self.search_surface.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.chrome_surface(x) + self.search_surface(x) + + model = MultiSurfaceDenseModel(name="model") + model.compile(optimizer="adam") + model.build((1, 2)) + model.optimizer.build(model.trainable_variables) + model.shared_dense.kernel.assign(np.full((2, 2), 4.2, dtype=np.float32)) + model.shared_dense.bias.assign(np.full((2,), 1.8, dtype=np.float32)) + manager.save_model_variables(model, epoch=1) + manager.wait_until_finished() + manager.close() + return model + + def _build_new_surface_model(self): + """Builds a target model with an independent, differently named layer.""" + model = keras.Sequential( + [ + keras.Sequential( + [keras.layers.Dense(2, name="new_dense")], name="new_surface" + ) + ], + name="model", + ) + model.build((1, 2)) + return model + + def test_restore_partial_checkpoint_with_prefix_map(self): + source = self._save_multi_surface_dense_model() + shared_var_map = keras_utils.extract_shared_variable_map(source) + # The search surface's kernel is an alias; restoring through it also + # exercises shared variable resolution. + search_alias = next( + p + for p in shared_var_map + if "search_surface" in p and p.endswith("kernel") + ) + source_prefix = search_alias.removesuffix("/kernel") + + target = self._build_new_surface_model() + target_dense = target.layers[0].layers[0] + target_prefix = target_dense.kernel.path.removesuffix("/kernel") + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + target_dense.kernel.path: target_dense.kernel, + target_dense.bias.path: target_dense.bias, + } + } + + keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, + partial_vars, + epoch=1, + prefix_map={target_prefix: source_prefix}, + ) + + np.testing.assert_allclose(target_dense.kernel.value, np.full((2, 2), 4.2)) + np.testing.assert_allclose(target_dense.bias.value, np.full((2,), 1.8)) + + def test_restore_partial_checkpoint_prefix_map_missing_source_fails(self): + self._save_multi_surface_dense_model() + target = self._build_new_surface_model() + target_dense = target.layers[0].layers[0] + target_prefix = target_dense.kernel.path.removesuffix("/kernel") + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + target_dense.kernel.path: target_dense.kernel, + } + } + + with self.assertRaisesRegex(ValueError, "Missing paths in checkpoint"): + keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, + partial_vars, + epoch=1, + prefix_map={target_prefix: "model/no_such_surface"}, + ) + + def test_restore_partial_checkpoint_prefix_map_unused_prefix_warns(self): + self._save_multi_surface_dense_model() + target = self._build_new_surface_model() + target_dense = target.layers[0].layers[0] + target_prefix = target_dense.kernel.path.removesuffix("/kernel") + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + target_dense.kernel.path: target_dense.kernel, + target_dense.bias.path: target_dense.bias, + } + } + + with mock.patch.object(keras_utils.logging, "warning") as mock_warning: + keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, + partial_vars, + epoch=1, + prefix_map={ + target_prefix: "model/chrome_surface/shared_dense", + "model/typo_surface": "model/chrome_surface", + }, + ) + mock_warning.assert_called() + warning_found = False + for call in mock_warning.call_args_list: + for arg in call[0]: + if "matched no target variables" in str(arg): + warning_found = True + break + self.assertTrue(warning_found) + + np.testing.assert_allclose(target_dense.kernel.value, np.full((2, 2), 4.2)) + np.testing.assert_allclose(target_dense.bias.value, np.full((2,), 1.8)) + + def test_restore_partial_checkpoint_prefix_map_precedence(self): + manager = keras_utils.KerasOrbaxCheckpointManagerV3( + checkpoint_dir=self.checkpoint_dir, + max_to_keep=1, + save_interval_epochs=1, + ) + + class TwoSurfaceModel(keras.Model): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.surface_a = keras.Sequential( + [keras.layers.Dense(2, name="dense")], name="surface_a" + ) + self.surface_b = keras.Sequential( + [keras.layers.Dense(2, name="dense")], name="surface_b" + ) + + @property + def dense_a(self): + return self.surface_a.layers[0] + + @property + def dense_b(self): + return self.surface_b.layers[0] + + def build(self, input_shape): + self.surface_a.build(input_shape) + self.surface_b.build(input_shape) + super().build(input_shape) + + def call(self, x): + return self.surface_a(x) + self.surface_b(x) + + model = TwoSurfaceModel(name="model") + model.compile(optimizer="adam") + model.build((1, 2)) + model.optimizer.build(model.trainable_variables) + model.dense_a.kernel.assign(np.full((2, 2), 1.0, dtype=np.float32)) + model.dense_a.bias.assign(np.full((2,), 1.0, dtype=np.float32)) + model.dense_b.kernel.assign(np.full((2, 2), 9.0, dtype=np.float32)) + model.dense_b.bias.assign(np.full((2,), 9.0, dtype=np.float32)) + manager.save_model_variables(model, epoch=1) + manager.wait_until_finished() + manager.close() + + target = TwoSurfaceModel(name="model") + target.build((1, 2)) + target.dense_a.kernel.assign(np.zeros((2, 2), dtype=np.float32)) + target.dense_a.bias.assign(np.zeros((2,), dtype=np.float32)) + + partial_vars = { + keras_utils.TRAINABLE_VARIABLES_KEY: { + target.dense_a.kernel.path: target.dense_a.kernel, + target.dense_a.bias.path: target.dense_a.bias, + } + } + + # Remap surface_a -> surface_b via prefix_map, but override bias via + # transforms. + explicit_bias_transform = ocp.transform_utils.Transform( + original_key="trainable_variables/model/surface_a/dense/bias" + ) + + keras_utils.restore_partial_checkpoint( + self.checkpoint_dir, + partial_vars, + epoch=1, + prefix_map={"model/surface_a": "model/surface_b"}, + transforms={target.dense_a.bias.path: explicit_bias_transform}, + ) + + # Kernel was remapped to surface_b (value 9.0) even though surface_a + # exists in checkpoint (value 1.0). + np.testing.assert_allclose( + target.dense_a.kernel.value, np.full((2, 2), 9.0, dtype=np.float32) + ) + # Bias explicit transform took precedence over prefix_map, restoring + # surface_a (value 1.0). + np.testing.assert_allclose( + target.dense_a.bias.value, np.full((2,), 1.0, dtype=np.float32) + ) + if __name__ == "__main__": absltest.main()