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:
@@ -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__":
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user