FlameF0X commited on
Commit
2321a0f
·
verified ·
1 Parent(s): fdf33ba

Upload inference.py

Browse files
Files changed (1) hide show
  1. inference.py +34 -13
inference.py CHANGED
@@ -56,6 +56,8 @@ class ShellDConfig:
56
  beta_end: float = 0.02
57
  model_name: str = "ShellD"
58
  dropout: float = 0.0 # match train.py; 0 = no dropout (inference)
 
 
59
 
60
  @classmethod
61
  def from_json(cls, path: str) -> "ShellDConfig":
@@ -142,7 +144,7 @@ class Decoder(nn.Module):
142
  nn.Sequential(
143
  ResidualBlock(ch, out_ch),
144
  ResidualBlock(out_ch, out_ch),
145
- nn.Upsample(scale_factor=2, mode="nearest"),
146
  )
147
  )
148
  ch = out_ch
@@ -151,7 +153,7 @@ class Decoder(nn.Module):
151
  self.out_conv = nn.Sequential(
152
  ResidualBlock(ch, ch),
153
  nn.Conv2d(ch, 3, 3, padding=1),
154
- nn.Sigmoid(),
155
  )
156
 
157
  def forward(self, z: torch.Tensor) -> torch.Tensor:
@@ -329,6 +331,8 @@ class ShellDModel(nn.Module):
329
  self.vae = VAE(cfg)
330
  self.dit = DiT(cfg)
331
  self.text_encoder = None # loaded separately
 
 
332
 
333
  def encode_text(self, prompts: List[str], device: torch.device) -> torch.Tensor:
334
  assert self.text_encoder is not None, "Text encoder not loaded"
@@ -396,9 +400,13 @@ class DiffusionSchedule:
396
  # Start from random noise
397
  z = torch.randn(B, cfg.latent_dim, H, W, device=device)
398
 
399
- # For classifier-free guidance, we need an unconditional embedding (zeros)
400
  if cfg_scale != 1.0:
401
- uncond_emb = torch.zeros_like(text_emb)
 
 
 
 
402
 
403
  # Time step resampling for faster inference (DDPM-style, evenly spaced)
404
  step_indices = torch.linspace(0, cfg.num_timesteps - 1, num_steps, device=device, dtype=torch.long)
@@ -535,6 +543,13 @@ class ShellDInference:
535
  self.model.vae.load_state_dict(vae_sd)
536
  self.model.dit.load_state_dict(dit_sd)
537
 
 
 
 
 
 
 
 
538
  if txt_sd and self.model.text_encoder is not None:
539
  self.model.text_encoder.load_state_dict(txt_sd)
540
  print("Text encoder weights loaded from safetensors.")
@@ -579,12 +594,12 @@ class ShellDInference:
579
  seed=seed,
580
  )
581
 
582
- # Decode latent to image
583
- img_tensor = self.model.vae.decode(z) # [1, 3, 256, 256], values in [0, 1]
584
 
585
- # Convert to PIL
586
  img_np = img_tensor[0].permute(1, 2, 0).cpu().numpy() # [256, 256, 3]
587
- img_np = (img_np * 255).clip(0, 255).astype(np.uint8)
588
  img = Image.fromarray(img_np)
589
 
590
  if output_size is not None:
@@ -594,9 +609,9 @@ class ShellDInference:
594
 
595
  def _decode_latent(self, z: torch.Tensor) -> Image.Image:
596
  """Decode a latent tensor to a PIL Image (helper for streaming)."""
597
- img_tensor = self.model.vae.decode(z) # [B, 3, H, W] in [0,1]
598
  img_np = img_tensor[0].permute(1, 2, 0).cpu().numpy()
599
- img_np = (img_np * 255).clip(0, 255).astype(np.uint8)
600
  return Image.fromarray(img_np)
601
 
602
  @torch.no_grad()
@@ -625,7 +640,13 @@ class ShellDInference:
625
 
626
  # --- Encode prompt ---
627
  text_emb = self.model.encode_text([prompt], device) # [1, 1, 384]
628
- uncond_emb = torch.zeros_like(text_emb) if cfg_scale != 1.0 else None
 
 
 
 
 
 
629
 
630
  # --- Start from pure noise ---
631
  H = W = cfg.image_size // (2 ** cfg.ae_num_blocks)
@@ -710,12 +731,12 @@ class ShellDInference:
710
  )
711
 
712
  # Decode
713
- img_tensor = self.model.vae.decode(z) # [B, 3, 256, 256]
714
 
715
  images = []
716
  for i in range(B):
717
  img_np = img_tensor[i].permute(1, 2, 0).cpu().numpy()
718
- img_np = (img_np * 255).clip(0, 255).astype(np.uint8)
719
  images.append(Image.fromarray(img_np))
720
 
721
  return images
 
56
  beta_end: float = 0.02
57
  model_name: str = "ShellD"
58
  dropout: float = 0.0 # match train.py; 0 = no dropout (inference)
