from __future__ import annotations import math import os from dataclasses import asdict, dataclass, field from typing import Any from adam.model_profiles import ModelProfile from adam.models import SystemSnapshot @dataclass(slots=True) class SettingsRecommendation: profile_id: str epochs: int settings: dict[str, Any] = field(default_factory=dict) reasons: list[str] = field(default_factory=list) warnings: list[str] = field(default_factory=list) summary: str = "" estimated_vram_gb: float | None = None risk_level: str = "normal" confidence: str = "conservative" def to_dict(self) -> dict[str, Any]: return asdict(self) def _field_default(profile: ModelProfile, key: str, fallback: Any) -> Any: return profile.training.get(key, {}).get("default", fallback) def _clamp_to_schema(profile: ModelProfile, key: str, value: Any) -> Any: spec = profile.training.get(key, {}) kind = str(spec.get("type", "text")) try: if kind in {"int", "slider"}: numeric = int(value) return max(int(spec.get("min", numeric)), min(numeric, int(spec.get("max", numeric)))) if kind == "float": numeric = float(value) return max(float(spec.get("min", numeric)), min(numeric, float(spec.get("max", numeric)))) except (TypeError, ValueError): return spec.get("default", value) if kind == "choice": options = list(spec.get("options", [])) return value if value in options else (options[0] if options else value) return value def estimate_vram_gb(profile: ModelProfile, resolution: int | str, batch_size: int, base_model_gb: float = 0.0) -> float: """Broad VRAM estimate used only for warnings and conservative defaults.""" architecture = profile.architecture.casefold() pixels = (max(64, resolution) / 512) ** 2 if profile.id == "lora" or "lora" in architecture: base = max(6.0, base_model_gb * 1.8) return base + pixels * max(1, batch_size) * 1.2 if profile.id == "oasis" or "action_conditioned" in architecture: width, height = (resolution, resolution) if isinstance(resolution, str) and "x" in resolution: try: width, height = (int(part) for part in resolution.lower().split("x", 1)) except ValueError: width, height = (256, 144) pixels = (max(width, height) / 512) ** 2 return 4.0 + pixels * max(1, batch_size) * 2.4 if "flow" in architecture: return 2.8 + pixels * max(1, batch_size) * 1.0 if "diffusion" in architecture: return 2.2 + pixels * max(1, batch_size) * 0.9 return 3.0 + pixels * max(1, batch_size) * 0.8 def recommend_for_profile( profile: ModelProfile, *, dataset_items: int, dataset_path: str = "", resolution: int | str | None = None, snapshot: SystemSnapshot | None = None, base_model_gb: float = 0.0, ) -> SettingsRecommendation: images = max(10, int(dataset_items or 10)) raw_resolution = resolution or _field_default(profile, "resolution", 256) or 256 if isinstance(raw_resolution, str) and "x" in raw_resolution: resolution = int(raw_resolution.lower().split("x", 1)[0]) else: resolution = int(raw_resolution) reasons: list[str] = [] warnings: list[str] = [] vram_total = snapshot.vram_total_gb if snapshot and snapshot.vram_total_gb else None available_vram = ( max(0.0, snapshot.vram_total_gb - snapshot.vram_used_gb) if snapshot and snapshot.vram_total_gb else vram_total ) architecture = profile.architecture.casefold() if profile.id == "oasis": # Oasis learns labelled transitions rather than independent images. A real # dataset inspection below replaces this fallback whenever one is selected. epochs = 20 reasons.append("Oasis starts from a transition-step budget, not an image-exposure target.") else: target_exposures = 80_000 if profile.id == "lora" else 180_000 if "diffusion" in architecture else 120_000 max_epochs = 220 if profile.id == "lora" else 600 if "diffusion" in architecture else 300 epochs = max(10 if profile.id == "lora" else 25, min(max_epochs, round(target_exposures / images))) reasons.append( f"Epochs target roughly {target_exposures:,} image exposures, then clamp to the profile's safe range." ) batch_defaults = { 64: 16, 128: 12, 256: 4, 384: 2, 512: 1, 768: 1, 1024: 1, } if profile.id == "flow": batch_defaults.update({64: 12, 128: 8, 256: 4}) if profile.id == "inrflow": batch_defaults.update({64: 4, 128: 2, 256: 1}) if profile.id == "oasis": batch_defaults.update({128: 4, 256: 2, 384: 1, 512: 1}) if profile.id == "lora": batch_defaults.update({512: 2, 768: 1, 1024: 1}) nearest = min(batch_defaults, key=lambda size: abs(size - resolution)) batch_size = batch_defaults[nearest] reasons.append(f"Batch starts from the closest resolution preset ({nearest}px).") if available_vram is not None and available_vram < 8: batch_size = max(1, batch_size // 2) reasons.append("Available VRAM is below 8 GB, so batch size is reduced conservatively.") settings: dict[str, Any] = {} for key in ("resolution", "batch_size"): if key in profile.training: value = batch_size if key == "resolution": value = raw_resolution if profile.id == "oasis" else resolution settings[key] = _clamp_to_schema( profile, key, value, ) if "learning_rate" in profile.training: settings["learning_rate"] = _clamp_to_schema( profile, "learning_rate", 0.00002 if profile.id == "oasis" else 0.0001 if profile.id in {"ddpm", "lora", "inrflow"} else 0.0002, ) workers = max(1, min(8, (os.cpu_count() or 4) // 2)) for key in ("dataloader_num_workers", "workers"): if key in profile.training: settings[key] = _clamp_to_schema(profile, key, workers) for key, value in { "gradient_accumulation_steps": 1, "gradient_accumulation": 1, "mixed_precision": "fp16", "save_every": max(5, min(25, max(1, epochs // 10))), "preview_every": max(5, min(50, max(1, epochs // 10))), "training_intensity": 100, "gradient_checkpointing": resolution >= 384 or (available_vram is not None and available_vram < 8), "rank": 16, "alpha": 16, "frame_gap": 1, "sequence_context": 1, "preview_steps": 1 if profile.id == "oasis" else 50 if profile.id == "ddpm" else 10, }.items(): if key in profile.training: settings[key] = _clamp_to_schema(profile, key, value) if profile.id == "inrflow" and "query_points" in profile.training: settings["query_points"] = _clamp_to_schema( profile, "query_points", min(1024, resolution * resolution) ) reasons.append( "INRFlow starts with at most 1,024 decoded pixel queries per image to keep training memory practical." ) if profile.id == "oasis" and dataset_path: from adam.oasis_dataset import dataset_directories, inspect_oasis_dataset, oasis_pace pace = oasis_pace(dataset_path, frame_gap=int(settings.get("frame_gap", 1))) recommended_gap = pace["recommended_frame_gap"] capture_fps = pace["capture_fps"] if isinstance(recommended_gap, int) and isinstance(capture_fps, (int, float)): settings["frame_gap"] = _clamp_to_schema(profile, "frame_gap", recommended_gap) native_fps = float(capture_fps) / int(settings["frame_gap"]) reasons.append( f"The dataset records at {float(capture_fps):g} FPS, so prediction gap " f"{settings['frame_gap']} gives a native trained pace of {native_fps:g} AI FPS." ) report = inspect_oasis_dataset(dataset_path, frame_gap=int(settings.get("frame_gap", 1))) if report.ok and report.valid_transitions: transition_count = report.valid_transitions # Large datasets need bounded epochs and a rotating, balanced sample. # This keeps the recommendation in tens of thousands of updates rather # than silently turning 10K captured frames into a multi-day run. chunk_size = 5_000 if transition_count >= 7_500 else 0 transitions_per_epoch = min(transition_count, chunk_size) if chunk_size else transition_count optimizer_steps_per_epoch = math.ceil( transitions_per_epoch / max(1, int(settings.get("batch_size", batch_size))) / max(1, int(settings.get("gradient_accumulation", 1))) ) target_updates = 50_000 if transition_count >= 7_500 else 30_000 epochs = max(5, min(45, math.ceil(target_updates / max(1, optimizer_steps_per_epoch)))) for key, value in { "chunk_size": chunk_size, "chunk_mode": "balanced", "chunk_offset": 0, "balance_actions": True, "tf32": True, "contrast_every": 4, "contrast_samples": 2, "recovery_minutes": 30, "save_every": max(5, min(10, max(1, epochs // 4))), "preview_every": max(2, min(10, max(1, epochs // 5))), }.items(): if key in profile.training: settings[key] = _clamp_to_schema(profile, key, value) if len(dataset_directories(dataset_path)) > 1: for key, value in {"include_older_data": True, "replay_older_percent": 50.0}.items(): if key in profile.training: settings[key] = _clamp_to_schema(profile, key, value) reasons.append( f"{transition_count:,} valid transitions use " f"{transitions_per_epoch:,} transition(s) per epoch, about " f"{optimizer_steps_per_epoch:,} optimizer steps per epoch, and a " f"{target_updates:,}-step initial budget." ) if chunk_size: reasons.append( "A balanced 5,000-transition chunk keeps rare controls represented; " "increase the chunk offset on a later continuation to rotate the sample." ) idle_ratio = report.idle_rows / max(1, report.valid_rows) if idle_ratio < 0.05: warnings.append( f"Only {idle_ratio:.1%} of labelled frames are idle. Record more no-input gameplay " "so the world can stay stable when the player releases controls." ) rare_threshold = max(10, math.ceil(report.valid_rows * 0.01)) rare_controls = [ name for name, count in report.action_counts.items() if 0 < count < rare_threshold ] if rare_controls: warnings.append( "Rare recorded controls: " + ", ".join(rare_controls[:5]) + ". Action balancing is enabled, but more examples are still safer." ) estimated = estimate_vram_gb(profile, resolution, int(settings.get("batch_size", batch_size)), base_model_gb) if available_vram is not None and estimated > available_vram * 0.9: warnings.append( f"Estimated VRAM need is about {estimated:.1f} GB, above the conservative {available_vram * 0.9:.1f} GB working limit." ) if "batch_size" in settings and int(settings["batch_size"]) > 1: settings["batch_size"] = max(1, int(settings["batch_size"]) // 2) estimated = estimate_vram_gb(profile, resolution, int(settings["batch_size"]), base_model_gb) warnings.append(f"Batch size was reduced to {settings['batch_size']} for a safer first run.") reasons.append("The initial VRAM estimate was high, so ADAM reduced the batch before applying the recipe.") if images < 20: warnings.append("Dataset is very small; expect overfitting unless this is just a smoke test.") reasons.append("Very small datasets get a warning because quality usually depends more on data cleanup than long training.") risk_level = "risky" if warnings else "normal" memory_note = ( f" using about {available_vram:.1f} GB available VRAM" if available_vram is not None else " without detected VRAM" ) if profile.id == "oasis": summary = ( f"Recommended {epochs:,} Oasis epochs, batch {settings.get('batch_size', batch_size)}, " f"prediction gap {settings.get('frame_gap', 1)}{memory_note}. " "The recipe uses labelled transitions and is a starting point, not a guarantee." ) else: summary = ( f"Recommended {epochs:,} epochs for {images:,} item(s), " f"batch {settings.get('batch_size', batch_size)} at {resolution}px{memory_note}. " "Treat this as a starting recipe, not a guarantee." ) return SettingsRecommendation( profile_id=profile.id, epochs=epochs, settings=settings, reasons=reasons, warnings=warnings, summary=summary, estimated_vram_gb=estimated, risk_level=risk_level, )