From a5f857ddb0093d96c9252230fee65e4adcbdf8f0 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Mon, 13 Jul 2026 18:48:56 -0600 Subject: [PATCH] Added patch from Fatalis to fix lokr offloading with convrot --- toolkit/models/lokr.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/toolkit/models/lokr.py b/toolkit/models/lokr.py index 48fa42a..253def7 100644 --- a/toolkit/models/lokr.py +++ b/toolkit/models/lokr.py @@ -327,6 +327,25 @@ class LokrModule(ToolkitModuleMixin, nn.Module): orig_dtype = x.dtype + # for OstrisLinear (convrot-quantized), use the quantizer's own forward + # for the base computation and add the LoKr delta separately. this avoids + # materializing the full dequantized weight every forward, which bypasses + # the hardware fp4/int8 GEMM path and pegs the CPU with repeated + # dequantization + kron + matmul for every training step. + if getattr(self.org_module[0], "is_ostris_quantized", False): + base_out = self.org_forward(x) + lokr_weight = self.get_weight().to(dtype=orig_dtype) + multiplier = self.network_ref().torch_multiplier + multiplier = torch.mean(multiplier) + # bias is handled by org_forward (quantizer includes it) + delta_out = self.op( + x.to(dtype=lokr_weight.dtype), + lokr_weight.view(self.shape), + None, + **self.extra_args + ) + return (base_out + delta_out * multiplier).to(orig_dtype) + orig_weight = self.get_orig_weight(x.device) lokr_weight = self.get_weight(orig_weight).to(dtype=orig_weight.dtype) multiplier = self.network_ref().torch_multiplier