Download training/train-launcher.py from akatz-ai/MiniMax-H3-Character-Swap-LoRA: direct link, hf CLI and curl.
- Browser
- Download file 3.6 kB
-
https://huggingface.co/akatz-ai/MiniMax-H3-Character-Swap-LoRA/resolve/main/training/train-launcher.py
- Command line
-
hf download hf://akatz-ai/MiniMax-H3-Character-Swap-LoRA/training/train-launcher.py
-
curl -L -o train-launcher.py https://huggingface.co/akatz-ai/MiniMax-H3-Character-Swap-LoRA/resolve/main/training/train-launcher.py
3.6 kB
| import os,sys,runpy,traceback,json,time | |
| config_path=os.path.abspath(sys.argv[1]) | |
| root=os.environ['AI_TOOLKIT_ROOT'] | |
| os.chdir(root);sys.path.insert(0,root) | |
| from extensions_built_in.sd_trainer.SDTrainer import SDTrainer | |
| import torch | |
| from torch.utils.checkpoint import checkpoint | |
| from extensions_built_in.diffusion_models.minimax_h3.src.transformer import MiniMaxH3Mlp | |
| mlp_forward=MiniMaxH3Mlp.forward | |
| def chunked_mlp(self,x): | |
| if x.shape[-2] <= 1024: | |
| return mlp_forward(self,x) | |
| return torch.cat([checkpoint(mlp_forward,self,c,use_reentrant=False) if torch.is_grad_enabled() else mlp_forward(self,c) for c in x.split(1024,dim=-2)],dim=-2) | |
| # Check values and gradients on the tokenwise operation before loading weights. | |
| torch.manual_seed(123) | |
| m=MiniMaxH3Mlp(16,32).double() | |
| a=torch.randn(1,2051,16,dtype=torch.float64,requires_grad=True) | |
| b=a.detach().clone().requires_grad_(True) | |
| y=mlp_forward(m,a); y.square().sum().backward() | |
| expected=[p.grad.clone() for p in m.parameters()]; m.zero_grad() | |
| z=chunked_mlp(m,b); z.square().sum().backward() | |
| assert torch.allclose(y,z,atol=1e-10,rtol=1e-10) | |
| assert torch.allclose(a.grad,b.grad,atol=1e-10,rtol=1e-10) | |
| assert all(torch.allclose(e,p.grad,atol=1e-9,rtol=1e-9) for e,p in zip(expected,m.parameters())) | |
| print('CHUNK_PARITY_PASS: outputs, input gradients, parameter gradients',flush=True) | |
| del m,a,b,y,z,expected | |
| MiniMaxH3Mlp.forward=chunked_mlp | |
| # Freeze evaluation conditioning with the full 73-frame video context. | |
| # Upstream cache_sample_prompts constructs GenerateImageConfig without num_frames, | |
| # which silently limits reference video conditioning to five frames. | |
| from PIL import Image | |
| from toolkit.config_modules import GenerateImageConfig | |
| def cache_eval_prompts(self): | |
| if self.train_config.disable_sampling: | |
| return | |
| vae_device=self.sd.vae.device | |
| self.sd.vae.to("cpu") | |
| torch.cuda.empty_cache() | |
| self.sd.sample_prompts_cache=[] | |
| for item in self.sample_config.samples: | |
| gen=GenerateImageConfig(prompt=item.prompt,width=item.width,height=item.height,num_frames=item.num_frames,fps=item.fps,ctrl_img_1=item.ctrl_img_1,ctrl_img_2=item.ctrl_img_2,output_path='/tmp/h3-unused.mp4') | |
| self.sd.prepare_sample_prompt_context(gen) | |
| with torch.no_grad(): | |
| emb=self.sd.get_prompt_embeds(item.prompt,control_images=[Image.open(item.ctrl_img_1).convert('RGB'),item.ctrl_img_2]).to('cpu') | |
| self.sd.sample_prompts_cache.append({'conditional':emb,'unconditional':emb}) | |
| self.sd.vae.to(vae_device) | |
| print('EVAL_PROMPTS_CACHED: full 73-frame reference context',flush=True) | |
| SDTrainer.cache_sample_prompts=cache_eval_prompts | |
| torch.set_num_threads(16) | |
| torch.set_num_interop_threads(4) | |
| original=SDTrainer.hook_train_loop | |
| def traced(self,batch): | |
| bs=batch if isinstance(batch,list) else [batch] | |
| info={'step':self.step_num,'regularization':[b.get_is_reg_list() for b in bs]} | |
| print('SMOKE_STEP_START '+json.dumps(info),flush=True) | |
| torch.cuda.synchronize(); started=time.monotonic() | |
| torch.cuda.reset_peak_memory_stats() | |
| try: | |
| result=original(self,batch) | |
| except Exception: | |
| traceback.print_exc() | |
| print('SMOKE_MEMORY '+json.dumps(dict(allocated=torch.cuda.max_memory_allocated(),reserved=torch.cuda.max_memory_reserved())),flush=True) | |
| raise ValueError('Smoke batch failed; aborting rather than skipping') | |
| torch.cuda.synchronize() | |
| print('SMOKE_STEP_PASS '+json.dumps(dict(**info,elapsed_s=time.monotonic()-started,peak_allocated=torch.cuda.max_memory_allocated(),peak_reserved=torch.cuda.max_memory_reserved())),flush=True) | |
| return result | |
| SDTrainer.hook_train_loop=traced | |
| sys.argv=[root+'/run.py',config_path] | |
| runpy.run_path(root+'/run.py',run_name='__main__') | |