Skip to content

jax backend and repo reorganize - #6

Merged
Vaibhavdixit02 merged 14 commits into
mainfrom
cursor/jax-backend-and-repo-reorganize
Apr 22, 2026
Merged

Vaibhavdixit02 merged 14 commits into
mainfrom
cursor/jax-backend-and-repo-reorganize

Conversation

@Vaibhavdixit02

Copy link
Copy Markdown
Owner

No description provided.

Vaibhavdixit02 and others added 14 commits April 14, 2026 20:15
…, move notebook

Symmetric directory structure for multi-backend support (torch/ and jax/).
Renames PyTorch modules into nanoasr/torch/ subpackage, prefixes test
files with test_torch_*, and moves the training notebook into notebooks/.

Made-with: Cursor
- Add nanoasr/jax/ subpackage: Conformer in Flax NNX, optax training
  loop, librosa mel, soundfile data loading — zero torch dependency
- Update all internal imports for torch/ subpackage move
- Add shared clean_text to vocab.py, remove duplication
- Add pyproject.toml [jax] optional deps and fix entry points
- Add train_jax.ipynb Colab TPU notebook
- Add test_jax_model.py (13 tests)

Made-with: Cursor
Colab pre-installs Flax 0.11.x which lacks nnx.List. Upgrade
the floor to 0.12 in pyproject.toml and force-upgrade in the
notebook install cell so the Conformer model builds correctly.

Made-with: Cursor
The prior fix only added wrt to optimizer.update(); the constructor
requires it too in Flax 0.11+.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
nnx.split(model) returns state that includes Rngs. In JAX's typed-PRNG
regime, those leaves cannot be converted to numpy via np.array() and
raise TypeError. Extract key_data on save and wrap_key_data on load so
the full state round-trips.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
- train() gains a max_steps arg that breaks the loop early; lets a CPU
  smoke run exercise the real entry point in seconds instead of minutes.
- tests/test_jax_smoke.py calls train() on dev-clean with depth=2,
  batch_size=2, max_steps=3 and asserts a checkpoint is written.

Run locally with:
    JAX_PLATFORMS=cpu pytest tests/test_jax_smoke.py -s

This run caught the PRNG-key checkpoint bug fixed in the previous commit.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
Each unique (mel_T, target_S) pair triggered a fresh XLA compile of the
Conformer train_step, and each compile peaks 15-20 GB of host RAM
during lowering. Two back-to-back batches with different shapes
reliably pushed Colab TPU hosts past 48 GB and OOM-killed the kernel.

compute_dataset_maxes() derives a single (max_mel_T, max_target_S) per
dataset (99th percentile of audio length, exact max of encoded target
length). make_loader now accepts pad_to= and max_audio_samples= so
every batch has identical shape; train_step compiles once for the
whole run. Outlier clips longer than the 99th percentile are dropped
to keep the pad ceiling reasonable.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
Without drop_last, the last batch of each epoch has fewer items than
batch_size, which is a second shape distinct from every full batch.
That's enough to trigger a second JIT compile of train_step, blowing
past host RAM on Colab even after the mel_T / target_S padding fix.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
@review-notebook-app

Copy link
Copy Markdown

Check out this pull request on  ReviewNB

See visual diffs & provide feedback on Jupyter Notebooks.


Powered by ReviewNB

@Vaibhavdixit02
Vaibhavdixit02 merged commit 1e5aeef into main Apr 22, 2026
1 check failed
@Vaibhavdixit02
Vaibhavdixit02 deleted the cursor/jax-backend-and-repo-reorganize branch April 22, 2026 22:37
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant