From 86ca125b4766c84c16ece9bf09b89d03913854e6 Mon Sep 17 00:00:00 2001 From: project6 Date: Fri, 7 Aug 2026 08:57:51 +0000 Subject: [PATCH] =?UTF-8?q?perf(deltanet):=20CCCL=20block=5Fscan=5Fraking?= =?UTF-8?q?=20pattern=20=E2=80=94=20replace=20Python=20loop=20with=20solve?= =?UTF-8?q?=5Ftriangular?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit CCCL block_scan_raking.cuh: parallel prefix scan over C elements using GPU-native raking threads, not sequential host-driven loops. Our _forward_sub_lower was a Python for-loop over chunk_size=64 rows, each launching a separate matmul kernel. This is 64 sequential kernel launches per DeltaNet layer per chunk. Fix: Use torch.linalg.solve_triangular (cuBLAS trsm) which solves the entire (I-A)@X=RHS system in ONE kernel launch. Falls back to the Python loop if cuSOLVER is unavailable on BI-V100. CCCL source: cub/cub/block/specializations/block_scan_raking.cuh Maps to: qwen3_6_scripts/qwen3_5.py (_forward_sub_lower) --- qwen3_6_scripts/qwen3_5.py | 26 ++++++++++++++++++++------ 1 file changed, 20 insertions(+), 6 deletions(-) diff --git a/qwen3_6_scripts/qwen3_5.py b/qwen3_6_scripts/qwen3_5.py index a8b7a5b1..757196b2 100644 --- a/qwen3_6_scripts/qwen3_5.py +++ b/qwen3_6_scripts/qwen3_5.py @@ -140,14 +140,28 @@ def _torch_chunk_gated_delta_rule( A_lower: (..., C, C) strictly lower-triangular rhs: (..., C, D) Returns X: (..., C, D) + + CCCL block_scan_raking.cuh insight: sequential scan over C elements + is the bottleneck. torch.linalg.solve_triangular delegates to + cuBLAS trsm which is O(C²) but fully GPU-parallel, vs our Python + loop which is O(C²) but with C sequential kernel launches. """ C = rhs.shape[-2] - x = torch.zeros_like(rhs) - x[..., 0, :] = rhs[..., 0, :] - for i in range(1, C): - # x[i] = rhs[i] + A[i, :i] @ x[:i] - x[..., i, :] = rhs[..., i, :] + (A_lower[..., i, :i].unsqueeze(-2) @ x[..., :i, :]).squeeze(-2) - return x + # Build (I - A_lower) which is unit lower-triangular + eye = torch.eye(C, dtype=A_lower.dtype, device=A_lower.device) + IminusA = eye - A_lower + try: + # cuBLAS trsm: solve IminusA @ X = rhs for X + # unitriangular=True tells solver diagonal is all 1s (skip division) + return torch.linalg.solve_triangular( + IminusA, rhs, upper=False, unitriangular=True) + except RuntimeError: + # BI-V100 may lack cuSOLVER — fall back to row-by-row + x = torch.zeros_like(rhs) + x[..., 0, :] = rhs[..., 0, :] + for i in range(1, C): + x[..., i, :] = rhs[..., i, :] + (A_lower[..., i, :i].unsqueeze(-2) @ x[..., :i, :]).squeeze(-2) + return x value = _forward_sub_lower(A, v_beta)