mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-11 12:21:09 +08:00
Code release.
This commit is contained in:
@@ -0,0 +1,53 @@
|
||||
from einops import rearrange
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import torch
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def make_reconstructions_from_batch(batch, save_dir, epoch, tokenizer):
|
||||
check_batch(batch)
|
||||
|
||||
original_frames = tensor_to_np_frames(rearrange(batch['observations'], 'b t c h w -> b t h w c'))
|
||||
all = [original_frames]
|
||||
|
||||
rec_frames = generate_reconstructions_with_tokenizer(batch, tokenizer)
|
||||
all.append(rec_frames)
|
||||
|
||||
for i, image in enumerate(map(Image.fromarray, np.concatenate(list(np.concatenate((original_frames, rec_frames), axis=-2)), axis=-3))):
|
||||
image.save(save_dir / f'epoch_{epoch:03d}_t_{i:03d}.png')
|
||||
|
||||
return
|
||||
|
||||
|
||||
def check_batch(batch):
|
||||
assert sorted(batch.keys()) == ['actions', 'ends', 'mask_padding', 'observations', 'rewards']
|
||||
b, t, _, _, _ = batch['observations'].shape # (B, T, C, H, W)
|
||||
assert batch['actions'].shape == batch['rewards'].shape == batch['ends'].shape == batch['mask_padding'].shape == (b, t)
|
||||
|
||||
|
||||
def tensor_to_np_frames(inputs):
|
||||
check_float_btw_0_1(inputs)
|
||||
return inputs.mul(255).cpu().numpy().astype(np.uint8)
|
||||
|
||||
|
||||
def check_float_btw_0_1(inputs):
|
||||
assert inputs.is_floating_point() and (inputs >= 0).all() and (inputs <= 1).all()
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def generate_reconstructions_with_tokenizer(batch, tokenizer):
|
||||
check_batch(batch)
|
||||
inputs = rearrange(batch['observations'], 'b t c h w -> (b t) c h w')
|
||||
outputs = reconstruct_through_tokenizer(inputs, tokenizer)
|
||||
b, t, _, _, _ = batch['observations'].size()
|
||||
outputs = rearrange(outputs, '(b t) c h w -> b t h w c', b=b, t=t)
|
||||
rec_frames = tensor_to_np_frames(outputs)
|
||||
return rec_frames
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def reconstruct_through_tokenizer(inputs, tokenizer):
|
||||
check_float_btw_0_1(inputs)
|
||||
reconstructions = tokenizer.encode_decode(inputs, should_preprocess=True, should_postprocess=True)
|
||||
return torch.clamp(reconstructions, 0, 1)
|
||||
Reference in New Issue
Block a user