Image-to-Image
Transformers
Safetensors
English
janus_ddpm
feature-extraction
project:genesis
project:ArtIMis
type:model
science:materials
science:physics-simulation
risk:general
custom-code
custom_code
Instructions to use 90879c/Janus-first2last with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use 90879c/Janus-first2last with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("image-to-image", model="90879c/Janus-first2last", trust_remote_code=True, device_map="auto")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("90879c/Janus-first2last", trust_remote_code=True, dtype="auto", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """Jenus conditional DDPM model with Hugging Face AutoModel support.""" | |
| import math | |
| import random as py_random | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import PreTrainedModel | |
| from .configuration_jenus_ddpm import JenusDDPMConfig | |
| class ResBlock(nn.Module): | |
| def __init__(self, in_ch, out_ch, time_emb_dim): | |
| super().__init__() | |
| self.norm1 = nn.GroupNorm(min(8, in_ch), in_ch) | |
| self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1) | |
| self.time_mlp = nn.Linear(time_emb_dim, out_ch) | |
| self.norm2 = nn.GroupNorm(min(8, out_ch), out_ch) | |
| self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1) | |
| self.shortcut = nn.Conv2d(in_ch, out_ch, 1) if in_ch != out_ch else nn.Identity() | |
| def forward(self, x, t_emb): | |
| h = self.conv1(F.silu(self.norm1(x))) | |
| h = h + self.time_mlp(F.silu(t_emb))[:, :, None, None] | |
| h = self.conv2(F.silu(self.norm2(h))) | |
| return h + self.shortcut(x) | |
| class AttentionBlock(nn.Module): | |
| def __init__(self, channels): | |
| super().__init__() | |
| self.norm = nn.GroupNorm(min(8, channels), channels) | |
| self.qkv = nn.Conv2d(channels, channels * 3, 1) | |
| self.proj = nn.Conv2d(channels, channels, 1) | |
| self.scale = channels ** -0.5 | |
| def forward(self, x): | |
| batch, channels, height, width = x.shape | |
| h = self.norm(x) | |
| qkv = self.qkv(h).reshape(batch, 3, channels, height * width) | |
| q, k, v = qkv[:, 0], qkv[:, 1], qkv[:, 2] | |
| attn = torch.softmax(torch.bmm(q.transpose(1, 2), k) * self.scale, dim=-1) | |
| h = torch.bmm(v, attn.transpose(1, 2)).reshape(batch, channels, height, width) | |
| return x + self.proj(h) | |
| class DownBlock(nn.Module): | |
| def __init__(self, in_ch, out_ch, time_emb_dim, has_attn=False): | |
| super().__init__() | |
| self.res = ResBlock(in_ch, out_ch, time_emb_dim) | |
| self.attn = AttentionBlock(out_ch) if has_attn else nn.Identity() | |
| self.down = nn.Conv2d(out_ch, out_ch, 3, stride=2, padding=1) | |
| def forward(self, x, t_emb): | |
| x = self.res(x, t_emb) | |
| skip = self.attn(x) | |
| return self.down(skip), skip | |
| class UpBlock(nn.Module): | |
| def __init__(self, in_ch, skip_ch, out_ch, time_emb_dim, has_attn=False): | |
| super().__init__() | |
| self.up = nn.ConvTranspose2d(in_ch, in_ch, 2, stride=2) | |
| self.res = ResBlock(in_ch + skip_ch, out_ch, time_emb_dim) | |
| self.attn = AttentionBlock(out_ch) if has_attn else nn.Identity() | |
| def forward(self, x, skip, t_emb): | |
| x = self.up(x) | |
| x = torch.cat([x, skip], dim=1) | |
| x = self.res(x, t_emb) | |
| return self.attn(x) | |
| def get_timestep_embedding(timesteps, embedding_dim): | |
| half_dim = embedding_dim // 2 | |
| emb = math.log(10000) / (half_dim - 1) | |
| emb = torch.exp(torch.arange(half_dim, device=timesteps.device) * -emb) | |
| emb = timesteps[:, None].float() * emb[None, :] | |
| return torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) | |
| class ConditionalUNet(nn.Module): | |
| def __init__(self, in_channels=2, out_channels=1, base_channels=64, time_emb_dim=256): | |
| super().__init__() | |
| self.time_emb_dim = time_emb_dim | |
| self.time_mlp = nn.Sequential( | |
| nn.Linear(time_emb_dim, time_emb_dim * 4), | |
| nn.SiLU(), | |
| nn.Linear(time_emb_dim * 4, time_emb_dim), | |
| ) | |
| ch = base_channels | |
| self.conv_in = nn.Conv2d(in_channels, ch, 3, padding=1) | |
| self.down1 = DownBlock(ch, ch * 2, time_emb_dim, has_attn=False) | |
| self.down2 = DownBlock(ch * 2, ch * 4, time_emb_dim, has_attn=False) | |
| self.down3 = DownBlock(ch * 4, ch * 4, time_emb_dim, has_attn=False) | |
| self.mid_res1 = ResBlock(ch * 4, ch * 4, time_emb_dim) | |
| self.mid_attn = AttentionBlock(ch * 4) | |
| self.mid_res2 = ResBlock(ch * 4, ch * 4, time_emb_dim) | |
| self.up1 = UpBlock(ch * 4, ch * 4, ch * 4, time_emb_dim, has_attn=False) | |
| self.up2 = UpBlock(ch * 4, ch * 4, ch * 2, time_emb_dim, has_attn=False) | |
| self.up3 = UpBlock(ch * 2, ch * 2, ch, time_emb_dim, has_attn=False) | |
| self.conv_out = nn.Sequential( | |
| nn.GroupNorm(8, ch), | |
| nn.SiLU(), | |
| nn.Conv2d(ch, out_channels, 3, padding=1), | |
| ) | |
| def forward(self, x, t, cond): | |
| x = torch.cat([x, cond], dim=1) | |
| t_emb = get_timestep_embedding(t, self.time_emb_dim) | |
| t_emb = self.time_mlp(t_emb) | |
| h = self.conv_in(x) | |
| h, skip1 = self.down1(h, t_emb) | |
| h, skip2 = self.down2(h, t_emb) | |
| h, skip3 = self.down3(h, t_emb) | |
| h = self.mid_res1(h, t_emb) | |
| h = self.mid_attn(h) | |
| h = self.mid_res2(h, t_emb) | |
| h = self.up1(h, skip3, t_emb) | |
| h = self.up2(h, skip2, t_emb) | |
| h = self.up3(h, skip1, t_emb) | |
| return self.conv_out(h) | |
| class DDPMScheduler: | |
| def __init__(self, num_timesteps=1000, beta_start=1e-4, beta_end=0.02): | |
| self.num_timesteps = num_timesteps | |
| self.betas = torch.linspace(beta_start, beta_end, num_timesteps) | |
| self.alphas = 1 - self.betas | |
| self.alphas_cumprod = torch.cumprod(self.alphas, dim=0) | |
| def to(self, device): | |
| self.betas = self.betas.to(device) | |
| self.alphas = self.alphas.to(device) | |
| self.alphas_cumprod = self.alphas_cumprod.to(device) | |
| return self | |
| class DDIMScheduler: | |
| def __init__(self, ddpm_scheduler, num_inference_steps=50): | |
| self.num_timesteps = ddpm_scheduler.num_timesteps | |
| self.alphas_cumprod = ddpm_scheduler.alphas_cumprod | |
| step_ratio = self.num_timesteps // num_inference_steps | |
| self.timesteps = torch.flip(torch.arange(0, num_inference_steps) * step_ratio, [0]) | |
| def step(self, model_output, t, t_prev, x_t, eta=0.0): | |
| alpha_t = self.alphas_cumprod[t].view(-1, 1, 1, 1) | |
| alpha_t_prev = ( | |
| self.alphas_cumprod[t_prev].view(-1, 1, 1, 1) | |
| if t_prev >= 0 | |
| else torch.ones_like(alpha_t) | |
| ) | |
| pred_x0 = (x_t - torch.sqrt(1 - alpha_t) * model_output) / torch.sqrt(alpha_t) | |
| pred_x0 = torch.clamp(pred_x0, -1, 1) | |
| sigma_t = eta * torch.sqrt((1 - alpha_t_prev) / (1 - alpha_t)) * torch.sqrt( | |
| 1 - alpha_t / alpha_t_prev | |
| ) | |
| pred_dir = torch.sqrt(1 - alpha_t_prev - sigma_t ** 2) * model_output | |
| x_prev = torch.sqrt(alpha_t_prev) * pred_x0 + pred_dir | |
| if eta > 0 and t_prev >= 0: | |
| x_prev = x_prev + sigma_t * torch.randn_like(x_t) | |
| return x_prev | |
| def to(self, device): | |
| self.alphas_cumprod = self.alphas_cumprod.to(device) | |
| self.timesteps = self.timesteps.to(device) | |
| return self | |
| class JenusDDPMModel(PreTrainedModel): | |
| """Hugging Face compatible Jenus conditional DDPM model.""" | |
| config_class = JenusDDPMConfig | |
| base_model_prefix = "jenus_ddpm" | |
| main_input_name = "x" | |
| _tied_weights_keys = [] | |
| all_tied_weights_keys = {} | |
| def __init__(self, config): | |
| super().__init__(config) | |
| self.unet = ConditionalUNet( | |
| in_channels=config.in_channels, | |
| out_channels=config.out_channels, | |
| base_channels=config.base_channels, | |
| time_emb_dim=config.time_emb_dim, | |
| ) | |
| self.ddpm_scheduler = None | |
| def forward(self, x, t, cond): | |
| return self.unet(x, t, cond) | |
| def to(self, *args, **kwargs): | |
| return super().to(*args, **kwargs) | |
| def device(self): | |
| return next(self.parameters()).device | |
| def _prepare_condition(self, raw_density): | |
| arr = raw_density.astype(np.float32) | |
| np.nan_to_num(arr, copy=False) | |
| raw_min, raw_max = float(arr.min()), float(arr.max()) | |
| img_size = tuple(self.config.img_size) | |
| t = torch.from_numpy(arr).unsqueeze(0).unsqueeze(0) | |
| t = F.interpolate(t, size=img_size, mode="bilinear", align_corners=False) | |
| tmin, tmax = t.min(), t.max() | |
| if tmax - tmin > 1e-6: | |
| t = 2.0 * (t - tmin) / (tmax - tmin) - 1.0 | |
| return t.to(self.device), raw_min, raw_max | |
| def _sample_once(self, condition, num_steps=50, eta=1.0, seed=None): | |
| if seed is None: | |
| seed = py_random.randint(0, 2**31 - 1) | |
| torch.manual_seed(seed) | |
| if self.device.type == "cuda": | |
| torch.cuda.manual_seed(seed) | |
| batch, _, height, width = condition.shape | |
| ddpm_scheduler = DDPMScheduler( | |
| num_timesteps=self.config.num_timesteps, | |
| beta_start=self.config.beta_start, | |
| beta_end=self.config.beta_end, | |
| ).to(self.device) | |
| ddim = DDIMScheduler(ddpm_scheduler, num_inference_steps=num_steps).to(self.device) | |
| x = torch.randn(batch, 1, height, width, device=self.device) | |
| timesteps = ddim.timesteps.long() | |
| self.eval() | |
| for i, t in enumerate(timesteps): | |
| t_prev = timesteps[i + 1] if i + 1 < len(timesteps) else torch.tensor(-1, device=self.device) | |
| t_batch = t.expand(batch).to(self.device) | |
| pred_noise = self.forward(x, t_batch, condition) | |
| x = ddim.step(pred_noise, t, t_prev, x, eta=eta) | |
| return x | |
| def predict(self, density, num_steps=50, eta=1.0, seed=None): | |
| orig_h, orig_w = density.shape | |
| condition, raw_min, raw_max = self._prepare_condition(density) | |
| x = self._sample_once(condition, num_steps=num_steps, eta=eta, seed=seed) | |
| pred_np = x[0, 0].cpu().float().numpy() | |
| pred_phys = (pred_np + 1.0) / 2.0 * (raw_max - raw_min) + raw_min | |
| pred_t = torch.from_numpy(pred_phys[None, None].astype(np.float32)) | |
| pred_t = F.interpolate(pred_t, size=(orig_h, orig_w), mode="bilinear", align_corners=False) | |
| return pred_t[0, 0].numpy() | |
| def predict_ensemble(self, density, n_samples=5, num_steps=50, eta=1.0): | |
| orig_h, orig_w = density.shape | |
| condition, raw_min, raw_max = self._prepare_condition(density) | |
| members = [] | |
| for _ in range(n_samples): | |
| x = self._sample_once(condition, num_steps=num_steps, eta=eta, seed=None) | |
| pred_np = x[0, 0].cpu().float().numpy() | |
| pred_phys = (pred_np + 1.0) / 2.0 * (raw_max - raw_min) + raw_min | |
| members.append(pred_phys) | |
| stack = np.stack(members) | |
| def resize(arr): | |
| t = torch.from_numpy(arr[None, None].astype(np.float32)) | |
| t = F.interpolate(t, size=(orig_h, orig_w), mode="bilinear", align_corners=False) | |
| return t[0, 0].numpy() | |
| return resize(stack.mean(axis=0)), resize(stack.std(axis=0)) | |