from __future__ import annotations import math import os from pathlib import Path from typing import Any IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} def dataset_image_count(raw_path: object) -> int: if not str(raw_path or "").strip(): return 0 path = Path(str(raw_path or "")).expanduser() if not path.is_dir(): return 0 try: return sum( 1 for item in path.rglob("*") if item.is_file() and item.suffix.casefold() in IMAGE_EXTENSIONS ) except OSError: return 0 def _available_vram_gb() -> float | None: try: import pynvml pynvml.nvmlInit() try: handle = pynvml.nvmlDeviceGetHandleByIndex(0) return pynvml.nvmlDeviceGetMemoryInfo(handle).total / (1024 ** 3) finally: pynvml.nvmlShutdown() except Exception: return None def _resolution_extent(value: object, fallback: int = 256) -> int: raw = str(value or fallback).strip().lower() if "x" in raw: try: return max(int(part.strip()) for part in raw.split("x", 1)) except ValueError: return fallback try: return int(raw) except ValueError: return fallback def recommend_training_settings( trainer: str, image_count: int, resolution: int, *, vram_gb: float | None = None, ) -> dict[str, Any]: """Return an explainable, conservative starting recipe for manual review.""" trainer = str(trainer).casefold() images = max(10, int(image_count)) resolution = max(64, min(512, _resolution_extent(resolution))) vram = _available_vram_gb() if vram_gb is None else vram_gb cpu_workers = max(2, min(8, (os.cpu_count() or 4) // 2)) if trainer == "lora": epochs = max(20, min(200, round(8_000 / images))) return { "epochs": epochs, "settings": {}, "summary": ( f"ORION chose {epochs} epochs from about {images:,} images. " "The connected LoRA trainer continues to own its detailed recipe." ), } exposure_target = 180_000 if trainer == "ddpm" else 120_000 minimum_epochs = 25 maximum_epochs = 600 if trainer == "ddpm" else 300 epochs = max(minimum_epochs, min(maximum_epochs, round(exposure_target / images))) batch_by_resolution = { 64: 16 if trainer == "ddpm" else 12, 128: 12 if trainer == "ddpm" else 8, 256: 4 if trainer == "ddpm" else 4, 384: 2, 512: 1, } nearest_resolution = min(batch_by_resolution, key=lambda size: abs(size - resolution)) batch = batch_by_resolution[nearest_resolution] if vram is not None and vram < 8: batch = max(1, batch // 2) gradient_checkpointing = resolution >= 384 or (vram is not None and vram < 8) settings: dict[str, Any] = { "batch_size": batch, "learning_rate": 0.0001 if trainer == "ddpm" else 0.0002, "gradient_accumulation_steps": 1, "dataloader_num_workers": cpu_workers, "mixed_precision": "fp16", "save_every": max(5, min(25, max(1, epochs // 10))), "preview_steps": 50 if trainer == "ddpm" else 10, "preview_every": max(5, min(50, max(1, epochs // 10))), "training_intensity": 100, "gradient_checkpointing": gradient_checkpointing, } if trainer == "flow": settings["gradient_accumulation"] = settings["gradient_accumulation_steps"] settings["workers"] = settings["dataloader_num_workers"] memory_note = ( f" using the detected {vram:.0f} GB GPU" if vram is not None else " without assuming a specific GPU" ) return { "epochs": epochs, "settings": settings, "summary": ( f"ORION chose about {epochs:,} epochs ({images * epochs:,} image exposures), " f"batch {batch} at {resolution}px{memory_note}. Review this starting recipe before training." ), } def review_training_plan(plan: Any) -> dict[str, Any]: """Give a conservative, explainable review without changing user settings.""" findings: list[dict[str, str]] = [] training_steps = [step for step in plan.steps if step.tool_id.endswith("_trainer")] projected_counts = { str(Path(str(step.arguments.get("output_dir", ""))).expanduser()): int( step.arguments.get("image_count", 0) or 0 ) for step in plan.steps if step.tool_id in {"dataset_collector", "youtube_video_collector"} and step.arguments.get("output_dir") } total_steps = 0 estimated_high_minutes = 0 for step in training_steps: args = step.arguments if step.tool_id == "wan_video_trainer": from adam.video_lora import clips_in count = len(clips_in(Path(str(args.get("dataset_dir", ""))))) epochs = int(args.get("epochs", 1)) findings.append({"level": "warning" if epochs > 100 or count < 20 else "ready", "message": ( f"Wan video: {count} clips, {epochs} epochs, frame buckets {args.get('target_frames', '25,49')}. " "Each clip can contribute multiple frame buckets. Inspect motion and captions; video runtime requires a measured run." )}) continue if step.tool_id == "oasis_trainer": from adam.oasis_dataset import inspect_oasis_dataset, oasis_pace gap = max(1, int(args.get("frame_gap", 1) or 1)) pace = oasis_pace(args.get("dataset_dir", ""), frame_gap=gap) recommended_gap = pace["recommended_frame_gap"] capture_fps = pace["capture_fps"] native_fps = pace["native_ai_fps"] if isinstance(capture_fps, (int, float)) and isinstance(native_fps, (int, float)): message = ( f"{step.title or 'Oasis'}: {float(capture_fps):g} FPS capture with " f"prediction gap {gap} trains at a native pace of {float(native_fps):g} AI FPS." ) if isinstance(recommended_gap, int) and recommended_gap != gap: message += f" Dataset metadata recommends gap {recommended_gap} for responsive control." findings.append({"level": "warning", "message": message}) else: findings.append({"level": "ready", "message": message}) report = inspect_oasis_dataset(args.get("dataset_dir", ""), frame_gap=gap) if report.ok and report.valid_transitions: batch = max(1, int(args.get("batch_size", 1) or 1)) accumulation = max(1, int(args.get("gradient_accumulation", 1) or 1)) requested_chunk = max(0, int(args.get("chunk_size", 0) or 0)) transitions_per_epoch = min(report.valid_transitions, requested_chunk) if requested_chunk else report.valid_transitions steps_per_epoch = math.ceil(transitions_per_epoch / batch / accumulation) epochs = max(1, int(args.get("epochs", 1) or 1)) optimizer_steps = steps_per_epoch * epochs total_steps += optimizer_steps label = step.title or "Oasis" findings.append({ "level": "ready", "message": ( f"{label}: {report.valid_transitions:,} valid transitions; " f"{transitions_per_epoch:,} used per epoch; about " f"{steps_per_epoch:,} optimizer steps per epoch." ), }) if report.valid_transitions >= 7_500 and not requested_chunk: findings.append({ "level": "warning", "message": ( f"{label}: every epoch uses all {report.valid_transitions:,} transitions. " "Use a balanced 5,000-transition chunk or explicitly confirm the longer run." ), }) if optimizer_steps >= 100_000: findings.append({ "level": "warning", "message": ( f"{label}: this plan schedules about {optimizer_steps:,} optimizer steps. " "Run the short benchmark and inspect rollout previews before committing." ), }) idle_ratio = report.idle_rows / max(1, report.valid_rows) if idle_ratio < 0.05: findings.append({ "level": "warning", "message": ( f"{label}: only {idle_ratio:.1%} of labelled frames are idle. " "Record more no-input gameplay to improve stable pauses." ), }) 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 and not bool(args.get("balance_actions", False)): findings.append({ "level": "warning", "message": ( "Rare controls are present (" + ", ".join(rare_controls[:5]) + "); turn on Balance rare actions or record more examples." ), }) continue dataset_key = str(Path(str(args.get("dataset_dir", ""))).expanduser()) images = dataset_image_count(args.get("dataset_dir")) or projected_counts.get(dataset_key, 0) epochs = max(1, int(args.get("epochs", 1) or 1)) batch = max(1, int(args.get("batch_size", 1) or 1)) accumulation = max( 1, int(args.get("gradient_accumulation_steps", args.get("gradient_accumulation", 1)) or 1), ) resolution = max(64, _resolution_extent(args.get("resolution", 256))) exposures = images * epochs if images else 0 optimizer_steps = math.ceil(images / batch / accumulation) * epochs if images else 0 total_steps += optimizer_steps # This remains deliberately broad: it is a planning guardrail, not a promise. if exposures: resolution_factor = (resolution / 256) ** 2 estimated_high_minutes += max(1, math.ceil(exposures * resolution_factor / batch * 0.003)) label = step.title or step.tool_id.replace("_", " ").title() if images >= 1_000 and epochs >= 300: suggested_epochs = max(25, min(150, round(180_000 / images))) findings.append({ "level": "warning", "message": ( f"{label}: {images:,} images × {epochs:,} epochs requests " f"{exposures:,} image exposures. This resembles a small-dataset preset. " f"Review the intent; about {suggested_epochs} epochs is a safer initial test." ), }) elif exposures >= 1_000_000: findings.append({ "level": "warning", "message": f"{label}: the plan exceeds 1,000,000 image exposures; confirm this is intentional.", }) if images >= 750 and "batch_size" in args and batch == 1 and resolution <= 256: findings.append({ "level": "warning", "message": ( f"{label}: batch size 1 at {resolution}px may leave substantial GPU capacity unused. " "Try a short test with a larger batch if VRAM allows." ), }) if epochs >= 1_000: findings.append({ "level": "warning", "message": f"{label}: {epochs:,} epochs is unusually long and deserves explicit review.", }) if not training_steps: return {} if not findings: findings.append({ "level": "ready", "message": "No obviously accidental training settings were found. Estimates are still approximate.", }) level = "warning" if any(item["level"] == "warning" for item in findings) else "ready" return { "agent": "ORION", "level": level, "headline": "Review recommended" if level == "warning" else "Plan looks reasonable", "findings": findings, "estimated_optimizer_steps": total_steps, "estimated_high_minutes": estimated_high_minutes, "settings_changed": False, } def apply_orion_review(plan: Any) -> dict[str, Any]: review = review_training_plan(plan) plan.orion_review = review if not review: return review lines = [f"ORION — {review['headline']}"] lines.extend(f"• {item['message']}" for item in review["findings"]) if review["estimated_optimizer_steps"]: lines.append(f"• Estimated optimizer steps: about {review['estimated_optimizer_steps']:,}") plan.summary += "\n\n" + "\n".join(lines) if review["level"] == "warning": plan.requires_confirmation = True reason = "ORION found unusual training settings. Review his findings before starting." if reason not in plan.confirmation_reason: plan.confirmation_reason = "; ".join(filter(None, [plan.confirmation_reason, reason])) return review