In [None]:
using_colab = 'google.colab' in str(get_ipython())
print("Using Colab:", using_colab)
if using_colab:
    !pip install -U "xformers==0.0.33" --index-url https://download.pytorch.org/whl/cu126
!pip install diffusers


In [None]:
import torch
from data.spair import SPairDataset
from torch.utils.data import DataLoader
import os
from pathlib import Path

from utils.utils_featuremaps import PreComputedFeaturemaps

base_dir = os.path.abspath(os.path.curdir)

if using_colab:
    base_dir = os.path.join(os.path.abspath(os.path.curdir), 'AML-polito')


def collate_single(batch_list):
    return batch_list[0]

save_dir = Path(base_dir) / "data" / "features"
dataset_size = 'large'  # 'small' or 'large'

# Load dataset and construct dataloader

test_dataset = SPairDataset(datatype='test', dataset_size=dataset_size)

test_dataloader = DataLoader(test_dataset, num_workers=4, batch_size=1, collate_fn=collate_single)
print("Dataset loaded")

In [None]:
import torch

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("Device:", device)
torch.backends.cuda.matmul.allow_tf32 = True



In [None]:
import torch
from torchvision import transforms

def resize_image_sd(
    sample: torch.Tensor,
    img_output_size: tuple[int, int] = (768, 768),  # (H, W)
    ensemble_size: int = 8,
) -> torch.Tensor:
    """
    sample: (C,H,W) con valori 0..255 (uint8 o float)
    return: (E,C,H',W') normalizzata in [-1,1]
    """
    if sample.ndim != 3:
        raise ValueError(f"`sample` deve essere (C,H,W). Trovato shape={tuple(sample.shape)}")

    resize = transforms.Resize(
        img_output_size,
        interpolation=transforms.InterpolationMode.BICUBIC,
        antialias=True
    )

    img = resize(sample)                      # (C,H',W')
    img = img.to(torch.float32)
    img = (img / 255.0 - 0.5) * 2.0          # [-1,1]
    img = img.unsqueeze(0).repeat(ensemble_size, 1, 1, 1)  # (E,C,H',W')
    return img


def resize_keypoints(
    keypoint: torch.Tensor,
    orig_size: tuple[int, int],              # (H, W) originali
    new_size: tuple[int, int],               # (H, W) dopo resize
    ensemble_size: int = 8,
    repeat: bool = True,
) -> torch.Tensor:
    """
    keypoint: (N,2) in pixel, ordine (x,y)
    return: (E,N,2) se repeat=True, altrimenti (N,2)
    """
    if keypoint.ndim != 2 or keypoint.shape[-1] != 2:
        raise ValueError(f"`keypoint` deve essere (N,2). Trovato shape={tuple(keypoint.shape)}")

    orig_h, orig_w = orig_size
    new_h, new_w = new_size

    sx = new_w / orig_w
    sy = new_h / orig_h

    scale = keypoint.new_tensor([sx, sy])   # (sx, sy) -> (x,y)
    kp = keypoint.to(torch.float32) * scale # (N,2)

    if repeat:
        kp = kp.unsqueeze(0).repeat(ensemble_size, 1, 1)  # (E,N,2)

    return kp

In [None]:
from diffusers import StableDiffusionPipeline
import torch
import torch.nn as nn
import matplotlib.pyplot as plt
import numpy as np
from typing import Any, Callable, Dict, List, Optional, Union, Tuple
from diffusers.models.unets.unet_2d_condition import UNet2DConditionModel
from diffusers import DDIMScheduler
import gc
import os


