Skip to content
KernelIndex
Search⌘K

submission 418024

shiyegao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 852 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-418024?include=source"
interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA H100
1.57ms
#20 of 71
2026-02-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d953e31dd740063673c4bdf89849475ede426a98a0f9214f66a111b8b28697cd
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memory__shared__ half tile_l[32][33];
vector-width = float4float4 v = reinterpret_cast<const float4*>(xrow)[c4];

Kernel source

submission.py852 lines
from __future__ import annotations

from typing import Any, Dict, Tuple

__PRECISION_NOTE__ = "fp16_gemm_fp32_accum"

_EXT = None
_KERNEL_CACHE: dict[tuple, dict[str, Any]] = {}
_WEIGHT_CACHE: dict[tuple, dict[str, Any]] = {}


def _load_ext():
    global _EXT
    if _EXT is not None:
        return _EXT

    import hashlib
    import os

    from torch.utils.cpp_extension import load_inline

    this_dir = os.path.dirname(os.path.abspath(__file__))
    build_dir = os.path.join(this_dir, ".torch_ext_build")
    os.makedirs(build_dir, exist_ok=True)

    os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "9.0")

    cpp_src = r"""
    #include <torch/extension.h>
    torch::Tensor trimul_forward(
        torch::Tensor x,
        torch::Tensor mask_h,
        torch::Tensor norm_w,
        torch::Tensor norm_b,
        torch::Tensor w_stack,
        torch::Tensor out_norm_w,
        torch::Tensor out_norm_b,
        torch::Tensor w_to_out,
        torch::Tensor work_u8,
        torch::Tensor x_norm_h,
        torch::Tensor y0_h,
        torch::Tensor left_packed,
        torch::Tensor right_packed,
        torch::Tensor out_gate_packed,
        torch::Tensor out_packed,
        torch::Tensor out2_h,
        torch::Tensor y_out_f);

    PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
      m.def("forward", &trimul_forward, "trimul outgoing forward (CUDA)");
    }
    """

    cuda_src = r"""
    #include <torch/extension.h>
    #include <cuda.h>
    #include <cuda_runtime.h>
    #include <cuda_fp16.h>
    #include <cublasLt.h>

    #ifndef CHECK_CUDA
    #define CHECK_CUDA(x) TORCH_CHECK((x).is_cuda(), #x " must be a CUDA tensor")
    #endif
    #ifndef CHECK_CONTIGUOUS
    #define CHECK_CONTIGUOUS(x) TORCH_CHECK((x).is_contiguous(), #x " must be contiguous")
    #endif
    #ifndef CHECK_DTYPE
    #define CHECK_DTYPE(x, dt) TORCH_CHECK((x).dtype() == (dt), #x " dtype mismatch")
    #endif

    static __device__ __forceinline__ float _warp_sum(float v) {
      unsigned mask = 0xffffffffu;
      v += __shfl_down_sync(mask, v, 16);
      v += __shfl_down_sync(mask, v, 8);
      v += __shfl_down_sync(mask, v, 4);
      v += __shfl_down_sync(mask, v, 2);
      v += __shfl_down_sync(mask, v, 1);
      return __shfl_sync(mask, v, 0);
    }

    static __device__ __forceinline__ float _sigmoid(float x) {
      return 1.0f / (1.0f + __expf(-x));
    }

    __global__ void ln_fwd_fp16(
        const float* __restrict__ x,
        half* __restrict__ y,
        const float* __restrict__ w,
        const float* __restrict__ b,
        int rows,
        int dim) {
      int tid = threadIdx.x;
      int warp = tid >> 5;
      int lane = tid & 31;
      int row = (blockIdx.x * (blockDim.x >> 5)) + warp;
      if (row >= rows) return;

      const float* xrow = x + (long long)row * dim;
      half* yrow = y + (long long)row * dim;

      float sum = 0.0f;
      float sum2 = 0.0f;

      for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {
        float4 v = reinterpret_cast<const float4*>(xrow)[c4];
        sum += v.x + v.y + v.z + v.w;
        sum2 += v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w;
      }
      sum = _warp_sum(sum);
      sum2 = _warp_sum(sum2);
      float mean = sum / (float)dim;
      float var = fmaxf(sum2 / (float)dim - mean * mean, 0.0f);
      float inv = rsqrtf(var + 1e-5f);

      for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {
        float4 v = reinterpret_cast<const float4*>(xrow)[c4];
        int c = c4 << 2;
        float4 ww = make_float4(w[c + 0], w[c + 1], w[c + 2], w[c + 3]);
        float4 bb = make_float4(b[c + 0], b[c + 1], b[c + 2], b[c + 3]);
        float4 o;
        o.x = (v.x - mean) * inv * ww.x + bb.x;
        o.y = (v.y - mean) * inv * ww.y + bb.y;
        o.z = (v.z - mean) * inv * ww.z + bb.z;
        o.w = (v.w - mean) * inv * ww.w + bb.w;
        reinterpret_cast<half2*>(yrow)[c4 * 2 + 0] = __floats2half2_rn(o.x, o.y);
        reinterpret_cast<half2*>(yrow)[c4 * 2 + 1] = __floats2half2_rn(o.z, o.w);
      }
    }

    __global__ void pack_lr_og(
        const half* __restrict__ y0,
        const half* __restrict__ mask,
        half* __restrict__ left_p,
        half* __restrict__ right_p,
        half* __restrict__ og_p,
        int rows,
        int n,
        int hidden) {
      __shared__ half tile_l[32][33];
      __shared__ half tile_r[32][33];
      __shared__ half tile_o[32][33];

      int x = (blockIdx.x << 5) + threadIdx.x;
      int y = (blockIdx.y << 5) + threadIdx.y;

      #pragma unroll
      for (int i = 0; i < 4; ++i) {
        int row = y + (i << 3);
        if (x < hidden && row < rows) {
          float m = __half2float(mask[row]);
          long long base = (long long)row * (5LL * hidden) + x;
          float lp = __half2float(y0[base + 0LL * hidden]);
          float rp = __half2float(y0[base + 1LL * hidden]);
          float lg = __half2float(y0[base + 2LL * hidden]);
          float rg = __half2float(y0[base + 3LL * hidden]);
          float og = __half2float(y0[base + 4LL * hidden]);

          float l = lp * _sigmoid(lg) * m;
          float r = rp * _sigmoid(rg) * m;
          float o = _sigmoid(og);
          tile_l[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(l);
          tile_r[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(r);
          tile_o[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(o);
        }
      }
      __syncthreads();

      int xt = (blockIdx.y << 5) + threadIdx.x;
      int yt = (blockIdx.x << 5) + threadIdx.y;

      #pragma unroll
      for (int i = 0; i < 4; ++i) {
        int h = yt + (i << 3);
        int row = xt;
        if (h < hidden && row < rows) {
          int b = row / (n * n);
          int local = row - b * (n * n);
          int ii = local / n;
          int jj = local - ii * n;
          long long out_idx = ((long long)(b * hidden + h) * n + ii) * n + jj;
          left_p[out_idx] = tile_l[threadIdx.x][threadIdx.y + (i << 3)];
          right_p[out_idx] = tile_r[threadIdx.x][threadIdx.y + (i << 3)];
          og_p[out_idx] = tile_o[threadIdx.x][threadIdx.y + (i << 3)];
        }
      }
    }

    __global__ void ln2_gate_pack_h128(
        const half* __restrict__ out_p,
        const half* __restrict__ og_p,
        const float* __restrict__ w,
        const float* __restrict__ b,
        half* __restrict__ out2,
        int rows_in_b) {
      int tx = (int)threadIdx.x;
      int ty = (int)threadIdx.y;
      int row_base = (int)blockIdx.x << 5;
      int bid = (int)blockIdx.y;

      int row = row_base + tx;

      __shared__ half sv[128][33];
      __shared__ half sg[128][33];

      float sum = 0.0f;
      float sum2 = 0.0f;

      #pragma unroll
      for (int h_base = 0; h_base < 128; h_base += 32) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
          int h = h_base + ty + (i << 3);
          if (row < rows_in_b) {
            long long idx = ((long long)(bid * 128 + h) * rows_in_b) + row;
            half hv = out_p[idx];
            half hg = og_p[idx];
            sv[h][tx] = hv;
            sg[h][tx] = hg;
            float v = __half2float(hv);
            sum += v;
            sum2 += v * v;
          }
        }
      }

      __shared__ float sh_sum[8][32];
      __shared__ float sh_sum2[8][32];
      __shared__ float sh_mean[32];
      __shared__ float sh_inv[32];

      sh_sum[ty][tx] = sum;
      sh_sum2[ty][tx] = sum2;
      __syncthreads();

      if (ty == 0 && row < rows_in_b) {
        float s = 0.0f;
        float s2 = 0.0f;
        #pragma unroll
        for (int t = 0; t < 8; ++t) {
          s += sh_sum[t][tx];
          s2 += sh_sum2[t][tx];
        }
        float mean = s * (1.0f / 128.0f);
        float var = fmaxf(s2 * (1.0f / 128.0f) - mean * mean, 0.0f);
        sh_mean[tx] = mean;
        sh_inv[tx] = rsqrtf(var + 1e-5f);
      }
      __syncthreads();

      #pragma unroll
      for (int h_base = 0; h_base < 128; h_base += 32) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
          int row_off = ty + (i << 3);
          int out_row = row_base + row_off;
          int h = h_base + tx;
          if (out_row < rows_in_b) {
            float v = __half2float(sv[h][row_off]);
            float g = __half2float(sg[h][row_off]);
            float mean = sh_mean[row_off];
            float inv = sh_inv[row_off];
            float nv = (v - mean) * inv * w[h] + b[h];
            out2[((long long)(bid * rows_in_b + out_row) * 128) + h] = __float2half_rn(nv * g);
          }
        }
      }
    }

    __global__ void ln2_gate_pack_tiled(
        const half* __restrict__ out_p,
        const half* __restrict__ og_p,
        const float* __restrict__ w,
        const float* __restrict__ b,
        half* __restrict__ out2,
        int rows_in_b,
        int hidden) {
      int tx = (int)threadIdx.x;
      int ty = (int)threadIdx.y;
      int row_base = (int)blockIdx.x << 5;
      int bid = (int)blockIdx.y;

      int row = row_base + tx;

      float sum = 0.0f;
      float sum2 = 0.0f;

      if (row < rows_in_b) {
        for (int h = ty; h < hidden; h += 8) {
          long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;
          float v = __half2float(out_p[idx]);
          sum += v;
          sum2 += v * v;
        }
      }

      __shared__ float sh_sum[8][32];
      __shared__ float sh_sum2[8][32];
      __shared__ float sh_mean[32];
      __shared__ float sh_inv[32];
      sh_sum[ty][tx] = sum;
      sh_sum2[ty][tx] = sum2;
      __syncthreads();

      if (ty == 0 && row < rows_in_b) {
        float s = 0.0f;
        float s2 = 0.0f;
        #pragma unroll
        for (int t = 0; t < 8; ++t) {
          s += sh_sum[t][tx];
          s2 += sh_sum2[t][tx];
        }
        float mean = s / (float)hidden;
        float var = fmaxf(s2 / (float)hidden - mean * mean, 0.0f);
        sh_mean[tx] = mean;
        sh_inv[tx] = rsqrtf(var + 1e-5f);
      }
      __syncthreads();

      __shared__ half tile_v[32][33];
      __shared__ half tile_g[32][33];

      for (int h_base = 0; h_base < hidden; h_base += 32) {
        #pragma unroll
        for (int i = 0; i < 4; ++i) {
          int h = h_base + ty + (i << 3);
          if (row < rows_in_b && h < hidden) {
            long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;
            tile_v[ty + (i << 3)][tx] = out_p[idx];
            tile_g[ty + (i << 3)][tx] = og_p[idx];
          }
        }
        __syncthreads();

        #pragma unroll
        for (int i = 0; i < 4; ++i) {
          int row_off = ty + (i << 3);
          int out_row = row_base + row_off;
          int h = h_base + tx;
          if (out_row < rows_in_b && h < hidden) {
            float v = __half2float(tile_v[tx][row_off]);
            float g = __half2float(tile_g[tx][row_off]);
            float mean = sh_mean[row_off];
            float inv = sh_inv[row_off];
            float nv = (v - mean) * inv * w[h] + b[h];
            out2[((long long)(bid * rows_in_b + out_row) * hidden) + h] = __float2half_rn(nv * g);
          }
        }
        __syncthreads();
      }
    }

    struct LtPlanKey {
      int m, n, k;
      int batch;
      int op_b;
      int a_type, b_type, c_type, d_type;
    };

    struct LtPlan {
      bool valid;
      LtPlanKey key;
      cublasLtMatmulAlgo_t algo;
      size_t work_bytes;
      cublasLtMatmulDesc_t op_desc;
      cublasLtMatrixLayout_t a_desc;
      cublasLtMatrixLayout_t b_desc;
      cublasLtMatrixLayout_t c_desc;
      cublasLtMatrixLayout_t d_desc;
    };

    static cublasLtHandle_t _lt = nullptr;
    static LtPlan _plan_g1 = {false};
    static LtPlan _plan_g2 = {false};
    static LtPlan _plan_ct = {false};

    static inline bool _key_eq(const LtPlanKey& a, const LtPlanKey& b) {
      return a.m==b.m && a.n==b.n && a.k==b.k && a.batch==b.batch && a.op_b==b.op_b &&
             a.a_type==b.a_type && a.b_type==b.b_type && a.c_type==b.c_type && a.d_type==b.d_type;
    }

    static inline void _lt_init() {
      if (_lt) return;
      auto st = cublasLtCreate(&_lt);
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtCreate failed");
    }

    static inline void _lt_plan_destroy(LtPlan* plan) {
      if (!plan->valid) return;
      if (plan->a_desc) cublasLtMatrixLayoutDestroy(plan->a_desc);
      if (plan->b_desc) cublasLtMatrixLayoutDestroy(plan->b_desc);
      if (plan->c_desc) cublasLtMatrixLayoutDestroy(plan->c_desc);
      if (plan->d_desc) cublasLtMatrixLayoutDestroy(plan->d_desc);
      if (plan->op_desc) cublasLtMatmulDescDestroy(plan->op_desc);
      plan->a_desc = nullptr;
      plan->b_desc = nullptr;
      plan->c_desc = nullptr;
      plan->d_desc = nullptr;
      plan->op_desc = nullptr;
      plan->valid = false;
    }

    static inline cublasLtMatrixLayout_t _lt_make_layout(
        cudaDataType type,
        int rows,
        int cols,
        int ld,
        int batch,
        long long stride,
        cublasLtOrder_t order) {
      cublasLtMatrixLayout_t out = nullptr;
      auto st = cublasLtMatrixLayoutCreate(&out, type, rows, cols, ld);
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatrixLayoutCreate failed");
      st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_ORDER, &order, sizeof(order));
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set order failed");
      if (batch > 1) {
        st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch, sizeof(batch));
        TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set batch failed");
        st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &stride, sizeof(stride));
        TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set stride failed");
      }
      return out;
    }

    static inline void _lt_pick_algo(
        LtPlan* plan,
        const LtPlanKey& key,
        cublasOperation_t op_a,
        cublasOperation_t op_b,
        cudaDataType a_type,
        cudaDataType b_type,
        cudaDataType c_type,
        cudaDataType d_type,
        int lda, int ldb, int ldc, int ldd,
        int batch,
        long long stride_a,
        long long stride_b,
        long long stride_c,
        long long stride_d,
        size_t work_bytes) {
      _lt_init();

      if (plan->valid) {
        _lt_plan_destroy(plan);
      }

      auto st = cublasLtMatmulDescCreate(&plan->op_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F);
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "matmul desc create failed");
      st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &op_a, sizeof(op_a));
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transa failed");
      st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &op_b, sizeof(op_b));
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transb failed");

      cublasLtOrder_t order = CUBLASLT_ORDER_ROW;
      int a_rows = (op_a == CUBLAS_OP_N) ? key.m : key.k;
      int a_cols = (op_a == CUBLAS_OP_N) ? key.k : key.m;
      int b_rows = (op_b == CUBLAS_OP_N) ? key.k : key.n;
      int b_cols = (op_b == CUBLAS_OP_N) ? key.n : key.k;
      plan->a_desc = _lt_make_layout(a_type, a_rows, a_cols, lda, batch, stride_a, order);
      plan->b_desc = _lt_make_layout(b_type, b_rows, b_cols, ldb, batch, stride_b, order);
      plan->c_desc = _lt_make_layout(c_type, key.m, key.n, ldc, batch, stride_c, order);
      plan->d_desc = _lt_make_layout(d_type, key.m, key.n, ldd, batch, stride_d, order);

      cublasLtMatmulPreference_t pref = nullptr;
      st = cublasLtMatmulPreferenceCreate(&pref);
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref create failed");
      st = cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &work_bytes, sizeof(work_bytes));
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref set failed");

      cublasLtMatmulHeuristicResult_t heur;
      int got = 0;
      st = cublasLtMatmulAlgoGetHeuristic(_lt, plan->op_desc, plan->a_desc, plan->b_desc, plan->c_desc, plan->d_desc, pref, 1, &heur, &got);
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS && got > 0, "no cublasLt heuristic algo");

      plan->valid = true;
      plan->key = key;
      plan->algo = heur.algo;
      plan->work_bytes = work_bytes;

      cublasLtMatmulPreferenceDestroy(pref);
    }

    static inline void _lt_matmul(
        LtPlan* plan,
        const LtPlanKey& key,
        cublasOperation_t op_a,
        cublasOperation_t op_b,
        const void* a,
        const void* b,
        const void* c,
        void* d,
        cudaDataType a_type,
        cudaDataType b_type,
        cudaDataType c_type,
        cudaDataType d_type,
        int lda, int ldb, int ldc, int ldd,
        int batch,
        long long stride_a,
        long long stride_b,
        long long stride_c,
        long long stride_d,
        void* work,
        size_t work_bytes) {
      _lt_init();
      if (!plan->valid || !_key_eq(plan->key, key)) {
        _lt_pick_algo(plan, key, op_a, op_b, a_type, b_type, c_type, d_type, lda, ldb, ldc, ldd, batch, stride_a, stride_b, stride_c, stride_d, work_bytes);
      }

      float alpha = 1.0f;
      float beta = 0.0f;
      auto st = cublasLtMatmul(_lt, plan->op_desc, &alpha, a, plan->a_desc, b, plan->b_desc, &beta, c, plan->c_desc, d, plan->d_desc, &plan->algo, work, plan->work_bytes, 0);
      TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatmul failed");
    }

    torch::Tensor trimul_forward(
        torch::Tensor x,
        torch::Tensor mask_h,
        torch::Tensor norm_w,
        torch::Tensor norm_b,
        torch::Tensor w_stack,
        torch::Tensor out_norm_w,
        torch::Tensor out_norm_b,
        torch::Tensor w_to_out,
        torch::Tensor work_u8,
        torch::Tensor x_norm_h,
        torch::Tensor y0_h,
        torch::Tensor left_packed,
        torch::Tensor right_packed,
        torch::Tensor out_gate_packed,
        torch::Tensor out_packed,
        torch::Tensor out2_h,
        torch::Tensor y_out_f) {
      CHECK_CUDA(x);
      CHECK_CUDA(mask_h);
      CHECK_CUDA(norm_w);
      CHECK_CUDA(norm_b);
      CHECK_CUDA(w_stack);
      CHECK_CUDA(out_norm_w);
      CHECK_CUDA(out_norm_b);
      CHECK_CUDA(w_to_out);
      CHECK_CUDA(work_u8);
      CHECK_CUDA(x_norm_h);
      CHECK_CUDA(y0_h);
      CHECK_CUDA(left_packed);
      CHECK_CUDA(right_packed);
      CHECK_CUDA(out_gate_packed);
      CHECK_CUDA(out_packed);
      CHECK_CUDA(out2_h);
      CHECK_CUDA(y_out_f);

      CHECK_CONTIGUOUS(x);
      CHECK_CONTIGUOUS(mask_h);
      CHECK_CONTIGUOUS(norm_w);
      CHECK_CONTIGUOUS(norm_b);
      CHECK_CONTIGUOUS(w_stack);
      CHECK_CONTIGUOUS(out_norm_w);
      CHECK_CONTIGUOUS(out_norm_b);
      CHECK_CONTIGUOUS(w_to_out);
      CHECK_CONTIGUOUS(work_u8);
      CHECK_CONTIGUOUS(x_norm_h);
      CHECK_CONTIGUOUS(y0_h);
      CHECK_CONTIGUOUS(left_packed);
      CHECK_CONTIGUOUS(right_packed);
      CHECK_CONTIGUOUS(out_gate_packed);
      CHECK_CONTIGUOUS(out_packed);
      CHECK_CONTIGUOUS(out2_h);
      CHECK_CONTIGUOUS(y_out_f);

      CHECK_DTYPE(x, torch::kFloat32);
      CHECK_DTYPE(mask_h, torch::kFloat16);
      CHECK_DTYPE(norm_w, torch::kFloat32);
      CHECK_DTYPE(norm_b, torch::kFloat32);
      CHECK_DTYPE(w_stack, torch::kFloat16);
      CHECK_DTYPE(out_norm_w, torch::kFloat32);
      CHECK_DTYPE(out_norm_b, torch::kFloat32);
      CHECK_DTYPE(w_to_out, torch::kFloat16);
      CHECK_DTYPE(work_u8, torch::kUInt8);
      CHECK_DTYPE(x_norm_h, torch::kFloat16);
      CHECK_DTYPE(y0_h, torch::kFloat16);
      CHECK_DTYPE(left_packed, torch::kFloat16);
      CHECK_DTYPE(right_packed, torch::kFloat16);
      CHECK_DTYPE(out_gate_packed, torch::kFloat16);
      CHECK_DTYPE(out_packed, torch::kFloat16);
      CHECK_DTYPE(out2_h, torch::kFloat16);
      CHECK_DTYPE(y_out_f, torch::kFloat32);

      TORCH_CHECK(x.dim() == 4, "x must be [bs,N,N,dim]");
      int bs = (int)x.size(0);
      int n = (int)x.size(1);
      int dim = (int)x.size(3);
      TORCH_CHECK((int)x.size(2) == n, "x must be square on N");
      TORCH_CHECK((int)mask_h.size(0) == bs && (int)mask_h.size(1) == n && (int)mask_h.size(2) == n, "mask shape");

      int hidden5 = (int)w_stack.size(1);
      TORCH_CHECK(hidden5 % 5 == 0, "w_stack second dim must be 5*hidden");
      int hidden = hidden5 / 5;

      int rows = bs * n * n;
      int rows_in_b = n * n;

      TORCH_CHECK((int)x_norm_h.size(0) == rows && (int)x_norm_h.size(1) == dim, "x_norm_h shape");
      TORCH_CHECK((dim & 3) == 0, "dim must be multiple of 4");
      int warps = 8;
      dim3 block1(32 * warps, 1, 1);
      dim3 grid1((rows + warps - 1) / warps, 1, 1);
      ln_fwd_fp16<<<grid1, block1>>>(
          (const float*)x.data_ptr<float>(),
          (half*)x_norm_h.data_ptr<at::Half>(),
          (const float*)norm_w.data_ptr<float>(),
          (const float*)norm_b.data_ptr<float>(),
          rows, dim);

      TORCH_CHECK((int)y0_h.size(0) == rows && (int)y0_h.size(1) == 5 * hidden, "y0_h shape");
      LtPlanKey k1;
      k1.m = rows; k1.n = 5 * hidden; k1.k = dim; k1.batch = 1; k1.op_b = 0;
      k1.a_type = (int)CUDA_R_16F; k1.b_type = (int)CUDA_R_16F; k1.c_type = (int)CUDA_R_16F; k1.d_type = (int)CUDA_R_16F;
      _lt_matmul(
          &_plan_g1, k1,
          CUBLAS_OP_N, CUBLAS_OP_N,
          x_norm_h.data_ptr<at::Half>(),
          w_stack.data_ptr<at::Half>(),
          y0_h.data_ptr<at::Half>(),
          y0_h.data_ptr<at::Half>(),
          CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,
          dim, 5 * hidden, 5 * hidden, 5 * hidden,
          1, 0, 0, 0, 0,
          work_u8.data_ptr(), (size_t)work_u8.numel());

      TORCH_CHECK((int)left_packed.size(0) == bs * hidden && (int)left_packed.size(1) == n && (int)left_packed.size(2) == n, "left_packed shape");
      TORCH_CHECK(left_packed.sizes() == right_packed.sizes(), "right_packed shape");
      TORCH_CHECK(left_packed.sizes() == out_gate_packed.sizes(), "out_gate_packed shape");
      dim3 block2(32, 8, 1);
      dim3 grid2((hidden + 31) / 32, (rows + 31) / 32, 1);
      pack_lr_og<<<grid2, block2>>>(
          (const half*)y0_h.data_ptr<at::Half>(),
          (const half*)mask_h.data_ptr<at::Half>(),
          (half*)left_packed.data_ptr<at::Half>(),
          (half*)right_packed.data_ptr<at::Half>(),
          (half*)out_gate_packed.data_ptr<at::Half>(),
          rows, n, hidden);

      TORCH_CHECK(out_packed.sizes() == left_packed.sizes(), "out_packed shape");
      int batch_ct = bs * hidden;
      LtPlanKey kc;
      kc.m = n; kc.n = n; kc.k = n; kc.batch = batch_ct; kc.op_b = 1;
      kc.a_type = (int)CUDA_R_16F; kc.b_type = (int)CUDA_R_16F; kc.c_type = (int)CUDA_R_16F; kc.d_type = (int)CUDA_R_16F;
      long long stride_mat = (long long)n * n;
      _lt_matmul(
          &_plan_ct, kc,
          CUBLAS_OP_N, CUBLAS_OP_T,
          left_packed.data_ptr<at::Half>(),
          right_packed.data_ptr<at::Half>(),
          out_packed.data_ptr<at::Half>(),
          out_packed.data_ptr<at::Half>(),
          CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,
          n, n, n, n,
          batch_ct,
          stride_mat, stride_mat, stride_mat, stride_mat,
          work_u8.data_ptr(), (size_t)work_u8.numel());

      TORCH_CHECK((int)out2_h.size(0) == rows && (int)out2_h.size(1) == hidden, "out2_h shape");
      TORCH_CHECK((int)out_norm_w.numel() == hidden && (int)out_norm_b.numel() == hidden, "out_norm weight/bias shape");
      dim3 block3(32, 8, 1);
      dim3 grid3((rows_in_b + 31) / 32, bs, 1);
      if (hidden == 128) {
        ln2_gate_pack_h128<<<grid3, block3>>>(
            (const half*)out_packed.data_ptr<at::Half>(),
            (const half*)out_gate_packed.data_ptr<at::Half>(),
            (const float*)out_norm_w.data_ptr<float>(),
            (const float*)out_norm_b.data_ptr<float>(),
            (half*)out2_h.data_ptr<at::Half>(),
            rows_in_b);
      } else {
        ln2_gate_pack_tiled<<<grid3, block3>>>(
            (const half*)out_packed.data_ptr<at::Half>(),
            (const half*)out_gate_packed.data_ptr<at::Half>(),
            (const float*)out_norm_w.data_ptr<float>(),
            (const float*)out_norm_b.data_ptr<float>(),
            (half*)out2_h.data_ptr<at::Half>(),
            rows_in_b, hidden);
      }

      TORCH_CHECK((int)y_out_f.size(0) == rows && (int)y_out_f.size(1) == dim, "y_out_f shape");
      LtPlanKey k2;
      k2.m = rows; k2.n = dim; k2.k = hidden; k2.batch = 1; k2.op_b = 0;
      k2.a_type = (int)CUDA_R_16F; k2.b_type = (int)CUDA_R_16F; k2.c_type = (int)CUDA_R_32F; k2.d_type = (int)CUDA_R_32F;
      _lt_matmul(
          &_plan_g2, k2,
          CUBLAS_OP_N, CUBLAS_OP_N,
          out2_h.data_ptr<at::Half>(),
          w_to_out.data_ptr<at::Half>(),
          y_out_f.data_ptr<float>(),
          y_out_f.data_ptr<float>(),
          CUDA_R_16F, CUDA_R_16F, CUDA_R_32F, CUDA_R_32F,
          hidden, dim, dim, dim,
          1, 0, 0, 0, 0,
          work_u8.data_ptr(), (size_t)work_u8.numel());

      return y_out_f;
    }
    """

    name = "trimul_ext_" + hashlib.sha256((cpp_src + cuda_src).encode("utf-8")).hexdigest()[:16]
    _EXT = load_inline(
        name=name,
        cpp_sources=[cpp_src],
        cuda_sources=[cuda_src],
        functions=None,
        extra_cflags=["-O3", "-std=c++17"],
        extra_cuda_cflags=["-O3", "--use_fast_math", "-std=c++17"],
        extra_ldflags=["-lcublasLt", "-lcublas"],
        with_cuda=True,
        build_directory=build_dir,
        verbose=False,
    )
    return _EXT


