eyaler/custom-diffusion
0
1from __future__ import annotations2 3import os4import pathlib5import shlex6import shutil7import subprocess8 9import gradio as gr10import PIL.Image11import torch12import json13 14os.environ['PYTHONPATH'] = f'custom-diffusion:{os.getenv("PYTHONPATH", "")}'15 16 17def pad_image(image: PIL.Image.Image) -> PIL.Image.Image:18 w, h = image.size19 if w == h:20 return image21 elif w > h:22 new_image = PIL.Image.new(image.mode, (w, w), (0, 0, 0))23 new_image.paste(image, (0, (w - h) // 2))24 return new_image25 else:26 new_image = PIL.Image.new(image.mode, (h, h), (0, 0, 0))27 new_image.paste(image, ((h - w) // 2, 0))28 return new_image29 30 31class Trainer:32 def __init__(self):33 self.is_running = False34 self.is_running_message = 'Another training is in progress.'35 36 self.output_dir = pathlib.Path('results')37 self.instance_data_dir = self.output_dir / 'training_data'38 self.class_data_dir = self.output_dir / 'regularization_data'39 40 def check_if_running(self) -> dict:41 if self.is_running:42 return gr.update(value=self.is_running_message)43 else:44 return gr.update(value='No training is running.')45 46 def cleanup_dirs(self) -> None:47 shutil.rmtree(self.output_dir, ignore_errors=True)48 49 def prepare_dataset(self, concept_images_collection: list, concept_prompt_collection: list, class_prompt_collection: list, resolution: int) -> None:50 self.instance_data_dir.mkdir(parents=True)51 concepts_list = []52 53 for i in range(len(concept_images_collection)):54 concept_dir = self.instance_data_dir / f'{i}'55 class_dir = self.class_data_dir / f'{i}'56 concept_dir.mkdir(parents=True)57 concept_images = concept_images_collection[i]58 59 concepts_list.append(60 {61 "instance_prompt": concept_prompt_collection[i],62 "class_prompt": class_prompt_collection[i],63 "instance_data_dir": f'{concept_dir}',64 "class_data_dir": f'{class_dir}'65 }66 )67 68 for i, temp_path in enumerate(concept_images):69 image = PIL.Image.open(temp_path.name)70 image = pad_image(image)71 # image = image.resize((resolution, resolution))72 image = image.convert('RGB')73 out_path = concept_dir / f'{i:03d}.jpg'74 image.save(out_path, format='JPEG', quality=100)75 76 print(concepts_list)77 json.dump(concepts_list, open( f'{self.output_dir}/temp.json' , 'w') )78 79 80 def run(81 self,82 base_model: str,83 resolution_s: str,84 n_steps: int,85 learning_rate: float,86 train_text_encoder: bool,87 modifier_token: bool,88 gradient_accumulation: int,89 batch_size: int,90 use_8bit_adam: bool,91 gradient_checkpointing: bool,92 gen_images: bool,93 num_reg_images: int,94 *inputs, 95 ) -> tuple[dict, list[pathlib.Path]]:96 if not torch.cuda.is_available():97 raise gr.Error('CUDA is not available.')98 99 num_concept = 0100 for i in range(len(inputs) // 3):101 if inputs[i] != None:102 num_concept +=1103 104 print(num_concept, inputs)105 concept_images_collection = inputs[: num_concept]106 concept_prompt_collection = inputs[3: 3 + num_concept]107 class_prompt_collection = inputs[6: 6+num_concept]108 if self.is_running:109 return gr.update(value=self.is_running_message), []110 111 if concept_images_collection is None:112 raise gr.Error('You need to upload images.')113 if not concept_prompt_collection:114 raise gr.Error('The concept prompt is missing.')115 116 resolution = int(resolution_s)117 118 self.cleanup_dirs()119 self.prepare_dataset(concept_images_collection, concept_prompt_collection, class_prompt_collection, resolution)120 torch.cuda.empty_cache()121 command = f'''122 accelerate launch custom-diffusion/src/diffuser_training.py \123 --pretrained_model_name_or_path={base_model} \124 --output_dir={self.output_dir} \125 --concepts_list={f'{self.output_dir}/temp.json'} \126 --with_prior_preservation --prior_loss_weight=1.0 \127 --resolution={resolution} \128 --train_batch_size={batch_size} \129 --gradient_accumulation_steps={gradient_accumulation} \130 --learning_rate={learning_rate} \131 --lr_scheduler="constant" \132 --lr_warmup_steps=0 \133 --max_train_steps={n_steps} \134 --num_class_images={num_reg_images} \135 --initializer_token="ktn+pll+ucd" \136 --scale_lr --hflip 137 '''138 if modifier_token:139 tokens = '+'.join([f'<new{i+1}>' for i in range(num_concept)])140 command += f' --modifier_token {tokens}'141 142 if not gen_images:143 command += ' --real_prior'144 if use_8bit_adam:145 command += ' --use_8bit_adam'146 if train_text_encoder:147 command += f' --train_text_encoder'148 if gradient_checkpointing:149 command += f' --gradient_checkpointing'150 151 with open(self.output_dir / 'train.sh', 'w') as f:152 command_s = ' '.join(command.split())153 f.write(command_s)154 155 self.is_running = True156 res = subprocess.run(shlex.split(command))157 self.is_running = False158 159 if res.returncode == 0:160 result_message = 'Training Completed!'161 else:162 result_message = 'Training Failed!'163 weight_paths = sorted(self.output_dir.glob('*.bin'))164 return gr.update(value=result_message), weight_paths165 