Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 6 additions & 4 deletions flax/nnx/spmd.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,11 +47,11 @@ def _add_axis(x: tp.Any):
metadata = x.get_metadata()
if 'sharding_names' in metadata and metadata['sharding_names']:
sharding = metadata['sharding_names']
x.sharding_names = insert_field(sharding, index, axis_name)
x.set_metadata(sharding_names=insert_field(sharding, index, axis_name))

for k, v in other_meta.items():
if hasattr(x, k) and (t := getattr(x, k)) and isinstance(t, tuple):
setattr(x, k, insert_field(t, index, v))
x.set_metadata(k, insert_field(t, index, v))

assert isinstance(x, variablelib.Variable)
x.add_axis(index, axis_name)
Expand All @@ -75,11 +75,13 @@ def remove_field(fields, index, value):
def _remove_axis(x: tp.Any):
if isinstance(x, variablelib.Variable):
if hasattr(x, 'sharding_names') and x.sharding_names is not None:
x.sharding_names = remove_field(x.sharding_names, index, axis_name)
x.set_metadata(
sharding_names=remove_field(x.sharding_names, index, axis_name)
)

for k, v in other_meta.items():
if hasattr(x, k) and (t := getattr(x, k)) and isinstance(t, tuple):
setattr(x, k, remove_field(t, index, v))
x.set_metadata(k, remove_field(t, index, v))

x.remove_axis(index, axis_name)
return x
Expand Down
43 changes: 32 additions & 11 deletions flax/nnx/variablelib.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
from flax import errors
from flax.core import spmd as core_spmd
from flax.nnx import filterlib, reprlib, tracers, visualization
from flax.typing import Missing, PathParts, SizeBytes
from flax.typing import MISSING, Missing, PathParts, SizeBytes
import jax.tree_util as jtu
import jax.numpy as jnp
from jax._src.state.types import AbstractRef
Expand Down Expand Up @@ -172,7 +172,14 @@ class VariableMetadata(tp.Generic[A]):
metadata: tp.Mapping[str, tp.Any] = dataclasses.field(default_factory=dict)


class Variable(tp.Generic[A], reprlib.Representable):
class VariableMeta(type):
def __new__(cls, cls_name, bases, attrs):
if '__slots__' not in attrs:
attrs['__slots__'] = ()
return super().__new__(cls, cls_name, bases, attrs)


class Variable(tp.Generic[A], reprlib.Representable, metaclass=VariableMeta):
"""The base class for all ``Variable`` types. Create custom ``Variable``
types by subclassing this class. Numerous NNX graph functions can filter
for specific ``Variable`` types, for example, :func:`split`, :func:`state`,
Expand Down Expand Up @@ -353,47 +360,61 @@ def has_ref(self) -> bool:
@tp.overload
def get_metadata(self) -> dict[str, tp.Any]: ...
@tp.overload
def get_metadata(self, name: str) -> tp.Any: ...
def get_metadata(self, name: str | None = None):
def get_metadata(self, name: str, default: tp.Any = MISSING) -> tp.Any: ...
def get_metadata(
self, name: str | None = None, default: tp.Any = MISSING
) -> tp.Any:
"""Get metadata for the Variable.

Args:
name: The key of the metadata element to get. If not provided, returns
the full metadata dictionary.
default: The default value to return if the metadata key is not found. If
not provided and the key is not found, raises a KeyError.
"""
if name is None:
return self._var_metadata
if name not in self._var_metadata and not isinstance(default, Missing):
return default
return self._var_metadata[name]

@tp.overload
def set_metadata(self, metadata: dict[str, tp.Any], /) -> None: ...
@tp.overload
def set_metadata(self, name: str, value: tp.Any, /) -> None: ...
@tp.overload
def set_metadata(self, **metadata: tp.Any) -> None: ...
def set_metadata(self, *args, **kwargs) -> None:
"""Set metadata for the Variable.

`set_metadata` can be called in two ways:
`set_metadata` can be called in 3 ways:

1. By passing a dictionary of metadata as the first argument, this will replace
the entire Variable's metadata.
2. By using keyword arguments, these will be merged into the existing Variable's
metadata.
2. By passing a name and value as the first two arguments, this will set
the metadata entry for the given name to the given value.
3. By using keyword arguments, this will update the Variable's metadata
with the provided key-value pairs.
"""
if not self._trace_state.is_valid():
raise errors.TraceContextError(
f'Cannot mutate {type(self).__name__} from a different trace level'
)
if not (bool(args) ^ bool(kwargs)):
if args and kwargs:
raise TypeError(
'set_metadata takes either a single dict argument or keyword arguments'
'Cannot mix positional and keyword arguments in set_metadata'
)
if len(args) == 1:
self._var_metadata = args[0]
self._var_metadata = dict(args[0])
elif len(args) == 2:
name, value = args
self._var_metadata[name] = value
elif kwargs:
self._var_metadata.update(kwargs)
else:
raise TypeError(
f'set_metadata takes either 1 argument or 1 or more keyword arguments, got args={args}, kwargs={kwargs}'
f'set_metadata takes either 1 or 2 arguments, or at least 1 keyword argument, '
f'got args={args}, kwargs={kwargs}'
)

def copy_from(self, other: Variable[A]) -> None:
Expand Down
4 changes: 4 additions & 0 deletions tests/nnx/variable_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,10 @@ def test_get_set_metadata(self):
self.assertEqual(v.get_metadata(), {'b': 3, 'c': 4})
self.assertEqual(v.get_metadata('b'), 3)
self.assertEqual(v.get_metadata('c'), 4)
c = v.get_metadata('c')
self.assertEqual(c, 4)
x = v.get_metadata('x', default=10)
self.assertEqual(x, 10)


if __name__ == '__main__':
Expand Down
Loading