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

Method fans

axlearn/common/base_layer.py:208–227  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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.

Callers 2

test_fanMethod · 0.80

Calls

no outgoing calls

Tested by 2

test_fanMethod · 0.64