Returns a dictionary with keys 'fan_in', 'fan_out', and 'fan_avg' containing the fan values for this parameter. The calculation is consistent with jax's initializers: Indices without an explicit axis type specified are treated as both in and out axes. Batch axes are
(self)
| 206 | weight_decay_scale: Optional[float] = None |
| 207 | |
| 208 | def fans(self) -> dict[str, float]: |
| 209 | """Returns a dictionary with keys 'fan_in', 'fan_out', and 'fan_avg' containing |
| 210 | the fan values for this parameter. |
| 211 | |
| 212 | The calculation is consistent with jax's initializers: Indices without |
| 213 | an explicit axis type specified are treated as both in and out axes. |
| 214 | Batch axes are ignored. |
| 215 | """ |
| 216 | sizes = {} |
| 217 | for axis_type in self.fan_axes._fields: # pylint: disable=protected-access |
| 218 | axes = getattr(self.fan_axes, axis_type) |
| 219 | if isinstance(axes, int): |
| 220 | axes = [axes] |
| 221 | sizes[axis_type] = math.prod(self.shape[axis] for axis in axes) |
| 222 | unbatched_size = math.prod(self.shape) / sizes["batch_axis"] |
| 223 | result = dict( |
| 224 | fan_in=unbatched_size / sizes["out_axis"], fan_out=unbatched_size / sizes["in_axis"] |
| 225 | ) |
| 226 | result["fan_avg"] = (result["fan_in"] + result["fan_out"]) / 2 |
| 227 | return result |
| 228 | |
| 229 | |
| 230 | # Legacy type alias. For new code, use Nested[ParameterSpec] from axlearn.common.utils. |
no outgoing calls