vllm.models.glm5next.nvidia.ops.third_party.kda.kernels
¶
Functions:
-
chunk_kda_scaled_dot_kkt_fwd–Compute beta * K * K^T.
-
chunk_kda_with_fused_gate–Run chunk KDA from raw gate projection using fused gate+cumsum.
-
fused_kda_gate–Forward pass for KDA gate:
chunk_kda_scaled_dot_kkt_fwd(q, k, gk=None, beta=None, scale=None, cu_seqlens=None, chunk_indices=None, chunk_size=FLA_CHUNK_SIZE, output_dtype=torch.float32)
¶
Compute beta * K * K^T.
Parameters:
-
(q¶Tensor) –The query tensor of shape
[B, T, H, K]. -
(k¶Tensor) –The key tensor of shape
[B, T, H, K]. -
(beta¶Tensor, default:None) –The beta tensor of shape
[B, T, H]. -
(gk¶Tensor, default:None) –The cumulative sum of the gate tensor of shape
[B, T, H, K]applied to the key tensor. Default:None. -
(scale¶float, default:None) –Scale applied to the query-key products. Default:
None. -
(cu_seqlens¶Tensor, default:None) –The cumulative sequence lengths of the input tensor. Default: None
-
(chunk_indices¶Tensor, default:None) –Precomputed chunk indices for
cu_seqlens. Default:None. -
(chunk_size¶int, default:FLA_CHUNK_SIZE) –The chunk size. Default: 64.
-
(output_dtype¶dtype, default:float32) –The dtype of the output tensor. Default:
torch.float32
Returns:
Source code in vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py
408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 | |
chunk_kda_with_fused_gate(q, k, v, raw_g, beta, A_log, g_bias, scale=None, initial_state=None, output_final_state=False, use_qk_l2norm_in_kernel=False, cu_seqlens=None, safe_gate=False, lower_bound=-5.0, **kwargs)
¶
Run chunk KDA from raw gate projection using fused gate+cumsum.
Source code in vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py
fused_kda_gate(g, A, head_k_dim, g_bias=None, beta=1.0, threshold=20.0, safe_gate=False, lower_bound=-5.0)
¶
Forward pass for KDA gate: input g: [..., HD] param A: [H] or [1, 1, H, 1] beta: softplus beta parameter (softplus branch only) threshold: softplus threshold parameter (softplus branch only) safe_gate: when False (default) compute y = -exp(A)softplus(g+g_bias); when True compute the bounded y = lower_boundsigmoid(exp(A)(g+g_bias)) lower_bound: floor for the safe_gate branch (default -5.0) return : [..., H, D]