WIP working on convrot offloading

This commit is contained in:
Jaret Burkett
2026-07-11 15:37:28 -06:00
parent 1d1e21177a
commit 4625406093
5 changed files with 176 additions and 12 deletions

View File

@@ -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)

View File

@@ -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)

View File

@@ -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

View File

@@ -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,

View File

@@ -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