Skip to content

Latest commit

 

History

History
78 lines (57 loc) · 1.98 KB

File metadata and controls

78 lines (57 loc) · 1.98 KB

Development Notes

Auto-generated Code

Never manually edit files in src/jaxls/_py310/. These are auto-generated by transpile_py310.py. If the transpiler produces incorrect output, fix the transpiler script instead.

Commands

# Type checking.
uv run --extra dev --extra examples pyright ./src ./examples ./tests

# Linting and formatting.
uv run --extra dev --extra examples ruff check --fix .
uv run --extra dev --extra examples ruff format .

# Run tests.
uv run --extra dev pytest tests/

# Transpile for Python 3.10/3.11 compatibility.
uv run --extra dev ./transpile_py310.py

# Build documentation.
uv run --extra docs sphinx-build -b dirhtml docs_source/source docs

# Run benchmark on example notebooks.
uv run --extra dev --extra docs python benchmark.py
uv run --extra dev --extra docs python benchmark.py --category robotics  # Filter by category

# Benchmark + regression suite (BA / examples / pyroki / float32). See
# benchmarks/README.md. Run from the repo root.
uv run --extra dev --extra docs python -m benchmarks.suite              # full run + report vs committed baseline
uv run --extra dev --extra docs python -m benchmarks.suite --quick      # fast GPU-only inner loop
uv run --extra dev --extra docs python -m benchmarks.suite --gate       # exit 1 on regression (CI)

Style Guidelines

Use jdc.copy_and_mutate instead of jdc.replace

# Bad - **kwargs in replace() can't be type-checked.
new_obj = jdc.replace(obj, field=value)

# Good - preserves types.
with jdc.copy_and_mutate(obj) as new_obj:
    new_obj.field = value

Avoid truthy/falsey checks

# Bad.
if my_list:
if my_value:

# Good.
if len(my_list) > 0:
if my_value is not None:
if my_value != 0:

JIT compatibility

All code paths must work with traced arrays. Array values cannot affect control flow.

# Bad - breaks tracing.
if array_value > 0:
    return x
else:
    return y

# Good - use jax.lax.cond or jnp.where.
return jnp.where(array_value > 0, x, y)