diff --git a/zeromodels/base/base_config.py b/zeromodels/base/base_config.py index efbb9437..c3d66579 100644 --- a/zeromodels/base/base_config.py +++ b/zeromodels/base/base_config.py @@ -78,14 +78,45 @@ def field_names(cls): def unknown_keys(cls, data): """Keys in a repo spec ``data`` this config does not recognize. - A typed config's ``from_dict`` keeps only annotated fields, so a key a newer - zeromodels added (e.g. ``rope_scaling_factor``, nested under ``text_config``) - is dropped. This finds those by round-trip: a key present in ``data`` but - absent from ``to_dict(from_dict(data))`` was not consumed. It recurses into - the nested sub-config dicts the repo format wraps fields in, and never flags - a legitimate key (a recognized key always round-trips back). + Match the input routes used by ``from_dict``, not its serialized output: + serialization relocates flat fields and omits absent optional towers and + some glue fields. Only declared config blocks are traversed; dictionary + values of ordinary fields (e.g. ``rope_scaling``) are opaque here. This + checks field names, not value types. """ - known = cls.from_dict(data).to_dict() + + def constructor_schema(config_cls): + schema = dict.fromkeys(config_cls.field_names()) + for key, sub_cls in config_cls.sub_configs.items(): + schema[key] = constructor_schema(sub_cls) + return schema + + known = dict.fromkeys(cls.field_names()) + if cls.sub_configs: + if any(key in data for key in cls.sub_configs): + known = constructor_schema(cls) + else: + for key, sub_cls in cls.sub_configs.items(): + for field in sub_cls.field_names(): + known[cls._sub_flat_name(key, field)] = None + else: + main_key = cls._main_key() + if main_key in data or "text_config" in data or "vision_config" in data: + fields = set(cls.field_names()) + if main_key not in cls.config_groups: + known[main_key] = dict.fromkeys(fields) + for key, prefix in cls.config_groups.items(): + extras = set(cls.group_extras.get(key, ())) + members = { + field[len(prefix) :] + for field in fields + if field.startswith(prefix) + and field[len(prefix) :] not in extras + } + members.update(extras & fields) + known[key] = dict.fromkeys(members) + if cls.model_type is not None: + known["model_type"] = None def diff(spec, seen): out = set() diff --git a/zeromodels/models/sam2/sam2_layers.py b/zeromodels/models/sam2/sam2_layers.py index a4735f51..d798310b 100644 --- a/zeromodels/models/sam2/sam2_layers.py +++ b/zeromodels/models/sam2/sam2_layers.py @@ -1590,7 +1590,9 @@ def _recompute_full_pos(self): dtype = "float64" if self.pos_embed.dtype == "float64" else "float32" matrices = [] for out_size, in_size in ((h, self.bg_size[0]), (w, self.bg_size[1])): - center = (ops.arange(out_size, dtype=dtype) + 0.5) * (in_size / out_size) - 0.5 + center = (ops.arange(out_size, dtype=dtype) + 0.5) * ( + in_size / out_size + ) - 0.5 start = ops.floor(center) fraction = center - start a = -0.75 @@ -1605,14 +1607,18 @@ def _recompute_full_pos(self): for offset, tap in zip((-1, 0, 1, 2), taps): index = ops.clip(ops.cast(start, "int32") + offset, 0, in_size - 1) # Accumulate duplicate taps when border indices are clamped. - matrix = matrix + ops.one_hot(index, in_size, dtype=dtype) * tap[:, None] + matrix = ( + matrix + ops.one_hot(index, in_size, dtype=dtype) * tap[:, None] + ) matrices.append(matrix) pos = ops.cast(self.pos_embed, dtype) pos = ops.einsum("wj,bijc->biwc", matrices[1], pos) pos = ops.einsum("hi,biwc->bhwc", matrices[0], pos) tile_h = h // self.window_size tile_w = w // self.window_size - window_pos = ops.tile(ops.cast(self.pos_embed_window, dtype), (1, tile_h, tile_w, 1)) + window_pos = ops.tile( + ops.cast(self.pos_embed_window, dtype), (1, tile_h, tile_w, 1) + ) self._full_pos.assign(pos + window_pos) def load_own_variables(self, store):