comdoleger commited on
Commit
8b89843
·
verified ·
1 Parent(s): 16bcf78

Upload jobs/TrainJob.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. jobs/TrainJob.py +44 -0
jobs/TrainJob.py ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import os
3
+
4
+ from jobs import BaseJob
5
+ from toolkit.kohya_model_util import load_models_from_stable_diffusion_checkpoint
6
+ from collections import OrderedDict
7
+ from typing import List
8
+ from jobs.process import BaseExtractProcess, TrainFineTuneProcess
9
+ from datetime import datetime
10
+
11
+
12
+ process_dict = {
13
+ 'vae': 'TrainVAEProcess',
14
+ 'slider': 'TrainSliderProcess',
15
+ 'slider_old': 'TrainSliderProcessOld',
16
+ 'lora_hack': 'TrainLoRAHack',
17
+ 'rescale_sd': 'TrainSDRescaleProcess',
18
+ 'esrgan': 'TrainESRGANProcess',
19
+ 'reference': 'TrainReferenceProcess',
20
+ }
21
+
22
+
23
+ class TrainJob(BaseJob):
24
+
25
+ def __init__(self, config: OrderedDict):
26
+ super().__init__(config)
27
+ self.training_folder = self.get_conf('training_folder', required=True)
28
+ self.is_v2 = self.get_conf('is_v2', False)
29
+ self.device = self.get_conf('device', 'cpu')
30
+ # self.gradient_accumulation_steps = self.get_conf('gradient_accumulation_steps', 1)
31
+ # self.mixed_precision = self.get_conf('mixed_precision', False) # fp16
32
+ self.log_dir = self.get_conf('log_dir', None)
33
+
34
+ # loads the processes from the config
35
+ self.load_processes(process_dict)
36
+
37
+
38
+ def run(self):
39
+ super().run()
40
+ print("")
41
+ print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
42
+
43
+ for process in self.process:
44
+ process.run()