diff --git a/toolkit/lora_special.py b/toolkit/lora_special.py index 170442e..2b733f9 100644 --- a/toolkit/lora_special.py +++ b/toolkit/lora_special.py @@ -69,8 +69,16 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module): torch.nn.Module.__init__(self) self.lora_name = lora_name self.orig_module_ref = weakref.ref(org_module) - self.scalar = torch.tensor(1.0, device=org_module.weight.device) - + # read the device off a param/buffer directly: OstrisLinear.weight is a + # property that dequantizes the whole weight just to answer .device + org_tensor = next( + (t for t in org_module._parameters.values() if t is not None), + next((t for t in org_module._buffers.values() if t is not None), None), + ) + self.scalar = torch.tensor( + 1.0, device=org_tensor.device if org_tensor is not None else None + ) + # if is ara lora module, mark it on the layer so memory manager can handle it if is_ara: org_module.ara_lora_ref = weakref.ref(self) @@ -182,8 +190,9 @@ class FullModule(ToolkitModuleMixin, torch.nn.Module): # trainable delta, zero initialized so an untrained layer is a no-op (zero diff) # dequantize first so the delta is full precision and shaped like the real (unpacked) weight - self.weight_is_quantized = _is_quantized_tensor(org_module.weight) - ref_weight = _dequantize_if_needed(org_module.weight) + org_weight = org_module.weight # single access: dequantizes on OstrisLinear + self.weight_is_quantized = _is_quantized_tensor(org_weight) + ref_weight = _dequantize_if_needed(org_weight) self.diff = torch.nn.Parameter(torch.zeros_like(ref_weight)) # some modules (e.g. Embedding) have no bias attribute at all org_bias = getattr(org_module, 'bias', None) diff --git a/toolkit/lycoris_special.py b/toolkit/lycoris_special.py index be2250b..e3da1e1 100644 --- a/toolkit/lycoris_special.py +++ b/toolkit/lycoris_special.py @@ -79,7 +79,12 @@ class LoConSpecialModule(ToolkitModuleMixin, LoConModule, ExtractableModuleMixin self.lora_up = nn.Linear(lora_dim, out_dim, bias=use_bias) else: raise NotImplementedError - self.shape = org_module.weight.shape + # avoid the weight property on quantized OstrisLinear: it dequantizes the + # whole weight just to answer .shape + if getattr(org_module, "is_ostris_quantized", False): + self.shape = torch.Size((org_module.out_features, org_module.in_features)) + else: + self.shape = org_module.weight.shape if dropout: self.dropout = nn.Dropout(dropout) diff --git a/toolkit/memory_management/manager.py b/toolkit/memory_management/manager.py index ddf0345..c048e94 100644 --- a/toolkit/memory_management/manager.py +++ b/toolkit/memory_management/manager.py @@ -1,5 +1,10 @@ import torch -from .manager_modules import LinearLayerMemoryManager, ConvLayerMemoryManager, _DEVICE_STATE +from .manager_modules import ( + LinearLayerMemoryManager, + ConvLayerMemoryManager, + OstrisLinearLayerMemoryManager, + _DEVICE_STATE, +) import random LINEAR_MODULES = [ @@ -105,10 +110,16 @@ class MemoryManager: if skip: module._memory_manager.unmanaged_modules.append(child_module) else: - # linear - LinearLayerMemoryManager.attach( - child_module, module._memory_manager - ) + # linear; OstrisLinear bounces its quantized buffers instead + # of a dequantized weight (module.weight is a property) + if getattr(child_module, "is_ostris_quantized", False): + OstrisLinearLayerMemoryManager.attach( + child_module, module._memory_manager + ) + else: + LinearLayerMemoryManager.attach( + child_module, module._memory_manager + ) # attach to ARA as well if hasattr(child_module, "ara_lora_ref"): ara = child_module.ara_lora_ref() @@ -195,7 +206,9 @@ class MemoryManager: child.forward = original_forward for param_name in ("weight", "bias"): - param = getattr(child, param_name, None) + # read _parameters directly: OstrisLinear.weight is a property that + # materializes a full dequantized weight on access + param = child._parameters.get(param_name, None) if param is None or not isinstance(param, torch.nn.Parameter): continue try: @@ -211,6 +224,20 @@ class MemoryManager: except Exception: pass + if getattr(child, "is_ostris_quantized", False): + # move quantized buffers home and unpin them (clone drops pinning) + for buf_name, buf in list(child._buffers.items()): + if buf is None: + continue + try: + if buf.device.type != "cpu": + buf = buf.to("cpu") + if buf.is_pinned(): + buf = buf.clone() + child._buffers[buf_name] = buf + except Exception: + pass + del child._layer_memory_manager if hasattr(child, "_memory_management_device"): del child._memory_management_device diff --git a/toolkit/memory_management/manager_modules.py b/toolkit/memory_management/manager_modules.py index 4ea5b0c..44f3b28 100644 --- a/toolkit/memory_management/manager_modules.py +++ b/toolkit/memory_management/manager_modules.py @@ -637,6 +637,124 @@ class LinearLayerMemoryManager(BaseLayerMemoryManager): self.module._memory_management_device = self.manager.process_device +class OstrisLinearLayerMemoryManager(BaseLayerMemoryManager): + """Offload manager for OstrisLinear (custom-quantized) layers. + + The generic linear bounce is wrong for these: module.weight is a property that + fully dequantizes on access, so bouncing it ships a full-precision weight over + PCIe every forward and bypasses the quantizer's hardware kernels. Instead this + keeps the (much smaller) quantized buffers pinned on CPU, stages them H2D into + the same forward ring the float path uses, swaps them onto the module, and runs + the quantizer's own forward on device — so fp4/int8 GEMM paths and the STE + training path work unchanged under offloading. Buffers are read live off the + module each forward (not cached) so requantize_ during merge/reset stays valid. + """ + + def __init__( + self, + module: nn.Module, + manager: "MemoryManager", + ): + super().__init__(module, manager) + + # 1) Move quantized buffers + bias to CPU and pin for fast async H2D + with torch.no_grad(): + for name, buf in list(module._buffers.items()): + if buf is None: + continue + if buf.device.type != "cpu": + buf = buf.to("cpu") + if torch.cuda.is_available() and not buf.is_pinned(): + try: + buf = buf.pin_memory() + except RuntimeError: + pass + module._buffers[name] = buf + bias = module._parameters.get("bias", None) + if bias is not None: + bias.data = _ensure_cpu_pinned(bias.data).detach() + + # 2) Hijack forward + if hasattr(self.module, "ara_lora_ref"): + # ARA, we need to replace the lora forward + self._original_forward = getattr(self.module.ara_lora_ref(), "org_forward") + else: + self._original_forward = getattr(self.module, "forward") + + def _mm_forward(x, *args, **kwargs): + # ensure we only use expected signature (Linear: x) + if args or kwargs: + return self._original_forward(x, *args, **kwargs) + + module = self.module + device = self.manager.process_device + if device.type != "cuda": + return self._original_forward(x) + + cpu_bufs = { + n: b + for n, b in module._buffers.items() + if b is not None and b.device.type == "cpu" + } + bias = module._parameters.get("bias", None) + bias_cpu = ( + bias.data + if bias is not None and bias.data.device.type == "cpu" + else None + ) + if not cpu_bufs and bias_cpu is None: + # already resident on device + return self._original_forward(x) + + state = _get_device_state(device) + d = state["depth"] + idx = state["forward_clk"] + state["forward_clk"] = (idx + 1) % d + ts = state["transfer_stream"] + # the guard makes current_stream() resolve to the process device and + # keeps that device's context active for the quantizer's triton + # kernels (nothing sets the global current device, so it is 0 even + # when training on another gpu) + with torch.cuda.device(device): + with torch.cuda.stream(ts): + ts.wait_event(state["fwd_slot_free"][idx]) + gpu_bufs = { + n: b.to(device, non_blocking=True) for n, b in cpu_bufs.items() + } + gpu_bias = ( + bias_cpu.to(device, non_blocking=True) + if bias_cpu is not None + else None + ) + state["w_buffers"][idx] = gpu_bufs + state["b_buffers"][idx] = gpu_bias + state["fwd_slot_ready"][idx].record() + torch.cuda.current_stream().wait_event(state["fwd_slot_ready"][idx]) + + # swap the quantized state onto the device, run the quantizer's own + # forward, then swap the pinned CPU state back + for n, t in gpu_bufs.items(): + module._buffers[n] = t + if gpu_bias is not None: + bias.data = gpu_bias + try: + out = self._original_forward(x) + finally: + for n, t in cpu_bufs.items(): + module._buffers[n] = t + if bias_cpu is not None: + bias.data = bias_cpu + _release_forward_slot(state, idx) + return out + + if hasattr(self.module, "ara_lora_ref"): + self.module.ara_lora_ref().org_forward = _mm_forward + else: + self.module.forward = _mm_forward + + self.module._memory_management_device = self.manager.process_device + + class ConvLayerMemoryManager(BaseLayerMemoryManager): def __init__( self, diff --git a/toolkit/models/lokr.py b/toolkit/models/lokr.py index 74de635..48fa42a 100644 --- a/toolkit/models/lokr.py +++ b/toolkit/models/lokr.py @@ -104,7 +104,12 @@ class LokrModule(ToolkitModuleMixin, nn.Module): self.use_w2 = False self.can_merge_in = True - self.shape = org_module.weight.shape + # avoid the weight property on quantized OstrisLinear: it dequantizes the + # whole weight just to answer .shape + if getattr(org_module, "is_ostris_quantized", False): + self.shape = torch.Size((org_module.out_features, org_module.in_features)) + else: + self.shape = org_module.weight.shape if org_module.__class__.__name__ == 'Conv2d': in_dim = org_module.in_channels k_size = org_module.kernel_size