CoolFace
Apppublic

avinjcy/custom-diffusion

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
trainer.py165 linesDownload Raw Back to root
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