Skip to content
Success

Changes

Summary

  1. Use JAX in tuner (#1383) (details)
  2. Add JSON case definitions and Dash case editor (#1389) (details)
  3. Run JAX statistics consistency with an explicit rounding tolerance (#1385) (details)
  4. Run timestep sweep tests with JAX (#1384) (details)
  5. Add JAX comparison plots (details)
Commit 9a1b8b23408d3a5abe78ebae7dd1390755fa4fec by noreply
Use JAX in tuner (#1383)

This makes the JAX work through the new tuner, same sort of way the fortran does. Probably needs a little more testing but seems to be working properly now.

* Support reusable JAX driver state and runtime batches

* Add standalone JAX loss evaluation and consistency coverage

* Keep standalone loss CLI coverage separate from tuning

* Integrate JAX with managed tuner jobs and add bounded Jenkins coverage

* Remove trailing whitespace from extracted loss CLI tests

* Run JAX loss help checks in the JAX pytest suite

* Keep JAX lifecycle coverage with native driver tests and document its origins

* Document loss test scenarios, native counterparts and JAX extensions

* Keep the native loss executable path intact in the test description

* Explain managed tuning test ownership and native backend contracts

* Mirror the native driver lifecycle test and compare complete JAX statistics

* Run JAX driver extensions as an independent step in the lifecycle stage

* Update the C8 comparison mutation for indexed parameter loading

* Refresh JAX tuner integration for merged loss frontends

Use the relocated standalone loss entry point and retain CLI checks only in their shared owner. Give each Jenkins tuning run a separate job directory so repeated builds can succeed.

* Record invalid direct tuner requests as terminal errors

Let the existing request and launcher error handlers report malformed JSON and non-string JAX options instead of crashing during bootstrap before writing job results.

* Preserve winning parameter columns in tuner reruns

Retain request physics overrides while excluding initialized values for tuned parameters. Window and full-case reruns now use the winning candidate values, including multiple columns and grouped override names.

* Use the selected JAX runtime throughout Dash tuning

Persist backend and per-job device settings through native and typed broker requests, saved revisions, continuation and result reruns. Reuse the launcher, shared device controls and rerun override resolver; supervise browser replay processes in the broker.

* Allow run IDs in broker replay completion metadata

Avoid a selector argument collision that killed the replay watcher before recording terminal state. Preserve run IDs and cancellation metadata while reporting completed jobs accurately.

* Keep JAX device setup inside the launcher and backends

Expose a standard-library selection interface and forward per-job device and preallocation settings through the SCM and tuner launch paths. Use the same backend configuration for discovery, provenance, and execution; remove GPU environment handling from Dash and the tuner.

Validate with Dash/shared suites, focused launcher checks, a CPU tuning and replay workflow, and synthetic two-GPU discovery.

* Forward device settings through the existing JAX selection option

Extend the launcher-owned -options grammar with device=UUID and prealloc_gpu_mem=true|false, while retaining xla_prealloc compatibility. Remove the transient device-specific options from SCM and tuner scripts so they forward one opaque -jax value.

Validate saved settings, discovery and replays through the same public interface; verify CPU tuning metrics remain identical.
The file was modifiedjenkins_tests/clubb_tuner/Jenkinsfile (diff)
The file was addedtuner/pytests/auto_llm_generated_pytests/test_jax_backend.py
The file was modifiedtuner/tuning_scheduler.py (diff)
The file was addedtuner/pytests/auto_llm_generated_pytests/test_runtime_device_settings.py
The file was addedtuner/pytests/auto_llm_generated_pytests/test_jax_launch_errors.py
The file was modifieddash_app/compile_tab/callbacks.py (diff)
The file was modifieddash_app/profile_tab/runtime.py (diff)
The file was modifieddash_app/pytests/test_jax_device.py (diff)
The file was modifiedtuner/status.py (diff)
The file was modifiedclubb_jax/README.md (diff)
The file was addeddash_app/pytests/auto_llm_generated_pytests/test_jax_tuning_runtime.py
The file was modifieddash_app/pytests/test_run_broker_simplification.py (diff)
The file was addeddash_app/pytests/auto_llm_generated_pytests/test_tune_replay_completion.py
The file was modifiedtuner/README.md (diff)
The file was modifiedtuner/tuning_worker.py (diff)
The file was modifiedclubb_jax/run_jax.py (diff)
The file was modifieddash_app/compile_tab/build_selector.py (diff)
The file was modifiedclubb_jax/pytests/auto_llm_generated_pytests/test_jax_cli_options.py (diff)
The file was modifiedtuner/request.py (diff)
The file was modifieddash_app/tune_tab/runtime.py (diff)
The file was modifiedrun_scripts/README.md (diff)
The file was modifiedclubb_jax/backends/cuda.py (diff)
The file was addedclubb_jax/pytests/auto_llm_generated_pytests/test_device_selection.py
The file was modifieddash_app/DEVELOPMENT.md (diff)
The file was modifieddash_app/services/models.py (diff)
The file was modifiedtuner/tune_clubb.py (diff)
The file was modifiedutilities/create_case_namelist.py (diff)
The file was modifieddash_app/shared/activity.py (diff)
The file was removeddash_app/shared/jax_device.py
The file was addedtuner/loss_backend.py
The file was modifieddash_app/README.md (diff)
The file was addedtuner/pytests/auto_llm_generated_pytests/test_tuner_top_overrides.py
The file was modifiedrun_scripts/run_tuner_job.py (diff)
The file was modifieddash_app/pytests/test_agent_services.py (diff)
The file was modifieddash_app/tune_tab/callbacks_runs.py (diff)
The file was modifiedtuner/job_runtime.py (diff)
The file was modifieddash_app/shared/actions.py (diff)
The file was modifieddash_app/run_tab/runtime.py (diff)
Commit 5f8b57f69c01c59fc3c4457178ae85d57894c1d9 by noreply
Add JSON case definitions and Dash case editor (#1389)

This is mainly a dash update right now, but it adds json versions of our model files. Basically just 2 new files for that - each json stores default values then one entry for each case with overrides it is not default yet and the current system of defining cases is unchanged.


* Add JSON case definitions and Dash case editor

* Simplify named case lookup and remove queued snapshots

* Remove model validation from JSON case conversion
The file was modifieddash_app/profile_tab/runtime.py (diff)
The file was addeddash_app/pytests/auto_llm_generated_pytests/test_case_editor.py
The file was modifieddash_app/assets/11_tab_run_theme.css (diff)
The file was modifieddash_app/run_tab/tab.py (diff)
The file was modifiedutilities/README.md (diff)
The file was addeddash_app/run_tab/cases.py
The file was addedutilities/pytests/auto_llm_generated_pytests/test_case_json_to_namelist.py
The file was addeddash_app/assets/12_case_editor.css
The file was modifieddash_app/persistence.py (diff)
The file was addedinput/case_setups/case_definitions.json
The file was modifieddash_app/run_tab/discovery.py (diff)
The file was modifieddash_app/run_tab/callbacks_selection.py (diff)
The file was addedutilities/case_json_to_namelist.py
The file was modifieddash_app/README.md (diff)
The file was addeddash_app/assets/34_case_editor.js
The file was modifiedinput/case_setups/README (diff)
The file was modifiedrun_scripts/run_tuner.py (diff)
The file was modifieddash_app/shared/actions.py (diff)
The file was modifiedrun_scripts/run_scm_loss.py (diff)
The file was modified.gitignore (diff)
The file was modifiedutilities/create_case_namelist.py (diff)
The file was modifiedtuner/case_defaults.py (diff)
The file was modifieddash_app/run_tab/layout.py (diff)
Commit 166ec211468128c37d8bb4b917ffac337bf80247 by noreply
Run JAX statistics consistency with an explicit rounding tolerance (#1385)

This adds a jax jenkins stage for the statistics consistency test. This is the test that runs clubb with different batches sizes and averaging methods, e.g
- start by: run once with 60s output intervals all the way through
- test: running with 300s output intervals should give the same stats as manually averaging 5 steps of 60s output
- test: outputing only a "window" of stats data (like from timesteps 100-200 instead of whole run), should produce a window idendical to the 60s full output run
- test: running 16 columns at once should produce the same as running 16 columns in batches of 4 at a time

In fortran these comparisons should be BFB (other than the manually averaging one, which is only close), but in jax, changing things like the batch size can cause optimizations to happen differently, which breaks BFBness and seems to cause differences larger than expected by fortran. For example, the threshold for the manual averaging test in fortran is 1e-12, but for the jax we need 1e-9 even for the batch test. I added a tolerance, `jax_tolerance`, that I left named after jax for now since that is the only current reason that flag exists.

This difficulty with jax wasn't unexpected - we also have to compile the fortran with `-O0` to get this test working, also to disable optimizations. There just seems to be less control over how jax does that, so I went with a "relax the tolerance for this specific instance" approach.

* Support reusable JAX driver state and runtime batches

* Expose JAX in the strict statistics consistency workflow

* Keep JAX lifecycle coverage with native driver tests and document its origins

* Reference native driver and statistics owners in the JAX test workflow

* Keep full Fortran statistics source paths in the test description

* Mirror the native driver lifecycle test and compare complete JAX statistics

* Run JAX driver extensions as an independent step in the lifecycle stage

* Update the C8 comparison mutation for indexed parameter loading

* Keep statistics windows within the capped native case duration

* Allow explicit JAX statistics tolerance for JIT rounding differences
The file was modifiedtests/run_stats_output_consistency.py (diff)
The file was modifiedjenkins_tests/clubb_stats_output_consistency/Jenkinsfile (diff)
The file was addedtests/pytests/auto_llm_generated_pytests/test_run_stats_output_consistency.py
The file was modifiedtests/README.md (diff)
Commit 341cac040a0d352fd5a4da6a5938a86d7d06e332 by noreply
Run timestep sweep tests with JAX (#1384)

Adding a jax run to the timestep test (now renamed timestep sweep because it's more descriptive).

This was really slow on JAX at first - one trick was to add a compilation cache, so that different runs that only use different timesteps didn't need to be recompiled. From my undestanding this is super easy with jax, basically just set `JAX_COMPILATION_CACHE_DIR` to a directory the process can save compiled parts is all it takes, then the compiled parts can live there even after the processes dies, and future ones can reuse it.

* Run existing timestep sweeps with the JAX runner

* Explain the shared native and JAX timestep sweep behavior

* Keep timestep checks aligned with current launcher and CLI conventions

Select CPU explicitly in the bounded JAX stage and resolve optional output roots through the shared owner. Support standard single-dash help, reject ambiguous option abbreviations, and cover opaque runtime forwarding while preserving native defaults and exploratory status.

Validate with 373 shared checks and four eight-step BOMEX/ATEX runs per backend at 60 and 600 seconds.

* Rename the timestep sweep and run shared JAX cases at native defaults

Select the 18 cases shared by the native sweep and JAX/Fortran comparison list. Remove Jenkins timestep and iteration overrides, update stage names and documentation, and preserve the sweep's stability-reporting behavior.

Validate the selection and rename with 374 shared checks before live branch testing.

* Reuse JAX compilations and parallelize case timestep sweeps

Keep a backend-specific persistent cache outside the checkout and preserve explicit cache overrides. Run four isolated case sweeps, prioritize measured expensive cases, and print each completed report as a block while preserving timestep order and first-failure semantics.

Validate cache reuse across fresh processes and cache clearing, scheduling/overlap/interruption contracts, 379 shared checks, and full 90-run cold-cache JAX/native matrices.
The file was addedclubb_jax/pytests/auto_llm_generated_pytests/test_compilation_cache.py
The file was modifiedclubb_jax/README.md (diff)
The file was addedtests/pytests/auto_llm_generated_pytests/test_run_timestep_sweep.py
The file was removedtests/run_timestep_tests.py
The file was modifiedtests/README.md (diff)
The file was addedtests/run_timestep_sweep.py
The file was modifiedjenkins_tests/clubb_timestep/Jenkinsfile (diff)
The file was modifiedclubb_jax/run_jax.py (diff)
The file was modifiedLLM_prompts/jenkins_workflow.md (diff)
The file was modifiedjenkins_tests/clubb_plots/Jenkinsfile (diff)