jax-ml/jax · error · runtime_error

Invalid pipeline stage

Error message

Invalid pipeline stage

What it means

register_xla_transform(name, callback, stage) maps stage to a PipelineStage enum (0 = pre-scheduler, 1 = post-scheduler). Any integer other than 0 or 1 throws 'Invalid pipeline stage'.

Source

Thrown at jaxlib/xla.cc:187

}  // namespace

NB_MODULE(_xla, m) {
  // Register a transform directly via RegisterHloXlaTransform.
  // This is used for in-process backends (e.g. CPU).
  m.def(
      "register_xla_transform",
      [](std::string name, int stage, nb::object callback) {
        xla::HloXlaTransform::PipelineStage pipeline_stage;
        switch (stage) {
          case 0:
            pipeline_stage = xla::HloXlaTransform::PipelineStage::kPreScheduler;
            break;
          case 1:
            pipeline_stage =
                xla::HloXlaTransform::PipelineStage::kPostScheduler;
            break;
          default:
            throw std::runtime_error("Invalid pipeline stage");
        }

        auto transform = std::make_shared<PyHloXlaTransform>(
            std::move(name), std::move(callback));
        xla::RegisterHloXlaTransform(pipeline_stage, std::move(transform));
      },
      nb::arg("name"), nb::arg("stage"), nb::arg("callback"));

  m.def(
      "clear_xla_transform",
      [](std::string name, int stage) {
        xla::HloXlaTransform::PipelineStage pipeline_stage;
        switch (stage) {
          case 0:
            pipeline_stage = xla::HloXlaTransform::PipelineStage::kPreScheduler;
            break;
          case 1:
            pipeline_stage =

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Use stage 0 (kPreScheduler) or 1 (kPostScheduler) only
  2. Prefer higher-level JAX APIs for staging rather than raw xla_extension registration

Example fix

# before
register_xla_transform('my_pass', cb, stage=2)

# after
register_xla_transform('my_pass', cb, stage=1)  # 0=pre-scheduler, 1=post-scheduler
Defensive patterns

Strategy: validation

Validate before calling

if stage not in (0, 1):
    raise ValueError('stage must be 0 (pre-scheduler) or 1 (post-scheduler)')
register_xla_transform(name, cb, stage)

Type guard

def is_valid_stage(s) -> bool:
    return s in (0, 1)

Prevention

When it happens

Trigger: Calling jaxlib.xla_extension.register_xla_transform with stage outside {0,1}, e.g. 2, -1, or an unmarshalled enum value.

Common situations: Passing a string stage name instead of an int; using a stage constant from a different JAX version; plugin/backend tooling written against a changed enum.

Understand the failure class

Background: Invalid enum value errors: "Unknown type", "Invalid scope", "must be one of" — when a string is not on the library's allowed list — this error's family across 23 libraries.

Related errors


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