Upload inference.py
Browse files- 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="
|
| 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.
|
| 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
|
| 400 |
if cfg_scale != 1.0:
|
| 401 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 [
|
| 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 *
|
| 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 [
|
| 598 |
img_np = img_tensor[0].permute(1, 2, 0).cpu().numpy()
|
| 599 |
-
img_np = (img_np *
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 *
|
| 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
|