class MyUNet2DConditionModel(UNet2DConditionModel):
    def forward(
            self,
            sample: torch.Tensor,
            timestep: Union[torch.Tensor, float, int],
            encoder_hidden_states: torch.Tensor,
            up_ft_indices,
            class_labels: Optional[torch.Tensor] = None,
            timestep_cond: Optional[torch.Tensor] = None,
            attention_mask: Optional[torch.Tensor] = None,
            cross_attention_kwargs: Optional[Dict[str, Any]] = None,
            added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
            down_block_additional_residuals: Optional[Tuple[torch.Tensor]] = None,
            mid_block_additional_residual: Optional[torch.Tensor] = None,
            down_intrablock_additional_residuals: Optional[Tuple[torch.Tensor]] = None,
            encoder_attention_mask: Optional[torch.Tensor] = None,
            return_dict: bool = True
    ):
        r"""
        Args:
            sample (`torch.FloatTensor`): (batch, channel, height, width) noisy inputs tensor
            timestep (`torch.FloatTensor` or `float` or `int`): (batch) timesteps
            encoder_hidden_states (`torch.FloatTensor`): (batch, sequence_length, feature_dim) encoder hidden states
            cross_attention_kwargs (`dict`, *optional*):
                A kwargs dictionary that if specified is passed along to the `AttnProcessor` as defined under
                `self.processor` in
                [diffusers.cross_attention](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/cross_attention.py).
        """
        # By default samples have to be AT least a multiple of the overall upsampling factor.
        # The overall upsampling factor is equal to 2 ** (# num of upsampling layears).
        # However, the upsampling interpolation output size can be forced to fit any upsampling size
        # on the fly if necessary.
        default_overall_up_factor = 2 ** self.num_upsamplers

        # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor`
        forward_upsample_size = False
        upsample_size = None

        if any(s % default_overall_up_factor != 0 for s in sample.shape[-2:]):
            # logger.info("Forward upsample size to force interpolation output size.")
            forward_upsample_size = True

        # prepare attention_mask
        if attention_mask is not None:
            attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0
            attention_mask = attention_mask.unsqueeze(1)

        # 0. center input if necessary
        if self.config.center_input_sample:
            sample = 2 * sample - 1.0

        # 1. time
        timesteps = timestep
        if not torch.is_tensor(timesteps):
            # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can
            # This would be a good case for the `match` statement (Python 3.10+)
            is_mps = sample.device.type == "mps"
            if isinstance(timestep, float):
                dtype = torch.float32 if is_mps else torch.float64
            else:
                dtype = torch.int32 if is_mps else torch.int64
            timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device)
        elif len(timesteps.shape) == 0:
            timesteps = timesteps[None].to(sample.device)

        # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
        timesteps = timesteps.expand(sample.shape[0])

        t_emb = self.time_proj(timesteps)

        # timesteps does not contain any weights and will always return f32 tensors
        # but time_embedding might actually be running in fp16. so we need to cast here.
        # there might be better ways to encapsulate this.
        t_emb = t_emb.to(dtype=self.dtype)

        emb = self.time_embedding(t_emb, timestep_cond)

        if self.class_embedding is not None:
            if class_labels is None:
                raise ValueError("class_labels should be provided when num_class_embeds > 0")

            if self.config.class_embed_type == "timestep":
                class_labels = self.time_proj(class_labels)

            class_emb = self.class_embedding(class_labels).to(dtype=self.dtype)
            emb = emb + class_emb

        # 2. pre-process
        sample = self.conv_in(sample)

        # 3. down
        down_block_res_samples = (sample,)
        for downsample_block in self.down_blocks:
            if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention:
                sample, res_samples = downsample_block(
                    hidden_states=sample,
                    temb=emb,
                    encoder_hidden_states=encoder_hidden_states,
                    attention_mask=attention_mask,
                    cross_attention_kwargs=cross_attention_kwargs,
                )
            else:
                sample, res_samples = downsample_block(hidden_states=sample, temb=emb)

            down_block_res_samples += res_samples

        # 4. mid
        if self.mid_block is not None:
            sample = self.mid_block(
                sample,
                emb,
                encoder_hidden_states=encoder_hidden_states,
                attention_mask=attention_mask,
                cross_attention_kwargs=cross_attention_kwargs,
            )

        # 5. up
        up_ft = {}
        for i, upsample_block in enumerate(self.up_blocks):

            if i > np.max(up_ft_indices):
                break

            is_final_block = i == len(self.up_blocks) - 1

            res_samples = down_block_res_samples[-len(upsample_block.resnets):]
            down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)]

            # if we have not reached the final block and need to forward the
            # upsample size, we do it here
            if not is_final_block and forward_upsample_size:
                upsample_size = down_block_res_samples[-1].shape[2:]

            if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention:
                sample = upsample_block(
                    hidden_states=sample,
                    temb=emb,
                    res_hidden_states_tuple=res_samples,
                    encoder_hidden_states=encoder_hidden_states,
                    cross_attention_kwargs=cross_attention_kwargs,
                    upsample_size=upsample_size,
                    attention_mask=attention_mask,
                )
            else:
                sample = upsample_block(
                    hidden_states=sample, temb=emb, res_hidden_states_tuple=res_samples, upsample_size=upsample_size
                )

            if i in up_ft_indices:
                up_ft[i] = sample.detach()

        output = {}
        output['up_ft'] = up_ft
        return output


class OneStepSDPipeline(StableDiffusionPipeline):
    @torch.no_grad()
    def __call__(
            self,
            img_tensor,
            t,
            up_ft_indices,
            prompt_embeds: Optional[torch.FloatTensor] = None,
    ):
        device = self._execution_device
        latents = self.vae.encode(img_tensor).latent_dist.sample() * self.vae.config.scaling_factor
        t = torch.tensor(t, dtype=torch.long, device=device)
        noise = torch.randn_like(latents).to(device)
        latents_noisy = self.scheduler.add_noise(latents, noise, t)
        print("forwarding unet...")
        unet_output = self.unet(
            latents_noisy,
            t,
            encoder_hidden_states=prompt_embeds,
            up_ft_indices=up_ft_indices,
        )
        return unet_output

