CoolFace
Apppublic

realfill-library/RealFill-Training-UI

sourceHugging Facemitupdated 2y agoView on Hugging Face
5likes
trainer.py173 linesDownload Raw Back to root
1from __future__ import annotations2 3import datetime4import os5import pathlib6import shlex7import shutil8import subprocess9 10import gradio as gr11import slugify12import torch13from PIL import Image14from huggingface_hub import HfApi15 16from app_upload import ModelUploader17from utils import save_model_card18 19URL_TO_JOIN_LIBRARY_ORG = 'https://huggingface.co/organizations/realfill-library/share/WctmaLvDHWxnuWoJxagTrzVXbGwxoqoJoG'20 21class Trainer:22    def __init__(self, hf_token: str | None = None):23        self.hf_token = hf_token24        self.api = HfApi(token=hf_token)25        self.model_uploader = ModelUploader(hf_token)26 27    def prepare_dataset(self, reference_images: list,28                        target_image: Image.Image, target_mask: Image.Image,29                        train_data_dir: pathlib.Path, output_dir: pathlib.Path) -> None:30        shutil.rmtree(train_data_dir, ignore_errors=True)31        train_data_dir.mkdir(parents=True)32        33        (train_data_dir / 'ref').mkdir(parents=True)34        (train_data_dir / 'target').mkdir(parents=True)35 36        for i, temp_path in enumerate(reference_images):37            image = Image.open(temp_path.name)38            image = image.convert('RGB')39            out_path = train_data_dir / 'ref' / f'{i:03d}.jpg'40            image.save(out_path, format='JPEG', quality=100)41 42        target_image = Image.open(target_image[0].name)43        target_image = target_image.convert('RGB')44        out_path = train_data_dir / 'target' / f'target.jpg'45        target_image.save(out_path, format='JPEG', quality=100)46        out_path = output_dir / f'target.jpg'47        target_image.save(out_path, format='JPEG', quality=100)48 49        target_mask = Image.open(target_mask[0].name)50        target_mask = target_mask.convert('L')51        out_path = train_data_dir / 'target' / f'mask.jpg'52        target_mask.save(out_path, format='JPEG', quality=100)53        out_path = output_dir / f'mask.jpg'54        target_mask.save(out_path, format='JPEG', quality=100)55 56    def join_library_org(self) -> None:57        subprocess.run(58            shlex.split(59                f'curl -X POST -H "Authorization: Bearer {self.hf_token}" -H "Content-Type: application/json" {URL_TO_JOIN_LIBRARY_ORG}'60            ))61 62    def run(63        self,64        reference_images: list | None,65        target_image: Image.Image | None,66        target_mask: Image.Image | None,67        output_model_name: str,68        overwrite_existing_model: bool,69        base_model: str,70        resolution_s: str,71        n_steps: int,72        unet_learning_rate: float,73        text_encoder_learning_rate: float,74        lora_rank: int,75        lora_dropout: float,76        lora_alpha: int,77        gradient_accumulation: int,78        seed: int,79        fp16: bool,80        use_8bit_adam: bool,81        checkpointing_steps: int,82        use_wandb: bool,83        validation_steps: int,84        upload_to_hub: bool,85        use_private_repo: bool,86        delete_existing_repo: bool,87        upload_to: str,88        remove_gpu_after_training: bool,89    ) -> str:90        if not torch.cuda.is_available():91            raise gr.Error('CUDA is not available.')92        if reference_images is None:93            raise gr.Error('You need to upload reference images.')94        if target_image is None:95            raise gr.Error('The instance prompt is missing.')96 97        resolution = int(resolution_s)98 99        if not output_model_name:100            timestamp = datetime.datetime.now().strftime('%Y-%m-%d-%H-%M-%S')101            output_model_name = f'realfill-{timestamp}'102        output_model_name = slugify.slugify(output_model_name)103 104        repo_dir = pathlib.Path(__file__).parent105        output_dir = repo_dir / 'experiments' / output_model_name106        if overwrite_existing_model or upload_to_hub:107            shutil.rmtree(output_dir, ignore_errors=True)108        output_dir.mkdir(parents=True)109 110        train_data_dir = repo_dir / 'training_data' / output_model_name111        self.prepare_dataset(reference_images, target_image, target_mask, train_data_dir, output_dir)112 113        if upload_to_hub:114            self.join_library_org()115 116        command = f'''117        python train_realfill.py \118          --pretrained_model_name_or_path={base_model}  \119          --train_data_dir={train_data_dir} \120          --output_dir={output_dir} \121          --resolution={resolution} \122          --train_batch_size=16 \123          --gradient_accumulation_steps={gradient_accumulation} --gradient_checkpointing \124          --unet_learning_rate={unet_learning_rate} \125          --text_encoder_learning_rate={text_encoder_learning_rate} \126          --lr_scheduler=constant \127          --lr_warmup_steps=100 \128          --set_grads_to_none \129          --max_train_steps={n_steps} \130          --checkpointing_steps={checkpointing_steps} \131          --validation_steps={validation_steps} \132          --lora_rank={lora_rank} \133          --lora_dropout={lora_dropout} \134          --lora_alpha={lora_alpha} \135          --seed={seed}136        '''137        if fp16:138            command += ' --mixed_precision fp16'139        if use_8bit_adam:140            command += ' --use_8bit_adam'141        if use_wandb:142            command += ' --report_to wandb'143 144        with open(output_dir / 'train.sh', 'w') as f:145            command_s = ' '.join(command.split())146            f.write(command_s)147        subprocess.run(shlex.split(command))148        save_model_card(save_dir=output_dir,149                        base_model=base_model,150                        target_image=output_dir / 'target.jpg',151                        target_mask=output_dir / 'mask.jpg')152 153        message = 'Training completed!'154        print(message)155 156        if upload_to_hub:157            upload_message = self.model_uploader.upload_model(158                folder_path=output_dir.as_posix(),159                repo_name=output_model_name,160                upload_to=upload_to,161                private=use_private_repo,162                delete_existing_repo=delete_existing_repo)163            print(upload_message)164            message = message + '\n' + upload_message165 166        if remove_gpu_after_training:167            space_id = os.getenv('SPACE_ID')168            if space_id:169                self.api.request_space_hardware(repo_id=space_id,170                                                hardware='cpu-basic')171 172        return message173