Files
OBLITERATUS/tests/test_offload_surgery.py
T

392 lines
13 KiB
Python

"""CPU-only contracts for transactional Accelerate offload surgery."""
from __future__ import annotations
from collections.abc import MutableMapping
import pytest
import torch
import torch.nn as nn
from accelerate.big_modeling import disk_offload
from accelerate.hooks import (
AlignDevicesHook,
ModelHook,
SequentialHook,
add_hook_to_module,
)
from accelerate.utils.modeling import get_state_dict_offloaded_model
import obliteratus.models.offload_surgery as offload_surgery
from obliteratus.abliterate import AbliterationPipeline
from obliteratus.models.offload_surgery import (
OffloadSurgeryError,
UnsupportedOffloadLayoutError,
UnsupportedOffloadedQuantizationError,
logical_module_device,
resolve_logical_parameter,
validate_offloaded_parameters,
)
def _direct_offload(module: nn.Module) -> AlignDevicesHook:
hook = AlignDevicesHook(execution_device="cpu", offload=True)
add_hook_to_module(module, hook)
return hook
def test_direct_align_hook_updates_backing_weight_and_preserves_meta_state():
module = nn.Linear(3, 2, bias=False)
original = module.weight.detach().clone()
hook = _direct_offload(module)
parameter_identity = id(module.weight)
transaction = resolve_logical_parameter(module)
transaction.commit(transaction.tensor + 2)
assert module.weight.device.type == "meta"
assert id(module.weight) == parameter_identity
torch.testing.assert_close(hook.weights_map["weight"], original + 2)
def test_sequential_hook_finds_nested_align_devices_hook():
module = nn.Linear(3, 2, bias=False)
original = module.weight.detach().clone()
align = AlignDevicesHook(execution_device="cpu", offload=True)
add_hook_to_module(module, SequentialHook(ModelHook(), align))
transaction = resolve_logical_parameter(module)
transaction.commit(transaction.tensor - 1)
assert module.weight.device.type == "meta"
torch.testing.assert_close(align.weights_map["weight"], original - 1)
def test_parent_prefixed_weight_map_resolves_nested_parameter():
block = nn.Module()
block.self_attn = nn.Module()
block.self_attn.o_proj = nn.Linear(3, 2, bias=False)
original = block.self_attn.o_proj.weight.detach().clone()
hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
place_submodules=True,
)
add_hook_to_module(block, hook)
transaction = resolve_logical_parameter(
block.self_attn.o_proj,
search_roots=(block,),
)
assert transaction.backing_key == "self_attn.o_proj.weight"
transaction.commit(transaction.tensor * 0.5)
assert block.self_attn.o_proj.weight.device.type == "meta"
torch.testing.assert_close(
hook.weights_map["self_attn.o_proj.weight"],
original * 0.5,
)
def test_disk_backed_prefixed_loader_uses_updated_authoritative_value(tmp_path):
model = nn.Sequential(nn.Linear(3, 2, bias=False))
original = model[0].weight.detach().clone()
disk_offload(model, tmp_path, execution_device=torch.device("cpu"))
transaction = resolve_logical_parameter(model[0], search_roots=(model,))
transaction.commit(transaction.tensor + 3)
assert model[0].weight.device.type == "meta"
hook = model[0]._hf_hook
torch.testing.assert_close(hook.weights_map["weight"], original + 3)
gathered = get_state_dict_offloaded_model(model)
assert all(tensor.device.type != "meta" for tensor in gathered.values())
torch.testing.assert_close(gathered["0.weight"], original + 3)
def test_save_reload_observes_surgery_instead_of_stale_backing_weight():
model = nn.Sequential(nn.Linear(3, 2, bias=False))
original = model[0].weight.detach().clone()
_direct_offload(model[0])
transaction = resolve_logical_parameter(model[0], search_roots=(model,))
transaction.commit(transaction.tensor - 4)
state_dict = get_state_dict_offloaded_model(model)
reloaded = nn.Sequential(nn.Linear(3, 2, bias=False))
reloaded.load_state_dict(state_dict)
torch.testing.assert_close(reloaded[0].weight, original - 4)
assert reloaded[0].weight.device.type == "cpu"
def test_tied_backing_values_retain_identity_and_consistent_values():
block = nn.Module()
block.first = nn.Linear(3, 3, bias=False)
block.second = nn.Linear(3, 3, bias=False)
shared = block.first.weight.detach().clone()
shared_meta = nn.Parameter(torch.empty_like(shared, device="meta"))
block.first.weight = shared_meta
block.second.weight = shared_meta
hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
weights_map={"first.weight": shared, "second.weight": shared},
place_submodules=True,
)
block._hf_hook = hook
live_identity = id(block.first.weight)
transaction = resolve_logical_parameter(block.first, search_roots=(block,))
transaction.commit(transaction.tensor + 1)
first = hook.weights_map["first.weight"]
second = hook.weights_map["second.weight"]
assert first is second
torch.testing.assert_close(first, shared + 1)
assert block.first.weight is block.second.weight
assert id(block.first.weight) == live_identity
class _FailOnceAliasMap(MutableMapping[str, torch.Tensor]):
def __init__(self, shared: torch.Tensor):
self.values = {"first.weight": shared, "second.weight": shared}
self.failed = False
def __getitem__(self, key: str) -> torch.Tensor:
return self.values[key]
def __setitem__(self, key: str, value: torch.Tensor) -> None:
if key == "second.weight" and not self.failed:
self.failed = True
raise OSError("simulated backing-store failure")
self.values[key] = value
def __delitem__(self, key: str) -> None:
del self.values[key]
def __iter__(self):
return iter(self.values)
def __len__(self) -> int:
return len(self.values)
def test_commit_failure_rolls_back_all_backing_aliases_and_leaves_meta_state():
block = nn.Module()
block.first = nn.Linear(3, 3, bias=False)
original = block.first.weight.detach().clone()
block.first.weight = nn.Parameter(torch.empty_like(original, device="meta"))
backing = _FailOnceAliasMap(original)
block._hf_hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
weights_map=backing,
place_submodules=True,
)
transaction = resolve_logical_parameter(block.first, search_roots=(block,))
with pytest.raises(OffloadSurgeryError, match="rolled back"):
transaction.commit(transaction.tensor + 5)
torch.testing.assert_close(backing["first.weight"], original)
torch.testing.assert_close(backing["second.weight"], original)
assert block.first.weight.device.type == "meta"
def test_quantized_offloaded_weight_fails_before_any_mutation():
module = nn.Linear(3, 2, bias=False, device="meta")
backing = {"weight": torch.ones(2, 3, dtype=torch.uint8)}
module._hf_hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
weights_map=backing,
)
with pytest.raises(UnsupportedOffloadedQuantizationError):
resolve_logical_parameter(module)
assert module.weight.device.type == "meta"
assert backing["weight"].dtype == torch.uint8
assert torch.equal(backing["weight"], torch.ones(2, 3, dtype=torch.uint8))
def test_read_only_unknown_mapping_fails_during_resolution():
module = nn.Linear(3, 2, bias=False, device="meta")
class MappingOnly:
def __getitem__(self, key):
return torch.ones(2, 3)
def __iter__(self):
return iter(("weight",))
def __len__(self):
return 1
# Registering through Mapping is deliberately avoided: this object models
# a future wrapper with read access but no supported write contract.
hook = AlignDevicesHook(execution_device="cpu", offload=True)
hook.weights_map = MappingOnly()
module._hf_hook = hook
with pytest.raises(UnsupportedOffloadLayoutError, match="read-only"):
resolve_logical_parameter(module)
def test_validation_resolves_all_meta_parameters_before_surgery():
block = nn.Module()
block.proj = nn.Linear(3, 2, bias=False)
hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
place_submodules=True,
)
add_hook_to_module(block, hook)
validate_offloaded_parameters(block)
hook.weights_map = {}
with pytest.raises(UnsupportedOffloadLayoutError, match="attempted keys"):
validate_offloaded_parameters(block)
def test_advanced_projection_updates_offloaded_logical_weight():
module = nn.Module()
module.proj = nn.Linear(3, 2, bias=False)
original = module.proj.weight.detach().clone()
hook = _direct_offload(module.proj)
direction = torch.tensor([[1.0], [0.0], [0.0]])
count = AbliterationPipeline._project_out_advanced(
module,
direction,
["proj"],
)
assert count == 1
assert module.proj.weight.device.type == "meta"
expected = original.clone()
expected[:, 0] = 0
torch.testing.assert_close(hook.weights_map["weight"], expected)
def test_advanced_projection_resolves_all_candidates_before_mutation():
block = nn.Module()
block.first = nn.Linear(3, 2, bias=False)
first_original = block.first.weight.detach().clone()
first_hook = _direct_offload(block.first)
block.second = nn.Linear(3, 2, bias=False, device="meta")
block.second._hf_hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
weights_map={"weight": torch.ones(2, 3, dtype=torch.uint8)},
)
with pytest.raises(UnsupportedOffloadedQuantizationError):
AbliterationPipeline._project_out_advanced(
block,
torch.tensor([[1.0], [0.0], [0.0]]),
["first", "second"],
)
torch.testing.assert_close(first_hook.weights_map["weight"], first_original)
assert block.first.weight.device.type == "meta"
def test_logical_module_device_uses_authoritative_backing_device():
module = nn.Linear(3, 2, bias=False)
_direct_offload(module)
assert logical_module_device(module) == torch.device("cpu")
def test_unknown_accelerate_major_version_fails_closed(monkeypatch):
module = nn.Linear(3, 2, bias=False)
original = module.weight.detach().clone()
hook = _direct_offload(module)
monkeypatch.setattr(offload_surgery, "version", lambda _name: "2.0.0")
with pytest.raises(UnsupportedOffloadLayoutError, match="outside the tested"):
resolve_logical_parameter(module)
assert module.weight.device.type == "meta"
torch.testing.assert_close(hook.weights_map["weight"], original)
def test_bias_projection_updates_parent_prefixed_backing_value():
block = nn.Module()
block.proj = nn.Linear(3, 3, bias=True)
original = block.proj.bias.detach().clone()
hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
place_submodules=True,
)
add_hook_to_module(block, hook)
count = AbliterationPipeline._project_bias(
block,
torch.tensor([[1.0], [0.0], [0.0]]),
["proj"],
offload_roots=(block,),
)
assert count == 1
assert block.proj.bias.device.type == "meta"
expected = original.clone()
expected[0] = 0
torch.testing.assert_close(hook.weights_map["proj.bias"], expected)
class _FailSecondParameterMap(MutableMapping[str, torch.Tensor]):
def __init__(self, first: torch.Tensor, second: torch.Tensor):
self.values = {"first.weight": first, "second.weight": second}
self.failed = False
def __getitem__(self, key: str) -> torch.Tensor:
return self.values[key]
def __setitem__(self, key: str, value: torch.Tensor) -> None:
if key == "second.weight" and not self.failed:
self.failed = True
raise OSError("simulated second-parameter failure")
self.values[key] = value
def __delitem__(self, key: str) -> None:
del self.values[key]
def __iter__(self):
return iter(self.values)
def __len__(self) -> int:
return len(self.values)
def test_multi_parameter_projection_rolls_back_prior_commits():
block = nn.Module()
block.first = nn.Linear(3, 2, bias=False)
block.second = nn.Linear(3, 2, bias=False)
first = block.first.weight.detach().clone()
second = block.second.weight.detach().clone()
block.first.weight = nn.Parameter(torch.empty_like(first, device="meta"))
block.second.weight = nn.Parameter(torch.empty_like(second, device="meta"))
backing = _FailSecondParameterMap(first, second)
block._hf_hook = AlignDevicesHook(
execution_device="cpu",
offload=True,
weights_map=backing,
place_submodules=True,
)
with pytest.raises(OffloadSurgeryError, match="rolled back"):
AbliterationPipeline._project_out_advanced(
block,
torch.tensor([[1.0], [0.0], [0.0]]),
["first", "second"],
offload_roots=(block,),
)
torch.testing.assert_close(backing["first.weight"], first)
torch.testing.assert_close(backing["second.weight"], second)
assert block.first.weight.device.type == "meta"
assert block.second.weight.device.type == "meta"