from __future__ import annotations
import textwrap
from typing import Any, Callable, List, Optional, Tuple, Union
import mlx.core as mx
from mlx.utils import tree_flatten, tree_unflatten
class Module(dict):
__call__: Callable
def __init__(self):
self._no_grad = set()
self._training = True
@property
def training(self):
return self._training
@property
def state(self):
return self
def _extra_repr(self) -> str:
return ""
def __repr__(self):
children = tree_flatten(self.children(), is_leaf=self.is_module)
value = f"{type(self).__name__}({self._extra_repr()}"
for k, v in children:
value += "\n"
value += textwrap.indent(f"({k}): {repr(v)}", prefix=" ")
if children:
value += "\n"
value += ")"
return value
def __getattr__(self, key: str):
if (value := self.get(key, None)) is not None:
return value
else:
super(Module, self).__getattribute__(key)
def __setattr__(self, key: str, val: Any):
if isinstance(val, (mx.array, dict, list, tuple)):
if hasattr(self, key) and key not in self:
delattr(self, key)
self[key] = val
else:
super(Module, self).__setattr__(key, val)
self.pop(key, None)
def __delattr__(self, name):
if (val := self.get(name, None)) is not None:
del self[name]
else:
super().__delattr__(name)
def load_weights(
self,
file_or_weights: Union[str, List[Tuple[str, mx.array]]],
strict: bool = True,
) -> Module:
weights = file_or_weights
if isinstance(weights, str):
weights = list(mx.load(weights).items())
if strict:
new_weights = dict(weights)
curr_weights = tree_flatten(self.parameters(), destination={})
if extras := (new_weights.keys() - curr_weights.keys()):
num_extra = len(extras)
extras = ",\n".join(sorted(extras))
raise ValueError(
f"Received {num_extra} parameters not in model: \n{extras}."
)
if missing := (curr_weights.keys() - new_weights.keys()):
num_missing = len(missing)
missing = ",\n".join(sorted(missing))
raise ValueError(f"Missing {num_missing} parameters: \n{missing}.")
for k, v in curr_weights.items():
v_new = new_weights[k]
if not isinstance(v_new, mx.array):
raise ValueError(
"Expected mx.array but received "
f"{type(v_new)} for parameter {k}"
)
if v_new.shape != v.shape:
raise ValueError(
f"Expected shape {v.shape} but received "
f"shape {v_new.shape} for parameter {k}"
)
if len(weights) != 0:
self.update(tree_unflatten(weights), strict=False)
return self
def save_weights(self, file: str):
params_dict = tree_flatten(self.parameters(), destination={})
if file.endswith(".npz"):
mx.savez(file, **params_dict)
elif file.endswith(".safetensors"):
mx.save_safetensors(file, params_dict)
else:
raise ValueError(
f"Unsupported file extension for {file}. Use '.npz' or '.safetensors'."
)
@staticmethod
def is_module(value):
return isinstance(value, Module)
@staticmethod
def valid_child_filter(module, key, value):
return isinstance(value, (dict, list))
@staticmethod
def valid_parameter_filter(module, key, value):
return isinstance(value, (dict, list, mx.array)) and not key.startswith("_")
@staticmethod
def trainable_parameter_filter(module, key, value):
return (
Module.valid_parameter_filter(module, key, value)
and key not in module._no_grad
)
def filter_and_map(
self,
filter_fn: Callable[[Module, str, Any], bool],
map_fn: Optional[Callable] = None,
is_leaf_fn: Optional[Callable[[Module, str, Any], bool]] = None,
):
map_fn = map_fn or (lambda x: x)
is_leaf_fn = is_leaf_fn or (
lambda m, k, v: not isinstance(v, (Module, dict, list))
)
return {
k: _unwrap(self, k, v, filter_fn, map_fn, is_leaf_fn)
for k, v in self.items()
if filter_fn(self, k, v)
}
def parameters(self):
return self.filter_and_map(self.valid_parameter_filter)
def trainable_parameters(self):
return self.filter_and_map(self.trainable_parameter_filter)
def children(self):
return self.filter_and_map(
self.valid_child_filter, is_leaf_fn=lambda m, k, v: isinstance(v, Module)
)
def leaf_modules(self):
def _is_leaf_module(m, k, v):
return isinstance(v, Module) and len(tree_flatten(v.children())) == 0
return self.filter_and_map(self.valid_child_filter, is_leaf_fn=_is_leaf_module)
def update(self, parameters: dict, strict: bool = True) -> Module:
def apply(dst, parameters):
if isinstance(parameters, dict):
for k in parameters:
if k in dst:
current_value = dst[k]
new_value = parameters[k]
if isinstance(current_value, mx.array):
if strict and not isinstance(new_value, mx.array):
raise ValueError(
f"Received invalid type: {type(new_value).__name__}."
)
dst[k] = new_value
else:
apply(current_value, new_value)
elif strict:
raise ValueError(f'Module does not have parameter named "{k}".')
elif isinstance(parameters, list):
for i in range(len(parameters)):
if i >= len(dst):
if strict:
raise ValueError(
f"List index {i} is out of bounds for "
f"destination of length {len(dst)}."
)
continue
current_value = dst[i]
new_value = parameters[i]
if isinstance(current_value, mx.array):
if strict and not isinstance(new_value, mx.array):
raise ValueError(
f"Received invalid type: {type(new_value).__name__}."
)
dst[i] = new_value
else:
apply(current_value, new_value)
elif strict:
raise ValueError(f"Received invalid type: {type(parameters).__name__}.")
apply(self, parameters)
return self
def apply(
self,
map_fn: Callable[[mx.array], mx.array],
filter_fn: Optional[Callable[[Module, str, Any], bool]] = None,
) -> Module:
filter_fn = filter_fn or Module.valid_parameter_filter
self.update(self.filter_and_map(filter_fn, map_fn))
return self
def update_modules(self, modules: dict, strict: bool = True) -> Module:
_update_modules(self, modules, strict)
return self
def apply_to_modules(self, apply_fn: Callable[[str, Module], Any]) -> Module:
module_stack = [("", self)]
while module_stack:
prefix, mod = module_stack.pop()
apply_fn(prefix, mod)
prefix = "." + prefix if prefix else ""
module_stack.extend(
tree_flatten(mod.children(), prefix=prefix, is_leaf=self.is_module)
)
return self
def modules(self):
modulelist = []
self.apply_to_modules(lambda k, m: modulelist.append(m))
return modulelist
def named_modules(self):
modulelist = []
self.apply_to_modules(lambda k, m: modulelist.append((k, m)))
return modulelist
def _validate_keys(self, keys, strict):
keys = keys if isinstance(keys, list) else [keys]
if strict:
for k in keys:
if k not in self:
raise KeyError(f"Module doesn't contain member {k}.")
return keys
def freeze(
self,
*,
recurse: bool = True,
keys: Optional[Union[str, List[str]]] = None,
strict: bool = False,
) -> Module:
def _freeze_impl(_, m):
local_keys = keys
if local_keys is None:
local_keys = tree_flatten(
m.filter_and_map(
lambda m, k, v: (not isinstance(v, Module))
and m.valid_parameter_filter(m, k, v)
)
)
local_keys = [k for (k, v) in local_keys]
local_keys = m._validate_keys(local_keys, strict)
m._no_grad.update(local_keys)
if recurse:
self.apply_to_modules(_freeze_impl)
else:
_freeze_impl("", self)
return self
def unfreeze(
self,
*,
recurse: bool = True,
keys: Optional[Union[str, List[str]]] = None,
strict: bool = False,
) -> Module:
def _unfreeze_impl(_, m):
if keys is None:
m._no_grad.clear()
else:
local_keys = m._validate_keys(keys, strict)
m._no_grad.difference_update(local_keys)
if recurse:
self.apply_to_modules(_unfreeze_impl)
else:
_unfreeze_impl("", self)
return self
def _set_training_mode(self, mode: bool) -> None:
self._training = mode
def train(self, mode: bool = True) -> Module:
self.apply_to_modules(lambda _, m: m._set_training_mode(mode))
return self
def eval(self) -> Module:
return self.train(False)
def set_dtype(
self,
dtype: mx.Dtype,
predicate: Optional[Callable[[mx.Dtype], bool]] = lambda x: mx.issubdtype(
x, mx.floating
),
):
if predicate is None:
predicate = lambda _: True
self.apply(lambda x: x.astype(dtype) if predicate(x.dtype) else x)
def _update_modules(dst, modules, strict):
if isinstance(modules, dict):
for k in modules:
if k in dst:
current_value = dst[k]
new_value = modules[k]
if Module.is_module(current_value) and Module.is_module(new_value):
dst[k] = new_value
elif isinstance(current_value, (dict, list)):
_update_modules(current_value, new_value, strict)
elif strict and new_value != {}:
raise ValueError(
f"Received invalid type: {type(new_value).__name__}."
)
elif strict:
raise ValueError(f'Module does not have sub-module named "{k}".')
elif isinstance(modules, list):
for i in range(len(modules)):
current_value = dst[i]
new_value = modules[i]
if Module.is_module(current_value) and Module.is_module(new_value):
dst[i] = new_value
elif isinstance(current_value, (dict, list)):
_update_modules(current_value, new_value, strict)
elif strict and new_value != {}:
raise ValueError(f"Received invalid type: {type(new_value).__name__}.")
elif strict:
raise ValueError(f"Received invalid type: {type(modules).__name__}.")
def _unwrap(model, value_key, value, filter_fn, map_fn, is_leaf_fn):
if is_leaf_fn(model, value_key, value):
return map_fn(value)
elif isinstance(value, Module):
return {
k: _unwrap(value, k, v, filter_fn, map_fn, is_leaf_fn)
for k, v in value.items()
if filter_fn(value, k, v)
}
elif isinstance(value, dict):
nd = {}
for k, v in value.items():
tk = f"{value_key}.{k}"
nd[k] = (
_unwrap(model, tk, v, filter_fn, map_fn, is_leaf_fn)
if filter_fn(model, tk, v)
else {}
)
return nd
elif isinstance(value, list):
nl = []
for i, vi in enumerate(value):
tk = f"{value_key}.{i}"
nl.append(
_unwrap(model, tk, vi, filter_fn, map_fn, is_leaf_fn)
if filter_fn(model, tk, vi)
else {}
)
return nl
raise RuntimeError("Unexpected leaf found while traversing the module")