WIP working on convrot offloading
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user