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
- Use stage 0 (kPreScheduler) or 1 (kPostScheduler) only
- 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
- Define named constants PRE_SCHEDULER=0, POST_SCHEDULER=1 instead of magic numbers
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
- Argument to register_custom_call_partitioner was not a pjrt_
- {full_name} must be a pytree prefix with bool leaves or a tu
- multi-platform lowering for buffer_callback
- Axes mentioned in `manual_axis_type` field of ShapedArray sh
- varying and unreduced cannot have common mesh axes. Got vary
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/60c64382c7af521f.
Report an issue: GitHub.