akatz-ai's picture
Publish only the final 1000-step character-swap LoRA
5823b1d verified
Raw History Blame Contribute Delete
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__')