def _prep_weights(weights: Dict[str, "torch.Tensor"], dim: int, hidden_dim: int, device: "torch.device") -> dict[str, Any]:
    import torch

    key = (
        int(weights["left_proj.weight"].data_ptr()),
        int(weights["right_proj.weight"].data_ptr()),
        int(weights["left_gate.weight"].data_ptr()),
        int(weights["right_gate.weight"].data_ptr()),
        int(weights["out_gate.weight"].data_ptr()),
        int(weights["to_out.weight"].data_ptr()),
        int(weights["norm.weight"].data_ptr()),
        int(weights["norm.bias"].data_ptr()),
        int(weights["to_out_norm.weight"].data_ptr()),
        int(weights["to_out_norm.bias"].data_ptr()),
        device.index if device.type == "cuda" else -1,
    )
    cached = _WEIGHT_CACHE.get(key)
    if cached is not None:
        return cached

    def _f32_vec(t: "torch.Tensor") -> "torch.Tensor":
        if t.device == device and t.dtype == torch.float32 and t.is_contiguous():
            return t
        return t.to(device=device, dtype=torch.float32).contiguous()

    def _to_half_t(w: "torch.Tensor") -> "torch.Tensor":
        return w.to(device=device, dtype=torch.float16).t().contiguous()

    w_stack = torch.cat(
        [
            _to_half_t(weights["left_proj.weight"]),
            _to_half_t(weights["right_proj.weight"]),
            _to_half_t(weights["left_gate.weight"]),
            _to_half_t(weights["right_gate.weight"]),
            _to_half_t(weights["out_gate.weight"]),
        ],
        dim=1,
    ).contiguous()

    w_to_out = _to_half_t(weights["to_out.weight"])

    out = {
        "norm_w": _f32_vec(weights["norm.weight"]),
        "norm_b": _f32_vec(weights["norm.bias"]),
        "out_norm_w": _f32_vec(weights["to_out_norm.weight"]),
        "out_norm_b": _f32_vec(weights["to_out_norm.bias"]),
        "w_stack": w_stack,
        "w_to_out": w_to_out,
    }
    _WEIGHT_CACHE[key] = out
    return out


