import torch from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.utils import GenerationBatchResult class _DummyEvent: def __init__(self): self.recorded = False def record(self): self.recorded = True class _HiddenStateSentinel: def to(self, *args, **kwargs): raise AssertionError("hidden_states should not be copied when disabled") def test_copy_to_cpu_skips_hidden_states_when_not_requested(): hidden_states = _HiddenStateSentinel() copy_done = _DummyEvent() result = GenerationBatchResult( logits_output=LogitsProcessorOutput( next_token_logits=None, hidden_states=hidden_states, ), next_token_ids=torch.tensor([1], dtype=torch.int64), ) result.copy_done = copy_done result.copy_to_cpu(return_logprob=False, return_hidden_states=False) assert result.logits_output.hidden_states is hidden_states assert copy_done.recorded