submission 44222
davidberard · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 323 lines, June 9 Researcher Reciprocity License v1.0.
v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-trimul-44222?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:06c27237cbfe1c5e6edb1c83fdfd8612837b32cce87231bc2e3700933b93d303
license declaredunknown
license concludedunknown
authorsdavidberard
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
configs.append(triton.Config({mma
accumulator1 = tl.dot(a, b1.T, accumulator1, allow_tf32=True)num-warps = 8
}, num_stages=4, num_warps=8))persistent-kernel
Persistent dual matrix multiplication: A @ B1.T and A @ B2.T using on-device TMA descriptors.stages = 4
}, num_stages=4, num_warps=8))Kernel source
v2.py323 lines
# from utils import make_match_reference, DisableCuDNNTF32
from task import input_t, output_t
import torch
from torch import nn, einsum
import math
import os
import triton
import triton.language as tl
# The flag below controls whether to allow TF32 on matmul. This flag defaults to False
# in PyTorch 1.12 and later.
torch.backends.cuda.matmul.allow_tf32 = True
# The flag below controls whether to allow TF32 on cuDNN. This flag defaults to True.
torch.backends.cudnn.allow_tf32 = True
# Set allocator for TMA descriptors (required for on-device TMA)
def alloc_fn(size: int, alignment: int, stream=None):
return torch.empty(size, device="cuda", dtype=torch.int8)
triton.set_allocator(alloc_fn)
os.environ['TRITON_PRINT_AUTOTUNING'] = '1'
os.environ['MLIR_ENABLE_DIAGNOSTICS'] = 'warnings,remarks'
# Reference code in PyTorch
class TriMul(nn.Module):
# Based on https://github.com/lucidrains/triangle-multiplicative-module/blob/main/triangle_multiplicative_module/triangle_multiplicative_module.py
def __init__(
self,
dim: int,
hidden_dim: int,
):
super().__init__()
self.norm = nn.LayerNorm(dim)
self.left_proj = nn.Linear(dim, hidden_dim, bias=False)
self.right_proj = nn.Linear(dim, hidden_dim, bias=False)
self.left_gate = nn.Linear(dim, hidden_dim, bias=False)
self.right_gate = nn.Linear(dim, hidden_dim, bias=False)
self.out_gate = nn.Linear(dim, hidden_dim, bias=False)
self.to_out_norm = nn.LayerNorm(hidden_dim)
self.to_out = nn.Linear(hidden_dim, dim, bias=False)
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
"""
x: [bs, seq_len, seq_len, dim]
mask: [bs, seq_len, seq_len]
Returns:
output: [bs, seq_len, seq_len, dim]
"""
batch_size, seq_len, _, dim = x.shape
x = self.norm(x)
left = self.left_proj(x)
right = self.right_proj(x)
mask = mask.unsqueeze(-1)
left = left * mask
right = right * mask
left_gate = self.left_gate(x).sigmoid()
right_gate = self.right_gate(x).sigmoid()
out_gate = self.out_gate(x).sigmoid()
left = left * left_gate
right = right * right_gate
out = einsum('... i k d, ... j k d -> ... i j d', left, right)
# This einsum is the same as the following:
# out = torch.zeros(batch_size, seq_len, seq_len, dim, device=x.device)
# # Compute using nested loops
# for b in range(batch_size):
# for i in range(seq_len):
# for j in range(seq_len):
# # Compute each output element
# for k in range(seq_len):
# out[b, i, j] += left[b, i, k, :] * right[b, j, k, :]
out = self.to_out_norm(out)
out = out * out_gate
return self.to_out(out)
def two_mm_kernel_configs():
configs = []
for BLOCK_M in [64, 128]:
for BLOCK_N in [64, 128, 256]:
for BLOCK_K in [32, 64, 128]:
configs.append(triton.Config({
'BLOCK_M': BLOCK_M,
'BLOCK_N': BLOCK_N,
'BLOCK_K': BLOCK_K,
'GROUP_SIZE_M': 8
}, num_stages=4, num_warps=8))
return configs
@triton.autotune(
two_mm_kernel_configs(), key=["M", "N", "K"]
)
@triton.jit
def two_mm_kernel(a_ptr, b1_ptr, b2_ptr, c1_ptr, c2_ptr, mask_ptr, M, N, K, stride_am, stride_ak, stride_b1k, stride_b1n, stride_b2k, stride_b2n, stride_c1m, stride_c1n, stride_c2m, stride_c2n, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, NUM_SMS: tl.constexpr):
# Persistent kernel using on-device TMA descriptors
start_pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
k_tiles = tl.cdiv(K, BLOCK_K)
num_tiles = num_pid_m * num_pid_n
# Create on-device TMA descriptors
a_desc = tl._experimental_make_tensor_descriptor(
a_ptr,
shape=[M, K],
strides=[stride_am, stride_ak],
block_shape=[BLOCK_M, BLOCK_K],
)
b1_desc = tl._experimental_make_tensor_descriptor(
b1_ptr,
shape=[N, K],
strides=[stride_b1n, stride_b1k],
block_shape=[BLOCK_N, BLOCK_K],
)
b2_desc = tl._experimental_make_tensor_descriptor(
b2_ptr,
shape=[N, K],
strides=[stride_b2n, stride_b2k],
block_shape=[BLOCK_N, BLOCK_K],
)
# tile_id_c is used in the epilogue to break the dependency between
# the prologue and the epilogue
tile_id_c = start_pid - NUM_SMS
num_pid_in_group = GROUP_SIZE_M * num_pid_n
# Persistent loop over tiles
for tile_id in tl.range(start_pid, num_tiles, NUM_SMS, flatten=False):
# Calculate PID for this tile using improved swizzling
group_id = tile_id // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (tile_id % group_size_m)
pid_n = (tile_id % num_pid_in_group) // group_size_m
# Calculate block offsets
offs_am = pid_m * BLOCK_M
offs_bn = pid_n * BLOCK_N
# Initialize accumulators for both outputs
accumulator1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
accumulator2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Main computation loop over K dimension
for ki in range(k_tiles):
offs_k = ki * BLOCK_K
# Load blocks from A, B1, B2 using on-device TMA
a = a_desc.load([offs_am, offs_k])
b1 = b1_desc.load([offs_bn, offs_k])
b2 = b2_desc.load([offs_bn, offs_k])
# Perform matrix multiplications: A @ B1.T and A @ B2.T using TF32
accumulator1 = tl.dot(a, b1.T, accumulator1, allow_tf32=True)
accumulator2 = tl.dot(a, b2.T, accumulator2, allow_tf32=True)
# Store results using separate tile_id_c for epilogue
tile_id_c += NUM_SMS
group_id = tile_id_c // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (tile_id_c % group_size_m)
pid_n = (tile_id_c % num_pid_in_group) // group_size_m
# Calculate output offsets and pointers
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# Create masks for bounds checking
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
# Calculate pointer addresses
c1_ptrs = c1_ptr + stride_c1m * offs_cm[:, None] + stride_c1n * offs_cn[None, :]
c2_ptrs = c2_ptr + stride_c2m * offs_cm[:, None] + stride_c2n * offs_cn[None, :]
mask = tl.load(mask_ptr + offs_cm, mask=(offs_cm < M))
# Broadcast mask to match accumulator dimensions [BLOCK_M, BLOCK_N]
mask_2d = mask[:, None] # Convert to [BLOCK_M, 1] then broadcast
accumulator1 = tl.where(mask_2d, accumulator1, 0)
accumulator2 = tl.where(mask_2d, accumulator2, 0)
# Convert to appropriate output dtype and store with normal tl.store
c1 = accumulator1.to(c1_ptr.dtype.element_ty)
c2 = accumulator2.to(c2_ptr.dtype.element_ty)
tl.store(c1_ptrs, c1, mask=c_mask)
tl.store(c2_ptrs, c2, mask=c_mask)
def two_mm(A, B1, B2, mask):
"""
Persistent dual matrix multiplication: A @ B1.T and A @ B2.T using on-device TMA descriptors.
Args:
A: [..., K] tensor (arbitrary leading dimensions)
B1: [N, K] matrix (will be transposed)
B2: [N, K] matrix (will be transposed)
Returns:
(C1, C2): Tuple of result tensors [..., N] with same leading dims as A
"""
# Check constraints
assert A.shape[-1] == B1.shape[1] == B2.shape[1], "Incompatible K dimensions"
assert A.dtype == B1.dtype == B2.dtype, "Incompatible dtypes"
# Get dimensions
original_shape = A.shape[:-1] # All dimensions except the last
K = A.shape[-1]
N = B1.shape[0]
dtype = A.dtype
# Flatten A to 2D for kernel processing
A_2d = A.view(-1, K) # [M, K] where M is product of all leading dims
M = A_2d.shape[0]
# Allocate outputs as 2D then reshape
C1_2d = torch.empty((M, N), device=A.device, dtype=dtype)
C2_2d = torch.empty((M, N), device=A.device, dtype=dtype)
# Get number of streaming multiprocessors
NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count
# Launch persistent kernel with limited number of blocks
grid = lambda META: (min(NUM_SMS, triton.cdiv(M, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"])),)
two_mm_kernel[grid](
A_2d, B1, B2, C1_2d, C2_2d, mask,
M, N, K,
A_2d.stride(0), A_2d.stride(1),
B1.stride(1), B1.stride(0), # Note: B1 is [N, K] but we access as transposed
B2.stride(1), B2.stride(0), # Note: B2 is [N, K] but we access as transposed
C1_2d.stride(0), C1_2d.stride(1),
C2_2d.stride(0), C2_2d.stride(1),
NUM_SMS=NUM_SMS
)
# Reshape outputs back to original shape + N dimension
output_shape = original_shape + (N,)
C1 = C1_2d.view(output_shape)
C2 = C2_2d.view(output_shape)
return C1, C2
def custom_kernel(data: input_t) -> output_t:
"""
Reference implementation of TriMul using PyTorch.
Args:
data: Tuple of (input: torch.Tensor, mask: torch.Tensor, weights: Dict[str, torch.Tensor], config: Dict)
- input: Input tensor of shape [batch_size, seq_len, seq_len, dim]
- mask: Mask tensor of shape [batch_size, seq_len, seq_len]
- weights: Dictionary containing model weights
- config: Dictionary containing model configuration parameters
"""
input_tensor, mask, weights, config = data
hidden_dim = config["hidden_dim"]
# trimul = TriMul(dim=config["dim"], hidden_dim=config["hidden_dim"]).to(input_tensor.device)
x = input_tensor
batch_size, seq_len, _, dim = x.shape
x = torch.nn.functional.layer_norm(x, (dim,), eps=1e-5, weight=weights['norm.weight'], bias=weights['norm.bias'])
left, right = two_mm(x, weights["left_proj.weight"], weights["right_proj.weight"], mask)
# left = torch.nn.functional.linear(x, weights['left_proj.weight'].to(torch.float16))
# right = torch.nn.functional.linear(x, weights['right_proj.weight'].to(torch.float16))
# left = left * mask.unsqueeze(-1)
# right = right * mask.unsqueeze(-1)
'''
left = left.to(torch.float32)
right = right.to(torch.float32)
x = x.to(torch.float32)
'''
left_gate = torch.nn.functional.linear(x, weights['left_gate.weight']).sigmoid()
right_gate = torch.nn.functional.linear(x, weights['right_gate.weight']).sigmoid()
out_gate = torch.nn.functional.linear(x, weights['out_gate.weight']).sigmoid()
left = left * left_gate
right = right * right_gate
out = einsum('... i k d, ... j k d -> ... i j d', left, right)
out = torch.nn.functional.layer_norm(out, (hidden_dim,), eps=1e-5, weight=weights['to_out_norm.weight'], bias=weights['to_out_norm.bias'])
out = out * out_gate
return torch.nn.functional.linear(out, weights['to_out.weight'])
'''
# Fill in the given weights of the model
trimul.norm.weight = nn.Parameter(weights['norm.weight'])
trimul.norm.bias = nn.Parameter(weights['norm.bias'])
trimul.left_proj.weight = nn.Parameter(weights['left_proj.weight'])
trimul.right_proj.weight = nn.Parameter(weights['right_proj.weight'])
trimul.left_gate.weight = nn.Parameter(weights['left_gate.weight'])
trimul.right_gate.weight = nn.Parameter(weights['right_gate.weight'])
trimul.out_gate.weight = nn.Parameter(weights['out_gate.weight'])
trimul.to_out_norm.weight = nn.Parameter(weights['to_out_norm.weight'])
trimul.to_out_norm.bias = nn.Parameter(weights['to_out_norm.bias'])
trimul.to_out.weight = nn.Parameter(weights['to_out.weight'])
output = trimul(input_tensor, mask)
return output
'''scrolls · 323 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 44207.
⋯ 3 unchanged linesimport torchfrom torch import nn, einsumimport math+ import os+ import triton+ import triton.language as tl+# The flag below controls whether to allow TF32 on matmul. This flag defaults to False# in PyTorch 1.12 and later.torch.backends.cuda.matmul.allow_tf32 = True⋯ 1 unchanged lines# The flag below controls whether to allow TF32 on cuDNN. This flag defaults to True.torch.backends.cudnn.allow_tf32 = True+ # Set allocator for TMA descriptors (required for on-device TMA)+ def alloc_fn(size: int, alignment: int, stream=None):+ return torch.empty(size, device="cuda", dtype=torch.int8)++ triton.set_allocator(alloc_fn)++ os.environ['TRITON_PRINT_AUTOTUNING'] = '1'+ os.environ['MLIR_ENABLE_DIAGNOSTICS'] = 'warnings,remarks'+# Reference code in PyTorchclass TriMul(nn.Module):# Based on https://github.com/lucidrains/triangle-multiplicative-module/blob/main/triangle_multiplicative_module/triangle_multiplicative_module.py⋯ 58 unchanged linesout = out * out_gatereturn self.to_out(out)+ def two_mm_kernel_configs():+ configs = []+ for BLOCK_M in [64, 128]:+ for BLOCK_N in [64, 128, 256]:+ for BLOCK_K in [32, 64, 128]:+ configs.append(triton.Config({+ 'BLOCK_M': BLOCK_M,+ 'BLOCK_N': BLOCK_N,+ 'BLOCK_K': BLOCK_K,+ 'GROUP_SIZE_M': 8+ }, num_stages=4, num_warps=8))+ return configs+ @triton.autotune(+ two_mm_kernel_configs(), key=["M", "N", "K"]+ )+ @triton.jit+ def two_mm_kernel(a_ptr, b1_ptr, b2_ptr, c1_ptr, c2_ptr, mask_ptr, M, N, K, stride_am, stride_ak, stride_b1k, stride_b1n, stride_b2k, stride_b2n, stride_c1m, stride_c1n, stride_c2m, stride_c2n, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, NUM_SMS: tl.constexpr):+ # Persistent kernel using on-device TMA descriptors+ start_pid = tl.program_id(axis=0)+ num_pid_m = tl.cdiv(M, BLOCK_M)+ num_pid_n = tl.cdiv(N, BLOCK_N)+ k_tiles = tl.cdiv(K, BLOCK_K)+ num_tiles = num_pid_m * num_pid_n++ # Create on-device TMA descriptors+ a_desc = tl._experimental_make_tensor_descriptor(+ a_ptr,+ shape=[M, K],+ strides=[stride_am, stride_ak],+ block_shape=[BLOCK_M, BLOCK_K],+ )+ b1_desc = tl._experimental_make_tensor_descriptor(+ b1_ptr,+ shape=[N, K],+ strides=[stride_b1n, stride_b1k],+ block_shape=[BLOCK_N, BLOCK_K],+ )+ b2_desc = tl._experimental_make_tensor_descriptor(+ b2_ptr,+ shape=[N, K],+ strides=[stride_b2n, stride_b2k],+ block_shape=[BLOCK_N, BLOCK_K],+ )++ # tile_id_c is used in the epilogue to break the dependency between+ # the prologue and the epilogue+ tile_id_c = start_pid - NUM_SMS+ num_pid_in_group = GROUP_SIZE_M * num_pid_n++ # Persistent loop over tiles+ for tile_id in tl.range(start_pid, num_tiles, NUM_SMS, flatten=False):+ # Calculate PID for this tile using improved swizzling+ group_id = tile_id // num_pid_in_group+ first_pid_m = group_id * GROUP_SIZE_M+ group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)+ pid_m = first_pid_m + (tile_id % group_size_m)+ pid_n = (tile_id % num_pid_in_group) // group_size_m++ # Calculate block offsets+ offs_am = pid_m * BLOCK_M+ offs_bn = pid_n * BLOCK_N++ # Initialize accumulators for both outputs+ accumulator1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)+ accumulator2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)++ # Main computation loop over K dimension+ for ki in range(k_tiles):+ offs_k = ki * BLOCK_K+ # Load blocks from A, B1, B2 using on-device TMA+ a = a_desc.load([offs_am, offs_k])+ b1 = b1_desc.load([offs_bn, offs_k])+ b2 = b2_desc.load([offs_bn, offs_k])++ # Perform matrix multiplications: A @ B1.T and A @ B2.T using TF32+ accumulator1 = tl.dot(a, b1.T, accumulator1, allow_tf32=True)+ accumulator2 = tl.dot(a, b2.T, accumulator2, allow_tf32=True)++ # Store results using separate tile_id_c for epilogue+ tile_id_c += NUM_SMS+ group_id = tile_id_c // num_pid_in_group+ first_pid_m = group_id * GROUP_SIZE_M+ group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)+ pid_m = first_pid_m + (tile_id_c % group_size_m)+ pid_n = (tile_id_c % num_pid_in_group) // group_size_m++ # Calculate output offsets and pointers+ offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)+ offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)++ # Create masks for bounds checking+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)++ # Calculate pointer addresses+ c1_ptrs = c1_ptr + stride_c1m * offs_cm[:, None] + stride_c1n * offs_cn[None, :]+ c2_ptrs = c2_ptr + stride_c2m * offs_cm[:, None] + stride_c2n * offs_cn[None, :]++ mask = tl.load(mask_ptr + offs_cm, mask=(offs_cm < M))++ # Broadcast mask to match accumulator dimensions [BLOCK_M, BLOCK_N]+ mask_2d = mask[:, None] # Convert to [BLOCK_M, 1] then broadcast+ accumulator1 = tl.where(mask_2d, accumulator1, 0)+ accumulator2 = tl.where(mask_2d, accumulator2, 0)++ # Convert to appropriate output dtype and store with normal tl.store+ c1 = accumulator1.to(c1_ptr.dtype.element_ty)+ c2 = accumulator2.to(c2_ptr.dtype.element_ty)++ tl.store(c1_ptrs, c1, mask=c_mask)+ tl.store(c2_ptrs, c2, mask=c_mask)++ def two_mm(A, B1, B2, mask):+ """+ Persistent dual matrix multiplication: A @ B1.T and A @ B2.T using on-device TMA descriptors.++ Args:+ A: [..., K] tensor (arbitrary leading dimensions)+ B1: [N, K] matrix (will be transposed)+ B2: [N, K] matrix (will be transposed)++ Returns:+ (C1, C2): Tuple of result tensors [..., N] with same leading dims as A+ """+ # Check constraints+ assert A.shape[-1] == B1.shape[1] == B2.shape[1], "Incompatible K dimensions"+ assert A.dtype == B1.dtype == B2.dtype, "Incompatible dtypes"++ # Get dimensions+ original_shape = A.shape[:-1] # All dimensions except the last+ K = A.shape[-1]+ N = B1.shape[0]+ dtype = A.dtype++ # Flatten A to 2D for kernel processing+ A_2d = A.view(-1, K) # [M, K] where M is product of all leading dims+ M = A_2d.shape[0]++ # Allocate outputs as 2D then reshape+ C1_2d = torch.empty((M, N), device=A.device, dtype=dtype)+ C2_2d = torch.empty((M, N), device=A.device, dtype=dtype)++ # Get number of streaming multiprocessors+ NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count+++ # Launch persistent kernel with limited number of blocks+ grid = lambda META: (min(NUM_SMS, triton.cdiv(M, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"])),)++ two_mm_kernel[grid](+ A_2d, B1, B2, C1_2d, C2_2d, mask,+ M, N, K,+ A_2d.stride(0), A_2d.stride(1),+ B1.stride(1), B1.stride(0), # Note: B1 is [N, K] but we access as transposed+ B2.stride(1), B2.stride(0), # Note: B2 is [N, K] but we access as transposed+ C1_2d.stride(0), C1_2d.stride(1),+ C2_2d.stride(0), C2_2d.stride(1),+ NUM_SMS=NUM_SMS+ )++ # Reshape outputs back to original shape + N dimension+ output_shape = original_shape + (N,)+ C1 = C1_2d.view(output_shape)+ C2 = C2_2d.view(output_shape)++ return C1, C2+def custom_kernel(data: input_t) -> output_t:"""Reference implementation of TriMul using PyTorch.⋯ 16 unchanged linesx = torch.nn.functional.layer_norm(x, (dim,), eps=1e-5, weight=weights['norm.weight'], bias=weights['norm.bias'])- left = torch.nn.functional.linear(x, weights['left_proj.weight'])- right = torch.nn.functional.linear(x, weights['right_proj.weight'])+ left, right = two_mm(x, weights["left_proj.weight"], weights["right_proj.weight"], mask)+ # left = torch.nn.functional.linear(x, weights['left_proj.weight'].to(torch.float16))+ # right = torch.nn.functional.linear(x, weights['right_proj.weight'].to(torch.float16))- left = left * mask.unsqueeze(-1)- right = right * mask.unsqueeze(-1)+ # left = left * mask.unsqueeze(-1)+ # right = right * mask.unsqueeze(-1)+ '''+ left = left.to(torch.float32)+ right = right.to(torch.float32)+ x = x.to(torch.float32)+ '''+left_gate = torch.nn.functional.linear(x, weights['left_gate.weight']).sigmoid()right_gate = torch.nn.functional.linear(x, weights['right_gate.weight']).sigmoid()out_gate = torch.nn.functional.linear(x, weights['out_gate.weight']).sigmoid()
scrolls · 226 diff lines total
Best evidence level for this revision: reported
JSON