File size: 4,665 Bytes
e0265b9 f8c73f9 e0265b9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 | from __future__ import annotations
from pathlib import Path
import json
import pytest
from adam.assets import AssetRegistry
from adam.commands import CommandValidationError, TrainingCommand
class FakeConfig:
def __init__(self, values: dict) -> None:
self.values = values
def get(self, key: str, default=None):
return self.values.get(key, default)
def test_asset_registry_prefers_exact_friendly_name(tmp_path: Path) -> None:
exact = tmp_path / "Mario"
similar = tmp_path / "Mario 2"
exact.mkdir()
similar.mkdir()
registry = AssetRegistry(tmp_path)
registry.register(kind="dataset", name="Mario", path=str(exact))
registry.register(kind="dataset", name="Mario 2", path=str(similar))
assert [item.name for item in registry.find("dataset", "Mario")] == ["Mario"]
def test_asset_discovery_removes_models_whose_paths_were_deleted(tmp_path: Path) -> None:
model = tmp_path / "LoRA output" / "Crystal_Biter"
model.mkdir(parents=True)
registry = AssetRegistry(tmp_path)
registry.register(kind="model", name="Crystal_Biter", path=str(model), trainer="lora")
model.rmdir()
registry.discover(FakeConfig({"tool_folders": {}}))
assert registry.find("model", "Crystal_Biter", trainer="lora") == []
def test_lora_discovery_registers_weights_in_nested_trainer_output(tmp_path: Path) -> None:
trainer = tmp_path / "LoRATrainer"
weight = trainer / "output" / "Named run" / "adapter" / "My_LoRA.safetensors"
weight.parent.mkdir(parents=True)
weight.write_bytes(b"weights")
(weight.parent.parent / "model_info.json").write_text(
'{"trigger_word": "my_lora"}', encoding="utf-8"
)
registry = AssetRegistry(tmp_path)
registry.discover(FakeConfig({"tool_folders": {"lora_trainer": str(trainer)}}))
model = registry.find("model", "My LoRA", trainer="lora")[0]
assert Path(model.path) == weight.resolve()
assert model.metadata == {"trigger_word": "my_lora"}
def test_lora_discovery_excludes_intermediate_epoch_checkpoints(tmp_path: Path) -> None:
trainer = tmp_path / "LoRATrainer"
output = trainer / "output" / "Named run"
output.mkdir(parents=True)
(output / "My_LoRA.safetensors").write_bytes(b"final")
(output / "My_LoRA_epoch_0050.safetensors").write_bytes(b"checkpoint")
(output / "checkpoint-e50_s100.safetensors").write_bytes(b"checkpoint")
registry = AssetRegistry(tmp_path)
# Simulate an index written by an older ADAM release.
registry.register(
kind="model", name="My_LoRA_epoch_0050",
path=str(output / "My_LoRA_epoch_0050.safetensors"), trainer="lora",
)
registry.discover(FakeConfig({"tool_folders": {"lora_trainer": str(trainer)}}))
assert [item.name for item in registry.assets if item.trainer == "lora"] == ["My_LoRA"]
def test_training_command_rejects_uncontrolled_fields() -> None:
with pytest.raises(CommandValidationError, match="Unsupported command fields"):
TrainingCommand.from_dict(
{
"action": "train",
"trainer": "ddpm",
"dataset": "dataset",
"model_name": "model",
"epochs": 10,
"shell_command": "unsafe",
}
)
def test_training_command_accepts_only_safe_ddpm_options() -> None:
command = TrainingCommand.from_dict({
"action": "train", "trainer": "ddpm", "dataset": "dataset", "model_name": "model", "epochs": 10,
"training_options": {"resolution": 256, "batch_size": 2, "learning_rate": 0.0001},
})
assert command.training_options["resolution"] == 256
def test_flow_discovery_recovers_dataset_from_adam_job_history(tmp_path: Path) -> None:
flow_root = tmp_path / "Flow"
dataset = tmp_path / "Dataset"
model = flow_root / "output_flow_models" / "Model"
dataset.mkdir()
(model / "unet").mkdir(parents=True)
(model / "unet" / "config.json").write_text("{}", encoding="utf-8")
(model / "flow_model_info.json").write_text(
'{"model_type":"rectified_flow","name":"Friendly Flow","resolution":128}',
encoding="utf-8",
)
(tmp_path / "data").mkdir()
(tmp_path / "data" / "jobs.json").write_text(json.dumps({"jobs": [{"plan": {"steps": [{
"tool_id": "flow_trainer", "arguments": {
"output_dir": str(model), "dataset_dir": str(dataset),
},
}]}}]}), encoding="utf-8")
registry = AssetRegistry(tmp_path)
registry.discover(FakeConfig({"tool_folders": {"flow_trainer": str(flow_root)}}))
flow = registry.find("model", "Friendly Flow", trainer="flow")[0]
assert flow.dataset_id
|