{"record":{"id":"5b75c194588e49be","repo":"jax-ml/jax","slug":"invalid-mode-for-variance-scaling-initializer-mo","errorCode":null,"errorMessage":"invalid mode for variance scaling initializer: {mode}","messagePattern":"invalid mode for variance scaling initializer: (.+?)","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/nn/initializers.py","lineNumber":346,"sourceCode":"    out_axis: axis or sequence of axes of the output dimension in the weights\n      array.\n    batch_axis: axis or sequence of axes in the weight array that should be\n      ignored.\n    dtype: the dtype of the weights.\n  \"\"\"\n  def init(key: Array,\n           shape: core.Shape,\n           dtype: DTypeLikeInexact | None = dtype,\n           out_sharding: OutShardingType = None) -> Array:\n    shape = core.canonicalize_shape(shape)\n    dtype = dtypes.default_float_dtype() if dtype is None else dtype\n    fan_in, fan_out = _compute_fans(shape, in_axis, out_axis, batch_axis)\n    if mode == \"fan_in\": denominator = fan_in\n    elif mode == \"fan_out\": denominator = fan_out\n    elif mode == \"fan_avg\": denominator = (fan_in + fan_out) / 2\n    elif mode == \"fan_geo_avg\": denominator = (fan_in * fan_out) ** 0.5\n    else:\n      raise ValueError(\n        f\"invalid mode for variance scaling initializer: {mode}\")\n    variance = jnp.array(scale / denominator, dtype=dtype)\n\n    if distribution == \"truncated_normal\":\n      if dtypes.issubdtype(dtype, np.floating):\n        # constant is stddev of standard normal truncated to (-2, 2)\n        stddev = jnp.sqrt(variance) / jnp.array(.87962566103423978, dtype)\n        return random.truncated_normal(key, -2, 2, shape, dtype,\n                                       out_sharding=out_sharding) * stddev\n      else:\n        # constant is stddev of complex standard normal truncated to 2\n        stddev = jnp.sqrt(variance) / jnp.array(.95311164380491208, dtype)\n        return _complex_truncated_normal(key, 2, shape, dtype) * stddev\n    elif distribution == \"normal\":\n      return random.normal(key, shape, dtype,\n                           out_sharding=out_sharding) * jnp.sqrt(variance)\n    elif distribution == \"uniform\":\n      if dtypes.issubdtype(dtype, np.floating):","sourceCodeStart":328,"sourceCodeEnd":364,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/nn/initializers.py#L328-L364","documentation":"jax.nn.initializers.variance_scaling selects the denominator via the `mode` string: 'fan_in', 'fan_out', 'fan_avg', 'fan_geo_avg'. Any other string raises this ValueError.","triggerScenarios":"Calling variance_scaling(scale, mode='FAN_IN') (case-sensitive), mode='avg', mode=None, or a typo.","commonSituations":"Porting configs from TF/PyTorch where mode names or casing differ; hand-written config strings from YAML/JSON typos.","solutions":["Use exactly one of: 'fan_in', 'fan_out', 'fan_avg', 'fan_geo_avg' (lowercase)","For Kaiming-style init use 'fan_in'; for Xavier-style balance use 'fan_avg' or 'fan_geo_avg'"],"exampleFix":"// before\ninit = jax.nn.initializers.variance_scaling(2.0, 'fan-in', 'truncated_normal')\n\n// after\ninit = jax.nn.initializers.variance_scaling(2.0, 'fan_in', 'truncated_normal')","handlingStrategy":"validation","validationCode":"MODES = ('fan_in', 'fan_out', 'fan_avg', 'fan_geo_avg')\nassert mode in MODES, f'mode must be one of {MODES}'","typeGuard":"def is_valid_mode(m: str) -> bool: return m in ('fan_in','fan_out','fan_avg','fan_geo_avg')","tryCatchPattern":null,"preventionTips":["Use Literal['fan_in','fan_out','fan_avg','fan_geo_avg'] in config dataclasses","Validate config strings at load time, not at init time"],"tags":["jax","nn","initializers","enum-argument"],"backgroundTag":"invalid-enum-argument","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}