google / google/tunix

Single-item batches truncate sampler token_buffer instead of populating it

Open
#809 0 comments 0 reactions 0 assignees View on GitHub
type:bug
Dominant language
Python
Stars
2.5k
Forks
345
Avg merge
1d 7h
Merged PRs (30d)
240

Description

I am doing the Tunix Kaggle competition but run into bugs when setting the TRAIN_MICRO_BATCH_SIZE to 1 in this[example notebook](https://www.kaggle.com/code/windmaple/grpo-demo-gemma2-2b).

https://github.com/google/tunix/blob/1fe360d47961abd1011d427adeb336a429857fbb/tunix/generate/sampler.py#L328-L330

### Description
When invoking `tunix.generate.sampler.Sampler` with `batch_size=1`, the generated output is always truncated after just a few tokens; the call never reaches `` even though the same prompt works fine with `batch_size=2`. Digging into the sampler code shows that:
- `token_buffer` is initialized in `init_sample_state()` as `(batch_size, total_sampling_steps)` (e.g., (1, 896)).
- `_sample()` is supposed to write each new token via `token_buffer = token_buffer.at[:, decoding_step + 1].set(next_token_candidate)`.
- During the JAX `while_loop` (`_decode_fn`), `_sample_step` repeatedly calls `_sample()`, but when `batch_size=1` the buffer never grows beyond the first 2–4 entries.
- By the time token extraction runs (`__call__`, lines 735-782), `token_buffer` has shape `(1, 2)` (or `(1, 4)`), so `token_buffer[max_prompt_length:]` only contains a few tokens and `find_first_eos_idx` returns immediately, resulting in truncated outputs. Running the same logic with `batch_size=2` produces a normal-length buffer and completion.

I added debug host-callbacks around `_sample_step` to log `token_buffer.shape` before/after each iteration and the logs confirm that the `.at[:, decoding_step + 1].set(...)` update silently stops writing when there’s a single example. This makes the sampler unusable when `TRAIN_MICRO_BATCH_SIZE=1`, yet it works fine for larger batches.

### Steps to reproduce
1. Instantiate the sampler used in `gemma2-grpo-original-demo.ipynb`, e.g. the Tunix transformer + tokenizer + cache config.
2. Call `sampler(input_strings=[prompt], ...)` with `max_generation_steps` large enough to reach ``, using greedy sampling (`temperature=0.0001`, `top_k=1`).
3. Inspect `result.tokens[0]` / `result.text[0]` – the output is only a few tokens long and often doesn’t end with ``.
4. Repeat step 2 with `input_strings=[prompt, prompt]` (batch_size=2) and observe a normal-length completion.

### Expected behavior
`token_buffer` should stay `(1, total_sampling_steps)` and `_sample()` should keep writing tokens until EOS, just like it does when `batch_size=2`.

### Actual behavior
`token_buffer` stops at ~2-4 slots, so the extractor only returns the first few tokens and the sample is truncated. This happens inside `_sample()` / `_sample_step`, meaning the loop never expands the buffer for single-item batches.

### Proposed fix
Investigate how `_sample()` / `_sample_step` behaves under JAX when `batch_size=1`. The `.at` update may be optimized out or shape-inferred incorrectly. Ensuring the `token_buffer` update runs (and that `token_buffer` remains `(batch_size, total_sampling_steps)`) should fix the bug. If necessary, add a test for `batch_size=1` to catch regressions.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.