submission 34397
zilli · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 189 lines, June 9 Researcher Reciprocity License v1.0.
trimul.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-34397?include=source"interfacepython
Compatibility
measured onAMD Instinct MI300X
declared hardwareAMD Instinct MI300X
architecturesgfx942
dtypesfp32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:4e12c47e13010a14e6acb8d02c16dbb8132bdb4154950bfa2a81b15bfb2a4a82
license declaredunknown
license concludedunknown
authorszilli
imported2026-08-15
Kernel source
trimul.py189 lines
#!POPCORN leaderboard trimul
#!POPCORN gpus MI300
# This is a submission template for popcorn leaderboard 'trimul'.
# Your task is as follows:
# > For a more complete description, see: https://tinyurl.com/gpumode-trimul
# > You will be implementing a Triangle Multiplicative Update (TriMul) module that is a core operation
# > for AlphaFold3, Chai, Protenix, and other protein structure prediction models in BioML.
# >
# > The TriMul operator operates over a 4D tensor of shape [B, N, N, C].
# >
# > Your task:
# > - Implement the "outgoing" version of the TriMul operator from the AlphaFold3 paper.
# > - You will not have to compute or store gradients for this version. You will only need to implement the forward pass.
# >
# > Input:
# > - `data`: Tuple of (input: torch.Tensor, weights: Dict[str, torch.Tensor], config: Dict)
# > - input: Input tensor of shape [bs, seq_len, seq_len, dim]
# > - mask: Mask tensor of shape [bs, seq_len, seq_len]
# > - weights: Dictionary containing model weights
# > - config: Dictionary containing model configuration parameters
# >
# > Output:
# > - Tuple containing:
# > - output: Processed tensor [bs, seq_len, seq_len, dim]
# The deadline for this leaderboard is 2025-09-30 00:00:00+00:00
# You can automatically route this file to specific GPUs by adding a line
# `#!POPCORN gpus <GPUs>` to the header of this file.
# Happy hacking!
# amd_trimul_bf16.py
import os
os.environ['PYTORCH_ROCM_ARCH'] = os.getenv('PYTORCH_ROCM_ARCH', 'gfx942')
from task import input_t, output_t
import torch
from torch import nn
from torch.utils.cpp_extension import load_inline
CPP = r"""
#include <torch/extension.h>
void trimul_contract(at::Tensor A_rm, at::Tensor B_rm, at::Tensor C_rm);
""";
HIP = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
#include <hipblas/hipblas.h>
#include <ATen/hip/HIPContext.h> // at::hip::getCurrentHIPStream()
// Row-major trick: C_rm = A_rm(NxK) * B_rm(NxK)^T
// Compute as column-major: C_cm = (B_cm)^T @ A_cm (same buffers, no extra transposes)
static inline void hipblasCheck(hipblasStatus_t st, const char* what) {
TORCH_CHECK(st == HIPBLAS_STATUS_SUCCESS, what, " (hipBLAS status=", (int)st, ")");
}
void trimul_contract(at::Tensor A_rm, at::Tensor B_rm, at::Tensor C_rm) {
TORCH_CHECK(A_rm.is_cuda() && B_rm.is_cuda() && C_rm.is_cuda(), "Use HIP tensors");
TORCH_CHECK(A_rm.dim()==3 && B_rm.dim()==3 && C_rm.dim()==3, "A,B,C must be [BD,N,K],[BD,N,K],[BD,N,N]");
TORCH_CHECK(A_rm.scalar_type()==at::kBFloat16 && B_rm.scalar_type()==at::kBFloat16, "A,B must be bf16");
TORCH_CHECK(C_rm.scalar_type()==at::kFloat, "C must be fp32");
TORCH_CHECK(A_rm.is_contiguous() && B_rm.is_contiguous() && C_rm.is_contiguous(), "A,B,C must be contiguous");
const int BD = (int)A_rm.size(0);
const int N = (int)A_rm.size(1);
const int K = (int)A_rm.size(2);
// Column-major dims/strides (interpret row-major [N,K] as col-major [K,N])
const int m = N, n = N, k = K;
const int lda = K; // rows of B_cm
const int ldb = K; // rows of A_cm
const int ldc = N; // rows of C_cm
const long long strideA = (long long)K * N; // B_cm
const long long strideB = (long long)K * N; // A_cm
const long long strideC = (long long)N * N; // C_cm
static hipblasHandle_t handle = nullptr;
if (!handle) hipblasCheck(hipblasCreate(&handle), "hipblasCreate");
hipStream_t stream = at::hip::getCurrentHIPStream();
hipblasCheck(hipblasSetStream(handle, stream), "hipblasSetStream");
// Types for Ex_v2
const hipblasDatatype_t Atype = HIPBLAS_R_16B; // bf16
const hipblasDatatype_t Btype = HIPBLAS_R_16B; // bf16
const hipblasDatatype_t Ctype = HIPBLAS_R_32F; // fp32 output
const hipblasComputeType_t computeType = HIPBLAS_COMPUTE_32F; // fp32 accumulate
const hipblasGemmAlgo_t algo = HIPBLAS_GEMM_DEFAULT;
const float alpha = 1.0f, beta = 0.0f;
const void* A_ptr = (const void*)B_rm.data_ptr<at::BFloat16>(); // A_hip = B_cm
const void* B_ptr = (const void*)A_rm.data_ptr<at::BFloat16>(); // B_hip = A_cm
void* C_ptr = (void*)C_rm.data_ptr<float>(); // C_hip = C_cm
hipblasStatus_t st = hipblasGemmStridedBatchedEx(
handle,
HIPBLAS_OP_T, HIPBLAS_OP_N, // (B_cm)^T @ A_cm
m, n, k,
&alpha,
A_ptr, Atype, lda, strideA, // A = B_cm
B_ptr, Btype, ldb, strideB, // B = A_cm
&beta,
C_ptr, Ctype, ldc, strideC, // C_cm
BD,
computeType,
algo);
hipblasCheck(st, "hipblasGemmStridedBatchedEx");
}
""";
module = load_inline(
name="trimul_bf16_hipblas",
cpp_sources=[CPP],
cuda_sources=[HIP],
functions=['trimul_contract'],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx942", "-O3", "-std=c++17"],
extra_ldflags=["-lhipblas"],
)
class TriMulHIPBLAS(nn.Module):
def __init__(self, dim, hidden_dim, assume_mask_all_ones=False):
super().__init__()
self.assume_mask_all_ones = assume_mask_all_ones
self.norm = nn.LayerNorm(dim)
self.left_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.right_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.left_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.right_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.out_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)
self.to_out_norm = nn.LayerNorm(hidden_dim)
self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)
def forward(self, x, mask):
x = self.norm(x) # fp32
left = self.left_proj(x) # fp32
right = self.right_proj(x) # fp32
if not self.assume_mask_all_ones:
m = mask.unsqueeze(-1)
left = left * m
right = right * m
lg = self.left_gate(x).sigmoid()
rg = self.right_gate(x).sigmoid()
og = self.out_gate(x).sigmoid()
left = left * lg
right = right * rg
# Pack (B,N,K,D) -> (BD,N,K). contiguous() makes the one copy we need.
B, N, K, D = left.shape
A_rm = left.permute(0,3,1,2).contiguous().view(B*D, N, K).to(torch.bfloat16)
B_rm = right.permute(0,3,1,2).contiguous().view(B*D, N, K).to(torch.bfloat16)
C_rm = torch.empty((B*D, N, N), device=x.device, dtype=torch.float32)
module.trimul_contract(A_rm, B_rm, C_rm) # hipBLAS call
out = C_rm.view(B, D, N, N).permute(0, 2, 3, 1).contiguous() # [B,N,N,D]
out = self.to_out_norm(out) * og
return self.to_out(out)
def custom_kernel(data: input_t) -> output_t:
input_tensor, mask, weights, config = data
trimul = TriMulHIPBLAS(
dim=config["dim"],
hidden_dim=config["hidden_dim"],
assume_mask_all_ones=config.get("assume_mask_all_ones", False),
).to(input_tensor.device).eval()
# load weights (fp32, no grads)
for name, mod in [
("norm.weight", "norm.weight"),
("norm.bias", "norm.bias"),
("left_proj.weight", "left_proj.weight"),
("right_proj.weight", "right_proj.weight"),
("left_gate.weight", "left_gate.weight"),
("right_gate.weight", "right_gate.weight"),
("out_gate.weight", "out_gate.weight"),
("to_out_norm.weight", "to_out_norm.weight"),
("to_out_norm.bias", "to_out_norm.bias"),
("to_out.weight", "to_out.weight"),
]:
m, p = mod.split('.')
getattr(trimul, m).__setattr__(p, nn.Parameter(weights[name].to(torch.float32), requires_grad=False))
with torch.inference_mode():
return trimul(input_tensor, mask).to(torch.float32)
scrolls · 189 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 34380.
- \x2321504f50434f524e206c6561646572626f617264207472696d756c0a2321504f50434f524e2067707573204d493330300a0a2320546869732069732061207375626d697373696f6e2074656d706c61746520666f7220706f70636f726e206c6561646572626f61726420277472696d756c272e0a2320596f7572207461736b20697320617320666f6c6c6f77733a0a23203e20466f722061206d6f726520636f6d706c657465206465736372697074696f6e2c207365653a2068747470733a2f2f74696e7975726c2e636f6d2f6770756d6f64652d7472696d756c0a23203e20596f752077696c6c20626520696d706c656d656e74696e67206120547269616e676c65204d756c7469706c696361746976652055706461746520285472694d756c29206d6f64756c652074686174206973206120636f7265206f7065726174696f6e0a23203e20666f7220416c706861466f6c64332c20436861692c2050726f74656e69782c20616e64206f746865722070726f7465696e207374727563747572652070726564696374696f6e206d6f64656c7320696e2042696f4d4c2e0a23203e0a23203e20546865205472694d756c206f70657261746f72206f70657261746573206f76657220612034442074656e736f72206f66207368617065205b422c204e2c204e2c20435d2e0a23203e0a23203e20596f7572207461736b3a0a23203e202d20496d706c656d656e742074686520226f7574676f696e67222076657273696f6e206f6620746865205472694d756c206f70657261746f722066726f6d2074686520416c706861466f6c64332070617065722e0a23203e202d20596f752077696c6c206e6f74206861766520746f20636f6d70757465206f722073746f7265206772616469656e747320666f7220746869732076657273696f6e2e20596f752077696c6c206f6e6c79206e65656420746f20696d706c656d656e742074686520666f727761726420706173732e0a23203e0a23203e20496e7075743a0a23203e202d206064617461603a205475706c65206f662028696e7075743a20746f7263682e54656e736f722c20776569676874733a20446963745b7374722c20746f7263682e54656e736f725d2c20636f6e6669673a2044696374290a23203e2020202d20696e7075743a20496e7075742074656e736f72206f66207368617065205b62732c207365715f6c656e2c207365715f6c656e2c2064696d5d0a23203e2020202d206d61736b3a204d61736b2074656e736f72206f66207368617065205b62732c207365715f6c656e2c207365715f6c656e5d0a23203e2020202d20776569676874733a2044696374696f6e61727920636f6e7461696e696e67206d6f64656c20776569676874730a23203e2020202d20636f6e6669673a2044696374696f6e61727920636f6e7461696e696e67206d6f64656c20636f6e66696775726174696f6e20706172616d65746572730a23203e0a23203e204f75747075743a0a23203e202d205475706c6520636f6e7461696e696e673a0a23203e2020202d206f75747075743a2050726f6365737365642074656e736f72205b62732c207365715f6c656e2c207365715f6c656e2c2064696d5d0a232054686520646561646c696e6520666f722074686973206c6561646572626f61726420697320323032352d30392d33302030303a30303a30302b30303a30300a0a2320596f752063616e206175746f6d61746963616c6c7920726f75746520746869732066696c6520746f207370656369666963204750557320627920616464696e672061206c696e650a2320602321504f50434f524e2067707573203c475055733e6020746f2074686520686561646572206f6620746869732066696c652e0a23204861707079206861636b696e67210a0a696d706f727420746f7263680a66726f6d20746f72636820696d706f7274206e6e2c2065696e73756d0a66726f6d207461736b20696d706f727420696e7075745f742c206f75747075745f740a0a636c617373205472694d756c286e6e2e4d6f64756c65293a0a20202020646566205f5f696e69745f5f280a202020202020202073656c662c0a202020202020202064696d3a20696e742c0a202020202020202068696464656e5f64696d3a20696e742c0a20202020293a0a2020202020202020737570657228292e5f5f696e69745f5f28290a0a202020202020202073656c662e6e6f726d203d206e6e2e4c617965724e6f726d2864696d290a0a202020202020202073656c662e6c6566745f70726f6a203d206e6e2e4c696e6561722864696d2c2068696464656e5f64696d2c20626961733d46616c73652c2064747970653d746f7263682e666c6f61743332290a202020202020202073656c662e72696768745f70726f6a203d206e6e2e4c696e6561722864696d2c2068696464656e5f64696d2c20626961733d46616c73652c2064747970653d746f7263682e666c6f61743332290a0a202020202020202073656c662e6c6566745f67617465203d206e6e2e4c696e6561722864696d2c2068696464656e5f64696d2c20626961733d46616c73652c2064747970653d746f7263682e666c6f61743332290a202020202020202073656c662e72696768745f67617465203d206e6e2e4c696e6561722864696d2c2068696464656e5f64696d2c20626961733d46616c73652c2064747970653d746f7263682e666c6f61743332290a202020202020202073656c662e6f75745f67617465203d206e6e2e4c696e6561722864696d2c2068696464656e5f64696d2c20626961733d46616c73652c2064747970653d746f7263682e666c6f61743332290a0a202020202020202073656c662e746f5f6f75745f6e6f726d203d206e6e2e4c617965724e6f726d2868696464656e5f64696d290a202020202020202073656c662e746f5f6f7574203d206e6e2e4c696e6561722868696464656e5f64696d2c2064696d2c20626961733d46616c73652c2064747970653d746f7263682e666c6f61743332290a0a2020202064656620666f72776172642873656c662c20783a20746f7263682e54656e736f722c206d61736b3a20746f7263682e54656e736f7229202d3e20746f7263682e54656e736f723a0a20202020202020202222220a2020202020202020783a205b62732c207365715f6c656e2c207365715f6c656e2c2064696d5d0a20202020202020206d61736b3a205b62732c207365715f6c656e2c207365715f6c656e5d0a0a202020202020202052657475726e733a0a2020202020202020202020206f75747075743a205b62732c207365715f6c656e2c207365715f6c656e2c2064696d5d0a20202020202020202222220a202020202020202062617463685f73697a652c207365715f6c656e2c205f2c2064696d203d20782e73686170650a0a202020202020202078203d2073656c662e6e6f726d2878290a202020202020202078203d20782e746f28746f7263682e666c6f61743332290a0a20202020202020206c656674203d2073656c662e6c6566745f70726f6a28782e746f28746f7263682e666c6f6174333229290a20202020202020207269676874203d2073656c662e72696768745f70726f6a28782e746f28746f7263682e666c6f6174333229290a0a20202020202020206d61736b203d206d61736b2e756e73717565657a65282d31290a20202020202020206c656674203d206c656674202a206d61736b0a20202020202020207269676874203d207269676874202a206d61736b0a0a20202020202020206c6566745f67617465203d2073656c662e6c6566745f6761746528782e746f28746f7263682e666c6f6174333229292e7369676d6f696428290a202020202020202072696768745f67617465203d2073656c662e72696768745f6761746528782e746f28746f7263682e666c6f6174333229292e7369676d6f696428290a20202020202020206f75745f67617465203d2073656c662e6f75745f6761746528782e746f28746f7263682e666c6f6174333229292e7369676d6f696428290a0a20202020202020206c656674203d206c656674202a206c6566745f676174650a20202020202020207269676874203d207269676874202a2072696768745f676174650a0a20202020202020206f7574203d2065696e73756d28272e2e2e2069206b20642c202e2e2e206a206b2064202d3e202e2e2e2069206a2064272c206c6566742e746f28746f7263682e62666c6f61743136292c2072696768742e746f28746f7263682e62666c6f6174313629290a20202020202020202320546869732065696e73756d206973207468652073616d652061732074686520666f6c6c6f77696e673a0a202020202020202023206f7574203d20746f7263682e7a65726f732862617463685f73697a652c207365715f6c656e2c207365715f6c656e2c2064696d2c206465766963653d782e646576696365290a0a202020202020202023202320436f6d70757465207573696e67206e6573746564206c6f6f70730a20202020202020202320666f72206220696e2072616e67652862617463685f73697a65293a0a2020202020202020232020202020666f72206920696e2072616e6765287365715f6c656e293a0a202020202020202023202020202020202020666f72206a20696e2072616e6765287365715f6c656e293a0a202020202020202023202020202020202020202020202320436f6d707574652065616368206f757470757420656c656d656e740a20202020202020202320202020202020202020202020666f72206b20696e2072616e6765287365715f6c656e293a0a20202020202020202320202020202020202020202020202020206f75745b622c20692c206a5d202b3d206c6566745b622c20692c206b2c203a5d202a2072696768745b622c206a2c206b2c203a5d0a0a20202020202020206f7574203d206f75742e746f28746f7263682e666c6f61743332290a20202020202020206f7574203d2073656c662e746f5f6f75745f6e6f726d286f7574290a20202020202020206f7574203d206f7574202a206f75745f676174650a202020202020202072657475726e2073656c662e746f5f6f7574286f7574290a0a0a64656620637573746f6d5f6b65726e656c28646174613a20696e7075745f7429202d3e206f75747075745f743a0a202020202222220a202020205265666572656e636520696d706c656d656e746174696f6e206f66205472694d756c207573696e67205079546f7263682e0a0a20202020417267733a0a2020202020202020646174613a205475706c65206f662028696e7075743a20746f7263682e54656e736f722c206d61736b3a20746f7263682e54656e736f722c20776569676874733a20446963745b7374722c20746f7263682e54656e736f725d2c20636f6e6669673a2044696374290a2020202020202020202020202d20696e7075743a20496e7075742074656e736f72206f66207368617065205b62617463685f73697a652c207365715f6c656e2c207365715f6c656e2c2064696d5d0a2020202020202020202020202d206d61736b3a204d61736b2074656e736f72206f66207368617065205b62617463685f73697a652c207365715f6c656e2c207365715f6c656e5d0a2020202020202020202020202d20776569676874733a2044696374696f6e61727920636f6e7461696e696e67206d6f64656c20776569676874730a2020202020202020202020202d20636f6e6669673a2044696374696f6e61727920636f6e7461696e696e67206d6f64656c20636f6e66696775726174696f6e20706172616d65746572730a202020202222220a20202020696e7075745f74656e736f722c206d61736b2c20776569676874732c20636f6e666967203d20646174610a202020207472696d756c203d205472694d756c28636f6e6669675b2264696d225d2c20636f6e6669675b2268696464656e5f64696d225d292e746f28696e7075745f74656e736f722e646576696365290a0a20202020232046696c6c20696e2074686520676976656e2077656967687473206f6620746865206d6f64656c0a202020207472696d756c2e6e6f726d2e776569676874203d206e6e2e506172616d6574657228776569676874735b276e6f726d2e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e6c6566745f70726f6a2e776569676874203d206e6e2e506172616d6574657228776569676874735b276c6566745f70726f6a2e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e72696768745f70726f6a2e776569676874203d206e6e2e506172616d6574657228776569676874735b2772696768745f70726f6a2e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e6c6566745f676174652e776569676874203d206e6e2e506172616d6574657228776569676874735b276c6566745f676174652e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e72696768745f676174652e776569676874203d206e6e2e506172616d6574657228776569676874735b2772696768745f676174652e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e6f75745f676174652e776569676874203d206e6e2e506172616d6574657228776569676874735b276f75745f676174652e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e746f5f6f75745f6e6f726d2e776569676874203d206e6e2e506172616d6574657228776569676874735b27746f5f6f75745f6e6f726d2e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e746f5f6f75742e776569676874203d206e6e2e506172616d6574657228776569676874735b27746f5f6f75742e776569676874275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e6e6f726d2e62696173203d206e6e2e506172616d6574657228776569676874735b276e6f726d2e62696173275d2e746f28746f7263682e666c6f6174333229290a202020207472696d756c2e746f5f6f75745f6e6f726d2e62696173203d206e6e2e506172616d6574657228776569676874735b27746f5f6f75745f6e6f726d2e62696173275d2e746f28746f7263682e666c6f6174333229290a0a202020206f7574707574203d207472696d756c28696e7075745f74656e736f722c206d61736b292e746f28746f7263682e666c6f61743332290a0a2020202072657475726e206f75747075740aNo newline at end of file+ #!POPCORN leaderboard trimul+ #!POPCORN gpus MI300++ # This is a submission template for popcorn leaderboard 'trimul'.+ # Your task is as follows:+ # > For a more complete description, see: https://tinyurl.com/gpumode-trimul+ # > You will be implementing a Triangle Multiplicative Update (TriMul) module that is a core operation+ # > for AlphaFold3, Chai, Protenix, and other protein structure prediction models in BioML.+ # >+ # > The TriMul operator operates over a 4D tensor of shape [B, N, N, C].+ # >+ # > Your task:+ # > - Implement the "outgoing" version of the TriMul operator from the AlphaFold3 paper.+ # > - You will not have to compute or store gradients for this version. You will only need to implement the forward pass.+ # >+ # > Input:+ # > - `data`: Tuple of (input: torch.Tensor, weights: Dict[str, torch.Tensor], config: Dict)+ # > - input: Input tensor of shape [bs, seq_len, seq_len, dim]+ # > - mask: Mask tensor of shape [bs, seq_len, seq_len]+ # > - weights: Dictionary containing model weights+ # > - config: Dictionary containing model configuration parameters+ # >+ # > Output:+ # > - Tuple containing:+ # > - output: Processed tensor [bs, seq_len, seq_len, dim]+ # The deadline for this leaderboard is 2025-09-30 00:00:00+00:00++ # You can automatically route this file to specific GPUs by adding a line+ # `#!POPCORN gpus <GPUs>` to the header of this file.+ # Happy hacking!++ # amd_trimul_bf16.py+ import os+ os.environ['PYTORCH_ROCM_ARCH'] = os.getenv('PYTORCH_ROCM_ARCH', 'gfx942')+ from task import input_t, output_t+ import torch+ from torch import nn+ from torch.utils.cpp_extension import load_inline++ CPP = r"""+ #include <torch/extension.h>+ void trimul_contract(at::Tensor A_rm, at::Tensor B_rm, at::Tensor C_rm);+ """;++ HIP = r"""+ #include <hip/hip_runtime.h>+ #include <torch/extension.h>+ #include <hipblas/hipblas.h>+ #include <ATen/hip/HIPContext.h> // at::hip::getCurrentHIPStream()++ // Row-major trick: C_rm = A_rm(NxK) * B_rm(NxK)^T+ // Compute as column-major: C_cm = (B_cm)^T @ A_cm (same buffers, no extra transposes)++ static inline void hipblasCheck(hipblasStatus_t st, const char* what) {+ TORCH_CHECK(st == HIPBLAS_STATUS_SUCCESS, what, " (hipBLAS status=", (int)st, ")");+ }++ void trimul_contract(at::Tensor A_rm, at::Tensor B_rm, at::Tensor C_rm) {+ TORCH_CHECK(A_rm.is_cuda() && B_rm.is_cuda() && C_rm.is_cuda(), "Use HIP tensors");+ TORCH_CHECK(A_rm.dim()==3 && B_rm.dim()==3 && C_rm.dim()==3, "A,B,C must be [BD,N,K],[BD,N,K],[BD,N,N]");+ TORCH_CHECK(A_rm.scalar_type()==at::kBFloat16 && B_rm.scalar_type()==at::kBFloat16, "A,B must be bf16");+ TORCH_CHECK(C_rm.scalar_type()==at::kFloat, "C must be fp32");+ TORCH_CHECK(A_rm.is_contiguous() && B_rm.is_contiguous() && C_rm.is_contiguous(), "A,B,C must be contiguous");++ const int BD = (int)A_rm.size(0);+ const int N = (int)A_rm.size(1);+ const int K = (int)A_rm.size(2);++ // Column-major dims/strides (interpret row-major [N,K] as col-major [K,N])+ const int m = N, n = N, k = K;+ const int lda = K; // rows of B_cm+ const int ldb = K; // rows of A_cm+ const int ldc = N; // rows of C_cm+ const long long strideA = (long long)K * N; // B_cm+ const long long strideB = (long long)K * N; // A_cm+ const long long strideC = (long long)N * N; // C_cm++ static hipblasHandle_t handle = nullptr;+ if (!handle) hipblasCheck(hipblasCreate(&handle), "hipblasCreate");+ hipStream_t stream = at::hip::getCurrentHIPStream();+ hipblasCheck(hipblasSetStream(handle, stream), "hipblasSetStream");++ // Types for Ex_v2+ const hipblasDatatype_t Atype = HIPBLAS_R_16B; // bf16+ const hipblasDatatype_t Btype = HIPBLAS_R_16B; // bf16+ const hipblasDatatype_t Ctype = HIPBLAS_R_32F; // fp32 output+ const hipblasComputeType_t computeType = HIPBLAS_COMPUTE_32F; // fp32 accumulate+ const hipblasGemmAlgo_t algo = HIPBLAS_GEMM_DEFAULT;++ const float alpha = 1.0f, beta = 0.0f;+ const void* A_ptr = (const void*)B_rm.data_ptr<at::BFloat16>(); // A_hip = B_cm+ const void* B_ptr = (const void*)A_rm.data_ptr<at::BFloat16>(); // B_hip = A_cm+ void* C_ptr = (void*)C_rm.data_ptr<float>(); // C_hip = C_cm++ hipblasStatus_t st = hipblasGemmStridedBatchedEx(+ handle,+ HIPBLAS_OP_T, HIPBLAS_OP_N, // (B_cm)^T @ A_cm+ m, n, k,+ &alpha,+ A_ptr, Atype, lda, strideA, // A = B_cm+ B_ptr, Btype, ldb, strideB, // B = A_cm+ &beta,+ C_ptr, Ctype, ldc, strideC, // C_cm+ BD,+ computeType,+ algo);++ hipblasCheck(st, "hipblasGemmStridedBatchedEx");+ }+ """;++ module = load_inline(+ name="trimul_bf16_hipblas",+ cpp_sources=[CPP],+ cuda_sources=[HIP],+ functions=['trimul_contract'],+ verbose=False,+ extra_cuda_cflags=["--offload-arch=gfx942", "-O3", "-std=c++17"],+ extra_ldflags=["-lhipblas"],+ )++ class TriMulHIPBLAS(nn.Module):+ def __init__(self, dim, hidden_dim, assume_mask_all_ones=False):+ super().__init__()+ self.assume_mask_all_ones = assume_mask_all_ones+ self.norm = nn.LayerNorm(dim)+ self.left_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.right_proj = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.left_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.right_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.out_gate = nn.Linear(dim, hidden_dim, bias=False, dtype=torch.float32)+ self.to_out_norm = nn.LayerNorm(hidden_dim)+ self.to_out = nn.Linear(hidden_dim, dim, bias=False, dtype=torch.float32)++ def forward(self, x, mask):+ x = self.norm(x) # fp32+ left = self.left_proj(x) # fp32+ right = self.right_proj(x) # fp32++ if not self.assume_mask_all_ones:+ m = mask.unsqueeze(-1)+ left = left * m+ right = right * m++ lg = self.left_gate(x).sigmoid()+ rg = self.right_gate(x).sigmoid()+ og = self.out_gate(x).sigmoid()+ left = left * lg+ right = right * rg++ # Pack (B,N,K,D) -> (BD,N,K). contiguous() makes the one copy we need.+ B, N, K, D = left.shape+ A_rm = left.permute(0,3,1,2).contiguous().view(B*D, N, K).to(torch.bfloat16)+ B_rm = right.permute(0,3,1,2).contiguous().view(B*D, N, K).to(torch.bfloat16)++ C_rm = torch.empty((B*D, N, N), device=x.device, dtype=torch.float32)+ module.trimul_contract(A_rm, B_rm, C_rm) # hipBLAS call++ out = C_rm.view(B, D, N, N).permute(0, 2, 3, 1).contiguous() # [B,N,N,D]+ out = self.to_out_norm(out) * og+ return self.to_out(out)++ def custom_kernel(data: input_t) -> output_t:+ input_tensor, mask, weights, config = data+ trimul = TriMulHIPBLAS(+ dim=config["dim"],+ hidden_dim=config["hidden_dim"],+ assume_mask_all_ones=config.get("assume_mask_all_ones", False),+ ).to(input_tensor.device).eval()++ # load weights (fp32, no grads)+ for name, mod in [+ ("norm.weight", "norm.weight"),+ ("norm.bias", "norm.bias"),+ ("left_proj.weight", "left_proj.weight"),+ ("right_proj.weight", "right_proj.weight"),+ ("left_gate.weight", "left_gate.weight"),+ ("right_gate.weight", "right_gate.weight"),+ ("out_gate.weight", "out_gate.weight"),+ ("to_out_norm.weight", "to_out_norm.weight"),+ ("to_out_norm.bias", "to_out_norm.bias"),+ ("to_out.weight", "to_out.weight"),+ ]:+ m, p = mod.split('.')+ getattr(trimul, m).__setattr__(p, nn.Parameter(weights[name].to(torch.float32), requires_grad=False))++ with torch.inference_mode():+ return trimul(input_tensor, mask).to(torch.float32)
scrolls · 190 diff lines total
Best evidence level for this revision: reported
JSON