-
Notifications
You must be signed in to change notification settings - Fork 493
Expand file tree
/
Copy pathaddr_.py
More file actions
99 lines (83 loc) · 2.73 KB
/
Copy pathaddr_.py
File metadata and controls
99 lines (83 loc) · 2.73 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
# Generated by KernelGen: https://github.com/flagos-ai/KernelGen
import logging
import torch
import triton
import triton.language as tl
from flag_gems.runtime import torch_device_fn
from flag_gems.utils import libentry
logger = logging.getLogger(__name__)
@libentry()
@triton.jit(do_not_specialize=["beta", "alpha"])
def addr_inplace_kernel(
input_ptr,
vec1_ptr,
vec2_ptr,
beta,
alpha,
M,
N,
stride_input_m,
stride_input_n,
stride_vec1,
stride_vec2,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
mask_m = offs_m < M
mask_n = offs_n < N
vec1_ptrs = vec1_ptr + offs_m * stride_vec1
vec2_ptrs = vec2_ptr + offs_n * stride_vec2
vec1 = tl.load(vec1_ptrs, mask=mask_m, other=0.0).to(tl.float32)
vec2 = tl.load(vec2_ptrs, mask=mask_n, other=0.0).to(tl.float32)
input_ptrs = (
input_ptr + offs_m[:, None] * stride_input_m + offs_n[None, :] * stride_input_n
)
mask_2d = mask_m[:, None] & mask_n[None, :]
input_val = tl.load(input_ptrs, mask=mask_2d, other=0.0).to(tl.float32)
# Fused computation: beta * input + alpha * outer(vec1, vec2)
result = beta * input_val + alpha * (vec1[:, None] * vec2[None, :])
tl.store(input_ptrs, result, mask=mask_2d)
def addr_(input, vec1, vec2, *, beta=1, alpha=1):
logger.debug("GEMS ADDR_")
assert input.dtype in (
torch.float16,
torch.bfloat16,
torch.float32,
torch.float64,
), f"addr_: unsupported dtype {input.dtype}, expected floating point"
if vec1.dim() != 1 or vec2.dim() != 1:
raise ValueError("addr_: expected 1-D vectors")
M, N = input.shape
if vec1.shape[0] != M or vec2.shape[0] != N:
raise ValueError(
f"addr_: vec1 size {vec1.shape[0]} must match input rows {M}, "
f"vec2 size {vec2.shape[0]} must match input cols {N}"
)
# Single fused kernel dispatch: reads input, computes, writes back in-place
BLOCK_SIZE_M = 32 # Tile size for rows
BLOCK_SIZE_N = 32 # Tile size for columns
grid = lambda META: (
triton.cdiv(M, BLOCK_SIZE_M),
triton.cdiv(N, BLOCK_SIZE_N),
)
with torch_device_fn.device(input.device):
addr_inplace_kernel[grid](
input,
vec1,
vec2,
beta,
alpha,
M,
N,
input.stride(0),
input.stride(1),
vec1.stride(0),
vec2.stride(0),
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
)
return input