SD3-Long-Captioner / test_engine.py
gokaygokay's picture
Add one-click detail pipeline, presets, variants and multi-step teacher
81564ac verified
Raw History Blame Contribute Delete
2.68 kB
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
from unittest.mock import patch
import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image
from engine import RestorationModel
class DummyModel:
def __call__(self, inputs, t, r, features):
latent, noise = inputs.float().chunk(2, dim=1)
# Exact one-step restoration is noise minus predicted velocity.
return noise - latent
def dummy():
obj = RestorationModel.__new__(RestorationModel)
obj.device = 'cpu'
obj.features = lambda pixels: None
obj.encode = lambda pixels: F.avg_pool2d(pixels,8)
obj.decode = lambda latent: F.interpolate(latent,scale_factor=8,mode='nearest')
obj.model = DummyModel()
return obj
def test_batching_and_shared_noise_are_invariant():
pixels = np.random.default_rng(5).integers(0,256,(512,1024,3),dtype=np.uint8)
image = Image.fromarray(pixels)
model = dummy()
with patch('torch.cuda.synchronize'):
first, stats = model.restore(image, seed=42, batch_size=1)
second, _ = model.restore(image, seed=42, batch_size=4)
np.testing.assert_array_equal(first, second)
assert stats['tiles'] == 3
expected = np.array(F.interpolate(F.avg_pool2d(torch.from_numpy(pixels.copy()).permute(2,0,1)[None].float(),8),scale_factor=8,mode='nearest')[0].permute(1,2,0))
assert np.max(np.abs(first.astype(float)-expected)) < 2
class ConstantTeacher:
def __init__(self): self.times=[]
def __call__(self, inputs, t, r, features):
self.times.append((float(t[0]),float(r[0])))
full,weak=inputs.chunk(2)
return torch.cat((torch.ones_like(full[:,:3])*2,torch.ones_like(weak[:,:3])))
def test_multistep_cfg_schedule_and_batch_invariance():
obj=dummy(); obj.mode='Detail'; obj.model=ConstantTeacher()
obj.features=lambda pixels:[pixels.mean((2,3))]
image=Image.new('RGB',(1024,512),(100,150,200))
with patch('torch.cuda.synchronize'):
first,stats=obj.restore(image,seed=42,batch_size=1,steps=5,guidance=.5)
second,_=obj.restore(image,seed=42,batch_size=4,steps=5,guidance=.5)
np.testing.assert_array_equal(first,second)
assert stats['steps']==5 and stats['model_evaluations']==30
assert abs(obj.model.times[0][0]-1)<1e-6 and abs(obj.model.times[-1][1])<1e-6
generator=torch.Generator().manual_seed(42)
noise=torch.randn((1,3,64,128),generator=generator)
expected=F.interpolate((noise-1.5).clamp(-1,1),scale_factor=8,mode='nearest')
expected=((expected[0].permute(1,2,0)+1)*127.5).round().to(torch.uint8).numpy()
np.testing.assert_array_equal(first,expected)