AI_Development_Automation_Manager / tests /test_oasis_integration.py
SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
16.2 kB
from __future__ import annotations
import json
import threading
from pathlib import Path
from PIL import Image
from adam.config import ConfigManager
from adam.executor import ToolContext
from adam.oasis_dataset import validate_oasis_dataset
from adam.planner import Planner
from adam.registry import ToolRegistry, ToolSpec
from adam.tools import oasis_adapter
def _write_oasis_dataset(root: Path, *, frames: int = 4) -> Path:
dataset = root / "Minecraft Action Dataset"
frames_dir = dataset / "frames"
frames_dir.mkdir(parents=True)
rows = []
for index in range(frames):
filename = f"frame_{index:08d}.png"
Image.new("RGB", (256, 144), (index * 20, 50, 80)).save(frames_dir / filename)
rows.append({
"session_id": "session-a",
"frame_index": index,
"filename": filename,
"timestamp_seconds": index / 10,
"w": 1 if index == 1 else 0,
"a": 0,
"s": 0,
"d": 1 if index == 2 else 0,
"jump": 1 if index == 3 else 0,
"mouse_dx": 0.0,
"mouse_dy": 0.0,
"zoom": 0.0,
})
(dataset / "actions.jsonl").write_text(
"\n".join(json.dumps(row) for row in rows),
encoding="utf-8",
)
(dataset / "dataset_info.json").write_text(
json.dumps({"capture_fps": 10, "output_resolution": "256x144"}),
encoding="utf-8",
)
return dataset
def _write_registry(root: Path) -> None:
config = root / "config"
config.mkdir()
(config / "tools.json").write_text(json.dumps({"tools": []}), encoding="utf-8")
def _write_oasis_model(root: Path, name: str = "Smoke") -> Path:
model = root / "Oasis-Game-Trainer" / "output_action_flow_models" / name
(model / "unet").mkdir(parents=True)
(model / "unet" / "config.json").write_text("{}", encoding="utf-8")
(model / "action_flow_model_info.json").write_text(
json.dumps({"model_type": "action_conditioned_rectified_flow_video", "model_name": name}),
encoding="utf-8",
)
return model
def test_oasis_plugin_is_registered() -> None:
registry = ToolRegistry(Path.cwd())
tool = registry.get("oasis_trainer")
player = registry.get("oasis_player")
assert tool.model_trainers == ()
assert "resume_training" in tool.capabilities
assert registry.model_plugins.training_schema("oasis")["resolution"]["default"] == "256x144"
assert player.model_trainers == ("oasis",)
def test_oasis_dataset_validation_accepts_legacy_actions_jsonl(tmp_path: Path) -> None:
dataset = _write_oasis_dataset(tmp_path)
report = validate_oasis_dataset(str(dataset), frame_gap=1)
assert report.ok
assert report.valid_rows == 4
assert report.valid_transitions == 3
assert report.action_counts["w"] == 1
def test_oasis_dataset_validation_rejects_missing_labels(tmp_path: Path) -> None:
dataset = _write_oasis_dataset(tmp_path)
rows = (dataset / "actions.jsonl").read_text(encoding="utf-8").splitlines()
first = json.loads(rows[0])
del first["jump"]
rows[0] = json.dumps(first)
(dataset / "actions.jsonl").write_text("\n".join(rows), encoding="utf-8")
report = validate_oasis_dataset(str(dataset), frame_gap=1)
assert not report.ok
assert any("missing action label" in error for error in report.errors)
def test_oasis_dataset_validation_accepts_legacy_action_suffix_mismatch(tmp_path: Path) -> None:
dataset = _write_oasis_dataset(tmp_path, frames=3)
frames_dir = dataset / "frames"
(frames_dir / "frame_00000001.png").rename(frames_dir / "frame_00000001_CAM.png")
report = validate_oasis_dataset(str(dataset), frame_gap=1)
assert report.ok
assert report.valid_rows == 3
assert any("missing legacy filename" in warning for warning in report.warnings)
def test_oasis_dataset_validation_skips_deleted_frames(tmp_path: Path) -> None:
dataset = _write_oasis_dataset(tmp_path, frames=5)
(dataset / "frames" / "frame_00000004.png").unlink()
report = validate_oasis_dataset(str(dataset), frame_gap=1)
assert report.ok
assert report.valid_rows == 4
assert report.valid_transitions == 3
assert any("skipping row" in warning for warning in report.warnings)
def test_oasis_training_request_creates_standard_plan(tmp_path: Path, monkeypatch) -> None:
_write_registry(tmp_path)
dataset = _write_oasis_dataset(tmp_path)
oasis_root = tmp_path / "Oasis-Game-Trainer"
oasis_root.mkdir()
(oasis_root / "roblox_action_flow_app.py").write_text("# oasis", encoding="utf-8")
config = ConfigManager(tmp_path)
config.settings["provider"] = "manual"
config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)}
monkeypatch.chdir(tmp_path)
planner = Planner(tmp_path, ToolRegistry(tmp_path), config)
plan = planner.plan(
f"Create an Oasis model called Old Minecraft Beta, use the {dataset} dataset, "
"train it for 5 epochs at 256x144. "
'[ADAM_TRAINING_OPTIONS:{"resolution":"256x144","batch_size":2,"workers":0}]'
)
assert plan.requires_confirmation is True
assert [step.tool_id for step in plan.steps] == ["oasis_trainer"]
assert plan.steps[0].arguments["model_name"] == "Old Minecraft Beta"
assert Path(plan.steps[0].arguments["output_dir"]).parent.name == "output_action_flow_models"
assert plan.steps[0].arguments["resolution"] == "256x144"
def test_oasis_player_request_finds_registered_model(tmp_path: Path) -> None:
_write_registry(tmp_path)
oasis_root = tmp_path / "Oasis-Game-Trainer"
_write_oasis_model(tmp_path, "Beta World")
config = ConfigManager(tmp_path)
config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)}
planner = Planner(tmp_path, ToolRegistry(tmp_path), config)
plan = planner.plan("Open Oasis Beta World with seed 42")
assert plan.requires_confirmation is False
assert [step.tool_id for step in plan.steps] == ["oasis_player"]
assert plan.steps[0].arguments["model_name"] == "Beta World"
assert plan.steps[0].arguments["seed"] == 42
def test_oasis_player_request_rejects_invalid_explicit_folder(tmp_path: Path) -> None:
_write_registry(tmp_path)
invalid = tmp_path / "not-a-model"
invalid.mkdir()
planner = Planner(tmp_path, ToolRegistry(tmp_path), ConfigManager(tmp_path))
plan = planner.plan(f"Launch Oasis from {invalid}")
assert plan.steps == []
assert "valid Oasis action model folder" in plan.summary
def test_oasis_fine_tune_accepts_existing_dataset_path(tmp_path: Path) -> None:
_write_registry(tmp_path)
bundle = tmp_path / "Roblox Dataset"
dataset = _write_oasis_dataset(bundle)
oasis_root = tmp_path / "Oasis-Game-Trainer"
_write_oasis_model(tmp_path, "Roblox Oasis V 2.3.6")
config = ConfigManager(tmp_path)
config.settings["provider"] = "manual"
config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)}
planner = Planner(tmp_path, ToolRegistry(tmp_path), config)
plan = planner.plan(
"Fine-tune Roblox Oasis V 2.3.6 for 30 epochs with Oasis Action World Model. "
"[ADAM_FINE_TUNE:"
+ json.dumps({
"dataset_mode": "existing",
"dataset_name": str(bundle),
"epochs": 30,
"image_count": 500,
"model_name": "Roblox Oasis V 2.3.6",
"new_subject": "",
"trainer": "oasis",
"training_options": {"resolution": "256x144", "workers": 0},
})
+ "]"
)
assert plan.requires_confirmation is True
step = plan.steps[0]
assert step.tool_id == "oasis_trainer"
assert step.arguments["dataset_dir"] == str(dataset.resolve())
assert step.arguments["resume_from"] == str(
(oasis_root / "output_action_flow_models" / "Roblox Oasis V 2.3.6").resolve()
)
def test_oasis_fine_tune_normalizes_path_shaped_model_name(tmp_path: Path) -> None:
_write_registry(tmp_path)
dataset = _write_oasis_dataset(tmp_path)
oasis_root = tmp_path / "Oasis-Game-Trainer"
model = _write_oasis_model(tmp_path, "Roblox Oasis V 2.3.6")
(model / "action_flow_model_info.json").write_text(
json.dumps({
"model_type": "action_conditioned_rectified_flow_video",
"model_name": str(model),
}),
encoding="utf-8",
)
config = ConfigManager(tmp_path)
config.settings["provider"] = "manual"
config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)}
planner = Planner(tmp_path, ToolRegistry(tmp_path), config)
plan = planner.plan(
f"Fine-tune {model} for 30 epochs with Oasis Action World Model. "
"[ADAM_FINE_TUNE:"
+ json.dumps({
"dataset_mode": "existing",
"dataset_name": str(dataset),
"epochs": 30,
"image_count": 500,
"model_name": str(model),
"new_subject": "",
"trainer": "oasis",
"training_options": {"resolution": "256x144", "workers": 0},
})
+ "]"
)
assert plan.requires_confirmation is True
assert plan.steps[0].arguments["model_name"] == "Roblox Oasis V 2.3.6"
def test_oasis_fine_tune_finds_model_when_payload_path_has_old_parent(tmp_path: Path) -> None:
_write_registry(tmp_path)
dataset = _write_oasis_dataset(tmp_path)
oasis_root = tmp_path / "FlowMatchImageGenerator" / "Oasis-Game-Trainer"
model = _write_oasis_model(tmp_path / "FlowMatchImageGenerator", "Roblox Oasis V 2.3.6")
old_parent_path = (
tmp_path
/ "FlowMatchImageGenerator"
/ "output_action_flow_models"
/ "Roblox Oasis V 2.3.6"
)
config = ConfigManager(tmp_path)
config.settings["provider"] = "manual"
config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)}
planner = Planner(tmp_path, ToolRegistry(tmp_path), config)
plan = planner.plan(
f"Fine-tune {old_parent_path} for 30 epochs with Oasis Action World Model. "
"[ADAM_FINE_TUNE:"
+ json.dumps({
"dataset_mode": "existing",
"dataset_name": str(dataset),
"epochs": 30,
"image_count": 500,
"model_name": str(old_parent_path),
"new_subject": "",
"trainer": "oasis",
"training_options": {"resolution": "256x144", "workers": 0},
})
+ "]"
)
assert plan.requires_confirmation is True
assert plan.steps[0].arguments["resume_from"] == str(model.resolve())
assert plan.steps[0].arguments["model_name"] == "Roblox Oasis V 2.3.6"
def test_oasis_fine_tune_expands_dataset_bundle_to_safe_list(tmp_path: Path) -> None:
_write_registry(tmp_path)
bundle = tmp_path / "Roblox Dataset"
first = _write_oasis_dataset(bundle / "one")
second = _write_oasis_dataset(bundle / "two")
oasis_root = tmp_path / "Oasis-Game-Trainer"
_write_oasis_model(tmp_path, "Roblox Oasis V 2.3.6")
config = ConfigManager(tmp_path)
config.settings["provider"] = "manual"
config.settings["tool_folders"] = {"oasis_trainer": str(oasis_root)}
planner = Planner(tmp_path, ToolRegistry(tmp_path), config)
plan = planner.plan(
"Fine-tune Roblox Oasis V 2.3.6 for 30 epochs with Oasis Action World Model. "
"[ADAM_FINE_TUNE:"
+ json.dumps({
"dataset_mode": "existing",
"dataset_name": str(bundle),
"epochs": 30,
"image_count": 500,
"model_name": "Roblox Oasis V 2.3.6",
"new_subject": "",
"trainer": "oasis",
"training_options": {"resolution": "256x144", "workers": 0},
})
+ "]"
)
dataset_dir = plan.steps[0].arguments["dataset_dir"]
assert dataset_dir == [str(first.resolve()), str(second.resolve())]
def test_oasis_adapter_builds_worker_command_and_registers_model(tmp_path: Path, monkeypatch) -> None:
dataset = _write_oasis_dataset(tmp_path)
oasis_root = tmp_path / "Oasis-Game-Trainer"
output = oasis_root / "output_action_flow_models" / "Smoke"
oasis_root.mkdir()
(oasis_root / "roblox_action_flow_app.py").write_text("# oasis", encoding="utf-8")
config = ConfigManager(tmp_path)
config.update({"tool_folders": {"oasis_trainer": str(oasis_root)}})
captured: dict[str, object] = {}
class FakeProcess:
pid = 123
returncode = 0
stdout = iter([
'ACTION_FLOW_EVENT:{"type":"start","transitions":3,"training":2,"validation":1,"device":"cpu"}\n',
'ACTION_FLOW_EVENT:{"type":"progress","epoch":1,"epochs":1,"update":1,"total_updates":1,"loss":0.5}\n',
'ACTION_FLOW_EVENT:{"type":"complete","output_dir":"x"}\n',
])
def poll(self):
return 0
def fake_popen(command, **kwargs):
captured["command"] = command
captured["cwd"] = kwargs.get("cwd")
(output / "unet").mkdir(parents=True)
(output / "unet" / "config.json").write_text("{}", encoding="utf-8")
(output / "action_flow_model_info.json").write_text(
json.dumps({"model_type": "action_conditioned_rectified_flow_video"}),
encoding="utf-8",
)
return FakeProcess()
monkeypatch.setattr(oasis_adapter.subprocess, "Popen", fake_popen)
context = ToolContext(
tmp_path,
"OASIS",
ToolSpec("oasis_trainer", "Oasis", "test", "Training", "train_oasis"),
threading.Event(),
threading.Event(),
lambda *_args, **_kwargs: None,
lambda *_args: None,
)
context.run_event.set()
result = oasis_adapter.train_oasis(
context,
dataset_dir=str(dataset),
model_name="Smoke",
epochs=1,
output_dir=str(output),
workers=0,
preview_enabled=False,
)
command = captured["command"]
assert "--train-worker" in command
assert "--dataset-dir" in command
assert str(dataset) in command
assert "--mixed-precision" in command
assert result["assets"][0]["trainer"] == "oasis"
def test_oasis_adapter_accepts_dataset_folder_list(tmp_path: Path, monkeypatch) -> None:
first = _write_oasis_dataset(tmp_path / "one")
second = _write_oasis_dataset(tmp_path / "two")
oasis_root = tmp_path / "Oasis-Game-Trainer"
output = oasis_root / "output_action_flow_models" / "Smoke"
oasis_root.mkdir()
(oasis_root / "roblox_action_flow_app.py").write_text("# oasis", encoding="utf-8")
ConfigManager(tmp_path).update({"tool_folders": {"oasis_trainer": str(oasis_root)}})
captured: dict[str, object] = {}
class FakeProcess:
pid = 123
returncode = 0
stdout = iter([
'ACTION_FLOW_EVENT:{"type":"complete","output_dir":"x"}\n',
])
def poll(self):
return 0
def fake_popen(command, **kwargs):
captured["command"] = command
(output / "unet").mkdir(parents=True)
(output / "unet" / "config.json").write_text("{}", encoding="utf-8")
(output / "action_flow_model_info.json").write_text(
json.dumps({"model_type": "action_conditioned_rectified_flow_video"}),
encoding="utf-8",
)
return FakeProcess()
monkeypatch.setattr(oasis_adapter.subprocess, "Popen", fake_popen)
context = ToolContext(
tmp_path,
"OASIS",
ToolSpec("oasis_trainer", "Oasis", "test", "Training", "train_oasis"),
threading.Event(),
threading.Event(),
lambda *_args, **_kwargs: None,
lambda *_args: None,
)
context.run_event.set()
result = oasis_adapter.train_oasis(
context,
dataset_dir=[str(first), str(second)],
model_name="Smoke",
epochs=1,
output_dir=str(output),
workers=0,
preview_enabled=False,
)
command = captured["command"]
dataset_arg = command[command.index("--dataset-dir") + 1]
assert dataset_arg == f"{first};{second}"
assert result["assets"][0]["dataset_path"] == f"{first};{second}"