Download jobs/TrainJob.py from comdoleger/ai-toolkit: direct link, hf CLI and curl.
- Browser
- Download file 1.4 kB
-
https://huggingface.co/comdoleger/ai-toolkit/resolve/983db42fb1f437c59fec079fbde7cd2733ac6b69/jobs/TrainJob.py
- Command line
-
hf download hf://comdoleger/ai-toolkit@983db42fb1f437c59fec079fbde7cd2733ac6b69/jobs/TrainJob.py
-
curl -L -o TrainJob.py https://huggingface.co/comdoleger/ai-toolkit/resolve/983db42fb1f437c59fec079fbde7cd2733ac6b69/jobs/TrainJob.py
1.4 kB
| import json | |
| import os | |
| from jobs import BaseJob | |
| from toolkit.kohya_model_util import load_models_from_stable_diffusion_checkpoint | |
| from collections import OrderedDict | |
| from typing import List | |
| from jobs.process import BaseExtractProcess, TrainFineTuneProcess | |
| from datetime import datetime | |
| process_dict = { | |
| 'vae': 'TrainVAEProcess', | |
| 'slider': 'TrainSliderProcess', | |
| 'slider_old': 'TrainSliderProcessOld', | |
| 'lora_hack': 'TrainLoRAHack', | |
| 'rescale_sd': 'TrainSDRescaleProcess', | |
| 'esrgan': 'TrainESRGANProcess', | |
| 'reference': 'TrainReferenceProcess', | |
| } | |
| class TrainJob(BaseJob): | |
| def __init__(self, config: OrderedDict): | |
| super().__init__(config) | |
| self.training_folder = self.get_conf('training_folder', required=True) | |
| self.is_v2 = self.get_conf('is_v2', False) | |
| self.device = self.get_conf('device', 'cpu') | |
| # self.gradient_accumulation_steps = self.get_conf('gradient_accumulation_steps', 1) | |
| # self.mixed_precision = self.get_conf('mixed_precision', False) # fp16 | |
| self.log_dir = self.get_conf('log_dir', None) | |
| # loads the processes from the config | |
| self.load_processes(process_dict) | |
| def run(self): | |
| super().run() | |
| print("") | |
| print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}") | |
| for process in self.process: | |
| process.run() | |