In [None]:
def encode_prompt_embeds(pipe, prompt: str, device: str = "cuda"):
    # Tokenizza anche la stringa vuota e produce embeddings validi [1, 77, dim]
    text_inputs = pipe.tokenizer(
        [prompt],
        padding="max_length",
        max_length=pipe.tokenizer.model_max_length,
        truncation=True,
        return_tensors="pt",
    )
    input_ids = text_inputs.input_ids.to(device)
    attention_mask = getattr(text_inputs, "attention_mask", None)
    if attention_mask is not None:
        attention_mask = attention_mask.to(device)

    with torch.no_grad():
        out = pipe.text_encoder(input_ids=input_ids, attention_mask=attention_mask)
        prompt_embeds = out[0]
    return prompt_embeds

In [None]:
import torch

class SDFeaturizer4Eval():
    def __init__(self, sd_id='Manojb/stable-diffusion-2-1-base', null_prompt='', cat_list=[]):
        unet = MyUNet2DConditionModel.from_pretrained(sd_id, subfolder="unet")
        onestep_pipe = OneStepSDPipeline.from_pretrained(sd_id, unet=unet, safety_checker=None)
        onestep_pipe.vae.decoder = None
        onestep_pipe.scheduler = DDIMScheduler.from_pretrained(sd_id, subfolder="scheduler")
        onestep_pipe = onestep_pipe.to("cuda")
        onestep_pipe.enable_attention_slicing()
        onestep_pipe.enable_xformers_memory_efficient_attention()
        null_prompt_embeds = encode_prompt_embeds(onestep_pipe, null_prompt, device="cuda")  # [1, 77, dim]

        self.null_prompt_embeds = null_prompt_embeds
        self.null_prompt = null_prompt
        self.pipe = onestep_pipe
        with torch.no_grad():
            cat2prompt_embeds = {}
            print("start encoding prompts for categories...")
            for cat in cat_list:
                prompt = f"a photo of a {cat}"
                prompt_embeds = encode_prompt_embeds(self.pipe, prompt, device="cuda")  # [1, 77, dim]
                cat2prompt_embeds[cat] = prompt_embeds
            print("encoded prompts for categories:", list(cat2prompt_embeds.keys()))
            self.cat2prompt_embeds = cat2prompt_embeds

        self.pipe.tokenizer = None
        self.pipe.text_encoder = None
        torch.cuda.empty_cache()

    @torch.no_grad()
    def forward(self,
                img_tensor,
                category=None,
                t=261,
                up_ft_index=1,
                ensemble_size=8):
        '''
        Args:
            img_tensor: should be a single torch tensor in the shape of [1, C, H, W] or [C, H, W]
            prompt: the prompt to use, a string
            t: the time step to use, should be an int in the range of [0, 1000]
            up_ft_index: which upsampling block of the U-Net to extract feature, you can choose [0, 1, 2, 3]
            ensemble_size: the number of repeated images used in the batch to extract features
        Return:
            unet_ft: a torch tensor in the shape of [1, c, h, w]
        '''

        if img_tensor is not None:
            img_tensor = img_tensor.cuda()  # ensem, c, h, w
        if category in self.cat2prompt_embeds:
            prompt_embeds = self.cat2prompt_embeds[category]
        else:
            prompt_embeds = self.null_prompt_embeds
        prompt_embeds = prompt_embeds.repeat(ensemble_size, 1, 1).cuda()
        unet_ft_all = self.pipe(
            img_tensor=img_tensor,
            t=t,
            up_ft_indices=[up_ft_index],
            prompt_embeds=prompt_embeds)
        unet_ft = unet_ft_all['up_ft'][up_ft_index]  # ensem, c, h, w
        unet_ft = unet_ft.mean(0, keepdim=True)  # 1,c,h,w
        return unet_ft

In [None]:
from pathlib import Path
from tqdm import tqdm

spair71k_categories = [
    "aeroplane",
    "bicycle",
    "bird",
    "boat",
    "bottle",
    "bus",
    "car",
    "cat",
    "chair",
    "cow",
    "dog",
    "horse",
    "motorbike",
    "person",
    "pottedplant",
    "sheep",
    "train",
    "tvmonitor",
]
save_dir = Path(base_dir) / "data" / "features"
print("Saving features to:", save_dir)

with torch.no_grad():
    dift = SDFeaturizer4Eval(cat_list=spair71k_categories)
    with PreComputedFeaturemaps(save_dir, device=device) as pcm:
        for img_tensor, img_size, img_category, img_name in tqdm(
                test_dataset.iter_images(),
                total=test_dataset.num_images(),
                desc="Generating embeddings"
        ):
            img_tensor = img_tensor.to(device)  # [1,3,H,W]
            img_transformed = resize_image_sd(img_tensor)

            featmap = dift.forward(img_tensor=img_transformed, category=img_category)

            pcm.save_featuremaps(featmap, img_category, img_name)