Made a fused GEMV kernel for convrot unpacking to increase speed further. Fix bug in test script that made train time add additional grads to bf16.

This commit is contained in:
Jaret Burkett
2026-07-15 10:24:52 -06:00
parent 691ddf434e
commit e28727d5cb
3 changed files with 631 additions and 114 deletions

View File

@@ -6,10 +6,13 @@ Compares bf16 against the custom OstrisLinear backends (convrot8, convrot4 for
now; add more qtypes to QTYPES as they land). now; add more qtypes to QTYPES as they land).
Measures, per qtype: Measures, per qtype:
- layer inference latency across DiT-representative shapes (vs bf16) - layer inference latency across DiT-representative shapes (vs bf16),
- layer training latency (forward + backward through the frozen layer) eager and torch.compile'd
- layer training latency (forward + backward through the frozen layer),
eager and torch.compile'd
- VRAM on a transformer-ish block stack: resident weights, peak during a - VRAM on a transformer-ish block stack: resident weights, peak during a
no-grad forward, peak during a train step no-grad forward, peak during a train step; forward/train-ckpt peaks also
under torch.compile
- accuracy drift vs bf16: output relative error per layer shape and - accuracy drift vs bf16: output relative error per layer shape and
accumulated through the block stack accumulated through the block stack
- weight reconstruction error and one-time quantize (conversion) time - weight reconstruction error and one-time quantize (conversion) time
@@ -46,7 +49,7 @@ VRAM_BLOCK_SHAPES = [(3072, 12288), (12288, 3072), (3072, 3072), (3072, 3072)]
VRAM_TOKENS = 4096 VRAM_TOKENS = 4096
QTYPES = [ QTYPES = [
"bf16", "qfloat8", "float8", "orbit4", "orbitvq4", "convrot8", "convrot4", "bf16", "qfloat8", "float8", "convrot8", "convrot4",
"convrotint7", "convrotint6", "convrotint5", "convrotint4", "convrotint3", "convrotint7", "convrotint6", "convrotint5", "convrotint4", "convrotint3",
"convrotint2", "convrotbitnet", "convrotcomfyw4a4", "convrotint2", "convrotbitnet", "convrotcomfyw4a4",
] ]
@@ -106,6 +109,10 @@ def make_layer(k: int, n: int, device) -> torch.nn.Linear:
lin = torch.nn.Linear(k, n, bias=True, dtype=torch.bfloat16, device=device) lin = torch.nn.Linear(k, n, bias=True, dtype=torch.bfloat16, device=device)
with torch.no_grad(): with torch.no_grad():
lin.weight.mul_(0.02) lin.weight.mul_(0.02)
# the train benches model lora-style training: base frozen, grads flow to
# the input only. without this, bf16/quanto accumulate weight grads that
# inflate every later vram measurement
lin.requires_grad_(False)
return lin return lin
@@ -120,6 +127,8 @@ def make_stack(device) -> torch.nn.ModuleList:
torch.nn.Linear(k, n, bias=True, dtype=torch.bfloat16, device=device) torch.nn.Linear(k, n, bias=True, dtype=torch.bfloat16, device=device)
for k, n in VRAM_BLOCK_SHAPES for k, n in VRAM_BLOCK_SHAPES
])) ]))
# frozen base (see make_layer)
blocks.requires_grad_(False)
return blocks return blocks
@@ -163,9 +172,59 @@ def run_speed(qtype: str, device, iters: int, results: dict):
t_train = bench(train_step, max(10, iters // 3), device) t_train = bench(train_step, max(10, iters // 3), device)
results[(qtype, "inf", (m, k, n))] = t_inf results[(qtype, "inf", (m, k, n))] = t_inf
results[(qtype, "train", (m, k, n))] = t_train results[(qtype, "train", (m, k, n))] = t_train
# compiled variants (compilation happens during bench warmup, so it
# isn't charged to the timing; a backend that won't compile records
# nothing and shows as '-')
lin_c = torch.compile(lin, dynamic=False)
try:
with torch.no_grad():
results[(qtype, "inf_comp", (m, k, n))] = bench(
lambda: lin_c(x), iters, device
)
except Exception as e:
print(f" [{qtype}] compiled inference failed for {m}x{k}->{n}: {e}")
def train_step_c():
xi = x.detach().requires_grad_(True)
lin_c(xi).sum().backward()
try:
results[(qtype, "train_comp", (m, k, n))] = bench(
train_step_c, max(10, iters // 3), device
)
except Exception as e:
print(f" [{qtype}] compiled train failed for {m}x{k}->{n}: {e}")
torch.cuda.empty_cache() torch.cuda.empty_cache()
def _stack_fwd_peak(blocks, x, device, base) -> int:
# warm up first so lazy-init allocations (and compilation) are not counted
# as steady-state peak
with torch.no_grad():
stack_forward(blocks, x)
torch.cuda.synchronize(device)
torch.cuda.reset_peak_memory_stats(device)
with torch.no_grad():
stack_forward(blocks, x)
torch.cuda.synchronize(device)
return torch.cuda.max_memory_allocated(device) - base
def _stack_train_peak(blocks, x, device, base, checkpoint) -> int:
# frozen base; grads flow to the input like lora training
def train_step():
xi = x.detach().requires_grad_(True)
stack_forward(blocks, xi, checkpoint).float().pow(2).mean().backward()
train_step()
torch.cuda.synchronize(device)
torch.cuda.reset_peak_memory_stats(device)
train_step()
torch.cuda.synchronize(device)
return torch.cuda.max_memory_allocated(device) - base
def run_vram(qtype: str, device, results: dict): def run_vram(qtype: str, device, results: dict):
torch.cuda.empty_cache() torch.cuda.empty_cache()
base = torch.cuda.memory_allocated(device) base = torch.cuda.memory_allocated(device)
@@ -179,30 +238,22 @@ def run_vram(qtype: str, device, results: dict):
x = torch.randn(VRAM_TOKENS, 3072, device=device, dtype=torch.bfloat16) x = torch.randn(VRAM_TOKENS, 3072, device=device, dtype=torch.bfloat16)
# no-grad forward peak (sampling); warm up first so lazy-init allocations are results[(qtype, "vram_fwd_peak")] = _stack_fwd_peak(blocks, x, device, base)
# not counted as steady-state peak # real training checkpoints, but the plain train peak still gets reported
with torch.no_grad(): results[(qtype, "vram_train_peak")] = _stack_train_peak(blocks, x, device, base, False)
stack_forward(blocks, x) results[(qtype, "vram_train_ckpt_peak")] = _stack_train_peak(blocks, x, device, base, True)
torch.cuda.synchronize(device)
torch.cuda.reset_peak_memory_stats(device)
with torch.no_grad():
stack_forward(blocks, x)
torch.cuda.synchronize(device)
results[(qtype, "vram_fwd_peak")] = torch.cuda.max_memory_allocated(device) - base
# train step peak (frozen base; grads flow to the input like lora training), # same peaks with every linear compiled (mirrors the trainer's block compile)
# with and without per-block gradient checkpointing (real training uses it) for b in blocks:
def train_step(checkpoint): for i in range(len(b)):
xi = x.detach().requires_grad_(True) b[i] = torch.compile(b[i], dynamic=False)
stack_forward(blocks, xi, checkpoint).float().pow(2).mean().backward() try:
results[(qtype, "vram_fwd_peak_comp")] = _stack_fwd_peak(blocks, x, device, base)
for key, ckpt in (("vram_train_peak", False), ("vram_train_ckpt_peak", True)): results[(qtype, "vram_train_ckpt_peak_comp")] = _stack_train_peak(
train_step(ckpt) blocks, x, device, base, True
torch.cuda.synchronize(device) )
torch.cuda.reset_peak_memory_stats(device) except Exception as e:
train_step(ckpt) print(f" [{qtype}] compiled vram measurement failed: {e}")
torch.cuda.synchronize(device)
results[(qtype, key)] = torch.cuda.max_memory_allocated(device) - base
blocks = x = None # release before the allocator accounting of the next run blocks = x = None # release before the allocator accounting of the next run
torch.cuda.empty_cache() torch.cuda.empty_cache()
@@ -252,18 +303,22 @@ def run_quality_and_quantize_time(qtype: str, device, results: dict):
def print_speed_table(title: str, kind: str, qts, results): def print_speed_table(title: str, kind: str, qts, results):
print(f"\n=== {title} (ms; speedup vs bf16) ===") # speedups always reference EAGER bf16, compiled kinds included, so the
# comp columns answer "what do I gain over plain bf16"
ref_kind = kind.removesuffix("_comp")
print(f"\n=== {title} (ms; speedup vs eager bf16) ===")
print(f"{'M x K -> N':<22}" + "".join(f"{qt:>18}" for qt in qts)) print(f"{'M x K -> N':<22}" + "".join(f"{qt:>18}" for qt in qts))
for shape in SPEED_SHAPES: for shape in SPEED_SHAPES:
m, k, n = shape m, k, n = shape
row = f"{f'{m} x {k} -> {n}':<22}" row = f"{f'{m} x {k} -> {n}':<22}"
ref = results.get(("bf16", kind, shape)) ref = results.get(("bf16", ref_kind, shape))
for qt in qts: for qt in qts:
t = results.get((qt, kind, shape)) t = results.get((qt, kind, shape))
if t is None: if t is None:
row += f"{'-':>18}" row += f"{'-':>18}"
continue continue
sp = f" ({ref / t:4.2f}x)" if ref and qt != "bf16" else " " * 8 is_self_ref = qt == "bf16" and kind == ref_kind
sp = f" ({ref / t:4.2f}x)" if ref and not is_self_ref else " " * 8
row += f"{t:8.3f}ms{sp}" row += f"{t:8.3f}ms{sp}"
print(row) print(row)
@@ -290,9 +345,15 @@ def main():
if qt != "bf16": if qt != "bf16":
get_ostris_quantizer(qt) get_ostris_quantizer(qt)
# many distinct module instances share one forward code object; the default
# per-code cache limit (8) would silently fall back to eager and corrupt the
# compiled columns
torch._dynamo.config.cache_size_limit = 4096
results = {} results = {}
for qt in args.qtypes: for qt in args.qtypes:
print(f"benchmarking {qt} ...") print(f"benchmarking {qt} ...")
torch._dynamo.reset() # drop the previous qtype's compiled artifacts
run_quality_and_quantize_time(qt, device, results) run_quality_and_quantize_time(qt, device, results)
run_drift(qt, device, results) run_drift(qt, device, results)
run_speed(qt, device, args.iters, results) run_speed(qt, device, args.iters, results)
@@ -300,17 +361,22 @@ def main():
qts = args.qtypes qts = args.qtypes
print_speed_table("layer latency, inference", "inf", qts, results) print_speed_table("layer latency, inference", "inf", qts, results)
print_speed_table("layer latency, inference (compiled)", "inf_comp", qts, results)
print_speed_table("layer latency, train fwd+bwd", "train", qts, results) print_speed_table("layer latency, train fwd+bwd", "train", qts, results)
print_speed_table("layer latency, train fwd+bwd (compiled)", "train_comp", qts, results)
print(f"\n=== vram on the block stack ({VRAM_BLOCKS} blocks, {VRAM_TOKENS} tokens) ===") print(f"\n=== vram on the block stack ({VRAM_BLOCKS} blocks, {VRAM_TOKENS} tokens) ===")
print(f"{'':<28}" + "".join(f"{qt:>18}" for qt in qts)) print(f"{'':<28}" + "".join(f"{qt:>18}" for qt in qts))
for key, label in (("vram_weights", "weights resident"), for key, label in (("vram_weights", "weights resident"),
("vram_fwd_peak", "peak, no-grad fwd"), ("vram_fwd_peak", "peak, no-grad fwd"),
("vram_train_peak", "peak, train step"), ("vram_train_peak", "peak, train step"),
("vram_train_ckpt_peak", "peak, train step (ckpt)")): ("vram_train_ckpt_peak", "peak, train step (ckpt)"),
("vram_fwd_peak_comp", "peak, no-grad fwd (comp)"),
("vram_train_ckpt_peak_comp", "peak, train ckpt (comp)")):
row = f"{label:<28}" row = f"{label:<28}"
for qt in qts: for qt in qts:
row += f"{gb(results[(qt, key)]):>18}" v = results.get((qt, key))
row += f"{gb(v):>18}" if v is not None else f"{'-':>18}"
print(row) print(row)
print("\n=== accuracy drift vs bf16 (output rel err, no-grad) ===") print("\n=== accuracy drift vs bf16 (output rel err, no-grad) ===")
@@ -335,25 +401,38 @@ def main():
# ---- clean per-qtype breakdown: speed (geomean over shapes) + accuracy ---- # ---- clean per-qtype breakdown: speed (geomean over shapes) + accuracy ----
def geomean_speedup(qt, kind): def geomean_speedup(qt, kind):
# every speedup references EAGER bf16 (compiled kinds included), so the
# comp columns answer "what do I gain over plain bf16"
ref_kind = kind.removesuffix("_comp")
logs = [] logs = []
for shape in SPEED_SHAPES: for shape in SPEED_SHAPES:
ref = results.get(("bf16", kind, shape)) ref = results.get(("bf16", ref_kind, shape))
t = results.get((qt, kind, shape)) t = results.get((qt, kind, shape))
if ref and t: if ref and t:
logs.append(math.log(ref / t)) logs.append(math.log(ref / t))
return math.exp(sum(logs) / len(logs)) if logs else float("nan") return math.exp(sum(logs) / len(logs)) if logs else None
def fmt_speed(v):
return f"{v:.2f}x" if v is not None else "-"
print("\n=== summary (speed = geomean speedup vs bf16; drift lower is better) ===") print("\n=== summary (speed = geomean speedup vs bf16; drift lower is better) ===")
print(f"{'':<12}{'inference':>18}{'train':>18}{'accuracy drift':>20}{'max vram':>16}") print(f"{'':<18}{'inference':>12}{'inference comp':>16}{'train':>12}{'train comp':>12}"
f"{'accuracy drift':>16}{'max vram':>12}{'max vram comp':>15}")
for qt in qts: for qt in qts:
# real training checkpoints, so the ckpt peak is the meaningful train # real training checkpoints, so the ckpt peak is the meaningful train
# number; the no-grad fwd peak still matters for sampling # number; the no-grad fwd peak still matters for sampling
max_vram = max(results[(qt, "vram_fwd_peak")], results[(qt, "vram_train_ckpt_peak")]) max_vram = max(results[(qt, "vram_fwd_peak")], results[(qt, "vram_train_ckpt_peak")])
print(f"{qt:<12}" fwd_c = results.get((qt, "vram_fwd_peak_comp"))
f"{geomean_speedup(qt, 'inf'):>17.2f}x" ckpt_c = results.get((qt, "vram_train_ckpt_peak_comp"))
f"{geomean_speedup(qt, 'train'):>17.2f}x" max_vram_comp = max(fwd_c, ckpt_c) if fwd_c is not None and ckpt_c is not None else None
f"{results[(qt, 'drift', STACK_KEY)]:>20.5f}" print(f"{qt:<18}"
f"{gb(max_vram):>16}") f"{fmt_speed(geomean_speedup(qt, 'inf')):>12}"
f"{fmt_speed(geomean_speedup(qt, 'inf_comp')):>16}"
f"{fmt_speed(geomean_speedup(qt, 'train')):>12}"
f"{fmt_speed(geomean_speedup(qt, 'train_comp')):>12}"
f"{results[(qt, 'drift', STACK_KEY)]:>16.5f}"
f"{gb(max_vram).strip():>12}"
f"{(gb(max_vram_comp).strip() if max_vram_comp is not None else '-'):>15}")
if __name__ == "__main__": if __name__ == "__main__":

View File

@@ -66,7 +66,7 @@ def get_convrot_quantizer(qtype: str):
if qtype == "convrotcomfyw4a4": if qtype == "convrotcomfyw4a4":
return ConvRotComfyW4A4Quantizer() return ConvRotComfyW4A4Quantizer()
if qtype.startswith("convrotint"): if qtype.startswith("convrotint"):
bits = int(qtype[len("convrotint"):]) bits = int(qtype[len("convrotint") :])
if 2 <= bits <= 8: if 2 <= bits <= 8:
return ConvRotIntNQuantizer(bits, rot_size=256) return ConvRotIntNQuantizer(bits, rot_size=256)
return None return None
@@ -147,11 +147,11 @@ def _optimal_nvfp4_scales(
reconstruction error. ~11% lower weight error than plain amax scaling; reconstruction error. ~11% lower weight error than plain amax scaling;
deterministic, and the storage/GEMM format is unchanged.""" deterministic, and the storage/GEMM format is unchanged."""
edges = _cached( edges = _cached(
_edges_cache, str(xb.device), lambda: torch.tensor(_E2M1_EDGES, device=xb.device) _edges_cache,
) str(xb.device),
vals = torch.tensor( lambda: torch.tensor(_E2M1_EDGES, device=xb.device),
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], device=xb.device
) )
vals = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], device=xb.device)
best_s = base.to(torch.float8_e4m3fn) best_s = base.to(torch.float8_e4m3fn)
best_e = None best_e = None
for frac in torch.linspace(0.70, 1.10, 9, dtype=torch.float64): for frac in torch.linspace(0.70, 1.10, 9, dtype=torch.float64):
@@ -214,19 +214,25 @@ def dequantize_nvfp4(
# single-pass triton path when available: the torch chain below is ~7 full-size # single-pass triton path when available: the torch chain below is ~7 full-size
# elementwise passes with fp32 intermediates, which made every convrot4 training # elementwise passes with fp32 intermediates, which made every convrot4 training
# backward pay a dequant cost comparable to the gradient matmul itself # backward pay a dequant cost comparable to the gradient matmul itself
if _triton_available() and packed.is_cuda and dtype in (torch.bfloat16, torch.float16, torch.float32): if (
_triton_available()
and packed.is_cuda
and dtype in (torch.bfloat16, torch.float16, torch.float32)
):
return _fp4_dequant_op( return _fp4_dequant_op(
packed, scales.view(torch.uint8), pts.reshape(1).view(torch.uint8), packed,
scales.view(torch.uint8),
pts.reshape(1).view(torch.uint8),
str(dtype).split(".")[-1], str(dtype).split(".")[-1],
) )
codes = torch.stack([packed & 15, packed >> 4], dim=-1).view(rows, K) codes = torch.stack([packed & 15, packed >> 4], dim=-1).view(rows, K)
# the lookup table is built inline (NOT module-cached): this function runs inside # the lookup table is built inline (NOT module-cached): this function runs inside
# custom-op backwards, which torch.compile traces with fake tensors where a # custom-op backwards, which torch.compile traces with fake tensors where a
# pre-existing real tensor is illegal; an in-trace constructed constant is fine # pre-existing real tensor is illegal; an in-trace constructed constant is fine
vals = torch.tensor( vals = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], device=packed.device)
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], device=packed.device mag = torch.index_select(vals, 0, (codes & 7).flatten().to(torch.int32)).view(
rows, K
) )
mag = torch.index_select(vals, 0, (codes & 7).flatten().to(torch.int32)).view(rows, K)
v = mag * torch.where((codes & 8) > 0, -1.0, 1.0) v = mag * torch.where((codes & 8) > 0, -1.0, 1.0)
v = v.view(rows, K // BLOCK, BLOCK) * (scales.float() * pts).unsqueeze(-1) v = v.view(rows, K // BLOCK, BLOCK) * (scales.float() * pts).unsqueeze(-1)
return v.view(rows, K).to(dtype) return v.view(rows, K).to(dtype)
@@ -340,9 +346,15 @@ def _launch_nvfp4_kernel(x, packed, scales, pts, blocked_scales: bool):
BLOCK_K = min(2048, 1 << (K - 1).bit_length()) BLOCK_K = min(2048, 1 << (K - 1).bit_length())
grid = (rows, -(-K // BLOCK_K)) grid = (rows, -(-K // BLOCK_K))
_get_kernel()[grid]( _get_kernel()[grid](
x, packed, scales, pts, x,
K, n_col_tiles, packed,
BLOCK_K=BLOCK_K, BLOCKED_SCALES=blocked_scales, num_warps=4, scales,
pts,
K,
n_col_tiles,
BLOCK_K=BLOCK_K,
BLOCKED_SCALES=blocked_scales,
num_warps=4,
) )
@@ -369,7 +381,8 @@ def _nvfp4_act_quant_op(x: torch.Tensor) -> list[torch.Tensor]:
# zero-init: rows are padded to 128-tiles and the pad region must be zero # zero-init: rows are padded to 128-tiles and the pad region must be zero
scales = torch.zeros( scales = torch.zeros(
(-(-rows_pad // 128)) * 128 * n_col_tiles * 4, (-(-rows_pad // 128)) * 128 * n_col_tiles * 4,
device=x.device, dtype=torch.float8_e4m3fn, device=x.device,
dtype=torch.float8_e4m3fn,
) )
_launch_nvfp4_kernel(x, packed, scales, pts, blocked_scales=True) _launch_nvfp4_kernel(x, packed, scales, pts, blocked_scales=True)
return [packed, scales.view(torch.uint8), pts] return [packed, scales.view(torch.uint8), pts]
@@ -382,7 +395,11 @@ def _nvfp4_act_quant_fake(x):
n_col_tiles = -(-(K // BLOCK) // 4) n_col_tiles = -(-(K // BLOCK) // 4)
return [ return [
torch.empty(rows_pad, K // 2, device=x.device, dtype=torch.uint8), torch.empty(rows_pad, K // 2, device=x.device, dtype=torch.uint8),
torch.empty((-(-rows_pad // 128)) * 128 * n_col_tiles * 4, device=x.device, dtype=torch.uint8), torch.empty(
(-(-rows_pad // 128)) * 128 * n_col_tiles * 4,
device=x.device,
dtype=torch.uint8,
),
torch.empty((), device=x.device, dtype=torch.float32), torch.empty((), device=x.device, dtype=torch.float32),
] ]
@@ -416,7 +433,6 @@ def quantize_nvfp4_fused(x: torch.Tensor, blocked_scales: bool = False):
return packed, scales, pts return packed, scales, pts
# ---------------- fp4 dequant kernel (backward hot path) ---------------- # ---------------- fp4 dequant kernel (backward hot path) ----------------
_dequant_kernel = None _dequant_kernel = None
@@ -431,7 +447,11 @@ def _get_dequant_kernel():
@triton.jit @triton.jit
def nvfp4_dequant_kernel( def nvfp4_dequant_kernel(
q_ptr, s_ptr, pts_ptr, out_ptr, K, q_ptr,
s_ptr,
pts_ptr,
out_ptr,
K,
BLOCK_B: tl.constexpr, BLOCK_B: tl.constexpr,
): ):
row = tl.program_id(0) row = tl.program_id(0)
@@ -443,7 +463,9 @@ def _get_dequant_kernel():
codes = tl.interleave(byte & 15, byte >> 4) # (2*BLOCK_B,), column order codes = tl.interleave(byte & 15, byte >> 4) # (2*BLOCK_B,), column order
m = (codes & 7).to(tl.float32) m = (codes & 7).to(tl.float32)
# arithmetic e2m1 decode ([0, .5, 1, 1.5, 2, 3, 4, 6]), exact # arithmetic e2m1 decode ([0, .5, 1, 1.5, 2, 3, 4, 6]), exact
mag = tl.where(m < 2, m * 0.5, tl.exp2(tl.floor(m / 2) - 1) * (1 + (m % 2) * 0.5)) mag = tl.where(
m < 2, m * 0.5, tl.exp2(tl.floor(m / 2) - 1) * (1 + (m % 2) * 0.5)
)
v = tl.where((codes & 8) > 0, -mag, mag) v = tl.where((codes & 8) > 0, -mag, mag)
n_s: tl.constexpr = (2 * BLOCK_B) // 16 n_s: tl.constexpr = (2 * BLOCK_B) // 16
offs_s = pid_k * n_s + tl.arange(0, n_s) offs_s = pid_k * n_s + tl.arange(0, n_s)
@@ -451,7 +473,11 @@ def _get_dequant_kernel():
vb = tl.reshape(v, (n_s, 16)) * (s.to(tl.float32) * pts)[:, None] vb = tl.reshape(v, (n_s, 16)) * (s.to(tl.float32) * pts)[:, None]
out = tl.reshape(vb, (2 * BLOCK_B,)) out = tl.reshape(vb, (2 * BLOCK_B,))
offs_v = pid_k * (2 * BLOCK_B) + tl.arange(0, 2 * BLOCK_B) offs_v = pid_k * (2 * BLOCK_B) + tl.arange(0, 2 * BLOCK_B)
tl.store(out_ptr + row * K + offs_v, out.to(out_ptr.dtype.element_ty), mask=offs_v < K) tl.store(
out_ptr + row * K + offs_v,
out.to(out_ptr.dtype.element_ty),
mask=offs_v < K,
)
_dequant_kernel = nvfp4_dequant_kernel _dequant_kernel = nvfp4_dequant_kernel
return _dequant_kernel return _dequant_kernel
@@ -467,7 +493,9 @@ def _fp4_dequant_op(
out_dtype: str, out_dtype: str,
) -> torch.Tensor: ) -> torch.Tensor:
rows, half = packed.shape rows, half = packed.shape
out = torch.empty(rows, half * 2, device=packed.device, dtype=getattr(torch, out_dtype)) out = torch.empty(
rows, half * 2, device=packed.device, dtype=getattr(torch, out_dtype)
)
kernel = _get_dequant_kernel() kernel = _get_dequant_kernel()
block_b = 1024 block_b = 1024
grid = (rows, -(-half // block_b)) grid = (rows, -(-half // block_b))
@@ -475,8 +503,10 @@ def _fp4_dequant_op(
packed.contiguous(), packed.contiguous(),
scales_u8.view(torch.float8_e4m3fn), scales_u8.view(torch.float8_e4m3fn),
pts_u8.view(torch.float32), pts_u8.view(torch.float32),
out, half * 2, out,
BLOCK_B=block_b, num_warps=4, half * 2,
BLOCK_B=block_b,
num_warps=4,
) )
return out return out
@@ -484,7 +514,9 @@ def _fp4_dequant_op(
@_fp4_dequant_op.register_fake @_fp4_dequant_op.register_fake
def _fp4_dequant_fake(packed, scales_u8, pts_u8, out_dtype): def _fp4_dequant_fake(packed, scales_u8, pts_u8, out_dtype):
rows, half = packed.shape rows, half = packed.shape
return torch.empty(rows, half * 2, device=packed.device, dtype=getattr(torch, out_dtype)) return torch.empty(
rows, half * 2, device=packed.device, dtype=getattr(torch, out_dtype)
)
# ---------------- backend ---------------- # ---------------- backend ----------------
@@ -547,7 +579,9 @@ def _fp4_linear_ste_op(
@_fp4_linear_ste_op.register_fake @_fp4_linear_ste_op.register_fake
def _fp4_linear_ste_fake(x2d, qdata, scales_u8, scales_blocked_u8, pts_u8, bias, out_dtype): def _fp4_linear_ste_fake(
x2d, qdata, scales_u8, scales_blocked_u8, pts_u8, bias, out_dtype
):
return torch.empty( return torch.empty(
x2d.shape[0], qdata.shape[0], device=x2d.device, dtype=getattr(torch, out_dtype) x2d.shape[0], qdata.shape[0], device=x2d.device, dtype=getattr(torch, out_dtype)
) )
@@ -562,9 +596,12 @@ def _fp4_linear_ste_backward(ctx, grad):
qdata, scales_u8, pts_u8 = ctx.saved_tensors qdata, scales_u8, pts_u8 = ctx.saved_tensors
out_f, in_half = qdata.shape out_f, in_half = qdata.shape
w = dequantize_nvfp4( w = dequantize_nvfp4(
qdata, scales_u8.view(torch.float8_e4m3fn), qdata,
scales_u8.view(torch.float8_e4m3fn),
pts_u8.view(torch.float32).reshape(()), pts_u8.view(torch.float32).reshape(()),
out_f, in_half * 2, grad.dtype, out_f,
in_half * 2,
grad.dtype,
) )
return grad @ w, None, None, None, None, None, None return grad @ w, None, None, None, None, None, None
@@ -633,7 +670,8 @@ class ConvRotQuantizer(OstrisQuantizer):
safe = torch.where(denom > 0, denom, torch.ones_like(denom)) safe = torch.where(denom > 0, denom, torch.ones_like(denom))
z = (w_rot.float().view(rows, K // BLOCK, BLOCK) / safe).clamp(-F4_MAX, F4_MAX) z = (w_rot.float().view(rows, K // BLOCK, BLOCK) / safe).clamp(-F4_MAX, F4_MAX)
edges = _cached( edges = _cached(
_edges_cache, str(w_rot.device), _edges_cache,
str(w_rot.device),
lambda: torch.tensor(_E2M1_EDGES, device=w_rot.device), lambda: torch.tensor(_E2M1_EDGES, device=w_rot.device),
) )
vals = torch.tensor( vals = torch.tensor(
@@ -669,7 +707,8 @@ class ConvRotQuantizer(OstrisQuantizer):
safe = torch.where(denom > 0, denom, torch.ones_like(denom)) safe = torch.where(denom > 0, denom, torch.ones_like(denom))
z = (w_rot.view(rows, K // BLOCK, BLOCK) / safe).clamp(-F4_MAX, F4_MAX) z = (w_rot.view(rows, K // BLOCK, BLOCK) / safe).clamp(-F4_MAX, F4_MAX)
edges = _cached( edges = _cached(
_edges_cache, str(w_rot.device), _edges_cache,
str(w_rot.device),
lambda: torch.tensor(_E2M1_EDGES, device=w_rot.device), lambda: torch.tensor(_E2M1_EDGES, device=w_rot.device),
) )
z = z.reshape(rows, K) z = z.reshape(rows, K)
@@ -699,8 +738,12 @@ class ConvRotQuantizer(OstrisQuantizer):
# with a straight-through analytic backward # with a straight-through analytic backward
x2d = rotate(x, rot).reshape(-1, in_f) x2d = rotate(x, rot).reshape(-1, in_f)
out = _fp4_linear_ste_op( out = _fp4_linear_ste_op(
x2d, module.cr_qdata, module.cr_scales, x2d,
module.cr_scales_blocked, module.cr_pts, module.bias, module.cr_qdata,
module.cr_scales,
module.cr_scales_blocked,
module.cr_pts,
module.bias,
str(x.dtype).split(".")[-1], str(x.dtype).split(".")[-1],
) )
return out.reshape(*x.shape[:-1], out_f) return out.reshape(*x.shape[:-1], out_f)
@@ -837,9 +880,7 @@ def _int8_act_quant_op(x: torch.Tensor, qmax: int) -> list[torch.Tensor]:
kernel, _ = _get_int8_kernels() kernel, _ = _get_int8_kernels()
# triton block shapes must be powers of 2; loads/stores are masked on offs < K # triton block shapes must be powers of 2; loads/stores are masked on offs < K
block_k = min(2048, 1 << (K - 1).bit_length()) block_k = min(2048, 1 << (K - 1).bit_length())
kernel[(rows,)]( kernel[(rows,)](x, q, scales, K, QMAX=qmax, BLOCK_K=block_k, num_warps=8)
x, q, scales, K, QMAX=qmax, BLOCK_K=block_k, num_warps=8
)
return [q, scales] return [q, scales]
@@ -927,7 +968,10 @@ def _int8_linear_ste_op(
aq, a_s = _int8_act_quant_padded(x2d, act_qmax) aq, a_s = _int8_act_quant_padded(x2d, act_qmax)
i32 = torch._int_mm(aq, qdata.t()) i32 = torch._int_mm(aq, qdata.t())
return _int8_epilogue( return _int8_epilogue(
i32[:m], a_s[:m], w_scales_u8.view(torch.float32), bias, i32[:m],
a_s[:m],
w_scales_u8.view(torch.float32),
bias,
getattr(torch, out_dtype), getattr(torch, out_dtype),
) )
@@ -1064,8 +1108,12 @@ class ConvRotInt8Quantizer(OstrisQuantizer):
"""Hardware STE linear for the training path. For int8 the saved qdata is """Hardware STE linear for the training path. For int8 the saved qdata is
the resident buffer itself, so autograd holds only a free reference.""" the resident buffer itself, so autograd holds only a free reference."""
return _int8_linear_ste_op( return _int8_linear_ste_op(
x2d, self._qdata(module), self._scales_u8(module), module.bias, x2d,
self.act_qmax, out_dtype, self._qdata(module),
self._scales_u8(module),
module.bias,
self.act_qmax,
out_dtype,
) )
def fake_quant_rotated_weight(self, module, w_rot: torch.Tensor) -> torch.Tensor: def fake_quant_rotated_weight(self, module, w_rot: torch.Tensor) -> torch.Tensor:
@@ -1099,6 +1147,11 @@ class ConvRotInt8Quantizer(OstrisQuantizer):
s = self._scales(module).unsqueeze(1) s = self._scales(module).unsqueeze(1)
module.cr8_qdata = torch.round(w_rot / s).clamp_(-127, 127).to(torch.int8) module.cr8_qdata = torch.round(w_rot / s).clamp_(-127, 127).to(torch.int8)
def _gemv_args(self, module):
"""(qdata, gratio fp32 or None, bits) for the fused decode gemv, or None
if this backend's storage isn't supported by it."""
return module.cr8_qdata, None, 8
def forward(self, module, x: torch.Tensor) -> torch.Tensor: def forward(self, module, x: torch.Tensor) -> torch.Tensor:
rot = self._rot(module) rot = self._rot(module)
in_f, out_f = module.in_features, module.out_features in_f, out_f = module.in_features, module.out_features
@@ -1115,7 +1168,8 @@ class ConvRotInt8Quantizer(OstrisQuantizer):
# utilization, computed twice for the amax and quant passes, cost # utilization, computed twice for the amax and quant passes, cost
# more than the activation round-trips they saved) # more than the activation round-trips they saved)
out = self._linear_ste( out = self._linear_ste(
module, rotate(x, rot).reshape(-1, in_f), module,
rotate(x, rot).reshape(-1, in_f),
str(x.dtype).split(".")[-1], str(x.dtype).split(".")[-1],
) )
return out.reshape(*x.shape[:-1], out_f) return out.reshape(*x.shape[:-1], out_f)
@@ -1129,6 +1183,25 @@ class ConvRotInt8Quantizer(OstrisQuantizer):
out = F.linear(x_ste, w, module.bias) out = F.linear(x_ste, w, module.bias)
return out.reshape(*x.shape[:-1], out_f) return out.reshape(*x.shape[:-1], out_f)
if m <= FUSED_GEMV_MAX_M and x.is_cuda and _triton_available():
# decode-size batches: one fused launch instead of the 4-kernel
# eager chain (bit-identical output, no unpacked weight transient)
args = self._gemv_args(module)
if args is not None:
qdata, gratio, bits = args
out = _int_gemv_op(
rotate(x, rot).reshape(-1, in_f),
qdata,
gratio,
self._scales_u8(module),
module.bias,
bits,
self.act_qmax,
out_f,
str(x.dtype).split(".")[-1],
)
return out.reshape(*x.shape[:-1], out_f)
if _int8_gemm_supported(x.device): if _int8_gemm_supported(x.device):
# row padding for _int_mm happens inside the act-quant op (compile # row padding for _int_mm happens inside the act-quant op (compile
# safety); slice the mm output back to m rows (a contiguous prefix) # safety); slice the mm output back to m rows (a contiguous prefix)
@@ -1136,7 +1209,9 @@ class ConvRotInt8Quantizer(OstrisQuantizer):
rotate(x, rot).reshape(-1, in_f), self.act_qmax rotate(x, rot).reshape(-1, in_f), self.act_qmax
) )
i32 = torch._int_mm(aq, self._qdata(module).t()) i32 = torch._int_mm(aq, self._qdata(module).t())
out = _int8_epilogue(i32[:m], a_s[:m], self._scales(module), module.bias, x.dtype) out = _int8_epilogue(
i32[:m], a_s[:m], self._scales(module), module.bias, x.dtype
)
return out.reshape(*x.shape[:-1], out_f) return out.reshape(*x.shape[:-1], out_f)
w = self._dequantize_rotated(module, x.dtype) w = self._dequantize_rotated(module, x.dtype)
@@ -1183,7 +1258,9 @@ def pack_intn_rows(q: torch.Tensor, bits: int) -> torch.Tensor:
return b.to(torch.uint8).reshape(rows, K // 8 * bits) return b.to(torch.uint8).reshape(rows, K // 8 * bits)
def unpack_intn_rows(packed: torch.Tensor, bits: int, rows: int, cols: int) -> torch.Tensor: def unpack_intn_rows(
packed: torch.Tensor, bits: int, rows: int, cols: int
) -> torch.Tensor:
"""Inverse of pack_intn_rows: (rows, cols) int8 codes in [-qmax, qmax].""" """Inverse of pack_intn_rows: (rows, cols) int8 codes in [-qmax, qmax]."""
qmax = (1 << (bits - 1)) - 1 qmax = (1 << (bits - 1)) - 1
b = packed.reshape(rows, cols // 8, bits).to(torch.int64) b = packed.reshape(rows, cols // 8, bits).to(torch.int64)
@@ -1312,8 +1389,13 @@ def _get_intn_kernel():
@triton.jit @triton.jit
def intn_unpack_kernel( def intn_unpack_kernel(
p_ptr, o_ptr, n_groups, p_ptr,
BITS: tl.constexpr, BPOW: tl.constexpr, QMAX: tl.constexpr, BLOCK: tl.constexpr, o_ptr,
n_groups,
BITS: tl.constexpr,
BPOW: tl.constexpr,
QMAX: tl.constexpr,
BLOCK: tl.constexpr,
): ):
# one row per group of 8 codes; 2d blocks keep the byte loads and int8 # one row per group of 8 codes; 2d blocks keep the byte loads and int8
# stores contiguous/coalesced (BPOW = BITS padded to a power of two for # stores contiguous/coalesced (BPOW = BITS padded to a power of two for
@@ -1330,7 +1412,9 @@ def _get_intn_kernel():
word = tl.sum(b << (8 * bi)[None, :].to(tl.int64), axis=1) word = tl.sum(b << (8 * bi)[None, :].to(tl.int64), axis=1)
j = tl.arange(0, 8) j = tl.arange(0, 8)
# mask after the shift kills any sign extension from int64 wrap # mask after the shift kills any sign extension from int64 wrap
v = ((word[:, None] >> (BITS * j)[None, :].to(tl.int64)) & ((1 << BITS) - 1)) - QMAX v = (
(word[:, None] >> (BITS * j)[None, :].to(tl.int64)) & ((1 << BITS) - 1)
) - QMAX
tl.store(o_ptr + g[:, None] * 8 + j[None, :], v.to(tl.int8), mask=gm[:, None]) tl.store(o_ptr + g[:, None] * 8 + j[None, :], v.to(tl.int8), mask=gm[:, None])
_intn_kernel = intn_unpack_kernel _intn_kernel = intn_unpack_kernel
@@ -1350,8 +1434,17 @@ def _get_intn_grouped_kernel():
@triton.jit @triton.jit
def intn_unpack_grouped_kernel( def intn_unpack_grouped_kernel(
p_ptr, r_ptr, o_ptr, n_groups, cols8, gdiv8, ngprow, p_ptr,
BITS: tl.constexpr, BPOW: tl.constexpr, QMAX: tl.constexpr, BLOCK: tl.constexpr, r_ptr,
o_ptr,
n_groups,
cols8,
gdiv8,
ngprow,
BITS: tl.constexpr,
BPOW: tl.constexpr,
QMAX: tl.constexpr,
BLOCK: tl.constexpr,
): ):
# like intn_unpack_kernel, plus a per-(row, k-group) ratio multiply that # like intn_unpack_kernel, plus a per-(row, k-group) ratio multiply that
# re-expresses the group-scaled codes on the row's int8 grid. the ratio # re-expresses the group-scaled codes on the row's int8 grid. the ratio
@@ -1370,7 +1463,9 @@ def _get_intn_grouped_kernel():
).to(tl.int64) ).to(tl.int64)
word = tl.sum(b << (8 * bi)[None, :].to(tl.int64), axis=1) word = tl.sum(b << (8 * bi)[None, :].to(tl.int64), axis=1)
j = tl.arange(0, 8) j = tl.arange(0, 8)
v = ((word[:, None] >> (BITS * j)[None, :].to(tl.int64)) & ((1 << BITS) - 1)) - QMAX v = (
(word[:, None] >> (BITS * j)[None, :].to(tl.int64)) & ((1 << BITS) - 1)
) - QMAX
vf = libdevice.rint(v.to(tl.float32) * ratio[:, None]) vf = libdevice.rint(v.to(tl.float32) * ratio[:, None])
vf = tl.minimum(tl.maximum(vf, -127.0), 127.0) vf = tl.minimum(tl.maximum(vf, -127.0), 127.0)
tl.store(o_ptr + g[:, None] * 8 + j[None, :], vf.to(tl.int8), mask=gm[:, None]) tl.store(o_ptr + g[:, None] * 8 + j[None, :], vf.to(tl.int8), mask=gm[:, None])
@@ -1410,7 +1505,9 @@ def _get_bitnet_kernel():
j < 1, 1, tl.where(j < 2, 3, tl.where(j < 3, 9, tl.where(j < 4, 27, 81))) j < 1, 1, tl.where(j < 2, 3, tl.where(j < 3, 9, tl.where(j < 4, 27, 81)))
) )
code = ((b[:, None] // p3[None, :]) % 3 - 1).to(tl.float32) code = ((b[:, None] // p3[None, :]) % 3 - 1).to(tl.float32)
ratio = tl.load(r_ptr + row[:, None] * ngprow + col // group, mask=cm, other=1.0) ratio = tl.load(
r_ptr + row[:, None] * ngprow + col // group, mask=cm, other=1.0
)
v = libdevice.rint(code * ratio) v = libdevice.rint(code * ratio)
v = tl.minimum(tl.maximum(v, -127.0), 127.0) v = tl.minimum(tl.maximum(v, -127.0), 127.0)
tl.store(o_ptr + row[:, None] * K + col, v.to(tl.int8), mask=cm) tl.store(o_ptr + row[:, None] * K + col, v.to(tl.int8), mask=cm)
@@ -1419,16 +1516,23 @@ def _get_bitnet_kernel():
return _bitnet_kernel return _bitnet_kernel
def _unpack_intn_impl(packed: torch.Tensor, bits: int, rows: int, cols: int) -> torch.Tensor: def _unpack_intn_impl(
packed: torch.Tensor, bits: int, rows: int, cols: int
) -> torch.Tensor:
if _triton_available() and packed.is_cuda: if _triton_available() and packed.is_cuda:
out = torch.empty(rows, cols, device=packed.device, dtype=torch.int8) out = torch.empty(rows, cols, device=packed.device, dtype=torch.int8)
n_groups = rows * cols // 8 n_groups = rows * cols // 8
BLOCK = 256 BLOCK = 256
kernel = _get_intn_kernel() kernel = _get_intn_kernel()
kernel[(-(-n_groups // BLOCK),)]( kernel[(-(-n_groups // BLOCK),)](
packed, out, n_groups, packed,
BITS=bits, BPOW=max(2, 1 << (bits - 1).bit_length()), out,
QMAX=(1 << (bits - 1)) - 1, BLOCK=BLOCK, num_warps=4, n_groups,
BITS=bits,
BPOW=max(2, 1 << (bits - 1).bit_length()),
QMAX=(1 << (bits - 1)) - 1,
BLOCK=BLOCK,
num_warps=4,
) )
return out return out
return unpack_intn_rows(packed, bits, rows, cols) return unpack_intn_rows(packed, bits, rows, cols)
@@ -1445,9 +1549,18 @@ def _unpack_intn_grouped_impl(
BLOCK = 256 BLOCK = 256
kernel = _get_intn_grouped_kernel() kernel = _get_intn_grouped_kernel()
kernel[(-(-n_groups // BLOCK),)]( kernel[(-(-n_groups // BLOCK),)](
packed, gratio, out, n_groups, cols // 8, group // 8, ngprow, packed,
BITS=bits, BPOW=max(2, 1 << (bits - 1).bit_length()), gratio,
QMAX=(1 << (bits - 1)) - 1, BLOCK=BLOCK, num_warps=4, out,
n_groups,
cols // 8,
group // 8,
ngprow,
BITS=bits,
BPOW=max(2, 1 << (bits - 1).bit_length()),
QMAX=(1 << (bits - 1)) - 1,
BLOCK=BLOCK,
num_warps=4,
) )
return out return out
return unpack_intn_rows_grouped(packed, gratio, bits, rows, cols) return unpack_intn_rows_grouped(packed, gratio, bits, rows, cols)
@@ -1456,7 +1569,9 @@ def _unpack_intn_grouped_impl(
# registered as custom ops so torch.compile treats the triton launches as opaque # registered as custom ops so torch.compile treats the triton launches as opaque
# nodes with known output shapes (see _nvfp4_act_quant_op) # nodes with known output shapes (see _nvfp4_act_quant_op)
@torch.library.custom_op("ostris::convrot_intn_unpack", mutates_args=()) @torch.library.custom_op("ostris::convrot_intn_unpack", mutates_args=())
def _intn_unpack_op(packed: torch.Tensor, bits: int, rows: int, cols: int) -> torch.Tensor: def _intn_unpack_op(
packed: torch.Tensor, bits: int, rows: int, cols: int
) -> torch.Tensor:
return _unpack_intn_impl(packed, bits, rows, cols) return _unpack_intn_impl(packed, bits, rows, cols)
@@ -1477,6 +1592,250 @@ def _intn_unpack_grouped_fake(packed, gratio, bits, rows, cols):
return torch.empty(rows, cols, device=packed.device, dtype=torch.int8) return torch.empty(rows, cols, device=packed.device, dtype=torch.int8)
# ---------------- fused decode GEMV (small m) ---------------------------------
#
# Single-launch decode path for the int backends. The eager inference path costs
# 4 launches + 3 custom-op dispatches per linear (act_quant, unpack, _int_mm,
# epilogue) and materializes the full unpacked int8 weight every forward; at
# generation batch sizes (m <= 16) that is host-launch-bound and the unpack
# write traffic dominates GPU time. This kernel reads the packed codes directly
# (unpack stays in registers) and does act-quant + int32 dot + scale/bias
# epilogue in one launch, with arithmetic that bit-matches the eager kernels:
# same rint/clamp/scale ops for act quant and grouped unpack, int32
# accumulation like torch._int_mm, same fp32 epilogue order.
FUSED_GEMV_MAX_M = 16
_int_gemv_kernel = None
def _get_int_gemv_kernel():
global _int_gemv_kernel
if _int_gemv_kernel is not None:
return _int_gemv_kernel
import triton
import triton.language as tl
from triton.language.extra import libdevice
@triton.jit
def _unpack_lane(
word,
ratio,
J: tl.constexpr,
BITS: tl.constexpr,
QMAX_W: tl.constexpr,
GROUPED: tl.constexpr,
):
# int8 codes of in-word position J; same arithmetic as the eager
# unpack kernels (rint(code * ratio) re-expression on the row grid)
code = ((word >> (BITS * J)) & ((1 << BITS) - 1)) - QMAX_W
if GROUPED:
cf = libdevice.rint(code.to(tl.float32) * ratio)
cf = tl.minimum(tl.maximum(cf, -127.0), 127.0)
return cf.to(tl.int8)
else:
return code.to(tl.int8)
@triton.jit
def int_gemv_kernel(
x_ptr,
w_ptr,
r_ptr,
ws_ptr,
b_ptr,
o_ptr,
M,
K,
N,
w_row_stride,
ngprow,
gdiv,
QMAX_A: tl.constexpr,
BITS: tl.constexpr,
PACKED: tl.constexpr,
GROUPED: tl.constexpr,
HAS_BIAS: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
pid = tl.program_id(0)
offs_n = pid * BLOCK_N + tl.arange(0, BLOCK_N)
mask_n = offs_n < N
offs_m = tl.arange(0, 16)
mask_m = offs_m < M
qmax_w: tl.constexpr = (1 << (BITS - 1)) - 1
# per-row activation amax -> scale (bit-matches int8_act_quant_kernel:
# max is exact under any blocking, the divisions see identical operands)
amax = tl.zeros((16,), tl.float32)
for k0 in range(0, K, BLOCK_K):
offs_k = k0 + tl.arange(0, BLOCK_K)
xv = tl.load(
x_ptr + offs_m[:, None] * K + offs_k[None, :],
mask=mask_m[:, None] & (offs_k[None, :] < K),
other=0.0,
).to(tl.float32)
amax = tl.maximum(amax, tl.max(tl.abs(xv), axis=1))
scale = tl.where(amax > 0, amax / QMAX_A, 1.0)
acc = tl.zeros((16, BLOCK_N), tl.int32)
if PACKED:
# 8 codes along K share one BITS-byte word (the whole word fits
# int32 at <= 4 bits, which halves the bitfield alu cost). the 8
# in-word positions are unpacked as separate lanes and stitched
# back into k order with an interleave tree, so each K-tile feeds
# ONE tensor-core dot instead of 8 skinny ones. int32 accumulation
# is exact under any order and the per-element unpack arithmetic is
# unchanged, so the output stays bit-identical to the eager path.
# the k-group scale ratio is constant across a word (group size is
# always a multiple of 8), so it loads once per word.
KG: tl.constexpr = BLOCK_K // 8
for k0 in range(0, K, BLOCK_K):
offs_g = k0 // 8 + tl.arange(0, KG)
mask_g = offs_g * 8 < K
mask_w = mask_n[:, None] & mask_g[None, :]
if BITS * 8 <= 32:
word = tl.zeros((BLOCK_N, KG), tl.int32)
else:
word = tl.zeros((BLOCK_N, KG), tl.int64)
for b in tl.static_range(BITS):
by = tl.load(
w_ptr
+ offs_n[:, None] * w_row_stride
+ offs_g[None, :] * BITS
+ b,
mask=mask_w,
other=0,
).to(word.dtype)
word += by << (8 * b)
if GROUPED:
ratio = tl.load(
r_ptr
+ offs_n[:, None] * ngprow
+ (offs_g * 8 // gdiv)[None, :],
mask=mask_w,
other=1.0,
)
else:
ratio = word # unused (DCE'd); any tensor satisfies the call
c0 = _unpack_lane(word, ratio, 0, BITS, qmax_w, GROUPED)
c1 = _unpack_lane(word, ratio, 1, BITS, qmax_w, GROUPED)
c2 = _unpack_lane(word, ratio, 2, BITS, qmax_w, GROUPED)
c3 = _unpack_lane(word, ratio, 3, BITS, qmax_w, GROUPED)
c4 = _unpack_lane(word, ratio, 4, BITS, qmax_w, GROUPED)
c5 = _unpack_lane(word, ratio, 5, BITS, qmax_w, GROUPED)
c6 = _unpack_lane(word, ratio, 6, BITS, qmax_w, GROUPED)
c7 = _unpack_lane(word, ratio, 7, BITS, qmax_w, GROUPED)
# interleave tree: (BLOCK_N, KG) j-lanes -> (BLOCK_N, BLOCK_K)
# with columns in k order (j cycling fastest within each word)
ev = tl.interleave(tl.interleave(c0, c4), tl.interleave(c2, c6))
od = tl.interleave(tl.interleave(c1, c5), tl.interleave(c3, c7))
wq_t = tl.interleave(ev, od)
offs_k = k0 + tl.arange(0, BLOCK_K)
xv = tl.load(
x_ptr + offs_m[:, None] * K + offs_k[None, :],
mask=mask_m[:, None] & (offs_k[None, :] < K),
other=0.0,
).to(tl.float32)
qa = libdevice.rint(xv / scale[:, None])
qa = tl.minimum(tl.maximum(qa, -1.0 * QMAX_A), 1.0 * QMAX_A)
acc = tl.dot(qa.to(tl.int8), tl.trans(wq_t), acc, out_dtype=tl.int32)
else:
for k0 in range(0, K, BLOCK_K):
offs_k = k0 + tl.arange(0, BLOCK_K)
mask_k = offs_k < K
xv = tl.load(
x_ptr + offs_m[:, None] * K + offs_k[None, :],
mask=mask_m[:, None] & mask_k[None, :],
other=0.0,
).to(tl.float32)
qa = libdevice.rint(xv / scale[:, None])
qa = tl.minimum(tl.maximum(qa, -1.0 * QMAX_A), 1.0 * QMAX_A)
wq = tl.load(
w_ptr + offs_n[None, :] * w_row_stride + offs_k[:, None],
mask=mask_k[:, None] & mask_n[None, :],
other=0,
)
acc = tl.dot(qa.to(tl.int8), wq, acc, out_dtype=tl.int32)
ws = tl.load(ws_ptr + offs_n, mask=mask_n, other=0.0)
out = acc.to(tl.float32) * (scale[:, None] * ws[None, :])
if HAS_BIAS:
out += tl.load(b_ptr + offs_n, mask=mask_n, other=0.0).to(tl.float32)[
None, :
]
tl.store(
o_ptr + offs_m[:, None] * N + offs_n[None, :],
out.to(o_ptr.dtype.element_ty),
mask=mask_m[:, None] & mask_n[None, :],
)
_int_gemv_kernel = int_gemv_kernel
return _int_gemv_kernel
@torch.library.custom_op("ostris::convrot_int_gemv", mutates_args=())
def _int_gemv_op(
x2d: torch.Tensor,
qdata: torch.Tensor,
gratio: Optional[torch.Tensor],
scales_u8: torch.Tensor,
bias: Optional[torch.Tensor],
bits: int,
act_qmax: int,
out_features: int,
out_dtype: str,
) -> torch.Tensor:
m, K = x2d.shape
N = out_features
out = torch.empty(m, N, device=x2d.device, dtype=getattr(torch, out_dtype))
ws = scales_u8.view(torch.float32)
# intn stores a uint8 bitstream; the int8 backend stores raw int8 codes
packed = qdata.dtype == torch.uint8
grouped = gratio is not None
ngprow = gratio.shape[1] if grouped else 1
gdiv = K // ngprow
kernel = _get_int_gemv_kernel()
# swept on RTX 5090 across qwen-sized decode shapes: small BLOCK_N keeps
# enough programs in flight at GEMV grid sizes; the packed branch prefers
# BN=16 (heavier per-program unpack alu)
BLOCK_N = 16 if packed else 32
BLOCK_K = 256
kernel[(-(-N // BLOCK_N),)](
x2d.contiguous(),
qdata,
gratio if grouped else ws,
ws,
bias if bias is not None else ws,
out,
m,
K,
N,
qdata.shape[1],
ngprow,
gdiv,
QMAX_A=act_qmax,
BITS=bits,
PACKED=packed,
GROUPED=grouped,
HAS_BIAS=bias is not None,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
num_warps=4,
num_stages=2,
)
return out
@_int_gemv_op.register_fake
def _int_gemv_fake(
x2d, qdata, gratio, scales_u8, bias, bits, act_qmax, out_features, out_dtype
):
return torch.empty(
x2d.shape[0], out_features, device=x2d.device, dtype=getattr(torch, out_dtype)
)
# training-path linear for the grouped widths: same STE scheme as # training-path linear for the grouped widths: same STE scheme as
# _int8_linear_ste_op, but autograd saves the PACKED codes and unpacks again in # _int8_linear_ste_op, but autograd saves the PACKED codes and unpacks again in
# the backward — otherwise every layer holds a full int8-size unpacked weight # the backward — otherwise every layer holds a full int8-size unpacked weight
@@ -1499,16 +1858,23 @@ def _intn_linear_ste_op(
aq, a_s = _int8_act_quant_padded(x2d, act_qmax) aq, a_s = _int8_act_quant_padded(x2d, act_qmax)
i32 = torch._int_mm(aq, qdata.t()) i32 = torch._int_mm(aq, qdata.t())
return _int8_epilogue( return _int8_epilogue(
i32[:m], a_s[:m], w_scales_u8.view(torch.float32), bias, i32[:m],
a_s[:m],
w_scales_u8.view(torch.float32),
bias,
getattr(torch, out_dtype), getattr(torch, out_dtype),
) )
@_intn_linear_ste_op.register_fake @_intn_linear_ste_op.register_fake
def _intn_linear_ste_fake(x2d, packed, gratio, w_scales_u8, bias, bits, act_qmax, out_dtype): def _intn_linear_ste_fake(
x2d, packed, gratio, w_scales_u8, bias, bits, act_qmax, out_dtype
):
return torch.empty( return torch.empty(
x2d.shape[0], w_scales_u8.view(torch.float32).numel(), x2d.shape[0],
device=x2d.device, dtype=getattr(torch, out_dtype), w_scales_u8.view(torch.float32).numel(),
device=x2d.device,
dtype=getattr(torch, out_dtype),
) )
@@ -1547,8 +1913,16 @@ def _unpack_bitnet_impl(
BLOCK = 256 BLOCK = 256
kernel = _get_bitnet_kernel() kernel = _get_bitnet_kernel()
kernel[(-(-n_bytes // BLOCK),)]( kernel[(-(-n_bytes // BLOCK),)](
packed, gratio, out, n_bytes, bpr, cols, cols // ngprow, ngprow, packed,
BLOCK=BLOCK, num_warps=4, gratio,
out,
n_bytes,
bpr,
cols,
cols // ngprow,
ngprow,
BLOCK=BLOCK,
num_warps=4,
) )
return out return out
return unpack_ternary_rows_grouped(packed, gratio, rows, cols) return unpack_ternary_rows_grouped(packed, gratio, rows, cols)
@@ -1582,16 +1956,23 @@ def _bitnet_linear_ste_op(
aq, a_s = _int8_act_quant_padded(x2d, act_qmax) aq, a_s = _int8_act_quant_padded(x2d, act_qmax)
i32 = torch._int_mm(aq, qdata.t()) i32 = torch._int_mm(aq, qdata.t())
return _int8_epilogue( return _int8_epilogue(
i32[:m], a_s[:m], w_scales_u8.view(torch.float32), bias, i32[:m],
a_s[:m],
w_scales_u8.view(torch.float32),
bias,
getattr(torch, out_dtype), getattr(torch, out_dtype),
) )
@_bitnet_linear_ste_op.register_fake @_bitnet_linear_ste_op.register_fake
def _bitnet_linear_ste_fake(x2d, packed, gratio, w_scales_u8, bias, act_qmax, out_dtype): def _bitnet_linear_ste_fake(
x2d, packed, gratio, w_scales_u8, bias, act_qmax, out_dtype
):
return torch.empty( return torch.empty(
x2d.shape[0], w_scales_u8.view(torch.float32).numel(), x2d.shape[0],
device=x2d.device, dtype=getattr(torch, out_dtype), w_scales_u8.view(torch.float32).numel(),
device=x2d.device,
dtype=getattr(torch, out_dtype),
) )
@@ -1666,13 +2047,24 @@ class ConvRotIntNQuantizer(ConvRotInt8Quantizer):
gratio = getattr(module, "crn_gratio", None) gratio = getattr(module, "crn_gratio", None)
if gratio is not None: if gratio is not None:
return _intn_unpack_grouped_op( return _intn_unpack_grouped_op(
module.crn_qdata, gratio.view(torch.float32), module.crn_qdata,
module.crn_bits, module.out_features, module.in_features, gratio.view(torch.float32),
module.crn_bits,
module.out_features,
module.in_features,
) )
return _intn_unpack_op( return _intn_unpack_op(
module.crn_qdata, module.crn_bits, module.out_features, module.in_features module.crn_qdata, module.crn_bits, module.out_features, module.in_features
) )
def _gemv_args(self, module):
gratio = getattr(module, "crn_gratio", None)
return (
module.crn_qdata,
gratio.view(torch.float32) if gratio is not None else None,
module.crn_bits,
)
def _scales_u8(self, module) -> torch.Tensor: def _scales_u8(self, module) -> torch.Tensor:
return module.crn_scales return module.crn_scales
@@ -1686,8 +2078,14 @@ class ConvRotIntNQuantizer(ConvRotInt8Quantizer):
# grouped: autograd saves only the packed codes (+ ratios), unpacked again # grouped: autograd saves only the packed codes (+ ratios), unpacked again
# in the backward # in the backward
return _intn_linear_ste_op( return _intn_linear_ste_op(
x2d, module.crn_qdata, gratio.view(torch.float32), x2d,
module.crn_scales, module.bias, module.crn_bits, self.act_qmax, out_dtype, module.crn_qdata,
gratio.view(torch.float32),
module.crn_scales,
module.bias,
module.crn_bits,
self.act_qmax,
out_dtype,
) )
def requantize_(self, module, fp_weight: torch.Tensor) -> None: def requantize_(self, module, fp_weight: torch.Tensor) -> None:
@@ -1755,19 +2153,31 @@ class ConvRotBitNetQuantizer(ConvRotIntNQuantizer):
gratio, rscales = _intn_group_ratio_and_rscales(gscales, self.bits) gratio, rscales = _intn_group_ratio_and_rscales(gscales, self.bits)
return pack_ternary_rows(q), rscales, gratio.contiguous() return pack_ternary_rows(q), rscales, gratio.contiguous()
def _gemv_args(self, module):
# base-3 5-codes-per-byte storage doesn't match the fused gemv's
# bitfield unpack; use the eager path
return None
def _pack(self, q: torch.Tensor) -> torch.Tensor: def _pack(self, q: torch.Tensor) -> torch.Tensor:
return pack_ternary_rows(q) return pack_ternary_rows(q)
def _qdata(self, module) -> torch.Tensor: def _qdata(self, module) -> torch.Tensor:
return _bitnet_unpack_op( return _bitnet_unpack_op(
module.crn_qdata, module.crn_gratio.view(torch.float32), module.crn_qdata,
module.out_features, module.in_features, module.crn_gratio.view(torch.float32),
module.out_features,
module.in_features,
) )
def _linear_ste(self, module, x2d: torch.Tensor, out_dtype: str) -> torch.Tensor: def _linear_ste(self, module, x2d: torch.Tensor, out_dtype: str) -> torch.Tensor:
return _bitnet_linear_ste_op( return _bitnet_linear_ste_op(
x2d, module.crn_qdata, module.crn_gratio.view(torch.float32), x2d,
module.crn_scales, module.bias, self.act_qmax, out_dtype, module.crn_qdata,
module.crn_gratio.view(torch.float32),
module.crn_scales,
module.bias,
self.act_qmax,
out_dtype,
) )
@@ -1889,4 +2299,3 @@ def convrot_qat_forward(module, x: torch.Tensor) -> torch.Tensor:
x_ste = x2d + (x_dq - x2d).detach() x_ste = x2d + (x_dq - x2d).detach()
out = F.linear(x_ste, w_ste, module.bias) out = F.linear(x_ste, w_ste, module.bias)
return out.reshape(*x.shape[:-1], out_f) return out.reshape(*x.shape[:-1], out_f)

View File

@@ -133,6 +133,7 @@ def quantize(
optimizer: Optional[Optimizer] = None, optimizer: Optional[Optimizer] = None,
include: Optional[Union[str, List[str]]] = None, include: Optional[Union[str, List[str]]] = None,
exclude: Optional[Union[str, List[str]]] = None, exclude: Optional[Union[str, List[str]]] = None,
quantize_device: Optional[torch.device] = None,
): ):
"""Quantize the specified model submodules """Quantize the specified model submodules
@@ -159,6 +160,10 @@ def quantize(
exclude (`Optional[Union[str, List[str]]]`): exclude (`Optional[Union[str, List[str]]]`):
Patterns constituting the denylist. If provided, module names must not match Patterns constituting the denylist. If provided, module names must not match
any patterns from the denylist. any patterns from the denylist.
quantize_device (`Optional[torch.device]`):
If provided, each module is moved to this device to quantize, then moved
back to the device its weights were on initially. Lets a CPU-resident
model (low vram) quantize layer-by-layer on the GPU.
""" """
if include is not None: if include is not None:
include = [include] if isinstance(include, str) else include include = [include] if isinstance(include, str) else include
@@ -175,7 +180,27 @@ def quantize(
# check if m is QLinear or QConv2d # check if m is QLinear or QConv2d
if m.__class__.__name__ in Q_MODULES: if m.__class__.__name__ in Q_MODULES:
continue continue
else: if (
isinstance(weights, aotype)
and not isinstance(m, torch.nn.Linear)
and (
quantize_device is not None
or include is not None
or exclude is not None
)
):
# torchao only quantizes nn.Linear; when a device round-trip or
# include/exclude filtering is in play, skip containers so each
# linear is handled individually (a container-level torchao call
# would quantize excluded children too)
continue
orig_device = None
if quantize_device is not None and next(m.children(), None) is None:
param = next(m.parameters(recurse=False), None)
if param is not None:
orig_device = param.device
m.to(quantize_device)
try:
if isinstance(weights, ostristype): if isinstance(weights, ostristype):
if isinstance(m, torch.nn.Linear): if isinstance(m, torch.nn.Linear):
convert_linear_to_ostris(m, weights.quantizer) convert_linear_to_ostris(m, weights.quantizer)
@@ -190,6 +215,10 @@ def quantize(
activations=activations, activations=activations,
optimizer=optimizer, optimizer=optimizer,
) )
finally:
if orig_device is not None:
# quanto replaces the module in its parent, so re-fetch by name
model.get_submodule(name).to(orig_device)
except Exception as e: except Exception as e:
print(f"Failed to quantize {name}: {e}") print(f"Failed to quantize {name}: {e}")
# raise e # raise e