File size: 13,627 Bytes
c61c435 f8c73f9 c61c435 f8c73f9 c61c435 f8c73f9 c61c435 f8c73f9 c61c435 f8c73f9 c61c435 f8c73f9 c61c435 f8c73f9 c61c435 f8c73f9 c61c435 f8c73f9 c61c435 | 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 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 | 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,
)
|