{"data":{"kind":"file","path":"README.md","version_id":"sgmqrir00qsnv4xdn7hmstq7","entry":{"name":"README.md","path":"README.md","is_directory":false,"size":5823,"modified_at":"2026-08-07T06:30:03.605000","content_hash":"5e34ec643fc6a086e54ea8d87875d01bf80b4939071f6050d06fd644115753ce"},"entries":[],"content":"# jax-env\r\n\r\n### Overview\r\n- **Environment ID**: `jax-env`\r\n- **Short description**: RL environment for JAX array/autodiff/linalg tasks using expected_output comparison\r\n- **Tags**: jax, autodiff, scientific-computing\r\n\r\n### Datasets\r\n- **Primary dataset(s)**: `eltociear/jax-tasks-v1`\r\n- **Source links**: [HuggingFace Dataset](https://huggingface.co/datasets/eltociear/jax-tasks-v1)\r\n- **Split sizes**: train (38 examples)\r\n\r\n### Task Categories\r\n\r\n| Category | Tasks | Description |\r\n|----------|-------|-------------|\r\n| array_ops | 14 | reshape, transpose, reductions, clip, sort/argsort, cumsum, matmul, row-normalise, `jnp.where`, functional index update (`.at[].set()`/`.at[].add()`), diff |\r\n| nn | 8 | softmax, log-softmax, relu, gelu, sigmoid, explicit dense layer, MSE, per-row cross-entropy |\r\n| autodiff | 7 | `jax.grad`, second derivative, `jax.vmap`, `jax.jacfwd`, `jax.value_and_grad`, `jax.jit` |\r\n| linalg | 6 | symmetric eigenvalues, Cholesky, solve, SVD, inverse, trace |\r\n| random | 3 | `PRNGKey`, `split`, `normal`, `uniform`, `permutation` with the seed pinned in the prompt |\r\n\r\n### Task\r\n- **Type**: Multi-turn tool use (default `max_turns=5`)\r\n- **Rubric overview**: Binary pass/fail using `numpy.allclose` (rtol 1e-6, atol 1e-8) against the reference array\r\n\r\n### Quickstart\r\n\r\n```bash\r\nuv run vf-eval jax-env -p prime -m openai/gpt-5.4-nano -s\r\n```\r\n\r\n### Environment Arguments\r\n\r\n| Arg | Type | Default | Description |\r\n|-----|------|---------|-------------|\r\n| `split` | str | `\"train\"` | Dataset split to use |\r\n| `dataset_name` | str | `\"eltociear/jax-tasks-v1\"` | HuggingFace dataset name |\r\n| `max_turns` | int | `5` | Maximum interaction turns per task |\r\n\r\n### Tools Available\r\n- `execute_code(code: str)` — run Python in the sandbox; input arrays are preloaded by name, `jax`/`jnp`/`np` imported, and `result` persists across turns\r\n- `bash(command: str)` — run shell commands in the sandbox\r\n\r\nThe sandbox installs `jax[cpu]`: every task is CPU-deterministic.\r\n\r\n### Two JAX-specific decisions\r\n\r\n1. **float32 throughout; `jax_enable_x64` deliberately left off.** JAX defaults to 32-bit and\r\n   silently downcasts. Building the answer key under x64 and grading under the default would\r\n   fail every float task for a reason unrelated to the model.\r\n2. **PRNG keys are explicit and the seed is stated in the prompt.** JAX's threefry stream is\r\n   reproducible, which is what makes those tasks gradeable — but only if the seed is pinned\r\n   rather than assumed.\r\n\r\n### Grading\r\n\r\nThe reference answer never enters the sandbox — it is compared host-side after the rollout.\r\nComparison uses a tolerance because linalg kernels differ in their last bits across builds;\r\nshape is compared exactly and an integer reference requires an integer answer, so the tolerance\r\ncannot launder a wrong result. A wrong answer scores `0.0`; a broken scorer scores `0.0` and\r\nsets `state[\"scoring_error\"]`.\r\n\r\n### Building the dataset\r\n\r\n```bash\r\npython build_tasks.py --verify          # check every task\r\npython build_tasks.py --out train.jsonl # regenerate\r\n```\r\n\r\nAll 38 pass on jax 0.11.0. The verifier earned its place here: on the first run it caught a\r\n\"sort ascending\" task whose input was already a `linspace` (an **identity task**) and a\r\nsecond-derivative task written as `grad(grad(f))`, which JAX rejects because the inner gradient\r\nof a sum over a vector is itself a vector and `grad` requires scalar output.\r\n\r\n## Environment arguments\r\n\r\n`load_environment()` deliberately exposes very little: the program's guidance is that an\r\nenvironment should have one correct way to be run, so nothing about the prompts, the parsing or\r\nthe grading is configurable.\r\n\r\n| Argument | Default | Meaning |\r\n|---|---|---|\r\n| `split` | `'train'` | Dataset split to load. |\r\n| `dataset_name` | `'eltociear/jax-tasks-v1'` | Hugging Face dataset of tasks. Change only to point at a fork. |\r\n| `max_turns` | `5` | Tool-use turns the model gets before the rollout ends. |\r\n| `**kwargs` | — | Passed through to the underlying `SandboxEnv`. |\r\n\r\n## Reward rubric\r\n\r\n| Reward function | Weight | What it returns |\r\n|---|---|---|\r\n| `correctness` | 1.0 | Binary: 1.0 when the answer matches the reference, else 0.0. |\r\n\r\nThere is no LLM judge and no partial credit. The score is computed on the host in\r\n`post_rollout` and read back by the rubric, so the reward is a deterministic function of the\r\nvalues the model left in `result`. A **harness** failure (dead sandbox, unreadable read-back)\r\nalso scores 0.0 but additionally sets `state[\"scoring_error\"]`, so an eval run can tell a\r\nbroken harness from a wrong answer instead of blaming the model.\r\n\r\n## Dependencies\r\n\r\n`datasets>=4.1.0`, `jax>=0.4.30`, `numpy>=1.26.0`, `verifiers>=0.1.8`\r\n\r\n## Sample `vf-eval` usage\r\n\r\n```bash\r\nuv run vf-install jax-env\r\nuv run vf-eval -s jax-env -m gpt-4.1 -n 5 -r 3\r\nuv run vf-tui                      # inspect the outputs/ folder it writes\r\n```\r\n\r\n### Known limitation\r\n\r\nThe Docker **sandbox transport** has not been executed — `docker run`, `pip install`, and the\r\nmodel's code running inside the container — because that needs a Docker runtime and an\r\ninference provider key, neither available where this was authored.\r\n\r\nEverything else is exercised. `environments/test_scoring_path.py` runs this environment's\r\n`post_rollout` and `Rubric` with the sandbox mocked out, asserting that a correct answer scores\r\n1.0, a well-formed wrong answer scores 0.0, and a dead sandbox scores 0.0 *and* sets\r\n`scoring_error` so a harness failure is never mistaken for a bad model. It also checks that\r\nwhat `build_tasks.py` emits is exactly what the scorer expects. The environment mirrors the\r\nstructure of `polars_env` (already accepted into the Environments Program) and imports cleanly\r\nagainst `verifiers` 0.2.1.\r\n","encoding":"utf-8","truncated":false,"total_bytes":5823},"status":null}