Files
Model-Optimizer/tests
yeyu-nvidiaandClaude Opus 5 ad1bad7817 fix(specdec): gather sharded hidden states in DFlash/DSpark AR generation (#2458)
### What does this PR do?

Type of change: Bug fix

Fixes the nightly
`tests/regression/torch/speculative/test_dflash.py::test_dflash_ar_validate`,
red every night since 2026-09-10 with:

```
WARNING: sample 0 (writing) failed: Expected all tensors to be on the same device,
but got tensors is on cuda:1, different from other tensors on cuda:0
(when checking argument in method wrapper_CUDA_cat)
RuntimeError: AR validation produced no results: all 3 samples failed.
```

**No PR caused this.** `ar_validate.py` loads with `device_map="auto"`,
and the regression runner (`…rtxpro6000-l-2-…`) has 2 GPUs, so even
Qwen3-0.6B gets sharded. `pseudo_speculative_generate()` then
concatenates hidden states from `target_layer_ids`, which spans early
*and* late layers — different devices once the base model is split:

```python
selected = [base_outputs.hidden_states[lid + hid_offset] for lid in self.target_layer_ids]
target_hidden = torch.cat(selected, dim=-1)   # cuda:0 ++ cuda:1 -> RuntimeError
```

The block tensors have the same problem from the other side: they follow
`input_ids.device`, while the draft module sits on the *last* base
layer's device (`_place_draft`).

What changed on 2026-09-10 was detection, not behaviour. #2288 made
`ar_validate.py` raise when every sample fails; before it, the report
block was guarded by `if results and …`, so zero results printed nothing
and exited **0**. The test had been green while validating nothing. Its
runtime is the tell: 23–26s in every green nightly, 21.6s in the first
red one, where it does no validation work at all.

The same defect is in #2288's own description, on an 8-GPU sharded
**EAGLE3** checkpoint: 80/80 samples dead on the identical error, job
exiting 0. So this is not DFlash-specific in principle — but EAGLE3's
generate path gathers separately (`pop_and_gather_aux_hiddens`) and
needs its own look, which is why this PR stops at the DFlash family.

### Fix

Gather everything the draft consumes onto the draft's device, and hand
results back on the caller's device since `validate_online` cats them
onto the running sequence. Every `.to()` is a no-op on a single device,
so single-GPU behaviour is bit-identical.

`HFDominoModel` inherits `HFDFlashModel.pseudo_speculative_generate`, so
it is fixed by the same change; `HFDSparkModel` overrides it and is
patched in parallel (including `prev_token` feeding `markov_step`).

### Usage

No API change.

```bash
python examples/speculative_decoding/scripts/ar_validate.py \
    --model_path <dflash ckpt> --osl 10 --num_samples 3 --steps 7
```

### Testing

- `tests/unit/torch/speculative/plugins/test_hf_dflash.py` and
`test_hf_dspark.py`: **110 passed, 1 skipped**.
- The skip is the new `TestShardedBaseGeneration`, which is gated on
`torch.cuda.device_count() >= 2`. It stands in for accelerate's dispatch
— the base forward stays whole, but its hidden states are spread across
two devices the way a real `device_map="auto"` split spreads them — and
asserts both returned tensors come back on the caller's device. **I have
no 2-GPU box to hand, so this test is unverified locally; it will first
execute in CI.** It reproduces the reported error against the unpatched
code by construction, but please treat that claim as unconfirmed until
the run goes green.
- The real verification is the next nightly: `test_dflash_ar_validate`
should pass *and* print an AR number, which it has not done in any log I
can find.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ✅ — device moves only; no-ops
when the base model is on one device.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅ — one multi-GPU regression
test (skipped without 2 GPUs).
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
❌ — bug in an unreleased-cycle path that only ever produced a silent
no-op; no user-visible behaviour to describe.
- Did you get Claude approval on this PR?: ❌ — not yet run.

### Additional Information

Worth a maintainer view: `device_map="auto"` is a poor default for AR
validation, which is a single-process step-by-step loop. Every sharded
run of it that I can find has failed. Pinning it to one device would be
a smaller blast radius than making every drafting path device-safe — but
it would also cap validation at models that fit on one GPU, so I have
not changed it here.

🤖 Generated with [Claude Code](https://claude.com/claude-code)


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **Bug Fixes**
* Improved speculative generation with sharded Hugging Face models
across multiple GPUs.
* Ensured generated draft tokens and outputs are returned on the
caller’s input device.
* Preserved Markov sampling behavior while improving device
coordination.

* **Tests**
* Added multi-GPU coverage for sharded model generation, including
output placement and draft shape validation.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Ye Yu <yeyu@nvidia.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
2026-09-17 11:28:57 -07:00
..