diff --git a/unirl/models/qwen3/ar.py b/unirl/models/qwen3/ar.py index 9a5862514..be2c329a0 100644 --- a/unirl/models/qwen3/ar.py +++ b/unirl/models/qwen3/ar.py @@ -316,6 +316,7 @@ def autoregress( past_key_values=model_kwargs.get("past_key_values"), attention_mask=model_kwargs.get("attention_mask"), use_cache=True, + logits_to_keep=1, ) with torch.no_grad(): out = transformer(**model_inputs, return_dict=True) diff --git a/unirl/models/qwen3_5/ar.py b/unirl/models/qwen3_5/ar.py index 77a9e0f5f..00cb29b39 100644 --- a/unirl/models/qwen3_5/ar.py +++ b/unirl/models/qwen3_5/ar.py @@ -320,6 +320,7 @@ def autoregress( "mm_token_type_ids": model_kwargs.get("mm_token_type_ids"), "next_sequence_length": next_sequence_length, "use_cache": True, + "logits_to_keep": 1, } if is_first_step: if pv is not None: @@ -336,6 +337,7 @@ def autoregress( "input_ids": cur_input_ids, "attention_mask": model_kwargs["attention_mask"], "use_cache": False, + "logits_to_keep": 1, } if pv is not None: model_inputs["pixel_values"] = pv diff --git a/unirl/models/qwen_vl/ar.py b/unirl/models/qwen_vl/ar.py index af81235c9..b26cedc62 100644 --- a/unirl/models/qwen_vl/ar.py +++ b/unirl/models/qwen_vl/ar.py @@ -203,6 +203,7 @@ def autoregress( "attention_mask": model_kwargs.get("attention_mask"), "cache_position": model_kwargs.get("cache_position"), "use_cache": True, + "logits_to_keep": 1, } if is_first_step: if "pixel_values" in model_kwargs: