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__')