def _get_scratch(device: "torch.device", bs: int, n: int, dim: int, hidden: int) -> dict[str, Any]:
    import torch

    key = (device.type, device.index if device.type == "cuda" else -1, bs, n, dim, hidden)
    cached = _KERNEL_CACHE.get(key)
    if cached is not None:
        return cached

    rows = bs * n * n
    work_bytes = 32 * 1024 * 1024
    scratch = {
        "work": torch.empty((work_bytes,), device=device, dtype=torch.uint8),
        "x_norm": torch.empty((rows, dim), device=device, dtype=torch.float16),
        "y0": torch.empty((rows, 5 * hidden), device=device, dtype=torch.float16),
        "left_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
        "right_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
        "og_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
        "out_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
        "out2": torch.empty((rows, hidden), device=device, dtype=torch.float16),
        "y_out": torch.empty((rows, dim), device=device, dtype=torch.float32),
    }
    _KERNEL_CACHE[key] = scratch
    return scratch


def _ensure_mask_half(mask: "torch.Tensor", device: "torch.device") -> "torch.Tensor":
    import torch

    if mask.dtype == torch.float16 and mask.is_contiguous() and mask.device == device:
        return mask
    return mask.to(device=device, dtype=torch.float16).contiguous()


def custom_kernel(data: Tuple["torch.Tensor", "torch.Tensor", Dict[str, "torch.Tensor"], Dict[str, Any]]) -> "torch.Tensor":
    import torch

    x, mask, weights, config = data
    dim = int(config["dim"])
    hidden_dim = int(config["hidden_dim"])

    if not x.is_cuda:
        raise RuntimeError("CUDA only")
    if x.dtype != torch.float32:
        x = x.to(dtype=torch.float32)
    if not x.is_contiguous():
        x = x.contiguous()

    bs, n0, n1, d0 = x.shape
    if n0 != n1 or d0 != dim:
        raise RuntimeError("shape mismatch")

    mask_h = _ensure_mask_half(mask, x.device)

    ext = _load_ext()
    wpack = _prep_weights(weights, dim=dim, hidden_dim=hidden_dim, device=x.device)
    scratch = _get_scratch(x.device, bs=bs, n=n0, dim=dim, hidden=hidden_dim)

    y_flat = ext.forward(
        x,
        mask_h,
        wpack["norm_w"],
        wpack["norm_b"],
        wpack["w_stack"],
        wpack["out_norm_w"],
        wpack["out_norm_b"],
        wpack["w_to_out"],
        scratch["work"],
        scratch["x_norm"],
        scratch["y0"],
        scratch["left_p"],
        scratch["right_p"],
        scratch["og_p"],
        scratch["out_p"],
        scratch["out2"],
        scratch["y_out"],
    )

    return y_flat.view(bs, n0, n0, dim)


