MCPcopy Create free account
hub / github.com/apple/axlearn / transform_factorization_spec

Method transform_factorization_spec

axlearn/common/pipeline.py:692–697  ·  view source on GitHub ↗
(
            spec: Optional[FactorizationSpec],
        )

Source from the content-addressed store, hash-verified

690 specs = VDict(**super().create_parameter_specs_recursively())
691
692 def transform_factorization_spec(
693 spec: Optional[FactorizationSpec],
694 ) -> Optional[FactorizationSpec]:
695 if spec is None:
696 return None
697 return FactorizationSpec(axes=[None] + list(spec.axes))
698
699 return jax.tree.map(
700 lambda spec: dataclasses.replace(

Callers

nothing calls this directly

Calls 1

FactorizationSpecClass · 0.90

Tested by

no test coverage detected