ADAM October 2026 source release: PixelRow, INRFlow, Wan Video, Oasis player and field guide
f8c73f9 verified Download adam/tools/ddpm_adapter.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 27.1 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/adam/tools/ddpm_adapter.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager/adam/tools/ddpm_adapter.py
-
curl -L -o ddpm_adapter.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/main/adam/tools/ddpm_adapter.py
27.1 kB
| """Safe execution bridge for the user's existing DDPM command-line trainer.""" | |
| from __future__ import annotations | |
| import json | |
| import importlib.util | |
| import math | |
| import queue | |
| import re | |
| import shutil | |
| import statistics | |
| import subprocess | |
| import sys | |
| import threading | |
| import time | |
| from pathlib import Path | |
| from adam.config import ConfigManager | |
| from adam.executor import ToolAdjustmentRequested, ToolCancelled, ToolContext, ToolExecutionError | |
| from adam.progressive_training import parse_stages, stage_batch_settings, stage_summary | |
| from adam.process_control import set_process_tree_paused, terminate_process_tree | |
| IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} | |
| FORCE_STOP_TIMEOUT_SECONDS = 30 | |
| def _image_files(folder: Path): | |
| """Yield training images below a dataset folder without loading them into memory.""" | |
| try: | |
| yield from ( | |
| path for path in folder.rglob("*") | |
| if path.is_file() and path.suffix.casefold() in IMAGE_EXTENSIONS | |
| ) | |
| except OSError: | |
| return | |
| def _training_image_folder(dataset: Path) -> tuple[Path, int]: | |
| """Choose the accepted-frame tree for video datasets, or the dataset itself. | |
| YouTube collections retain their provenance by storing accepted frames in | |
| ``frames/<source>/`` and rejected candidates separately. Passing the root | |
| to a recursive trainer would include rejected images, while only checking | |
| the root makes the collection appear empty. | |
| """ | |
| frames = dataset / "frames" | |
| if frames.is_dir(): | |
| frame_count = sum(1 for _ in _image_files(frames)) | |
| if frame_count: | |
| return frames, frame_count | |
| return dataset, sum(1 for _ in _image_files(dataset)) | |
| def _saved_unet_size(path: Path) -> tuple[int, int] | None: | |
| config_path = path / "unet" / "config.json" | |
| if not config_path.is_file(): | |
| return None | |
| try: | |
| data = json.loads(config_path.read_text(encoding="utf-8")) | |
| except (OSError, json.JSONDecodeError): | |
| return None | |
| sample_size = data.get("sample_size") | |
| try: | |
| if isinstance(sample_size, list) and len(sample_size) >= 2: | |
| return int(sample_size[0]), int(sample_size[1]) | |
| if sample_size is not None: | |
| size = int(sample_size) | |
| return size, size | |
| return None | |
| except (TypeError, ValueError): | |
| return None | |
| def _snap_dimension(value: float, *, multiple: int = 16) -> int: | |
| return max(64, min(512, int(round(value / multiple)) * multiple)) | |
| def _dataset_aspect_ratio(dataset: Path) -> float: | |
| from PIL import Image | |
| ratios: list[float] = [] | |
| for path in _image_files(dataset): | |
| try: | |
| with Image.open(path) as image: | |
| if image.width > 0 and image.height > 0: | |
| ratios.append(image.width / image.height) | |
| except OSError: | |
| continue | |
| return statistics.median(ratios) if ratios else 1.0 | |
| def _training_canvas(dataset: Path, resolution: int, aspect_ratio: str) -> tuple[int, int]: | |
| ratios = { | |
| "1:1 (Square)": 1.0, | |
| "16:9 (Widescreen)": 16 / 9, | |
| "9:16 (Portrait)": 9 / 16, | |
| "4:3 (Classic)": 4 / 3, | |
| "3:4 (Portrait Classic)": 3 / 4, | |
| "3:2 (Photo)": 3 / 2, | |
| "2:3 (Portrait Photo)": 2 / 3, | |
| } | |
| ratio = _dataset_aspect_ratio(dataset) if aspect_ratio == "Dataset (Auto)" else ratios.get(aspect_ratio) | |
| if ratio is None or ratio <= 0: | |
| raise ToolExecutionError("Choose a supported DDPM training aspect ratio.") | |
| if abs(ratio - 1.0) < 0.01: | |
| return resolution, resolution | |
| if ratio > 1: | |
| width, height = resolution, _snap_dimension(resolution / ratio) | |
| else: | |
| width, height = _snap_dimension(resolution * ratio), resolution | |
| return width, height | |
| def _latest_preview(folder: Path) -> Path | None: | |
| try: | |
| images = [ | |
| path for path in folder.rglob("*") | |
| if path.is_file() and path.suffix.lower() in IMAGE_EXTENSIONS | |
| and any(token in path.name.lower() for token in ("preview", "sample", "epoch")) | |
| ] | |
| return max(images, key=lambda path: path.stat().st_mtime) if images else None | |
| except OSError: | |
| return None | |
| def _safe_model_name(value: str) -> str: | |
| name = value.strip() | |
| if not name or len(name) > 96 or any(char in name for char in "<>:\\|?*\x00"): | |
| raise ToolExecutionError("Choose a short model name without filesystem-reserved characters.") | |
| return name | |
| def _parse_progress( | |
| context: ToolContext, payload: dict[str, object], epochs: int, output: Path, | |
| preview_enabled: bool, preview_every: int, preview_prompt: str, | |
| preview_seed: int, preview_steps: int, | |
| ) -> None: | |
| event = str(payload.get("event", "progress")) | |
| if event == "error": | |
| raise ToolExecutionError(str(payload.get("message", "DDPM trainer reported an error."))) | |
| if event == "step": | |
| current = int(payload.get("global_step", 0) or 0) | |
| total = int(payload.get("total_steps", 0) or 0) | |
| if total: | |
| context.progress( | |
| max(1, min(99, round(current * 100 / total))), | |
| f"Training step {current:,} of {total:,}", | |
| current_step=current, | |
| total_steps=total, | |
| epoch=int(payload.get("epoch", 0) or 0), | |
| total_epochs=epochs, | |
| unit="step", | |
| ) | |
| return | |
| if event == "epoch_end": | |
| epoch = int(payload.get("epoch", 0) or 0) | |
| context.progress( | |
| max(1, min(99, round(epoch * 100 / max(epochs, 1)))), | |
| f"Finished epoch {epoch} of {epochs}", | |
| epoch=epoch, | |
| total_epochs=epochs, | |
| unit="epoch", | |
| ) | |
| if preview_enabled and epoch and epoch % preview_every == 0: | |
| candidate = Path(str(payload.get("preview_path", ""))) if payload.get("preview_path") else _latest_preview(output) | |
| if candidate: | |
| context.preview(candidate, epoch=epoch, next_epoch=min(epochs, epoch + preview_every), | |
| prompt=preview_prompt, seed=preview_seed, steps=preview_steps) | |
| return | |
| if event == "done": | |
| context.progress(100, "DDPM training completed") | |
| def _train_ddpm_stage( | |
| context: ToolContext, | |
| dataset_dir: str, | |
| model_name: str, | |
| epochs: int, | |
| output_dir: str, | |
| resume_from: str = "", | |
| resolution: int = 128, batch_size: int = 1, learning_rate: float = 0.0001, | |
| gradient_accumulation_steps: int = 1, dataloader_num_workers: int = 4, | |
| mixed_precision: str = "fp16", save_every: int = 10, preview_steps: int = 50, | |
| training_intensity: int = 100, preview_enabled: bool = True, | |
| preview_every: int = 5, preview_prompt: str = "", preview_seed: int = 123456789, | |
| completed_epochs: int = 0, | |
| training_aspect_ratio: str = "Dataset (Auto)", resize_mode: str = "fit", | |
| ) -> dict[str, object]: | |
| """Run the registered DDPM project without shell interpolation or overwrites.""" | |
| trainer_root = Path(str(ConfigManager(context.root).get("tool_folders", {}).get("ddpm_trainer", ""))).expanduser() | |
| script = trainer_root / "train.py" | |
| dataset = Path(dataset_dir).expanduser().resolve() | |
| output = Path(output_dir).expanduser().resolve() | |
| model_name = _safe_model_name(model_name) | |
| if not script.is_file(): | |
| raise ToolExecutionError("DDPM train.py was not found. Re-scan the DDPM folder in Settings.") | |
| if not dataset.is_dir(): | |
| raise ToolExecutionError("The selected DDPM dataset folder no longer exists.") | |
| training_dataset, image_count = _training_image_folder(dataset) | |
| if image_count < 2: | |
| raise ToolExecutionError("The DDPM dataset needs at least two image files before training can start.") | |
| requested_epochs = int(epochs) | |
| if not 64 <= int(resolution) <= 512 or int(resolution) % 8 or not 1 <= int(batch_size) <= 64: | |
| raise ToolExecutionError("DDPM resolution must be a multiple of 8 (64–512) and batch size 1–64.") | |
| canvas_width, canvas_height = _training_canvas(training_dataset, int(resolution), str(training_aspect_ratio)) | |
| if resize_mode not in {"fit", "fill", "stretch"}: | |
| raise ToolExecutionError("DDPM resize mode must be fit, fill, or stretch.") | |
| if not 1e-7 <= float(learning_rate) <= 0.1 or not 1 <= int(gradient_accumulation_steps) <= 64: | |
| raise ToolExecutionError("DDPM learning rate or gradient accumulation is outside ADAM's safe range.") | |
| if not 0 <= int(dataloader_num_workers) <= 16 or not 1 <= int(save_every) <= 1000 or not 1 <= int(preview_steps) <= 500 or not 1 <= int(preview_every) <= 100_000 or not 10 <= int(training_intensity) <= 100 or mixed_precision not in {"fp16", "no"}: | |
| raise ToolExecutionError("DDPM training options are outside ADAM's safe range.") | |
| if requested_epochs == 0: | |
| epochs = 200 if image_count <= 100 else 100 | |
| context.log( | |
| f"Adaptive epoch policy selected {epochs} epochs for {image_count} collected images." | |
| ) | |
| elif 1 <= requested_epochs <= 100_000: | |
| epochs = requested_epochs | |
| else: | |
| raise ToolExecutionError("Epoch count must be between 1 and 100000, or 0 for ADAM's adaptive policy.") | |
| missing_packages = [ | |
| package | |
| for package in ("datasets", "diffusers", "transformers", "accelerate", "torch", "torchvision") | |
| if importlib.util.find_spec(package) is None | |
| ] | |
| if missing_packages: | |
| raise ToolExecutionError( | |
| "ADAM's Python environment is missing DDPM packages: " | |
| + ", ".join(missing_packages) | |
| + ". Close ADAM and open Launch ADAM.bat; it will install the needed DDPM requirements. " | |
| f"Current Python: {sys.executable}" | |
| ) | |
| output_root = (trainer_root / "output").resolve() | |
| try: | |
| output.relative_to(output_root) | |
| except ValueError as exc: | |
| raise ToolExecutionError("DDPM outputs must stay inside the registered DDPM output folder.") from exc | |
| resume = Path(resume_from).expanduser().resolve() if resume_from else None | |
| pretrained_model: Path | None = None | |
| if resume: | |
| # Accelerate checkpoints store the training UNet below ``unet/`` plus | |
| # optimizer and scheduler state; they are not standalone pipelines. | |
| accelerate_checkpoint = ( | |
| (resume / "unet" / "diffusion_pytorch_model.safetensors").is_file() | |
| and (resume / "optimizer.bin").is_file() | |
| and (resume / "scheduler.bin").is_file() | |
| ) | |
| standalone_checkpoint = (resume / "pytorch_model.bin").is_file() or (resume / "model.safetensors").is_file() | |
| source_model = resume.parent if resume.name.startswith("checkpoint-") else resume | |
| if not accelerate_checkpoint and not standalone_checkpoint: | |
| if not (source_model / "model_index.json").is_file(): | |
| raise ToolExecutionError("The saved DDPM model is incomplete and cannot be fine-tuned safely.") | |
| pretrained_model = source_model | |
| resume = None | |
| context.log( | |
| "The exact resume checkpoint is incomplete. Creating a new fine-tuned model " | |
| "from the saved DDPM pipeline instead." | |
| ) | |
| else: | |
| if not resume.is_dir() or not resume.name.startswith("checkpoint-"): | |
| raise ToolExecutionError("A valid DDPM checkpoint-* folder is required to resume.") | |
| checkpoint_size = _saved_unet_size(resume) | |
| target_size = (canvas_height, canvas_width) | |
| if checkpoint_size and checkpoint_size != target_size and (source_model / "model_index.json").is_file(): | |
| pretrained_model = source_model | |
| resume = None | |
| context.log( | |
| f"Changing DDPM canvas from {checkpoint_size[1]}x{checkpoint_size[0]} to " | |
| f"{canvas_width}x{canvas_height}. " | |
| "Starting a fresh fine-tune from the saved model weights instead of resuming the old optimizer schedule." | |
| ) | |
| if resume is not None: | |
| steps_per_epoch = max(1, math.ceil(image_count / int(batch_size))) | |
| checkpoint_step = int(resume.name.rsplit("-", 1)[-1]) | |
| prior_epochs = int(completed_epochs) if int(completed_epochs) > 0 else checkpoint_step // steps_per_epoch | |
| epochs = prior_epochs + int(epochs) | |
| context.log( | |
| f"Continuing after approximately {prior_epochs} completed epochs " | |
| f"for {int(epochs) - prior_epochs} additional epochs." | |
| ) | |
| if output.exists(): | |
| raise ToolExecutionError("The chosen DDPM output folder already exists; ADAM will not overwrite it.") | |
| output.mkdir(parents=True, exist_ok=False) | |
| if resume: | |
| # The connected trainer resolves --resume_from_checkpoint inside its | |
| # output directory. Copy only the checkpoint into this new branch so | |
| # it can restore optimizer state without touching the source model. | |
| resume_copy = output / resume.name | |
| shutil.copytree(resume, resume_copy) | |
| resume = resume_copy | |
| # The connected trainer asks Accelerate/TensorBoard to write directly to | |
| # output/logs/train. Create it up front because its writer does not always | |
| # create the nested directory on Windows. | |
| (output / "logs" / "train").mkdir(parents=True, exist_ok=True) | |
| stop_file = output / ".adam_stop_training.flag" | |
| command = [ | |
| sys.executable, str(script), "--train_data_dir", str(training_dataset), "--output_dir", str(output), | |
| "--model_name", model_name, "--resolution", str(int(resolution)), "--train_batch_size", str(int(batch_size)), | |
| "--resolution_width", str(canvas_width), "--resolution_height", str(canvas_height), "--resize_mode", resize_mode, | |
| "--num_epochs", str(int(epochs)), "--learning_rate", str(float(learning_rate)), "--mixed_precision", mixed_precision, | |
| "--ddpm_beta_schedule", "linear", "--tf32", "true", "--save_images_epochs", str(int(preview_every) if preview_enabled else int(epochs) + 1), | |
| "--save_model_epochs", str(int(save_every)), "--training_intensity", str(int(training_intensity)), "--dataloader_num_workers", str(int(dataloader_num_workers)), | |
| "--gradient_accumulation_steps", str(int(gradient_accumulation_steps)), "--preview_num_inference_steps", str(int(preview_steps)), | |
| "--preview_sampler", "DDIM", "--pin_memory", "true", "--stop_signal_file", str(stop_file), | |
| "--checkpointing_steps", str(max(1, math.ceil(image_count / int(batch_size)))), | |
| "--checkpoints_total_limit", "1", "--keep_latest_resume_checkpoint", "--gui_progress", | |
| ] | |
| if resume: | |
| command.extend(["--resume_from_checkpoint", resume.name]) | |
| if int(completed_epochs) > 0: | |
| command.extend(["--resume_completed_epochs", str(int(completed_epochs))]) | |
| if pretrained_model: | |
| command.extend(["--pretrained_model_path", str(pretrained_model)]) | |
| context.log( | |
| f"Starting real DDPM training with {image_count} images from {training_dataset} on a {canvas_width}x{canvas_height} " | |
| f"{resize_mode} canvas, batch {batch_size}, lr {learning_rate}." | |
| ) | |
| context.log(f"Output folder: {output}") | |
| process = subprocess.Popen(command, cwd=str(trainer_root), stdout=subprocess.PIPE, stderr=subprocess.STDOUT, | |
| text=True, encoding="utf-8", errors="replace", shell=False) | |
| lines: queue.Queue[str | None] = queue.Queue() | |
| def read_output() -> None: | |
| assert process.stdout is not None | |
| for line in process.stdout: | |
| lines.put(line.rstrip()) | |
| lines.put(None) | |
| threading.Thread(target=read_output, daemon=True).start() | |
| context.progress(1, "Starting DDPM trainer") | |
| stop_requested_at: float | None = None | |
| cancelled = False | |
| force_stop_sent = False | |
| suspended = False | |
| adjustment_requested = False | |
| stopped_details: dict[str, object] = {} | |
| while True: | |
| should_pause = not context.run_event.is_set() | |
| if should_pause != suspended: | |
| if set_process_tree_paused(process, should_pause): | |
| suspended = should_pause | |
| context.log("DDPM trainer paused safely." if suspended else "DDPM trainer resumed.") | |
| if context.adjustment_event and context.adjustment_event.is_set() and not adjustment_requested: | |
| if suspended: | |
| set_process_tree_paused(process, False) | |
| suspended = False | |
| stop_file.write_text("after_epoch", encoding="utf-8") | |
| adjustment_requested = True | |
| context.log("Settings change accepted; finishing this epoch and saving a resume checkpoint.") | |
| if context.cancel_event.is_set() and stop_requested_at is None: | |
| if suspended: | |
| set_process_tree_paused(process, False) | |
| suspended = False | |
| stop_file.touch(exist_ok=True) | |
| stop_requested_at = time.monotonic() | |
| cancelled = True | |
| context.log("Safe stop requested; waiting for DDPM to finish its current step.") | |
| if stop_requested_at and not force_stop_sent and time.monotonic() - stop_requested_at > FORCE_STOP_TIMEOUT_SECONDS: | |
| terminate_process_tree(process, timeout=3) | |
| force_stop_sent = True | |
| context.log("DDPM did not stop in time; terminating the trainer process.") | |
| try: | |
| line = lines.get(timeout=0.12) | |
| if line: | |
| if line.startswith("PROGRESS_JSON:"): | |
| try: | |
| payload = json.loads(line.split(":", 1)[1]) | |
| if str(payload.get("event", "")) == "stopped": | |
| stopped_details = payload | |
| if not cancelled: | |
| _parse_progress(context, payload, int(epochs), output, | |
| bool(preview_enabled), int(preview_every), preview_prompt, | |
| int(preview_seed), int(preview_steps)) | |
| except json.JSONDecodeError: | |
| context.log(line) | |
| else: | |
| context.log(line) | |
| except queue.Empty: | |
| pass | |
| if process.poll() is not None and lines.empty(): | |
| break | |
| if cancelled: | |
| raise ToolCancelled("DDPM training stopped by user.") | |
| if adjustment_requested: | |
| checkpoints = sorted(output.glob("checkpoint-*"), key=lambda path: int(path.name.rsplit("-", 1)[-1])) | |
| if not checkpoints: | |
| raise ToolExecutionError("Training stopped for adjustment, but no complete resume checkpoint was found.") | |
| raise ToolAdjustmentRequested( | |
| "DDPM training reached a safe epoch boundary.", | |
| { | |
| "checkpoint": str(checkpoints[-1]), | |
| "completed_epochs": int(stopped_details.get("completed_epochs", 0) or 0), | |
| "updates": dict(context.adjustment_request or {}), | |
| }, | |
| ) | |
| if process.returncode != 0: | |
| raise ToolExecutionError(f"DDPM trainer exited with code {process.returncode}. See the job log for details.") | |
| context.progress(100, "DDPM training completed") | |
| checkpoints = sorted( | |
| output.glob("checkpoint-*"), | |
| key=lambda path: int(path.name.rsplit("-", 1)[-1]) | |
| if path.name.rsplit("-", 1)[-1].isdigit() | |
| else -1, | |
| ) | |
| latest_checkpoint = str(checkpoints[-1]) if checkpoints else "" | |
| return { | |
| "output_folder": str(output), | |
| "model_name": model_name, | |
| "assets": [ | |
| { | |
| "kind": "model", | |
| "name": model_name, | |
| "path": str(output), | |
| "trainer": "ddpm", | |
| "dataset_path": str(dataset), | |
| "checkpoint": latest_checkpoint, | |
| "epochs": int(epochs), | |
| } | |
| ], | |
| } | |
| def _stage_context(context: ToolContext, *, stage_index: int, stage_count: int) -> ToolContext: | |
| """Map one stage's progress into the single job progress bar.""" | |
| def report(percent: int, message: str, **details: object) -> None: | |
| overall = round(((stage_index + max(0, min(percent, 100)) / 100) / stage_count) * 100) | |
| context.progress( | |
| overall, | |
| f"Stage {stage_index + 1}/{stage_count} · {message}", | |
| **details, | |
| ) | |
| return ToolContext( | |
| root=context.root, job_id=context.job_id, tool=context.tool, | |
| cancel_event=context.cancel_event, run_event=context.run_event, | |
| progress_callback=report, log_callback=context.log_callback, | |
| preview_callback=context.preview_callback, | |
| # A settings-change request currently restarts a single DDPM process. | |
| # Keep a curriculum stage atomic until that recovery path understands | |
| # its stage manifest. | |
| adjustment_event=None, adjustment_request=None, step_delay=context.step_delay, | |
| ) | |
| def _stage_output(root: Path, stage_number: int, resolution: int) -> Path: | |
| base = root / f"stage-{stage_number:02d}-{resolution}px" | |
| if not base.exists(): | |
| return base | |
| attempt = 2 | |
| while (candidate := root / f"{base.name}-retry-{attempt}").exists(): | |
| attempt += 1 | |
| return candidate | |
| def train_ddpm( | |
| context: ToolContext, | |
| dataset_dir: str, | |
| model_name: str, | |
| epochs: int, | |
| output_dir: str, | |
| resume_from: str = "", | |
| resolution: int = 128, batch_size: int = 1, learning_rate: float = 0.0001, | |
| gradient_accumulation_steps: int = 1, dataloader_num_workers: int = 4, | |
| mixed_precision: str = "fp16", save_every: int = 10, preview_steps: int = 50, | |
| training_intensity: int = 100, preview_enabled: bool = True, | |
| preview_every: int = 5, preview_prompt: str = "", preview_seed: int = 123456789, | |
| completed_epochs: int = 0, | |
| training_aspect_ratio: str = "Dataset (Auto)", resize_mode: str = "fit", | |
| progressive_stages: list[dict[str, object]] | None = None, | |
| progressive_auto_batch: bool = True, | |
| ) -> dict[str, object]: | |
| """Train once, or run a low-to-high resolution DDPM curriculum. | |
| Every completed stage is retained in a hidden sibling folder. The public | |
| output folder is created only after the final pipeline has completed, so a | |
| partial curriculum can never replace a usable completed model. | |
| """ | |
| if not progressive_stages: | |
| return _train_ddpm_stage( | |
| context, dataset_dir, model_name, epochs, output_dir, resume_from, | |
| resolution, batch_size, learning_rate, gradient_accumulation_steps, | |
| dataloader_num_workers, mixed_precision, save_every, preview_steps, | |
| training_intensity, preview_enabled, preview_every, preview_prompt, | |
| preview_seed, completed_epochs, training_aspect_ratio, resize_mode, | |
| ) | |
| stages = parse_stages(progressive_stages, trainer="ddpm", total_epochs=epochs) | |
| public_output = Path(output_dir).expanduser().resolve() | |
| if public_output.exists(): | |
| raise ToolExecutionError("The chosen DDPM output folder already exists; ADAM will not overwrite it.") | |
| stage_root = public_output.parent / f".{public_output.name}.progressive" | |
| state_path = stage_root / "progressive_state.json" | |
| stage_root.mkdir(parents=True, exist_ok=True) | |
| try: | |
| state = json.loads(state_path.read_text(encoding="utf-8")) | |
| except (OSError, json.JSONDecodeError): | |
| state = {"model_name": model_name, "stages": [], "completed": []} | |
| completed = state.get("completed", []) if isinstance(state.get("completed"), list) else [] | |
| completed_by_index = { | |
| int(item.get("index")): Path(str(item.get("output"))) | |
| for item in completed if isinstance(item, dict) and str(item.get("index", "")).isdigit() | |
| } | |
| final_resolution = stages[-1].resolution | |
| prior_model = resume_from | |
| final_result: dict[str, object] | None = None | |
| context.log( | |
| "Progressive DDPM schedule: " + stage_summary(stages) + ". " | |
| + ("Auto batch caps are enabled." if progressive_auto_batch else "Using the same batch settings at every stage.") | |
| ) | |
| for index, stage in enumerate(stages): | |
| completed_output = completed_by_index.get(index) | |
| if completed_output and completed_output.is_dir(): | |
| prior_model = str(completed_output) | |
| context.log(f"Stage {index + 1}/{len(stages)} already completed; using its saved weights.") | |
| continue | |
| stage_output = _stage_output(stage_root, index + 1, stage.resolution) | |
| stage_batch, stage_accumulation = stage_batch_settings( | |
| trainer="ddpm", stage_resolution=stage.resolution, final_resolution=final_resolution, | |
| final_batch_size=batch_size, base_accumulation=gradient_accumulation_steps, | |
| auto_batch=bool(progressive_auto_batch), | |
| ) | |
| context.log( | |
| f"Stage {index + 1}/{len(stages)}: {stage.resolution}px for {stage.epochs} epochs; " | |
| f"batch {stage_batch}, gradient accumulation {stage_accumulation}." | |
| ) | |
| result = _train_ddpm_stage( | |
| _stage_context(context, stage_index=index, stage_count=len(stages)), | |
| dataset_dir, model_name, stage.epochs, str(stage_output), prior_model, | |
| stage.resolution, stage_batch, learning_rate, stage_accumulation, | |
| dataloader_num_workers, mixed_precision, min(save_every, stage.epochs), preview_steps, | |
| training_intensity, preview_enabled, min(preview_every, stage.epochs), preview_prompt, | |
| preview_seed, 0, training_aspect_ratio, resize_mode, | |
| ) | |
| prior_model = str(stage_output) | |
| final_result = result | |
| completed.append({"index": index, "resolution": stage.resolution, "epochs": stage.epochs, "output": prior_model}) | |
| state.update({"stages": [{"resolution": item.resolution, "epochs": item.epochs} for item in stages], "completed": completed}) | |
| state_path.write_text(json.dumps(state, indent=2), encoding="utf-8") | |
| if not prior_model or not Path(prior_model).is_dir(): | |
| raise ToolExecutionError("Progressive DDPM training did not produce a final stage model.") | |
| shutil.move(prior_model, public_output) | |
| final_checkpoint = sorted( | |
| public_output.glob("checkpoint-*"), | |
| key=lambda path: int(path.name.rsplit("-", 1)[-1]) | |
| if path.name.rsplit("-", 1)[-1].isdigit() else -1, | |
| )[-1:] | |
| checkpoint = str(final_checkpoint[0]) if final_checkpoint else "" | |
| context.progress(100, "Progressive DDPM training completed") | |
| return { | |
| "output_folder": str(public_output), "model_name": model_name, | |
| "progressive_stages": [{"resolution": stage.resolution, "epochs": stage.epochs} for stage in stages], | |
| "assets": [{ | |
| "kind": "model", "name": model_name, "path": str(public_output), "trainer": "ddpm", | |
| "dataset_path": str(Path(dataset_dir).expanduser().resolve()), "checkpoint": checkpoint, | |
| "epochs": sum(stage.epochs for stage in stages), | |
| }], | |
| } | |