__all__ = ["custom_kernel"]
scrolls · 852 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 417990.

⋯ 1 unchanged lines
from typing import Any, Dict, Tuple
- import torch
- import torch.nn.functional as F
+ __PRECISION_NOTE__ = "fp16_gemm_fp32_accum"
- import cutlass
- import cutlass.cute as cute
- from cutlass.cute.runtime import make_ptr
- from cutlass.cutlass_dsl import for_generate, if_generate, yield_out
+ _EXT = None
+ _KERNEL_CACHE: dict[tuple, dict[str, Any]] = {}
+ _WEIGHT_CACHE: dict[tuple, dict[str, Any]] = {}
- _TMN = 32
- _TK = 32
- _THREADS = 256
+ def _load_ext():
+ global _EXT
+ if _EXT is not None:
+ return _EXT
+ import hashlib
+ import os
- class _BatchedGemmF32_32x32x32:
- def __init__(self) -> None:
- self.threads = _THREADS
+ from torch.utils.cpp_extension import load_inline
- @cute.jit
- def __call__(
- self,
- a_ptr: "cute.Pointer",
- b_ptr: "cute.Pointer",
- c_ptr: "cute.Pointer",
- problem: tuple,
- ):
- batch, n = problem
+ this_dir = os.path.dirname(os.path.abspath(__file__))
+ build_dir = os.path.join(this_dir, ".torch_ext_build")
+ os.makedirs(build_dir, exist_ok=True)
- stride_batch = n * n
- stride_i = n
+ os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "9.0")
- a = cute.make_tensor(
- a_ptr,
- cute.make_layout(
- (batch, n, n),
- stride=(stride_batch, stride_i, 1),
- ),
- )
- b = cute.make_tensor(
- b_ptr,
- cute.make_layout(
- (batch, n, n),
- stride=(stride_batch, stride_i, 1),
- ),
- )
- c = cute.make_tensor(
- c_ptr,
- cute.make_layout(
- (batch, n, n),
- stride=(stride_batch, stride_i, 1),
- ),
- )
+ cpp_src = r"""
+ #include <torch/extension.h>
+ torch::Tensor trimul_forward(
+ torch::Tensor x,
+ torch::Tensor mask_h,
+ torch::Tensor norm_w,
+ torch::Tensor norm_b,
+ torch::Tensor w_stack,
+ torch::Tensor out_norm_w,
+ torch::Tensor out_norm_b,
+ torch::Tensor w_to_out,
+ torch::Tensor work_u8,
+ torch::Tensor x_norm_h,
+ torch::Tensor y0_h,
+ torch::Tensor left_packed,
+ torch::Tensor right_packed,
+ torch::Tensor out_gate_packed,
+ torch::Tensor out_packed,
+ torch::Tensor out2_h,
+ torch::Tensor y_out_f);
- if (n & (_TMN - 1)) == 0:
- grid_x = n // _TMN
- grid_y = n // _TMN
- self.kernel_fast(a, b, c, n).launch(
- grid=[grid_x, grid_y, batch],
- block=[self.threads, 1, 1],
- )
- else:
- grid_x = (n + _TMN - 1) // _TMN
- grid_y = (n + _TMN - 1) // _TMN
- self.kernel_bnd(a, b, c, n).launch(
- grid=[grid_x, grid_y, batch],
- block=[self.threads, 1, 1],
- )
- return
+ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
+ m.def("forward", &trimul_forward, "trimul outgoing forward (CUDA)");
+ }
+ """
- @cute.kernel
- def kernel_fast(self, a: "cute.Tensor", b: "cute.Tensor", c: "cute.Tensor", n: int):
- tx, _, _ = cute.arch.thread_idx()
- bx, by, bz = cute.arch.block_idx()
+ cuda_src = r"""
+ #include <torch/extension.h>
+ #include <cuda.h>
+ #include <cuda_runtime.h>
+ #include <cuda_fp16.h>
+ #include <cublasLt.h>
- tid = tx
- lane_m = tid >> 4
- lane_n = tid & 15
+ #ifndef CHECK_CUDA
+ #define CHECK_CUDA(x) TORCH_CHECK((x).is_cuda(), #x " must be a CUDA tensor")
+ #endif
+ #ifndef CHECK_CONTIGUOUS
+ #define CHECK_CONTIGUOUS(x) TORCH_CHECK((x).is_contiguous(), #x " must be contiguous")
+ #endif
+ #ifndef CHECK_DTYPE
+ #define CHECK_DTYPE(x, dt) TORCH_CHECK((x).dtype() == (dt), #x " dtype mismatch")
+ #endif
- base_m = by * _TMN
- base_n = bx * _TMN
+ static __device__ __forceinline__ float _warp_sum(float v) {
+ unsigned mask = 0xffffffffu;
+ v += __shfl_down_sync(mask, v, 16);
+ v += __shfl_down_sync(mask, v, 8);
+ v += __shfl_down_sync(mask, v, 4);
+ v += __shfl_down_sync(mask, v, 2);
+ v += __shfl_down_sync(mask, v, 1);
+ return __shfl_sync(mask, v, 0);
+ }
- i0 = base_m + lane_m
- j0 = base_n + lane_n
- i1 = i0 + 16
- j1 = j0 + 16
+ static __device__ __forceinline__ float _sigmoid(float x) {
+ return 1.0f / (1.0f + __expf(-x));
+ }
- acc00 = cutlass.Float32(0.0)
- acc01 = cutlass.Float32(0.0)
- acc10 = cutlass.Float32(0.0)
- acc11 = cutlass.Float32(0.0)
+ __global__ void ln_fwd_fp16(
+ const float* __restrict__ x,
+ half* __restrict__ y,
+ const float* __restrict__ w,
+ const float* __restrict__ b,
+ int rows,
+ int dim) {
+ int tid = threadIdx.x;
+ int warp = tid >> 5;
+ int lane = tid & 31;
+ int row = (blockIdx.x * (blockDim.x >> 5)) + warp;
+ if (row >= rows) return;
- smem_stride = _TK + 1
- smem_a_elems = _TMN * smem_stride
- smem_b_elems = _TMN * smem_stride
- smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)
+ const float* xrow = x + (long long)row * dim;
+ half* yrow = y + (long long)row * dim;
- sA = cute.make_tensor(
- smem_ptr,
- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
- )
- sB = cute.make_tensor(
- smem_ptr + smem_a_elems,
- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
- )
+ float sum = 0.0f;
+ float sum2 = 0.0f;
- for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
- 0,
- n,
- _TK,
- iter_args=[acc00, acc01, acc10, acc11],
- ):
- base = tid << 2
- for t in range(4):
- idx = base + t
- row = idx >> 5
- col = idx & 31
- sA[row, col] = a[bz, base_m + row, k0 + col]
- sB[row, col] = b[bz, base_n + row, k0 + col]
+ for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {
+ float4 v = reinterpret_cast<const float4*>(xrow)[c4];
+ sum += v.x + v.y + v.z + v.w;
+ sum2 += v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w;
+ }
+ sum = _warp_sum(sum);
+ sum2 = _warp_sum(sum2);
+ float mean = sum / (float)dim;
+ float var = fmaxf(sum2 / (float)dim - mean * mean, 0.0f);
+ float inv = rsqrtf(var + 1e-5f);
- cute.arch.sync_threads()
+ for (int c4 = lane; c4 < (dim >> 2); c4 += 32) {
+ float4 v = reinterpret_cast<const float4*>(xrow)[c4];
+ int c = c4 << 2;
+ float4 ww = make_float4(w[c + 0], w[c + 1], w[c + 2], w[c + 3]);
+ float4 bb = make_float4(b[c + 0], b[c + 1], b[c + 2], b[c + 3]);
+ float4 o;
+ o.x = (v.x - mean) * inv * ww.x + bb.x;
+ o.y = (v.y - mean) * inv * ww.y + bb.y;
+ o.z = (v.z - mean) * inv * ww.z + bb.z;
+ o.w = (v.w - mean) * inv * ww.w + bb.w;
+ reinterpret_cast<half2*>(yrow)[c4 * 2 + 0] = __floats2half2_rn(o.x, o.y);
+ reinterpret_cast<half2*>(yrow)[c4 * 2 + 1] = __floats2half2_rn(o.z, o.w);
+ }
+ }
- for kk in range(_TK):
- a0 = sA[lane_m, kk]
- a1 = sA[lane_m + 16, kk]
- b0 = sB[lane_n, kk]
- b1 = sB[lane_n + 16, kk]
+ __global__ void pack_lr_og(
+ const half* __restrict__ y0,
+ const half* __restrict__ mask,
+ half* __restrict__ left_p,
+ half* __restrict__ right_p,
+ half* __restrict__ og_p,
+ int rows,
+ int n,
+ int hidden) {
+ __shared__ half tile_l[32][33];
+ __shared__ half tile_r[32][33];
+ __shared__ half tile_o[32][33];
- acc00 = acc00 + a0 * b0
- acc01 = acc01 + a0 * b1
- acc10 = acc10 + a1 * b0
- acc11 = acc11 + a1 * b1
+ int x = (blockIdx.x << 5) + threadIdx.x;
+ int y = (blockIdx.y << 5) + threadIdx.y;
- cute.arch.sync_threads()
- yield_out([acc00, acc01, acc10, acc11])
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ int row = y + (i << 3);
+ if (x < hidden && row < rows) {
+ float m = __half2float(mask[row]);
+ long long base = (long long)row * (5LL * hidden) + x;
+ float lp = __half2float(y0[base + 0LL * hidden]);
+ float rp = __half2float(y0[base + 1LL * hidden]);
+ float lg = __half2float(y0[base + 2LL * hidden]);
+ float rg = __half2float(y0[base + 3LL * hidden]);
+ float og = __half2float(y0[base + 4LL * hidden]);
- c[bz, i0, j0] = acc00_out
- c[bz, i0, j1] = acc01_out
- c[bz, i1, j0] = acc10_out
- c[bz, i1, j1] = acc11_out
+ float l = lp * _sigmoid(lg) * m;
+ float r = rp * _sigmoid(rg) * m;
+ float o = _sigmoid(og);
+ tile_l[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(l);
+ tile_r[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(r);
+ tile_o[threadIdx.y + (i << 3)][threadIdx.x] = __float2half_rn(o);
+ }
+ }
+ __syncthreads();
- @cute.kernel
- def kernel_bnd(self, a: "cute.Tensor", b: "cute.Tensor", c: "cute.Tensor", n: int):
- tx, _, _ = cute.arch.thread_idx()
- bx, by, bz = cute.arch.block_idx()
+ int xt = (blockIdx.y << 5) + threadIdx.x;
+ int yt = (blockIdx.x << 5) + threadIdx.y;
- tid = tx
- lane_m = tid >> 4
- lane_n = tid & 15
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ int h = yt + (i << 3);
+ int row = xt;
+ if (h < hidden && row < rows) {
+ int b = row / (n * n);
+ int local = row - b * (n * n);
+ int ii = local / n;
+ int jj = local - ii * n;
+ long long out_idx = ((long long)(b * hidden + h) * n + ii) * n + jj;
+ left_p[out_idx] = tile_l[threadIdx.x][threadIdx.y + (i << 3)];
+ right_p[out_idx] = tile_r[threadIdx.x][threadIdx.y + (i << 3)];
+ og_p[out_idx] = tile_o[threadIdx.x][threadIdx.y + (i << 3)];
+ }
+ }
+ }
- base_m = by * _TMN
- base_n = bx * _TMN
+ __global__ void ln2_gate_pack_h128(
+ const half* __restrict__ out_p,
+ const half* __restrict__ og_p,
+ const float* __restrict__ w,
+ const float* __restrict__ b,
+ half* __restrict__ out2,
+ int rows_in_b) {
+ int tx = (int)threadIdx.x;
+ int ty = (int)threadIdx.y;
+ int row_base = (int)blockIdx.x << 5;
+ int bid = (int)blockIdx.y;
- i0 = base_m + lane_m
- j0 = base_n + lane_n
- i1 = i0 + 16
- j1 = j0 + 16
+ int row = row_base + tx;
- acc00 = cutlass.Float32(0.0)
- acc01 = cutlass.Float32(0.0)
- acc10 = cutlass.Float32(0.0)
- acc11 = cutlass.Float32(0.0)
+ __shared__ half sv[128][33];
+ __shared__ half sg[128][33];
- smem_stride = _TK + 1
- smem_a_elems = _TMN * smem_stride
- smem_b_elems = _TMN * smem_stride
- smem_ptr = cute.arch.alloc_smem(cutlass.Float32, smem_a_elems + smem_b_elems, alignment=16)
+ float sum = 0.0f;
+ float sum2 = 0.0f;
- sA = cute.make_tensor(
- smem_ptr,
- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
- )
- sB = cute.make_tensor(
- smem_ptr + smem_a_elems,
- cute.make_layout((_TMN, _TK), stride=(smem_stride, 1)),
- )
+ #pragma unroll
+ for (int h_base = 0; h_base < 128; h_base += 32) {
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ int h = h_base + ty + (i << 3);
+ if (row < rows_in_b) {
+ long long idx = ((long long)(bid * 128 + h) * rows_in_b) + row;
+ half hv = out_p[idx];
+ half hg = og_p[idx];
+ sv[h][tx] = hv;
+ sg[h][tx] = hg;
+ float v = __half2float(hv);
+ sum += v;
+ sum2 += v * v;
+ }
+ }
+ }
- for k0, (acc00, acc01, acc10, acc11), (acc00_out, acc01_out, acc10_out, acc11_out) in for_generate(
- 0,
- n,
- _TK,
- iter_args=[acc00, acc01, acc10, acc11],
- ):
- base = tid << 2
- for t in range(4):
- idx = base + t
- row = idx >> 5
- col = idx & 31
+ __shared__ float sh_sum[8][32];
+ __shared__ float sh_sum2[8][32];
+ __shared__ float sh_mean[32];
+ __shared__ float sh_inv[32];
- gi_a = base_m + row
- gk_a = k0 + col
- gj_b = base_n + row
- gk_b = k0 + col
+ sh_sum[ty][tx] = sum;
+ sh_sum2[ty][tx] = sum2;
+ __syncthreads();
- def _ld_a():
- sA[row, col] = a[bz, gi_a, gk_a]
+ if (ty == 0 && row < rows_in_b) {
+ float s = 0.0f;
+ float s2 = 0.0f;
+ #pragma unroll
+ for (int t = 0; t < 8; ++t) {
+ s += sh_sum[t][tx];
+ s2 += sh_sum2[t][tx];
+ }
+ float mean = s * (1.0f / 128.0f);
+ float var = fmaxf(s2 * (1.0f / 128.0f) - mean * mean, 0.0f);
+ sh_mean[tx] = mean;
+ sh_inv[tx] = rsqrtf(var + 1e-5f);
+ }
+ __syncthreads();
- def _ze_a():
- sA[row, col] = cutlass.Float32(0.0)
+ #pragma unroll
+ for (int h_base = 0; h_base < 128; h_base += 32) {
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ int row_off = ty + (i << 3);
+ int out_row = row_base + row_off;
+ int h = h_base + tx;
+ if (out_row < rows_in_b) {
+ float v = __half2float(sv[h][row_off]);
+ float g = __half2float(sg[h][row_off]);
+ float mean = sh_mean[row_off];
+ float inv = sh_inv[row_off];
+ float nv = (v - mean) * inv * w[h] + b[h];
+ out2[((long long)(bid * rows_in_b + out_row) * 128) + h] = __float2half_rn(nv * g);
+ }
+ }
+ }
+ }
- def _ld_b():
- sB[row, col] = b[bz, gj_b, gk_b]
+ __global__ void ln2_gate_pack_tiled(
+ const half* __restrict__ out_p,
+ const half* __restrict__ og_p,
+ const float* __restrict__ w,
+ const float* __restrict__ b,
+ half* __restrict__ out2,
+ int rows_in_b,
+ int hidden) {
+ int tx = (int)threadIdx.x;
+ int ty = (int)threadIdx.y;
+ int row_base = (int)blockIdx.x << 5;
+ int bid = (int)blockIdx.y;
- def _ze_b():
- sB[row, col] = cutlass.Float32(0.0)
+ int row = row_base + tx;
- if_generate((gi_a < n) & (gk_a < n), _ld_a, _ze_a)
- if_generate((gj_b < n) & (gk_b < n), _ld_b, _ze_b)
+ float sum = 0.0f;
+ float sum2 = 0.0f;
- cute.arch.sync_threads()
+ if (row < rows_in_b) {
+ for (int h = ty; h < hidden; h += 8) {
+ long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;
+ float v = __half2float(out_p[idx]);
+ sum += v;
+ sum2 += v * v;
+ }
+ }
- for kk in range(_TK):
- a0 = sA[lane_m, kk]
- a1 = sA[lane_m + 16, kk]
- b0 = sB[lane_n, kk]
- b1 = sB[lane_n + 16, kk]
+ __shared__ float sh_sum[8][32];
+ __shared__ float sh_sum2[8][32];
+ __shared__ float sh_mean[32];
+ __shared__ float sh_inv[32];
+ sh_sum[ty][tx] = sum;
+ sh_sum2[ty][tx] = sum2;
+ __syncthreads();
- acc00 = acc00 + a0 * b0
- acc01 = acc01 + a0 * b1
- acc10 = acc10 + a1 * b0
- acc11 = acc11 + a1 * b1
+ if (ty == 0 && row < rows_in_b) {
+ float s = 0.0f;
+ float s2 = 0.0f;
+ #pragma unroll
+ for (int t = 0; t < 8; ++t) {
+ s += sh_sum[t][tx];
+ s2 += sh_sum2[t][tx];
+ }
+ float mean = s / (float)hidden;
+ float var = fmaxf(s2 / (float)hidden - mean * mean, 0.0f);
+ sh_mean[tx] = mean;
+ sh_inv[tx] = rsqrtf(var + 1e-5f);
+ }
+ __syncthreads();
- cute.arch.sync_threads()
- yield_out([acc00, acc01, acc10, acc11])
+ __shared__ half tile_v[32][33];
+ __shared__ half tile_g[32][33];
- def _st00():
- c[bz, i0, j0] = acc00_out
+ for (int h_base = 0; h_base < hidden; h_base += 32) {
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ int h = h_base + ty + (i << 3);
+ if (row < rows_in_b && h < hidden) {
+ long long idx = ((long long)(bid * hidden + h) * rows_in_b) + row;
+ tile_v[ty + (i << 3)][tx] = out_p[idx];
+ tile_g[ty + (i << 3)][tx] = og_p[idx];
+ }
+ }
+ __syncthreads();
- def _st01():
- c[bz, i0, j1] = acc01_out
+ #pragma unroll
+ for (int i = 0; i < 4; ++i) {
+ int row_off = ty + (i << 3);
+ int out_row = row_base + row_off;
+ int h = h_base + tx;
+ if (out_row < rows_in_b && h < hidden) {
+ float v = __half2float(tile_v[tx][row_off]);
+ float g = __half2float(tile_g[tx][row_off]);
+ float mean = sh_mean[row_off];
+ float inv = sh_inv[row_off];
+ float nv = (v - mean) * inv * w[h] + b[h];
+ out2[((long long)(bid * rows_in_b + out_row) * hidden) + h] = __float2half_rn(nv * g);
+ }
+ }
+ __syncthreads();
+ }
+ }
- def _st10():
- c[bz, i1, j0] = acc10_out
+ struct LtPlanKey {
+ int m, n, k;
+ int batch;
+ int op_b;
+ int a_type, b_type, c_type, d_type;
+ };
- def _st11():
- c[bz, i1, j1] = acc11_out
+ struct LtPlan {
+ bool valid;
+ LtPlanKey key;
+ cublasLtMatmulAlgo_t algo;
+ size_t work_bytes;
+ cublasLtMatmulDesc_t op_desc;
+ cublasLtMatrixLayout_t a_desc;
+ cublasLtMatrixLayout_t b_desc;
+ cublasLtMatrixLayout_t c_desc;
+ cublasLtMatrixLayout_t d_desc;
+ };
- if_generate((i0 < n) & (j0 < n), _st00)
- if_generate((i0 < n) & (j1 < n), _st01)
- if_generate((i1 < n) & (j0 < n), _st10)
- if_generate((i1 < n) & (j1 < n), _st11)
+ static cublasLtHandle_t _lt = nullptr;
+ static LtPlan _plan_g1 = {false};
+ static LtPlan _plan_g2 = {false};
+ static LtPlan _plan_ct = {false};
+ static inline bool _key_eq(const LtPlanKey& a, const LtPlanKey& b) {
+ return a.m==b.m && a.n==b.n && a.k==b.k && a.batch==b.batch && a.op_b==b.op_b &&
+ a.a_type==b.a_type && a.b_type==b.b_type && a.c_type==b.c_type && a.d_type==b.d_type;
+ }
- _CONTRACT = _BatchedGemmF32_32x32x32()
- _CONTRACT_C = None
+ static inline void _lt_init() {
+ if (_lt) return;
+ auto st = cublasLtCreate(&_lt);
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtCreate failed");
+ }
+ static inline void _lt_plan_destroy(LtPlan* plan) {
+ if (!plan->valid) return;
+ if (plan->a_desc) cublasLtMatrixLayoutDestroy(plan->a_desc);
+ if (plan->b_desc) cublasLtMatrixLayoutDestroy(plan->b_desc);
+ if (plan->c_desc) cublasLtMatrixLayoutDestroy(plan->c_desc);
+ if (plan->d_desc) cublasLtMatrixLayoutDestroy(plan->d_desc);
+ if (plan->op_desc) cublasLtMatmulDescDestroy(plan->op_desc);
+ plan->a_desc = nullptr;
+ plan->b_desc = nullptr;
+ plan->c_desc = nullptr;
+ plan->d_desc = nullptr;
+ plan->op_desc = nullptr;
+ plan->valid = false;
+ }
- def _get_contract_compiled():
- global _CONTRACT_C
- if _CONTRACT_C is not None:
- return _CONTRACT_C
+ static inline cublasLtMatrixLayout_t _lt_make_layout(
+ cudaDataType type,
+ int rows,
+ int cols,
+ int ld,
+ int batch,
+ long long stride,
+ cublasLtOrder_t order) {
+ cublasLtMatrixLayout_t out = nullptr;
+ auto st = cublasLtMatrixLayoutCreate(&out, type, rows, cols, ld);
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatrixLayoutCreate failed");
+ st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_ORDER, &order, sizeof(order));
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set order failed");
+ if (batch > 1) {
+ st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch, sizeof(batch));
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set batch failed");
+ st = cublasLtMatrixLayoutSetAttribute(out, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &stride, sizeof(stride));
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "layout set stride failed");
+ }
+ return out;
+ }
- a_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
- b_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
- c_ptr = make_ptr(cutlass.Float32, 0, cute.AddressSpace.gmem, assumed_align=16)
- _CONTRACT_C = cute.compile(
- _CONTRACT,
- a_ptr,
- b_ptr,
- c_ptr,
- (0, 0),
- options="--opt-level 3",
+ static inline void _lt_pick_algo(
+ LtPlan* plan,
+ const LtPlanKey& key,
+ cublasOperation_t op_a,
+ cublasOperation_t op_b,
+ cudaDataType a_type,
+ cudaDataType b_type,
+ cudaDataType c_type,
+ cudaDataType d_type,
+ int lda, int ldb, int ldc, int ldd,
+ int batch,
+ long long stride_a,
+ long long stride_b,
+ long long stride_c,
+ long long stride_d,
+ size_t work_bytes) {
+ _lt_init();
+
+ if (plan->valid) {
+ _lt_plan_destroy(plan);
+ }
+
+ auto st = cublasLtMatmulDescCreate(&plan->op_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F);
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "matmul desc create failed");
+ st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &op_a, sizeof(op_a));
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transa failed");
+ st = cublasLtMatmulDescSetAttribute(plan->op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &op_b, sizeof(op_b));
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "set transb failed");
+
+ cublasLtOrder_t order = CUBLASLT_ORDER_ROW;
+ int a_rows = (op_a == CUBLAS_OP_N) ? key.m : key.k;
+ int a_cols = (op_a == CUBLAS_OP_N) ? key.k : key.m;
+ int b_rows = (op_b == CUBLAS_OP_N) ? key.k : key.n;
+ int b_cols = (op_b == CUBLAS_OP_N) ? key.n : key.k;
+ plan->a_desc = _lt_make_layout(a_type, a_rows, a_cols, lda, batch, stride_a, order);
+ plan->b_desc = _lt_make_layout(b_type, b_rows, b_cols, ldb, batch, stride_b, order);
+ plan->c_desc = _lt_make_layout(c_type, key.m, key.n, ldc, batch, stride_c, order);
+ plan->d_desc = _lt_make_layout(d_type, key.m, key.n, ldd, batch, stride_d, order);
+
+ cublasLtMatmulPreference_t pref = nullptr;
+ st = cublasLtMatmulPreferenceCreate(&pref);
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref create failed");
+ st = cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &work_bytes, sizeof(work_bytes));
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "pref set failed");
+
+ cublasLtMatmulHeuristicResult_t heur;
+ int got = 0;
+ st = cublasLtMatmulAlgoGetHeuristic(_lt, plan->op_desc, plan->a_desc, plan->b_desc, plan->c_desc, plan->d_desc, pref, 1, &heur, &got);
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS && got > 0, "no cublasLt heuristic algo");
+
+ plan->valid = true;
+ plan->key = key;
+ plan->algo = heur.algo;
+ plan->work_bytes = work_bytes;
+
+ cublasLtMatmulPreferenceDestroy(pref);
+ }
+
+ static inline void _lt_matmul(
+ LtPlan* plan,
+ const LtPlanKey& key,
+ cublasOperation_t op_a,
+ cublasOperation_t op_b,
+ const void* a,
+ const void* b,
+ const void* c,
+ void* d,
+ cudaDataType a_type,
+ cudaDataType b_type,
+ cudaDataType c_type,
+ cudaDataType d_type,
+ int lda, int ldb, int ldc, int ldd,
+ int batch,
+ long long stride_a,
+ long long stride_b,
+ long long stride_c,
+ long long stride_d,
+ void* work,
+ size_t work_bytes) {
+ _lt_init();
+ if (!plan->valid || !_key_eq(plan->key, key)) {
+ _lt_pick_algo(plan, key, op_a, op_b, a_type, b_type, c_type, d_type, lda, ldb, ldc, ldd, batch, stride_a, stride_b, stride_c, stride_d, work_bytes);
+ }
+
+ float alpha = 1.0f;
+ float beta = 0.0f;
+ auto st = cublasLtMatmul(_lt, plan->op_desc, &alpha, a, plan->a_desc, b, plan->b_desc, &beta, c, plan->c_desc, d, plan->d_desc, &plan->algo, work, plan->work_bytes, 0);
+ TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, "cublasLtMatmul failed");
+ }
+
+ torch::Tensor trimul_forward(
+ torch::Tensor x,
+ torch::Tensor mask_h,
+ torch::Tensor norm_w,
+ torch::Tensor norm_b,
+ torch::Tensor w_stack,
+ torch::Tensor out_norm_w,
+ torch::Tensor out_norm_b,
+ torch::Tensor w_to_out,
+ torch::Tensor work_u8,
+ torch::Tensor x_norm_h,
+ torch::Tensor y0_h,
+ torch::Tensor left_packed,
+ torch::Tensor right_packed,
+ torch::Tensor out_gate_packed,
+ torch::Tensor out_packed,
+ torch::Tensor out2_h,
+ torch::Tensor y_out_f) {
+ CHECK_CUDA(x);
+ CHECK_CUDA(mask_h);
+ CHECK_CUDA(norm_w);
+ CHECK_CUDA(norm_b);
+ CHECK_CUDA(w_stack);
+ CHECK_CUDA(out_norm_w);
+ CHECK_CUDA(out_norm_b);
+ CHECK_CUDA(w_to_out);
+ CHECK_CUDA(work_u8);
+ CHECK_CUDA(x_norm_h);
+ CHECK_CUDA(y0_h);
+ CHECK_CUDA(left_packed);
+ CHECK_CUDA(right_packed);
+ CHECK_CUDA(out_gate_packed);
+ CHECK_CUDA(out_packed);
+ CHECK_CUDA(out2_h);
+ CHECK_CUDA(y_out_f);
+
+ CHECK_CONTIGUOUS(x);
+ CHECK_CONTIGUOUS(mask_h);
+ CHECK_CONTIGUOUS(norm_w);
+ CHECK_CONTIGUOUS(norm_b);
+ CHECK_CONTIGUOUS(w_stack);
+ CHECK_CONTIGUOUS(out_norm_w);
+ CHECK_CONTIGUOUS(out_norm_b);
+ CHECK_CONTIGUOUS(w_to_out);
+ CHECK_CONTIGUOUS(work_u8);
+ CHECK_CONTIGUOUS(x_norm_h);
+ CHECK_CONTIGUOUS(y0_h);
+ CHECK_CONTIGUOUS(left_packed);
+ CHECK_CONTIGUOUS(right_packed);
+ CHECK_CONTIGUOUS(out_gate_packed);
+ CHECK_CONTIGUOUS(out_packed);
+ CHECK_CONTIGUOUS(out2_h);
+ CHECK_CONTIGUOUS(y_out_f);
+
+ CHECK_DTYPE(x, torch::kFloat32);
+ CHECK_DTYPE(mask_h, torch::kFloat16);
+ CHECK_DTYPE(norm_w, torch::kFloat32);
+ CHECK_DTYPE(norm_b, torch::kFloat32);
+ CHECK_DTYPE(w_stack, torch::kFloat16);
+ CHECK_DTYPE(out_norm_w, torch::kFloat32);
+ CHECK_DTYPE(out_norm_b, torch::kFloat32);
+ CHECK_DTYPE(w_to_out, torch::kFloat16);
+ CHECK_DTYPE(work_u8, torch::kUInt8);
+ CHECK_DTYPE(x_norm_h, torch::kFloat16);
+ CHECK_DTYPE(y0_h, torch::kFloat16);
+ CHECK_DTYPE(left_packed, torch::kFloat16);
+ CHECK_DTYPE(right_packed, torch::kFloat16);
+ CHECK_DTYPE(out_gate_packed, torch::kFloat16);
+ CHECK_DTYPE(out_packed, torch::kFloat16);
+ CHECK_DTYPE(out2_h, torch::kFloat16);
+ CHECK_DTYPE(y_out_f, torch::kFloat32);
+
+ TORCH_CHECK(x.dim() == 4, "x must be [bs,N,N,dim]");
+ int bs = (int)x.size(0);
+ int n = (int)x.size(1);
+ int dim = (int)x.size(3);
+ TORCH_CHECK((int)x.size(2) == n, "x must be square on N");
+ TORCH_CHECK((int)mask_h.size(0) == bs && (int)mask_h.size(1) == n && (int)mask_h.size(2) == n, "mask shape");
+
+ int hidden5 = (int)w_stack.size(1);
+ TORCH_CHECK(hidden5 % 5 == 0, "w_stack second dim must be 5*hidden");
+ int hidden = hidden5 / 5;
+
+ int rows = bs * n * n;
+ int rows_in_b = n * n;
+
+ TORCH_CHECK((int)x_norm_h.size(0) == rows && (int)x_norm_h.size(1) == dim, "x_norm_h shape");
+ TORCH_CHECK((dim & 3) == 0, "dim must be multiple of 4");
+ int warps = 8;
+ dim3 block1(32 * warps, 1, 1);
+ dim3 grid1((rows + warps - 1) / warps, 1, 1);
+ ln_fwd_fp16<<<grid1, block1>>>(
+ (const float*)x.data_ptr<float>(),
+ (half*)x_norm_h.data_ptr<at::Half>(),
+ (const float*)norm_w.data_ptr<float>(),
+ (const float*)norm_b.data_ptr<float>(),
+ rows, dim);
+
+ TORCH_CHECK((int)y0_h.size(0) == rows && (int)y0_h.size(1) == 5 * hidden, "y0_h shape");
+ LtPlanKey k1;
+ k1.m = rows; k1.n = 5 * hidden; k1.k = dim; k1.batch = 1; k1.op_b = 0;
+ k1.a_type = (int)CUDA_R_16F; k1.b_type = (int)CUDA_R_16F; k1.c_type = (int)CUDA_R_16F; k1.d_type = (int)CUDA_R_16F;
+ _lt_matmul(
+ &_plan_g1, k1,
+ CUBLAS_OP_N, CUBLAS_OP_N,
+ x_norm_h.data_ptr<at::Half>(),
+ w_stack.data_ptr<at::Half>(),
+ y0_h.data_ptr<at::Half>(),
+ y0_h.data_ptr<at::Half>(),
+ CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,
+ dim, 5 * hidden, 5 * hidden, 5 * hidden,
+ 1, 0, 0, 0, 0,
+ work_u8.data_ptr(), (size_t)work_u8.numel());
+
+ TORCH_CHECK((int)left_packed.size(0) == bs * hidden && (int)left_packed.size(1) == n && (int)left_packed.size(2) == n, "left_packed shape");
+ TORCH_CHECK(left_packed.sizes() == right_packed.sizes(), "right_packed shape");
+ TORCH_CHECK(left_packed.sizes() == out_gate_packed.sizes(), "out_gate_packed shape");
+ dim3 block2(32, 8, 1);
+ dim3 grid2((hidden + 31) / 32, (rows + 31) / 32, 1);
+ pack_lr_og<<<grid2, block2>>>(
+ (const half*)y0_h.data_ptr<at::Half>(),
+ (const half*)mask_h.data_ptr<at::Half>(),
+ (half*)left_packed.data_ptr<at::Half>(),
+ (half*)right_packed.data_ptr<at::Half>(),
+ (half*)out_gate_packed.data_ptr<at::Half>(),
+ rows, n, hidden);
+
+ TORCH_CHECK(out_packed.sizes() == left_packed.sizes(), "out_packed shape");
+ int batch_ct = bs * hidden;
+ LtPlanKey kc;
+ kc.m = n; kc.n = n; kc.k = n; kc.batch = batch_ct; kc.op_b = 1;
+ kc.a_type = (int)CUDA_R_16F; kc.b_type = (int)CUDA_R_16F; kc.c_type = (int)CUDA_R_16F; kc.d_type = (int)CUDA_R_16F;
+ long long stride_mat = (long long)n * n;
+ _lt_matmul(
+ &_plan_ct, kc,
+ CUBLAS_OP_N, CUBLAS_OP_T,
+ left_packed.data_ptr<at::Half>(),
+ right_packed.data_ptr<at::Half>(),
+ out_packed.data_ptr<at::Half>(),
+ out_packed.data_ptr<at::Half>(),
+ CUDA_R_16F, CUDA_R_16F, CUDA_R_16F, CUDA_R_16F,
+ n, n, n, n,
+ batch_ct,
+ stride_mat, stride_mat, stride_mat, stride_mat,
+ work_u8.data_ptr(), (size_t)work_u8.numel());
+
+ TORCH_CHECK((int)out2_h.size(0) == rows && (int)out2_h.size(1) == hidden, "out2_h shape");
+ TORCH_CHECK((int)out_norm_w.numel() == hidden && (int)out_norm_b.numel() == hidden, "out_norm weight/bias shape");
+ dim3 block3(32, 8, 1);
+ dim3 grid3((rows_in_b + 31) / 32, bs, 1);
+ if (hidden == 128) {
+ ln2_gate_pack_h128<<<grid3, block3>>>(
+ (const half*)out_packed.data_ptr<at::Half>(),
+ (const half*)out_gate_packed.data_ptr<at::Half>(),
+ (const float*)out_norm_w.data_ptr<float>(),
+ (const float*)out_norm_b.data_ptr<float>(),
+ (half*)out2_h.data_ptr<at::Half>(),
+ rows_in_b);
+ } else {
+ ln2_gate_pack_tiled<<<grid3, block3>>>(
+ (const half*)out_packed.data_ptr<at::Half>(),
+ (const half*)out_gate_packed.data_ptr<at::Half>(),
+ (const float*)out_norm_w.data_ptr<float>(),
+ (const float*)out_norm_b.data_ptr<float>(),
+ (half*)out2_h.data_ptr<at::Half>(),
+ rows_in_b, hidden);
+ }
+
+ TORCH_CHECK((int)y_out_f.size(0) == rows && (int)y_out_f.size(1) == dim, "y_out_f shape");
+ LtPlanKey k2;
+ k2.m = rows; k2.n = dim; k2.k = hidden; k2.batch = 1; k2.op_b = 0;
+ k2.a_type = (int)CUDA_R_16F; k2.b_type = (int)CUDA_R_16F; k2.c_type = (int)CUDA_R_32F; k2.d_type = (int)CUDA_R_32F;
+ _lt_matmul(
+ &_plan_g2, k2,
+ CUBLAS_OP_N, CUBLAS_OP_N,
+ out2_h.data_ptr<at::Half>(),
+ w_to_out.data_ptr<at::Half>(),
+ y_out_f.data_ptr<float>(),
+ y_out_f.data_ptr<float>(),
+ CUDA_R_16F, CUDA_R_16F, CUDA_R_32F, CUDA_R_32F,
+ hidden, dim, dim, dim,
+ 1, 0, 0, 0, 0,
+ work_u8.data_ptr(), (size_t)work_u8.numel());
+
+ return y_out_f;
+ }
+ """
+
+ name = "trimul_ext_" + hashlib.sha256((cpp_src + cuda_src).encode("utf-8")).hexdigest()[:16]
+ _EXT = load_inline(
+ name=name,
+ cpp_sources=[cpp_src],
+ cuda_sources=[cuda_src],
+ functions=None,
+ extra_cflags=["-O3", "-std=c++17"],
+ extra_cuda_cflags=["-O3", "--use_fast_math", "-std=c++17"],
+ extra_ldflags=["-lcublasLt", "-lcublas"],
+ with_cuda=True,
+ build_directory=build_dir,
+ verbose=False,
)
- return _CONTRACT_C
+ return _EXT
- def _contract_outgoing(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
- bs, n, _, hidden = left.shape
+ def _prep_weights(weights: Dict[str, "torch.Tensor"], dim: int, hidden_dim: int, device: "torch.device") -> dict[str, Any]:
+ import torch
- if not (left.is_cuda and right.is_cuda):
- raise RuntimeError("该实现仅支持 CUDA 张量。")
- if left.dtype != torch.float32 or right.dtype != torch.float32:
- raise RuntimeError("该实现期望 left/right 为 float32。")
+ key = (
+ int(weights["left_proj.weight"].data_ptr()),
+ int(weights["right_proj.weight"].data_ptr()),
+ int(weights["left_gate.weight"].data_ptr()),
+ int(weights["right_gate.weight"].data_ptr()),
+ int(weights["out_gate.weight"].data_ptr()),
+ int(weights["to_out.weight"].data_ptr()),
+ int(weights["norm.weight"].data_ptr()),
+ int(weights["norm.bias"].data_ptr()),
+ int(weights["to_out_norm.weight"].data_ptr()),
+ int(weights["to_out_norm.bias"].data_ptr()),
+ device.index if device.type == "cuda" else -1,
+ )
+ cached = _WEIGHT_CACHE.get(key)
+ if cached is not None:
+ return cached
-
- left_t = left.permute(0, 3, 1, 2).contiguous()
- right_t = right.permute(0, 3, 1, 2).contiguous()
+ def _f32_vec(t: "torch.Tensor") -> "torch.Tensor":
+ if t.device == device and t.dtype == torch.float32 and t.is_contiguous():
+ return t
+ return t.to(device=device, dtype=torch.float32).contiguous()
- bh = bs * hidden
- a = left_t.view(bh, n, n)
- b = right_t.view(bh, n, n)
- c = torch.empty((bh, n, n), device=left.device, dtype=torch.float32)
+ def _to_half_t(w: "torch.Tensor") -> "torch.Tensor":
+ return w.to(device=device, dtype=torch.float16).t().contiguous()
- compiled = _get_contract_compiled()
- a_ptr = make_ptr(cutlass.Float32, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- b_ptr = make_ptr(cutlass.Float32, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- c_ptr = make_ptr(cutlass.Float32, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- compiled(a_ptr, b_ptr, c_ptr, (bh, n))
+ w_stack = torch.cat(
+ [
+ _to_half_t(weights["left_proj.weight"]),
+ _to_half_t(weights["right_proj.weight"]),
+ _to_half_t(weights["left_gate.weight"]),
+ _to_half_t(weights["right_gate.weight"]),
+ _to_half_t(weights["out_gate.weight"]),
+ ],
+ dim=1,
+ ).contiguous()
- out = c.view(bs, hidden, n, n).permute(0, 2, 3, 1).contiguous()
+ w_to_out = _to_half_t(weights["to_out.weight"])
+
+ out = {
+ "norm_w": _f32_vec(weights["norm.weight"]),
+ "norm_b": _f32_vec(weights["norm.bias"]),
+ "out_norm_w": _f32_vec(weights["to_out_norm.weight"]),
+ "out_norm_b": _f32_vec(weights["to_out_norm.bias"]),
+ "w_stack": w_stack,
+ "w_to_out": w_to_out,
+ }
+ _WEIGHT_CACHE[key] = out
return out
- @torch.inference_mode()
- def custom_kernel(data: Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor], Dict[str, Any]]) -> torch.Tensor:
- x, mask, weights, config = data
+ def _get_scratch(device: "torch.device", bs: int, n: int, dim: int, hidden: int) -> dict[str, Any]:
+ import torch
+ key = (device.type, device.index if device.type == "cuda" else -1, bs, n, dim, hidden)
+ cached = _KERNEL_CACHE.get(key)
+ if cached is not None:
+ return cached
+
+ rows = bs * n * n
+ work_bytes = 32 * 1024 * 1024
+ scratch = {
+ "work": torch.empty((work_bytes,), device=device, dtype=torch.uint8),
+ "x_norm": torch.empty((rows, dim), device=device, dtype=torch.float16),
+ "y0": torch.empty((rows, 5 * hidden), device=device, dtype=torch.float16),
+ "left_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
+ "right_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
+ "og_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
+ "out_p": torch.empty((bs * hidden, n, n), device=device, dtype=torch.float16),
+ "out2": torch.empty((rows, hidden), device=device, dtype=torch.float16),
+ "y_out": torch.empty((rows, dim), device=device, dtype=torch.float32),
+ }
+ _KERNEL_CACHE[key] = scratch
+ return scratch
+
+
+ def _ensure_mask_half(mask: "torch.Tensor", device: "torch.device") -> "torch.Tensor":
+ import torch
+
+ if mask.dtype == torch.float16 and mask.is_contiguous() and mask.device == device:
+ return mask
+ return mask.to(device=device, dtype=torch.float16).contiguous()
+
+
+ def custom_kernel(data: Tuple["torch.Tensor", "torch.Tensor", Dict[str, "torch.Tensor"], Dict[str, Any]]) -> "torch.Tensor":
+ import torch
+
+ x, mask, weights, config = data
dim = int(config["dim"])
hidden_dim = int(config["hidden_dim"])
if not x.is_cuda:
- raise RuntimeError("该实现要求 x 在 CUDA 上。")
-
+ raise RuntimeError("CUDA only")
if x.dtype != torch.float32:
x = x.to(dtype=torch.float32)
+ if not x.is_contiguous():
+ x = x.contiguous()
- try:
- torch.backends.cuda.matmul.allow_tf32 = True
- torch.backends.cudnn.allow_tf32 = True
- except Exception:
- pass
+ bs, n0, n1, d0 = x.shape
+ if n0 != n1 or d0 != dim:
+ raise RuntimeError("shape mismatch")
- x = F.layer_norm(x, (dim,), weights["norm.weight"], weights["norm.bias"], 1e-5)
+ mask_h = _ensure_mask_half(mask, x.device)
- left = F.linear(x, weights["left_proj.weight"], None)
- right = F.linear(x, weights["right_proj.weight"], None)
+ ext = _load_ext()
+ wpack = _prep_weights(weights, dim=dim, hidden_dim=hidden_dim, device=x.device)
+ scratch = _get_scratch(x.device, bs=bs, n=n0, dim=dim, hidden=hidden_dim)
- mask_f = mask.unsqueeze(-1)
- if mask_f.dtype != left.dtype:
- mask_f = mask_f.to(dtype=left.dtype)
- left.mul_(mask_f)
- right.mul_(mask_f)
+ y_flat = ext.forward(
+ x,
+ mask_h,
+ wpack["norm_w"],
+ wpack["norm_b"],
+ wpack["w_stack"],
+ wpack["out_norm_w"],
+ wpack["out_norm_b"],
+ wpack["w_to_out"],
+ scratch["work"],
+ scratch["x_norm"],
+ scratch["y0"],
+ scratch["left_p"],
+ scratch["right_p"],
+ scratch["og_p"],
+ scratch["out_p"],
+ scratch["out2"],
+ scratch["y_out"],
+ )
- left_gate = F.linear(x, weights["left_gate.weight"], None)
- right_gate = F.linear(x, weights["right_gate.weight"], None)
- out_gate = F.linear(x, weights["out_gate.weight"], None)
- left_gate.sigmoid_()
- right_gate.sigmoid_()
- out_gate.sigmoid_()
- left.mul_(left_gate)
- right.mul_(right_gate)
+ return y_flat.view(bs, n0, n0, dim)
- out = _contract_outgoing(left, right)
- out = F.layer_norm(out, (hidden_dim,), weights["to_out_norm.weight"], weights["to_out_norm.bias"], 1e-5)
- out.mul_(out_gate)
- out = F.linear(out, weights["to_out.weight"], None)
- return out
-
-
__all__ = ["custom_kernel"]
scrolls · 1115 diff lines total

Best evidence level for this revision: reported

JSON