59
+ kl_weight: float = 0.1 # VAE KL weight (not used in inference)
60
+ cfg_dropout_prob: float = 0.15 # CFG dropout prob (not used in inference)
61
 
62
  @classmethod
63
  def from_json(cls, path: str) -> "ShellDConfig":
 
144
  nn.Sequential(
145
  ResidualBlock(ch, out_ch),
146
  ResidualBlock(out_ch, out_ch),
147
+ nn.Upsample(scale_factor=2, mode="bilinear", align_corners=False),
148
  )
149
  )
150
  ch = out_ch
 
153
  self.out_conv = nn.Sequential(
154
  ResidualBlock(ch, ch),
155
  nn.Conv2d(ch, 3, 3, padding=1),
156
+ nn.Tanh(),
157
  )
158
 
159
  def forward(self, z: torch.Tensor) -> torch.Tensor:
 
331
  self.vae = VAE(cfg)
332
  self.dit = DiT(cfg)
333
  self.text_encoder = None # loaded separately
334
+ # Null text embedding for classifier-free guidance (loaded from checkpoint)
335
+ self.null_text_embed: Optional[torch.Tensor] = None
336
 
337
  def encode_text(self, prompts: List[str], device: torch.device) -> torch.Tensor:
338
  assert self.text_encoder is not None, "Text encoder not loaded"
 
400
  # Start from random noise
401
  z = torch.randn(B, cfg.latent_dim, H, W, device=device)
402
 
403
+ # For classifier-free guidance, we need the learned null embedding
404
  if cfg_scale != 1.0:
405
+ if model.null_text_embed is not None:
406
+ uncond_emb = model.null_text_embed.to(device).expand(B, -1, -1)
407
+ else:
408
+ # Fallback: zeros (works if model was trained with zero-dropped text)
409
+ uncond_emb = torch.zeros_like(text_emb)
410
 
411
  # Time step resampling for faster inference (DDPM-style, evenly spaced)
412
  step_indices = torch.linspace(0, cfg.num_timesteps - 1, num_steps, device=device, dtype=torch.long)
 
543
  self.model.vae.load_state_dict(vae_sd)
544
  self.model.dit.load_state_dict(dit_sd)
545
 
546
+ # Restore null text embedding for CFG
547
+ if "null_text_embed" in sd:
548
+ self.model.null_text_embed = sd["null_text_embed"].to(self.device)
549
+ print("Null text embedding loaded for CFG.")
550
+ else:
551
+ print("No null_text_embed in checkpoint — CFG will use zeros (may degrade quality).")
552
+
553
  if txt_sd and self.model.text_encoder is not None:
554
  self.model.text_encoder.load_state_dict(txt_sd)
555
  print("Text encoder weights loaded from safetensors.")
 
594
  seed=seed,
595
  )
596
 
597
+ # Decode latent to image (VAE decoder outputs [-1, 1] via Tanh)
598
+ img_tensor = self.model.vae.decode(z) # [1, 3, 256, 256], values in [-1, 1]
599
 
600
+ # Convert to PIL: [-1,1] → [0,255]
601
  img_np = img_tensor[0].permute(1, 2, 0).cpu().numpy() # [256, 256, 3]
602
+ img_np = ((img_np + 1.0) * 127.5).clip(0, 255).astype(np.uint8)
603
  img = Image.fromarray(img_np)
604
 
605
  if output_size is not None:
 
609
 
610
  def _decode_latent(self, z: torch.Tensor) -> Image.Image:
611
  """Decode a latent tensor to a PIL Image (helper for streaming)."""
612
+ img_tensor = self.model.vae.decode(z) # [B, 3, H, W] in [-1,1]
613
  img_np = img_tensor[0].permute(1, 2, 0).cpu().numpy()
614
+ img_np = ((img_np + 1.0) * 127.5).clip(0, 255).astype(np.uint8)
615
  return Image.fromarray(img_np)
616
 
617
  @torch.no_grad()
 
640
 
641
  # --- Encode prompt ---
642
  text_emb = self.model.encode_text([prompt], device) # [1, 1, 384]
643
+ if cfg_scale != 1.0:
644
+ if self.model.null_text_embed is not None:
645
+ uncond_emb = self.model.null_text_embed.to(device).expand(1, -1, -1)
646
+ else:
647
+ uncond_emb = torch.zeros_like(text_emb)
648
+ else:
649
+ uncond_emb = None
650
 
651
  # --- Start from pure noise ---
652
  H = W = cfg.image_size // (2 ** cfg.ae_num_blocks)
 
731
  )
732
 
733
  # Decode
734
+ img_tensor = self.model.vae.decode(z) # [B, 3, 256, 256] in [-1,1]
735
 
736
  images = []
737
  for i in range(B):
738
  img_np = img_tensor[i].permute(1, 2, 0).cpu().numpy()
739
+ img_np = ((img_np + 1.0) * 127.5).clip(0, 255).astype(np.uint8)
740
  images.append(Image.fromarray(img_np))
741
 
742
  return images