Download tests/test_oasis_integration.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 16.2 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/tests/test_oasis_integration.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager/tests/test_oasis_integration.py
-
curl -L -o test_oasis_integration.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/tests/test_oasis_integration.py
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}" | |