mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
interpolate pure-torch
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -82,13 +82,13 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1,
|
||||
for _i in range(pad): # fill starting frames
|
||||
buffer.put(frame)
|
||||
|
||||
I1 = f_pad(torch.from_numpy(np.transpose(frame, (2,0,1))).to(devices.device).unsqueeze(0).float() / 255.)
|
||||
I1 = f_pad(torch.from_numpy(np.transpose(frame, (2,0,1))).to(devices.device).unsqueeze(0).float() / 255.0)
|
||||
with torch.no_grad():
|
||||
with tqdm(total=len(images), desc='Interpolate', unit='frame') as pbar:
|
||||
for image in images:
|
||||
frame = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
|
||||
I0 = I1
|
||||
I1 = f_pad(torch.from_numpy(np.transpose(frame, (2,0,1))).to(devices.device).unsqueeze(0).float() / 255.)
|
||||
I1 = f_pad(torch.from_numpy(np.transpose(frame, (2,0,1))).to(devices.device).unsqueeze(0).float() / 255.0)
|
||||
I0_small = F.interpolate(I0, (32, 32), mode='bilinear', align_corners=False).to(torch.float32)
|
||||
I1_small = F.interpolate(I1, (32, 32), mode='bilinear', align_corners=False).to(torch.float32)
|
||||
ssim = ssim_matlab(I0_small[:, :3], I1_small[:, :3])
|
||||
@@ -104,7 +104,7 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1,
|
||||
else:
|
||||
output = execute(I0, I1, count-1)
|
||||
for mid in output:
|
||||
mid = (((mid[0] * 255.).byte().cpu().numpy().transpose(1, 2, 0)))
|
||||
mid = (((mid[0] * 255.0).byte().cpu().numpy().transpose(1, 2, 0)))
|
||||
buffer.put(mid[:h, :w])
|
||||
buffer.put(frame)
|
||||
pbar.update(1)
|
||||
@@ -116,3 +116,37 @@ def interpolate(images: list, count: int = 2, scale: float = 1.0, pad: int = 1,
|
||||
t1 = time.time()
|
||||
shared.log.info(f'Video interpolate: input={len(images)} frames={len(interpolated)} buffer={buffer.qsize()} duplicate={duplicate} width={w} height={h} interpolate={count} scale={scale} pad={pad} change={change} time={round(t1 - t0, 2)}')
|
||||
return interpolated
|
||||
|
||||
|
||||
def interpolate_nchw(images: list, count: int = 2, scale: float = 1.0):
|
||||
if images is None or len(images) < 2:
|
||||
return images
|
||||
if model is None:
|
||||
load()
|
||||
interpolated = []
|
||||
_n, _c, h, w = images.shape
|
||||
t0 = time.time()
|
||||
|
||||
def f_pad(img):
|
||||
return F.pad(img, padding).to(device=devices.device, dtype=devices.dtype) # pylint: disable=not-callable
|
||||
|
||||
tmp = max(128, int(128 / scale))
|
||||
ph = ((h - 1) // tmp + 1) * tmp
|
||||
pw = ((w - 1) // tmp + 1) * tmp
|
||||
padding = (0, pw - w, 0, ph - h)
|
||||
|
||||
I1 = f_pad(images[0].unsqueeze(0))
|
||||
with torch.no_grad():
|
||||
with tqdm(total=len(images), desc='Interpolate', unit='frame') as pbar:
|
||||
for frame in images:
|
||||
I0 = I1
|
||||
I1 = f_pad(frame.unsqueeze(0))
|
||||
for i in range(count-1):
|
||||
output = model.inference(I0, I1, (i+1) * 1. / (count), scale)
|
||||
interpolated.append(output)
|
||||
interpolated.append(I1)
|
||||
pbar.update(1)
|
||||
|
||||
t1 = time.time()
|
||||
shared.log.info(f'Video interpolate: input={len(images)} frames={len(interpolated)} width={w} height={h} interpolate={count} scale={scale} time={round(t1 - t0, 2)}')
|
||||
return interpolated
|
||||
|
||||
Reference in New Issue
Block a user