| import json,torch | |
| from pathlib import Path | |
| import yaml | |
| def cfg(r):return yaml.safe_load((Path(r)/'conf/config.yaml').read_text()) | |
| def frames(): | |
| y,x=torch.meshgrid(torch.arange(128),torch.arange(128),indexing='ij');return torch.stack([torch.exp(-((x-45-t*3)**2+(y-60-t*2)**2)/300)*20 for t in range(3)]) | |
| class STEPS(torch.nn.Module): | |
| def __init__(self,levels=8,members=24,steps=12):super().__init__();self.levels=levels;self.members=members;self.steps=steps;self.scale=torch.nn.Parameter(torch.ones(levels)) | |
| def forward(self,x): | |
| velocity=x[-1]-x[-2];base=x[-1];out=[] | |
| for m in range(self.members):out.append(torch.stack([(base+s*velocity+torch.randn_like(base)*.02*self.scale.mean()).clamp_min(0) for s in range(1,self.steps+1)])) | |
| return torch.stack(out) | |
| def write(p,o):p=Path(p);p.parent.mkdir(parents=True,exist_ok=True);p.write_text(json.dumps(o,indent=2)+'\n') | |