{"record":{"id":"3511cdd73b463652","repo":"jax-ml/jax","slug":"invalid-distribution-for-variance-scaling-initiali","errorCode":null,"errorMessage":"invalid distribution for variance scaling initializer: {distribution}","messagePattern":"invalid distribution for variance scaling initializer: (.+?)","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/nn/initializers.py","lineNumber":370,"sourceCode":"        # 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):\n        return random.uniform(key, shape, dtype, -1,\n                              out_sharding=out_sharding) * jnp.sqrt(3 * variance)\n      else:\n        return _complex_uniform(key, shape, dtype) * jnp.sqrt(variance)\n    else:\n      raise ValueError(f\"invalid distribution for variance scaling initializer: {distribution}\")\n\n  return init\n\n@export\ndef glorot_uniform(in_axis: int | Sequence[int] = -2,\n                   out_axis: int | Sequence[int] = -1,\n                   batch_axis: int | Sequence[int] = (),\n                   dtype: DTypeLikeInexact | None = None) -> Initializer:\n  \"\"\"Builds a Glorot uniform initializer (aka Xavier uniform initializer).\n\n  A `Glorot uniform initializer`_ is a specialization of\n  :func:`jax.nn.initializers.variance_scaling` where ``scale = 1.0``,\n  ``mode=\"fan_avg\"``, and ``distribution=\"uniform\"``.\n\n  Args:\n    in_axis: axis or sequence of axes of the input dimension in the weights\n      array.\n    out_axis: axis or sequence of axes of the output dimension in the weights","sourceCodeStart":352,"sourceCodeEnd":388,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/nn/initializers.py#L352-L388","documentation":"variance_scaling's second enum argument, `distribution`, supports 'truncated_normal', 'normal', and 'uniform'. Any other string raises this ValueError at init call time.","triggerScenarios":"Passing distribution='tn', 'gaussian', distribution=None, or 'truncated-normal' (hyphen typo).","commonSituations":"Config-driven initializer construction with typo'd strings; porting from libraries that use 'trunc_normal' naming.","solutions":["Use exactly 'truncated_normal', 'normal', or 'uniform'","If you need a different sampler, wrap the initializer and post-transform the samples"],"exampleFix":"// before\ninit = jax.nn.initializers.variance_scaling(1.0, 'fan_in', 'trunc_normal')\n\n// after\ninit = jax.nn.initializers.variance_scaling(1.0, 'fan_in', 'truncated_normal')","handlingStrategy":"validation","validationCode":"DISTS = ('truncated_normal', 'normal', 'uniform')\nassert distribution in DISTS, f'distribution must be one of {DISTS}'","typeGuard":"def is_valid_distribution(d: str) -> bool: return d in ('truncated_normal','normal','uniform')","tryCatchPattern":null,"preventionTips":["Use Literal typing for initializer enum args","Validate initializer configs once at parse 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"}