ZoomLDM-CDM-brca / run_demo_inference.py
BiliSakura's picture
Add files using upload-large-folder tool
b4865e1 verified
Raw
History Blame Contribute Delete
1.75 kB
#!/usr/bin/env python3
"""Run ZoomLDM-CDM-brca demo inference and save visualization to demo_images/."""
from pathlib import Path
import numpy as np
import torch
from diffusers import DiffusionPipeline
from PIL import Image, ImageDraw
def main() -> None:
repo = Path(__file__).resolve().parent
demo_dir = repo / "demo_images"
demo_dir.mkdir(exist_ok=True)
pipe = DiffusionPipeline.from_pretrained(
str(repo),
custom_pipeline=str(repo / "pipeline.py"),
trust_remote_code=True,
local_files_only=True,
).to("cuda")
magnification = torch.tensor([0], device="cuda")
out = pipe(
batch_size=1,
magnification=magnification,
num_inference_steps=50,
guidance_scale=1.0,
)
# out.samples: (1, 512, 65) -> visualize first sample as a heatmap-like image
sample = out.samples[0].detach().float().cpu().numpy()
sample = (sample - sample.min()) / (sample.max() - sample.min() + 1e-8)
sample_u8 = (sample * 255).astype(np.uint8) # (512, 65)
img = Image.fromarray(sample_u8, mode="L").resize((520, 512), Image.Resampling.NEAREST).convert("RGB")
img.save(demo_dir / "output.jpeg")
# Simple input card to make widget context explicit
card = Image.new("RGB", (520, 512), color=(245, 245, 245))
draw = ImageDraw.Draw(card)
draw.text((24, 24), "ZoomLDM-CDM-brca", fill=(20, 20, 20))
draw.text((24, 64), "Input: magnification class = 0", fill=(40, 40, 40))
draw.text((24, 100), "Output: sampled conditioning tensor visualization", fill=(40, 40, 40))
card.save(demo_dir / "input.jpeg")
print(f"Saved {demo_dir / 'input.jpeg'}")
print(f"Saved {demo_dir / 'output.jpeg'}")
if __name__ == "__main__":
main()