jax-ml/jax · error · ValueError

Invalid value "{default_str}" for JAX flag {name}

Error message

Invalid value "{default_str}" for JAX flag {name}

What it means

enum_class_state reads JAX_<FLAG> and tries enum_class(default_str). If the env var string does not map to any Enum value, it raises this ValueError (chained from the Enum's own ValueError) at import/definition time.

Source

Thrown at jax/_src/config.py:642

      the trace context.
    extra_validator: optional function to validate the value of the config
      option.

  Returns:
    A contextmanager to control the thread-local state value.
  """
  if not isinstance(default, enum_class):
    raise TypeError(
        f'Default value must be of type {enum_class}, got {default} '
        f"of type {getattr(type(default), '__name__', type(default))}"
    )
  name = name.lower()
  default_str = os.getenv(name.upper(), None)
  if default_str is not None:
    try:
      default = enum_class(default_str)
    except ValueError as e:
      raise ValueError(f"Invalid value \"{default_str}\" for JAX flag {name}") from e
  config._contextmanager_flags.add(name)

  def parser(new_val):
    if isinstance(new_val, str):
      return enum_class(new_val)
    if not isinstance(new_val, enum_class):
      raise TypeError(
          f'new enum value must be an instance of {enum_class}, got'
          f' {new_val} of type {type(new_val)}.'
      )
    if extra_validator is not None:
      extra_validator(new_val)
    return new_val

  s = State[EnumType](
      name,
      default,
      help,

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. List valid members: python -c 'from x import MyEnum; print([e.value for e in MyEnum])' and set the env var to one of them
  2. Unset the env var to fall back to the code default
  3. After jax upgrades, re-check flag values documented in the release notes

Example fix

# before
export JAX_MY_FLAG=old_value
# after
export JAX_MY_FLAG=new_value  # matches enum_class member
Defensive patterns

Strategy: validation

Validate before calling

import os
raw = os.getenv('JAX_MY_FLAG')
if raw is not None:
    enum_class(raw)  # raises early with a clear Enum error if invalid

Prevention

When it happens

Trigger: Exporting JAX_MY_FLAG=old_name after an enum value was renamed in a jax upgrade (e.g. profiler or dumping option renames).

Common situations: Upgrading jax where an enum-backed flag's member names changed, leaving stale env vars in shell profiles, Dockerfiles, or CI secrets.

Related errors


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/ccb7a5cb067a1603. Report an issue: GitHub.