mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-18 00:47:23 +02:00
232 lines
7.8 KiB
Python
232 lines
7.8 KiB
Python
"""Integration contracts between checkpoint helpers and the model pipeline."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
from unittest.mock import MagicMock
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
import obliteratus.abliterate as abliterate
|
|
from obliteratus.abliterate import AbliterationPipeline
|
|
|
|
|
|
pytestmark = pytest.mark.cpu
|
|
|
|
|
|
def _pipeline(tmp_path: Path) -> AbliterationPipeline:
|
|
pipeline = AbliterationPipeline(
|
|
model_name="test-model",
|
|
output_dir=str(tmp_path / "checkpoint"),
|
|
method="basic",
|
|
)
|
|
pipeline._on_log = lambda _message: None
|
|
pipeline._on_stage = lambda _event: None
|
|
pipeline.handle = MagicMock()
|
|
return pipeline
|
|
|
|
|
|
def test_rebirth_rejects_invalid_metadata_before_gathering_state(tmp_path):
|
|
pipeline = _pipeline(tmp_path)
|
|
destination = pipeline.output_dir
|
|
destination.mkdir()
|
|
sentinel = destination / "sentinel"
|
|
sentinel.write_text("old", encoding="utf-8")
|
|
pipeline._build_metadata = MagicMock(return_value={"invalid": object()})
|
|
pipeline._gather_state_dict = MagicMock()
|
|
|
|
with pytest.raises(TypeError):
|
|
pipeline._rebirth()
|
|
|
|
pipeline._gather_state_dict.assert_not_called()
|
|
assert sentinel.read_text(encoding="utf-8") == "old"
|
|
|
|
|
|
def test_rebirth_insufficient_capacity_preserves_checkpoint_and_runtime_state(
|
|
tmp_path,
|
|
monkeypatch,
|
|
):
|
|
pipeline = _pipeline(tmp_path)
|
|
destination = pipeline.output_dir
|
|
destination.mkdir()
|
|
sentinel = destination / "sentinel"
|
|
sentinel.write_text("old", encoding="utf-8")
|
|
pipeline._build_metadata = MagicMock(return_value={"schema": 1})
|
|
pipeline._gather_state_dict = MagicMock(
|
|
return_value={"weight": torch.ones(4, dtype=torch.float32)},
|
|
)
|
|
pipeline._write_local_checkpoint = MagicMock()
|
|
pipeline._free_gpu_memory = MagicMock()
|
|
pipeline._cleanup_offload_dir = MagicMock()
|
|
monkeypatch.setattr(
|
|
abliterate.shutil,
|
|
"disk_usage",
|
|
lambda _path: SimpleNamespace(free=16),
|
|
)
|
|
|
|
with pytest.raises(OSError, match="Insufficient disk space"):
|
|
pipeline._rebirth()
|
|
|
|
pipeline._write_local_checkpoint.assert_not_called()
|
|
pipeline._free_gpu_memory.assert_not_called()
|
|
pipeline._cleanup_offload_dir.assert_not_called()
|
|
assert sentinel.read_text(encoding="utf-8") == "old"
|
|
assert list(tmp_path.glob(".checkpoint.staging-*")) == []
|
|
|
|
|
|
def test_rebirth_ignores_non_os_disk_probe_failure_and_promotes_checkpoint(
|
|
tmp_path,
|
|
monkeypatch,
|
|
):
|
|
pipeline = _pipeline(tmp_path)
|
|
pipeline._build_metadata = MagicMock(return_value={"schema": 1})
|
|
state_dict = {"weight": torch.ones(1)}
|
|
pipeline._gather_state_dict = MagicMock(return_value=state_dict)
|
|
|
|
def write_checkpoint(checkpoint_dir, metadata_json, received_state_dict):
|
|
assert json.loads(metadata_json) == {"schema": 1}
|
|
assert received_state_dict is state_dict
|
|
(checkpoint_dir / "model.bin").write_bytes(b"saved")
|
|
|
|
pipeline._write_local_checkpoint = MagicMock(side_effect=write_checkpoint)
|
|
pipeline._free_gpu_memory = MagicMock()
|
|
pipeline._cleanup_offload_dir = MagicMock()
|
|
|
|
def fail_probe(_path):
|
|
raise RuntimeError("filesystem probe unavailable")
|
|
|
|
monkeypatch.setattr(abliterate.shutil, "disk_usage", fail_probe)
|
|
|
|
assert pipeline._rebirth() == pipeline.output_dir
|
|
assert (pipeline.output_dir / "model.bin").read_bytes() == b"saved"
|
|
pipeline._free_gpu_memory.assert_called_once_with()
|
|
pipeline._cleanup_offload_dir.assert_called_once_with()
|
|
|
|
|
|
def test_rebirth_pushes_only_after_local_checkpoint_promotion(tmp_path, monkeypatch):
|
|
import huggingface_hub
|
|
|
|
pipeline = _pipeline(tmp_path)
|
|
pipeline.push_to_hub = "auto"
|
|
pipeline.hub_token = "test-token"
|
|
pipeline.hub_community_org = "test-org"
|
|
pipeline._build_metadata = MagicMock(return_value={"schema": 1})
|
|
pipeline._gather_state_dict = MagicMock(return_value={"weight": torch.ones(1)})
|
|
pipeline._free_gpu_memory = MagicMock()
|
|
pipeline._cleanup_offload_dir = MagicMock()
|
|
|
|
def write_checkpoint(checkpoint_dir, _metadata_json, _state_dict):
|
|
(checkpoint_dir / "model.bin").write_bytes(b"saved")
|
|
|
|
pipeline._write_local_checkpoint = MagicMock(side_effect=write_checkpoint)
|
|
api = MagicMock()
|
|
api_factory = MagicMock(return_value=api)
|
|
monkeypatch.setattr(huggingface_hub, "HfApi", api_factory)
|
|
auto_name = MagicMock(return_value="test-org/test-model-OBLITERATED")
|
|
monkeypatch.setattr(abliterate, "auto_hub_repo_id", auto_name)
|
|
|
|
assert pipeline._rebirth() == pipeline.output_dir
|
|
|
|
assert (pipeline.output_dir / "model.bin").read_bytes() == b"saved"
|
|
api_factory.assert_called_once_with(token="test-token")
|
|
auto_name.assert_called_once_with(
|
|
"test-model",
|
|
api=api,
|
|
org="test-org",
|
|
)
|
|
api.create_repo.assert_called_once_with(
|
|
"test-org/test-model-OBLITERATED",
|
|
exist_ok=True,
|
|
)
|
|
api.upload_folder.assert_called_once_with(
|
|
folder_path=str(pipeline.output_dir),
|
|
repo_id="test-org/test-model-OBLITERATED",
|
|
commit_message="OBLITERATUS: abliterated test-model (basic)",
|
|
)
|
|
|
|
|
|
def test_rebirth_uses_fallback_token_for_explicit_hub_destination(
|
|
tmp_path,
|
|
monkeypatch,
|
|
):
|
|
import huggingface_hub
|
|
|
|
pipeline = _pipeline(tmp_path)
|
|
pipeline.push_to_hub = "test-org/explicit-repo"
|
|
pipeline._build_metadata = MagicMock(return_value={"schema": 1})
|
|
pipeline._gather_state_dict = MagicMock(return_value={"weight": torch.ones(1)})
|
|
pipeline._write_local_checkpoint = MagicMock(
|
|
side_effect=lambda path, *_args: (path / "model.bin").write_bytes(b"saved"),
|
|
)
|
|
pipeline._free_gpu_memory = MagicMock()
|
|
pipeline._cleanup_offload_dir = MagicMock()
|
|
api = MagicMock()
|
|
api_factory = MagicMock(return_value=api)
|
|
monkeypatch.setattr(huggingface_hub, "HfApi", api_factory)
|
|
monkeypatch.setenv("HF_TOKEN", "fallback-token")
|
|
|
|
pipeline._rebirth()
|
|
|
|
api_factory.assert_called_once_with(token="fallback-token")
|
|
api.create_repo.assert_called_once_with(
|
|
"test-org/explicit-repo",
|
|
exist_ok=True,
|
|
)
|
|
|
|
|
|
def test_write_local_checkpoint_strips_runtime_only_state_and_writes_metadata(
|
|
tmp_path,
|
|
monkeypatch,
|
|
):
|
|
import obliteratus.lora_ablation as lora_ablation
|
|
|
|
class Quantizer:
|
|
def __init__(self):
|
|
self.models = []
|
|
|
|
def remove_quantization_config(self, model):
|
|
self.models.append(model)
|
|
|
|
class Model:
|
|
def __init__(self):
|
|
self.hf_quantizer = Quantizer()
|
|
self._weight_conversions = {"legacy": "conversion"}
|
|
self.saved = None
|
|
|
|
def save_pretrained(self, path, **kwargs):
|
|
self.saved = (path, kwargs)
|
|
|
|
pipeline = _pipeline(tmp_path)
|
|
model = Model()
|
|
tokenizer = MagicMock()
|
|
pipeline.handle = SimpleNamespace(model=model, tokenizer=tokenizer)
|
|
pipeline._lora_adapters = {"layer": (torch.ones(1), torch.ones(1))}
|
|
checkpoint_dir = tmp_path / "staging"
|
|
adapter_path = checkpoint_dir / "lora"
|
|
save_adapters = MagicMock(return_value=adapter_path)
|
|
monkeypatch.setattr(lora_ablation, "save_lora_adapters", save_adapters)
|
|
state_dict = {"weight": torch.ones(1)}
|
|
metadata_json = '{"schema": 1}'
|
|
checkpoint_dir.mkdir()
|
|
|
|
pipeline._write_local_checkpoint(checkpoint_dir, metadata_json, state_dict)
|
|
|
|
assert model.hf_quantizer.models == [model]
|
|
assert not hasattr(model, "_weight_conversions")
|
|
assert model.saved == (
|
|
checkpoint_dir,
|
|
{
|
|
"state_dict": state_dict,
|
|
"max_shard_size": "2GB",
|
|
"save_original_format": False,
|
|
},
|
|
)
|
|
tokenizer.save_pretrained.assert_called_once_with(checkpoint_dir)
|
|
save_adapters.assert_called_once_with(pipeline._lora_adapters, checkpoint_dir)
|
|
assert (checkpoint_dir / "abliteration_metadata.json").read_text(
|
|
encoding="utf-8",
|
|
) == metadata_json
|