Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 25 additions & 11 deletions vllm_gaudi/ops/hpu_gdn_pytorch.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,11 +36,26 @@
_USE_EXACT_SOLVE = os.getenv("VLLM_GDN_EXACT_SOLVE", "0") == "1"


@torch._dynamo.disable
def _l2norm_tpc_last_dim(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
# input bf16, output fp32 to avoid precision issues
return torch.ops.hpu.l2_norm(x, epsilon=eps)


def _l2norm_pt_last_dim(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
x = x.to(torch.float32)
return (x / torch.sqrt(torch.sum(x * x, dim=-1, keepdim=True) + eps)).to(torch.float32)


def _l2norm_last_dim(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
return _l2norm_tpc_last_dim(x, eps)
# return _l2norm_pt_last_dim(x, eps)


# @torch._dynamo.disable
def _preprocess_qk_l2norm(q, k):
"""L2norm in eager mode — HPU torch.compile miscompiles l2norm."""
Comment on lines +54 to 56
q = _l2norm_last_dim(q.to(torch.float32))
k = _l2norm_last_dim(k.to(torch.float32))
q = _l2norm_last_dim(q).to(torch.float32)
k = _l2norm_last_dim(k).to(torch.float32)
return q, k


Expand Down Expand Up @@ -506,10 +521,6 @@ def _materialize_seq_ranges(cu_seqlens: torch.Tensor, total_tokens: int) -> list
return ranges


def _l2norm_last_dim(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
return x / torch.sqrt(torch.sum(x * x, dim=-1, keepdim=True) + eps)


def hpu_fused_gdn_gating(
A_log: torch.Tensor,
a: torch.Tensor,
Expand Down Expand Up @@ -612,15 +623,18 @@ def hpu_fused_recurrent_gated_delta_rule(

# Flatten token axis.
# Compute dtype controlled by VLLM_GDN_COMPUTE_FP32 env var (default: bf16)
qf = q.reshape(-1, H, Kdim).to(_GDN_COMPUTE_DTYPE)
kf = k.reshape(-1, H, Kdim).to(_GDN_COMPUTE_DTYPE)
qf = q.reshape(-1, H, Kdim)
kf = k.reshape(-1, H, Kdim)
Comment on lines 624 to +627
vf = v.reshape(-1, HV, Vdim).to(_GDN_COMPUTE_DTYPE)
gf = g.reshape(-1, HV).to(torch.float32)
bf = beta.reshape(-1, HV).to(_GDN_COMPUTE_DTYPE)

if use_qk_l2norm_in_kernel:
qf = _l2norm_last_dim(qf)
kf = _l2norm_last_dim(kf)
qf = _l2norm_last_dim(qf) # output is fp32
kf = _l2norm_last_dim(kf) # output is fp32
Comment on lines +633 to +634
else:
qf = qf.to(_GDN_COMPUTE_DTYPE)
kf = kf.to(_GDN_COMPUTE_DTYPE)

out_full = torch.zeros(num_seqs, HV, Vdim, dtype=v.dtype, device=device)

Expand Down