Image-to-Image
Diffusers
Safetensors
Core ML
StableDiffusionInpaintPipeline
image-editing
local-ai
clover-image
inpainting
stable-diffusion
Instructions to use neonforestmist/Clover-Image-Tiny-Inpaint with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use neonforestmist/Clover-Image-Tiny-Inpaint with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline from diffusers.utils import load_image # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("neonforestmist/Clover-Image-Tiny-Inpaint", dtype=torch.bfloat16, device_map="cuda") prompt = "Turn this cat into a dog" input_image = load_image("https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png") image = pipe(image=input_image, prompt=prompt).images[0] - Notebooks
- Google Colab
- Kaggle
| #!/usr/bin/env python3 | |
| """Train Clover Image Tiny v2 as a context-aware 9-channel inpainting model. | |
| The trainer warm-starts from the existing Clover inpainting checkpoint and | |
| distills the official SD 1.5 inpainting U-Net while retaining a ground-truth | |
| diffusion objective. Masks are synthesized on the fly with free-form, | |
| multi-region, object-like, and outpainting geometry. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import random | |
| import shutil | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from accelerate import Accelerator | |
| from datasets import load_dataset | |
| from diffusers import ( | |
| AutoencoderKL, | |
| DDPMScheduler, | |
| StableDiffusionInpaintPipeline, | |
| UNet2DConditionModel, | |
| get_scheduler, | |
| ) | |
| from PIL import Image, ImageOps | |
| from transformers import CLIPTextModel, CLIPTokenizer | |
| from inpainting.masks import apply_mask, mask_area_fraction, random_mask | |
| from inpainting.model import make_inpainting_unet | |
| from inpainting.objective import min_snr_weights, spatial_loss_weights, weighted_mse | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--pretrained_model_name_or_path", required=True) | |
| parser.add_argument("--revision") | |
| parser.add_argument("--initial_inpaint_model") | |
| parser.add_argument("--initial_inpaint_revision") | |
| parser.add_argument("--teacher_model_name_or_path") | |
| parser.add_argument("--teacher_revision") | |
| parser.add_argument("--dataset_name", required=True) | |
| parser.add_argument("--dataset_config_name") | |
| parser.add_argument("--dataset_revision") | |
| parser.add_argument("--dataset_split", default="train") | |
| parser.add_argument("--image_column", default="image") | |
| parser.add_argument("--caption_column", default="caption") | |
| parser.add_argument("--resolution", type=int, default=512) | |
| parser.add_argument("--train_batch_size", type=int, default=1) | |
| parser.add_argument("--max_train_steps", type=int, default=12000) | |
| parser.add_argument("--learning_rate", type=float, default=5e-6) | |
| parser.add_argument("--lr_scheduler", default="cosine") | |
| parser.add_argument("--lr_warmup_steps", type=int, default=500) | |
| parser.add_argument("--gradient_accumulation_steps", type=int, default=4) | |
| parser.add_argument("--gradient_checkpointing", action="store_true") | |
| parser.add_argument("--mixed_precision", choices=("no", "fp16", "bf16"), default="bf16") | |
| parser.add_argument("--seed", type=int, default=20260811) | |
| parser.add_argument("--mask_min_area", type=float, default=0.04) | |
| parser.add_argument("--mask_max_area", type=float, default=0.65) | |
| parser.add_argument("--caption_dropout_probability", type=float, default=0.10) | |
| parser.add_argument("--teacher_loss_weight", type=float, default=0.75) | |
| parser.add_argument("--ground_truth_loss_weight", type=float, default=0.25) | |
| parser.add_argument("--context_loss_weight", type=float, default=0.25) | |
| parser.add_argument("--masked_loss_weight", type=float, default=2.5) | |
| parser.add_argument("--boundary_loss_weight", type=float, default=2.0) | |
| parser.add_argument("--boundary_radius", type=int, default=2) | |
| parser.add_argument("--snr_gamma", type=float, default=5.0) | |
| parser.add_argument("--random_flip", action="store_true") | |
| parser.add_argument("--max_train_samples", type=int) | |
| parser.add_argument("--validation_samples", type=int, default=128) | |
| parser.add_argument("--num_workers", type=int, default=4) | |
| parser.add_argument("--checkpointing_steps", type=int, default=500) | |
| parser.add_argument("--checkpoints_total_limit", type=int, default=3) | |
| parser.add_argument("--resume_from_checkpoint") | |
| parser.add_argument("--output_dir", type=Path, required=True) | |
| parser.add_argument("--push_to_hub", action="store_true") | |
| parser.add_argument("--hub_model_id") | |
| return parser.parse_args() | |
| def model_kwargs(revision: str | None) -> dict[str, Any]: | |
| return {"revision": revision} if revision else {} | |
| def _frozen_weight_dtype(mixed_precision: str) -> torch.dtype: | |
| if mixed_precision == "fp16": | |
| return torch.float16 | |
| if mixed_precision == "bf16": | |
| return torch.bfloat16 | |
| return torch.float32 | |
| def _normalize_caption(value: Any) -> str: | |
| if isinstance(value, list): | |
| value = value[0] if value else "" | |
| return " ".join(str(value or "").split()) | |
| def _crop_image( | |
| image: Image.Image, | |
| *, | |
| resolution: int, | |
| rng: random.Random, | |
| random_flip: bool, | |
| ) -> Image.Image: | |
| image = ImageOps.exif_transpose(image).convert("RGB") | |
| width, height = image.size | |
| scale = rng.uniform(0.78, 1.0) | |
| crop_side = max(1, round(min(width, height) * scale)) | |
| left = rng.randint(0, max(0, width - crop_side)) | |
| top = rng.randint(0, max(0, height - crop_side)) | |
| image = image.crop((left, top, left + crop_side, top + crop_side)) | |
| image = image.resize((resolution, resolution), Image.Resampling.LANCZOS) | |
| if random_flip and rng.random() < 0.5: | |
| image = ImageOps.mirror(image) | |
| return image | |
| def _image_to_tensor(image: Image.Image) -> torch.Tensor: | |
| array = np.asarray(image.convert("RGB"), dtype=np.float32) / 127.5 - 1.0 | |
| return torch.from_numpy(array).permute(2, 0, 1).contiguous() | |
| def make_collate_fn( | |
| *, | |
| tokenizer: CLIPTokenizer, | |
| resolution: int, | |
| min_area: float, | |
| max_area: float, | |
| seed: int, | |
| caption_dropout_probability: float, | |
| random_flip: bool, | |
| ): | |
| worker_rng: random.Random | None = None | |
| def collate(examples: list[dict[str, Any]]) -> dict[str, Any]: | |
| nonlocal worker_rng | |
| if worker_rng is None: | |
| # torch.initial_seed differs across DataLoader workers and remains | |
| # stable for the lifetime of each worker process. | |
| worker_rng = random.Random(seed ^ torch.initial_seed()) | |
| images: list[torch.Tensor] = [] | |
| masked_images: list[torch.Tensor] = [] | |
| masks: list[torch.Tensor] = [] | |
| captions: list[str] = [] | |
| mask_areas: list[float] = [] | |
| for example in examples: | |
| image = example["image"] | |
| if not isinstance(image, Image.Image): | |
| image = Image.fromarray(np.asarray(image)) | |
| image = _crop_image( | |
| image, | |
| resolution=resolution, | |
| rng=worker_rng, | |
| random_flip=random_flip, | |
| ) | |
| mask = random_mask( | |
| (resolution, resolution), | |
| worker_rng, | |
| min_area=min_area, | |
| max_area=max_area, | |
| ) | |
| masked = apply_mask(image, mask) | |
| images.append(_image_to_tensor(image)) | |
| masked_images.append(_image_to_tensor(masked)) | |
| masks.append( | |
| torch.from_numpy(np.asarray(mask, dtype=np.float32) / 255.0).unsqueeze(0) | |
| ) | |
| caption = _normalize_caption(example["caption"]) | |
| if worker_rng.random() < caption_dropout_probability: | |
| caption = "" | |
| captions.append(caption) | |
| mask_areas.append(mask_area_fraction(mask)) | |
| tokenized = tokenizer( | |
| captions, | |
| max_length=tokenizer.model_max_length, | |
| padding="max_length", | |
| truncation=True, | |
| return_tensors="pt", | |
| ) | |
| return { | |
| "pixel_values": torch.stack(images), | |
| "masked_pixel_values": torch.stack(masked_images), | |
| "mask": torch.stack(masks), | |
| "input_ids": tokenized.input_ids, | |
| "mask_area": torch.tensor(mask_areas, dtype=torch.float32), | |
| } | |
| return collate | |
| def _load_student(args: argparse.Namespace) -> UNet2DConditionModel: | |
| if args.initial_inpaint_model: | |
| student = UNet2DConditionModel.from_pretrained( | |
| args.initial_inpaint_model, | |
| subfolder="unet", | |
| low_cpu_mem_usage=False, | |
| **model_kwargs(args.initial_inpaint_revision), | |
| ) | |
| if student.config.in_channels != 9: | |
| raise ValueError("initial inpainting U-Net must have nine input channels") | |
| return student | |
| return make_inpainting_unet( | |
| args.pretrained_model_name_or_path, | |
| revision=args.revision, | |
| ) | |
| def _save_pipeline( | |
| *, | |
| args: argparse.Namespace, | |
| unet: UNet2DConditionModel, | |
| output_dir: Path, | |
| ) -> None: | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| pipeline = StableDiffusionInpaintPipeline.from_pretrained( | |
| args.pretrained_model_name_or_path, | |
| revision=args.revision, | |
| unet=unet, | |
| ) | |
| pipeline.save_pretrained(output_dir, safe_serialization=True) | |
| (output_dir / "inpainting-config.json").write_text( | |
| json.dumps( | |
| { | |
| "recipe_version": 2, | |
| "unet_in_channels": 9, | |
| "mask_semantics": "white=regenerate, black=preserve", | |
| "base_model": args.pretrained_model_name_or_path, | |
| "base_revision": args.revision, | |
| "initial_inpaint_model": args.initial_inpaint_model, | |
| "initial_inpaint_revision": args.initial_inpaint_revision, | |
| "teacher_model": args.teacher_model_name_or_path, | |
| "teacher_revision": args.teacher_revision, | |
| "objective": "teacher distillation + ground-truth diffusion", | |
| "mask_distribution": "brush, multi-brush, box, ellipse, polygon, multi-region, outpaint", | |
| }, | |
| indent=2, | |
| ) | |
| + "\n" | |
| ) | |
| def _prune_checkpoints(checkpoint_root: Path, limit: int) -> None: | |
| checkpoints = sorted( | |
| (path for path in checkpoint_root.glob("checkpoint-*") if path.is_dir()), | |
| key=lambda path: int(path.name.rsplit("-", 1)[-1]), | |
| ) | |
| for path in checkpoints[:-limit]: | |
| shutil.rmtree(path) | |
| def _validate_args(args: argparse.Namespace) -> None: | |
| if not 0.0 < args.mask_min_area < args.mask_max_area < 1.0: | |
| raise ValueError("mask area bounds must satisfy 0 < min < max < 1") | |
| if not 0.0 <= args.caption_dropout_probability < 1.0: | |
| raise ValueError("caption dropout probability must be in [0, 1)") | |
| if args.teacher_loss_weight < 0 or args.ground_truth_loss_weight < 0: | |
| raise ValueError("loss weights must be non-negative") | |
| if args.teacher_loss_weight > 0 and not args.teacher_model_name_or_path: | |
| raise ValueError("teacher model is required when teacher loss is enabled") | |
| if args.teacher_loss_weight + args.ground_truth_loss_weight <= 0: | |
| raise ValueError("at least one training objective must be enabled") | |
| if args.push_to_hub and not args.hub_model_id: | |
| raise ValueError("--hub_model_id is required with --push_to_hub") | |
| if args.resolution % 8: | |
| raise ValueError("resolution must be divisible by 8") | |
| def main() -> None: | |
| args = parse_args() | |
| _validate_args(args) | |
| args.output_dir.mkdir(parents=True, exist_ok=True) | |
| checkpoint_root = args.output_dir / "checkpoints" | |
| metrics_path = args.output_dir / "training-metrics.jsonl" | |
| accelerator = Accelerator( | |
| gradient_accumulation_steps=args.gradient_accumulation_steps, | |
| mixed_precision=args.mixed_precision, | |
| ) | |
| torch.manual_seed(args.seed) | |
| random.seed(args.seed) | |
| np.random.seed(args.seed % (2**32)) | |
| kwargs = model_kwargs(args.revision) | |
| weight_dtype = _frozen_weight_dtype(args.mixed_precision) | |
| dataset = load_dataset( | |
| args.dataset_name, | |
| args.dataset_config_name, | |
| split=args.dataset_split, | |
| revision=args.dataset_revision, | |
| ) | |
| dataset = dataset.shuffle(seed=args.seed) | |
| if args.validation_samples < 0 or args.validation_samples >= len(dataset): | |
| raise ValueError("validation_samples must be smaller than the dataset") | |
| if args.validation_samples: | |
| dataset = dataset.select(range(len(dataset) - args.validation_samples)) | |
| if args.max_train_samples: | |
| dataset = dataset.select(range(min(args.max_train_samples, len(dataset)))) | |
| if args.image_column not in dataset.column_names: | |
| raise ValueError(f"Missing image column {args.image_column!r}: {dataset.column_names}") | |
| if args.caption_column not in dataset.column_names: | |
| raise ValueError(f"Missing caption column {args.caption_column!r}: {dataset.column_names}") | |
| dataset = dataset.rename_columns( | |
| {args.image_column: "image", args.caption_column: "caption"} | |
| ) | |
| tokenizer = CLIPTokenizer.from_pretrained( | |
| args.pretrained_model_name_or_path, | |
| subfolder="tokenizer", | |
| **kwargs, | |
| ) | |
| text_encoder = CLIPTextModel.from_pretrained( | |
| args.pretrained_model_name_or_path, | |
| subfolder="text_encoder", | |
| dtype=weight_dtype, | |
| **kwargs, | |
| ) | |
| vae = AutoencoderKL.from_pretrained( | |
| args.pretrained_model_name_or_path, | |
| subfolder="vae", | |
| torch_dtype=weight_dtype, | |
| **kwargs, | |
| ) | |
| student = _load_student(args) | |
| noise_scheduler = DDPMScheduler.from_pretrained( | |
| args.pretrained_model_name_or_path, | |
| subfolder="scheduler", | |
| **kwargs, | |
| ) | |
| teacher = None | |
| if args.teacher_loss_weight > 0: | |
| teacher = UNet2DConditionModel.from_pretrained( | |
| args.teacher_model_name_or_path, | |
| subfolder="unet", | |
| low_cpu_mem_usage=True, | |
| torch_dtype=weight_dtype, | |
| use_safetensors=True, | |
| variant="fp16", | |
| **model_kwargs(args.teacher_revision), | |
| ) | |
| if teacher.config.in_channels != 9: | |
| raise ValueError("teacher inpainting U-Net must have nine input channels") | |
| if args.gradient_checkpointing: | |
| student.enable_gradient_checkpointing() | |
| text_encoder.requires_grad_(False).eval() | |
| vae.requires_grad_(False).eval() | |
| if teacher is not None: | |
| teacher.requires_grad_(False).eval() | |
| collate_fn = make_collate_fn( | |
| tokenizer=tokenizer, | |
| resolution=args.resolution, | |
| min_area=args.mask_min_area, | |
| max_area=args.mask_max_area, | |
| seed=args.seed, | |
| caption_dropout_probability=args.caption_dropout_probability, | |
| random_flip=args.random_flip, | |
| ) | |
| dataloader = torch.utils.data.DataLoader( | |
| dataset, | |
| shuffle=True, | |
| collate_fn=collate_fn, | |
| batch_size=args.train_batch_size, | |
| num_workers=args.num_workers, | |
| pin_memory=True, | |
| persistent_workers=args.num_workers > 0, | |
| ) | |
| optimizer = torch.optim.AdamW( | |
| student.parameters(), | |
| lr=args.learning_rate, | |
| betas=(0.9, 0.999), | |
| weight_decay=1e-2, | |
| eps=1e-8, | |
| ) | |
| lr_scheduler = get_scheduler( | |
| args.lr_scheduler, | |
| optimizer=optimizer, | |
| num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, | |
| num_training_steps=args.max_train_steps * accelerator.num_processes, | |
| ) | |
| student, optimizer, dataloader, lr_scheduler = accelerator.prepare( | |
| student, optimizer, dataloader, lr_scheduler | |
| ) | |
| vae.to(accelerator.device) | |
| text_encoder.to(accelerator.device) | |
| if teacher is not None: | |
| teacher.to(accelerator.device) | |
| global_step = 0 | |
| if args.resume_from_checkpoint: | |
| checkpoint = Path(args.resume_from_checkpoint) | |
| if args.resume_from_checkpoint == "latest": | |
| candidates = sorted( | |
| checkpoint_root.glob("checkpoint-*"), | |
| key=lambda path: int(path.name.rsplit("-", 1)[-1]), | |
| ) | |
| if not candidates: | |
| raise ValueError("No checkpoint is available to resume") | |
| checkpoint = candidates[-1] | |
| accelerator.load_state(checkpoint) | |
| global_step = int(checkpoint.name.rsplit("-", 1)[-1]) | |
| accelerator.print(f"Resumed from {checkpoint} at step {global_step}") | |
| if accelerator.is_main_process: | |
| summary = vars(args).copy() | |
| summary["output_dir"] = str(args.output_dir) | |
| summary["dataset_size"] = len(dataset) | |
| summary["effective_batch_size"] = ( | |
| args.train_batch_size | |
| * args.gradient_accumulation_steps | |
| * accelerator.num_processes | |
| ) | |
| (args.output_dir / "training-summary.json").write_text( | |
| json.dumps(summary, indent=2, default=str) + "\n" | |
| ) | |
| data_iterator = iter(dataloader) | |
| while global_step < args.max_train_steps: | |
| try: | |
| batch = next(data_iterator) | |
| except StopIteration: | |
| data_iterator = iter(dataloader) | |
| batch = next(data_iterator) | |
| with accelerator.accumulate(student): | |
| pixel_values = batch["pixel_values"].to( | |
| accelerator.device, dtype=weight_dtype, non_blocking=True | |
| ) | |
| masked_pixel_values = batch["masked_pixel_values"].to( | |
| accelerator.device, dtype=weight_dtype, non_blocking=True | |
| ) | |
| mask = batch["mask"].to( | |
| accelerator.device, dtype=weight_dtype, non_blocking=True | |
| ) | |
| with torch.no_grad(): | |
| latents = vae.encode(pixel_values).latent_dist.sample() | |
| latents = latents * vae.config.scaling_factor | |
| masked_latents = vae.encode(masked_pixel_values).latent_dist.sample() | |
| masked_latents = masked_latents * vae.config.scaling_factor | |
| encoder_hidden_states = text_encoder( | |
| batch["input_ids"].to(accelerator.device) | |
| )[0] | |
| noise = torch.randn_like(latents) | |
| timesteps = torch.randint( | |
| 0, | |
| noise_scheduler.config.num_train_timesteps, | |
| (latents.shape[0],), | |
| device=latents.device, | |
| ).long() | |
| noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps) | |
| latent_mask = F.interpolate(mask, size=latents.shape[-2:], mode="nearest") | |
| model_input = torch.cat([noisy_latents, latent_mask, masked_latents], dim=1) | |
| model_pred = student( | |
| model_input, | |
| timesteps, | |
| encoder_hidden_states=encoder_hidden_states, | |
| ).sample | |
| prediction_type = noise_scheduler.config.prediction_type | |
| if prediction_type == "epsilon": | |
| ground_truth_target = noise | |
| elif prediction_type == "v_prediction": | |
| ground_truth_target = noise_scheduler.get_velocity(latents, noise, timesteps) | |
| else: | |
| raise ValueError(f"Unsupported prediction type: {prediction_type}") | |
| spatial_weights = spatial_loss_weights( | |
| latent_mask, | |
| context_weight=args.context_loss_weight, | |
| masked_weight=args.masked_loss_weight, | |
| boundary_weight=args.boundary_loss_weight, | |
| boundary_radius=args.boundary_radius, | |
| ) | |
| sample_weights = min_snr_weights( | |
| noise_scheduler.alphas_cumprod, | |
| timesteps, | |
| gamma=args.snr_gamma, | |
| prediction_type=prediction_type, | |
| ) | |
| ground_truth_loss = weighted_mse( | |
| model_pred, | |
| ground_truth_target, | |
| spatial_weights=spatial_weights, | |
| sample_weights=sample_weights, | |
| ) | |
| teacher_loss = torch.zeros((), device=accelerator.device) | |
| if teacher is not None: | |
| with torch.no_grad(): | |
| teacher_target = teacher( | |
| model_input, | |
| timesteps, | |
| encoder_hidden_states=encoder_hidden_states, | |
| ).sample | |
| teacher_loss = weighted_mse( | |
| model_pred, | |
| teacher_target, | |
| spatial_weights=spatial_weights, | |
| sample_weights=sample_weights, | |
| ) | |
| loss = ( | |
| args.ground_truth_loss_weight * ground_truth_loss | |
| + args.teacher_loss_weight * teacher_loss | |
| ) | |
| accelerator.backward(loss) | |
| if accelerator.sync_gradients: | |
| accelerator.clip_grad_norm_(student.parameters(), 1.0) | |
| optimizer.step() | |
| lr_scheduler.step() | |
| optimizer.zero_grad(set_to_none=True) | |
| if accelerator.sync_gradients: | |
| global_step += 1 | |
| metrics = { | |
| "step": global_step, | |
| "loss": loss.detach().float().item(), | |
| "teacher_loss": teacher_loss.detach().float().item(), | |
| "ground_truth_loss": ground_truth_loss.detach().float().item(), | |
| "lr": lr_scheduler.get_last_lr()[0], | |
| "mask_area": batch["mask_area"].float().mean().item(), | |
| } | |
| if accelerator.is_main_process and (global_step == 1 or global_step % 25 == 0): | |
| accelerator.print( | |
| f"step={global_step}/{args.max_train_steps} " | |
| f"loss={metrics['loss']:.5f} " | |
| f"teacher={metrics['teacher_loss']:.5f} " | |
| f"ground_truth={metrics['ground_truth_loss']:.5f} " | |
| f"mask={metrics['mask_area']:.3f} " | |
| f"lr={metrics['lr']:.3e}" | |
| ) | |
| with metrics_path.open("a") as handle: | |
| handle.write(json.dumps(metrics) + "\n") | |
| if global_step % args.checkpointing_steps == 0: | |
| accelerator.wait_for_everyone() | |
| checkpoint = checkpoint_root / f"checkpoint-{global_step}" | |
| accelerator.save_state(checkpoint, safe_serialization=True) | |
| if accelerator.is_main_process: | |
| _prune_checkpoints(checkpoint_root, args.checkpoints_total_limit) | |
| (args.output_dir / "progress.json").write_text( | |
| json.dumps(metrics, indent=2) + "\n" | |
| ) | |
| accelerator.print(f"Saved resumable checkpoint {checkpoint}") | |
| accelerator.wait_for_everyone() | |
| if accelerator.is_main_process: | |
| unwrapped = accelerator.unwrap_model(student).cpu() | |
| _save_pipeline(args=args, unet=unwrapped, output_dir=args.output_dir) | |
| final_metrics = { | |
| "completed_steps": global_step, | |
| "final_loss": loss.detach().float().item(), | |
| "final_teacher_loss": teacher_loss.detach().float().item(), | |
| "final_ground_truth_loss": ground_truth_loss.detach().float().item(), | |
| } | |
| (args.output_dir / "training-complete.json").write_text( | |
| json.dumps(final_metrics, indent=2) + "\n" | |
| ) | |
| if args.push_to_hub: | |
| pipeline = StableDiffusionInpaintPipeline.from_pretrained(args.output_dir) | |
| pipeline.push_to_hub(args.hub_model_id) | |
| accelerator.print(f"Saved Clover Image Tiny Inpaint v2 to {args.output_dir}") | |
| accelerator.end_training() | |
| if __name__ == "__main__": | |
| main() | |