Views
No views yet
1from models.unet import PixelArtUNet
2
3model = PixelArtUNet(
4 channels = [128, 256, 512, 1024],
5 num_residual_layers = 2,
6 t_embed_dim = 128,
7 midcoder_dropout_p=0.2
8).to(device)1from huggingface_hub import hf_hub_download
2from safetensors.torch import load_file
3
4repo_id = "mradovic38/sprite-flow"
5filename = "model.safetensors"
6file_path = hf_hub_download(repo_id=repo_id, filename=filename)
7checkpoint = load_file(file_path)
8model.load_state_dict(checkpoint)
9model.to(device)
10model.eval()1from sampling.conditional_probability_path import GaussianConditionalProbabilityPath
2from sampling.noise_scheduling import LinearAlpha, LinearBeta
3
4path = GaussianConditionalProbabilityPath(
5 p_data=None,
6 p_simple_shape=[4, 128, 128],
7 alpha=LinearAlpha(),
8 beta=LinearBeta()
9).to(device)
10path.eval()1import torch
2
3from diff_eq.ode_sde import UnguidedVectorFieldODE
4from diff_eq.simulator import EulerSimulator
5
6num_timesteps = 200 # example number of timesteps
7num_samples = 3 # example number of samples
8
9ts = torch.linspace(0, 1, num_timesteps).view(1, -1, 1, 1, 1).expand(num_samples, -1, 1, 1, 1).to(device)
10x0 = path.p_simple.sample(num_samples).to(device) # (num_samples, 4, 128, 128)
11ode = UnguidedVectorFieldODE(model)
12simulator = EulerSimulator(ode)
13x1 = simulator.simulate(x0, ts) # (num_samples, 4, 128, 128)1from utils.helpers import tensor_to_rgba_image, normalize_to_unit
2
3x1 = normalize_to_unit(x1) # [-1, 1] -> [0, 1]
4imgs = tensor_to_rgba_image(x1)