(
spec: Optional[FactorizationSpec],
)
| 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( |
nothing calls this directly
no test coverage detected