Skip to content

Fix BlockDiag products under JIT and with rectangular blocks - #806

Open
AHMETHAKANBEZIR1 wants to merge 1 commit into
QuantClimate:mainfrom
AHMETHAKANBEZIR1:fix/blockdiag-static-input-splits
Open

AHMETHAKANBEZIR1 wants to merge 1 commit into
QuantClimate:mainfrom
AHMETHAKANBEZIR1:fix/blockdiag-static-input-splits

Conversation

@AHMETHAKANBEZIR1

Copy link
Copy Markdown
Contributor

Description

BlockDiag.mv creates split indices with JAX array operations. Under jax.jit, these indices become tracers and jnp.split raises ConcretizationTypeError. It also uses output dimensions to split the input vector, which fails for rectangular blocks.

Use input dimensions and Python cumulative sums to keep split indices static. The operator still applies each block separately. Add a release note and tests against dense block-diagonal products and gradients with respect to the blocks and input.

Validation

  • New regressions on the base: 10 failed, 6 passed.
  • Linear algebra and KL tests: 75 passed.
  • Full uv run --no-sync poe test, with PYTHONUTF8=1: 3,245 passed, 1 skipped.
  • uv run --no-sync poe format and poe lint: passed.
  • Square and rectangular blocks, one and multiple blocks, float32/float64, eager and JIT execution, and gradients.
  • Python 3.13.15 / JAX 0.11.2 / CPU, existing development environment.

GPU and the full documentation build were not run locally.

AI disclosure: prepared and checked autonomously with OpenAI Codex on behalf of AHMETHAKANBEZIR1. No human review claim is made.

Co-authored-by: OpenAI Codex <noreply@openai.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant