From 2e8f2b0c58fb0be5e20d9e37bcba8f1ff521464d Mon Sep 17 00:00:00 2001 From: Ruixiang Wang Date: Mon, 27 Apr 2026 20:27:33 +0000 Subject: [PATCH] Fix pageable H2D copy in Gated DeltaNet PyTorch fallback --- .../models/olmo_hybrid/modeling_olmo_hybrid.py | 8 +++++--- src/transformers/models/qwen3_5/modeling_qwen3_5.py | 8 +++++--- .../models/qwen3_5_moe/modeling_qwen3_5_moe.py | 8 +++++--- src/transformers/models/qwen3_next/modeling_qwen3_next.py | 8 +++++--- src/transformers/models/qwen3_next/modular_qwen3_next.py | 8 +++++--- 5 files changed, 25 insertions(+), 15 deletions(-) diff --git a/src/transformers/models/olmo_hybrid/modeling_olmo_hybrid.py b/src/transformers/models/olmo_hybrid/modeling_olmo_hybrid.py index 563680286b64..5c76f0a8ca22 100644 --- a/src/transformers/models/olmo_hybrid/modeling_olmo_hybrid.py +++ b/src/transformers/models/olmo_hybrid/modeling_olmo_hybrid.py @@ -551,7 +551,7 @@ def torch_chunk_gated_delta_rule( value = attn @ v_beta k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1)) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) @@ -595,9 +595,11 @@ def torch_recurrent_gated_delta_rule( scale = 1 / (query.shape[-1] ** 0.5) query = query * scale - core_attn_out = torch.zeros(batch_size, num_heads, sequence_length, v_head_dim).to(value) + core_attn_out = torch.zeros( + batch_size, num_heads, sequence_length, v_head_dim, dtype=value.dtype, device=value.device + ) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) diff --git a/src/transformers/models/qwen3_5/modeling_qwen3_5.py b/src/transformers/models/qwen3_5/modeling_qwen3_5.py index 904d08a5570f..bad700952673 100644 --- a/src/transformers/models/qwen3_5/modeling_qwen3_5.py +++ b/src/transformers/models/qwen3_5/modeling_qwen3_5.py @@ -283,7 +283,7 @@ def torch_chunk_gated_delta_rule( value = attn @ v_beta k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1)) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) @@ -327,9 +327,11 @@ def torch_recurrent_gated_delta_rule( scale = 1 / (query.shape[-1] ** 0.5) query = query * scale - core_attn_out = torch.zeros(batch_size, num_heads, sequence_length, v_head_dim).to(value) + core_attn_out = torch.zeros( + batch_size, num_heads, sequence_length, v_head_dim, dtype=value.dtype, device=value.device + ) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) diff --git a/src/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py b/src/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py index bd8a49c969ca..d7b45a276412 100644 --- a/src/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py +++ b/src/transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py @@ -284,7 +284,7 @@ def torch_chunk_gated_delta_rule( value = attn @ v_beta k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1)) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) @@ -328,9 +328,11 @@ def torch_recurrent_gated_delta_rule( scale = 1 / (query.shape[-1] ** 0.5) query = query * scale - core_attn_out = torch.zeros(batch_size, num_heads, sequence_length, v_head_dim).to(value) + core_attn_out = torch.zeros( + batch_size, num_heads, sequence_length, v_head_dim, dtype=value.dtype, device=value.device + ) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) diff --git a/src/transformers/models/qwen3_next/modeling_qwen3_next.py b/src/transformers/models/qwen3_next/modeling_qwen3_next.py index 24f4d8a47b29..395f13d1420c 100644 --- a/src/transformers/models/qwen3_next/modeling_qwen3_next.py +++ b/src/transformers/models/qwen3_next/modeling_qwen3_next.py @@ -423,7 +423,7 @@ def torch_chunk_gated_delta_rule( value = attn @ v_beta k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1)) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) @@ -467,9 +467,11 @@ def torch_recurrent_gated_delta_rule( scale = 1 / (query.shape[-1] ** 0.5) query = query * scale - core_attn_out = torch.zeros(batch_size, num_heads, sequence_length, v_head_dim).to(value) + core_attn_out = torch.zeros( + batch_size, num_heads, sequence_length, v_head_dim, dtype=value.dtype, device=value.device + ) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) diff --git a/src/transformers/models/qwen3_next/modular_qwen3_next.py b/src/transformers/models/qwen3_next/modular_qwen3_next.py index ef55cbdda3f2..0bb527288bb9 100644 --- a/src/transformers/models/qwen3_next/modular_qwen3_next.py +++ b/src/transformers/models/qwen3_next/modular_qwen3_next.py @@ -262,7 +262,7 @@ def torch_chunk_gated_delta_rule( value = attn @ v_beta k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1)) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) ) @@ -306,9 +306,11 @@ def torch_recurrent_gated_delta_rule( scale = 1 / (query.shape[-1] ** 0.5) query = query * scale - core_attn_out = torch.zeros(batch_size, num_heads, sequence_length, v_head_dim).to(value) + core_attn_out = torch.zeros( + batch_size, num_heads, sequence_length, v_head_dim, dtype=value.dtype, device=value.device + ) last_recurrent_state = ( - torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim).to(value) + torch.zeros(batch_size, num_heads, k_head_dim, v_head_dim, dtype=value.dtype, device=value.device) if initial_state is None else initial_state.to(value) )