realfill-library/RealFill-Training-UI
5
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 