| """SANA LRM reward model wrapper implemented fully inside this folder.""" |
|
|
| import os |
| from pathlib import Path |
|
|
| import torch |
| from safetensors.torch import load_file |
| from torch import nn |
| from transformers import AutoModel, AutoTokenizer |
|
|
| try: |
| from diffusers import AutoencoderDC, FlowMatchEulerDiscreteScheduler, SanaTransformer2DModel |
| except Exception: |
| from diffusers import FlowMatchEulerDiscreteScheduler |
| from diffusers.models import AutoencoderDC, SanaTransformer2DModel |
|
|
|
|
| PROFILE_TO_MODEL_ID = { |
| "sana_600m_512": "Efficient-Large-Model/Sana_600M_512px_diffusers", |
| "sana_1600m_512": "Efficient-Large-Model/Sana_1600M_512px_diffusers", |
| "sana_sprint_0_6b_1024": "Efficient-Large-Model/Sana_Sprint_0.6B_1024px_diffusers", |
| "sana_sprint_1_6b_1024": "Efficient-Large-Model/Sana_Sprint_1.6B_1024px_diffusers", |
| } |
|
|
| PROFILE_TO_IMAGE_SIZE = { |
| "sana_600m_512": 512, |
| "sana_1600m_512": 512, |
| "sana_sprint_0_6b_1024": 1024, |
| "sana_sprint_1_6b_1024": 1024, |
| } |
|
|
|
|
| def _offline_mode_enabled() -> bool: |
| return os.getenv("HF_HUB_OFFLINE", "0").strip().lower() in {"1", "true", "yes", "on"} |
|
|
|
|
| def _hf_pretrained_kwargs(): |
| kwargs = {"local_files_only": _offline_mode_enabled()} |
| cache_dir = os.getenv("HF_HUB_CACHE") or os.getenv("HUGGINGFACE_HUB_CACHE") |
| if cache_dir: |
| kwargs["cache_dir"] = cache_dir |
| return kwargs |
|
|
|
|
| def _module_device(module: nn.Module) -> torch.device: |
| return next(module.parameters()).device |
|
|
|
|
| def _extract_hidden_states(outputs) -> torch.Tensor: |
| if hasattr(outputs, "last_hidden_state"): |
| return outputs.last_hidden_state |
| if isinstance(outputs, (list, tuple)): |
| return outputs[0] |
| return outputs |
|
|
|
|
| def _build_attention_mask(token_ids: torch.Tensor, pad_token_id: int) -> torch.Tensor: |
| if pad_token_id is None: |
| return torch.ones_like(token_ids, dtype=torch.long) |
| return (token_ids != pad_token_id).long() |
|
|
|
|
| def _masked_mean(hidden_states: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor: |
| mask = attention_mask.to(device=hidden_states.device, dtype=hidden_states.dtype).unsqueeze(-1) |
| denom = mask.sum(dim=1).clamp(min=1.0) |
| return (hidden_states * mask).sum(dim=1) / denom |
|
|
|
|
| def _infer_hidden_size(model: nn.Module) -> int: |
| for attr in ("hidden_size", "d_model", "projection_dim"): |
| value = getattr(model.config, attr, None) |
| if isinstance(value, int): |
| return value |
| raise ValueError("Could not infer hidden size from text encoder config") |
|
|
|
|
| class LRMRewardModelSana(nn.Module): |
| """Latent reward model wrapper for SANA.""" |
|
|
| def __init__( |
| self, |
| pretrained_model_name_or_path="Efficient-Large-Model/Sana_600M_512px_diffusers", |
| lrm_model_path=None, |
| model_profile="sana_600m_512", |
| guidance_scale=4.5, |
| max_sequence_length=300, |
| max_sequence_length_2=300, |
| projection_dim=1024, |
| logit_scale_init_value=2.6592, |
| device="cuda", |
| ): |
| super().__init__() |
|
|
| if model_profile not in PROFILE_TO_MODEL_ID: |
| raise ValueError( |
| f"Unknown model_profile={model_profile}. " |
| f"Available: {', '.join(sorted(PROFILE_TO_MODEL_ID.keys()))}" |
| ) |
|
|
| self.device = device |
| self.guidance_scale = guidance_scale |
| self.max_sequence_length = max_sequence_length |
| self.max_sequence_length_2 = max_sequence_length_2 |
| self.image_size = PROFILE_TO_IMAGE_SIZE[model_profile] |
|
|
| model_id = pretrained_model_name_or_path or PROFILE_TO_MODEL_ID[model_profile] |
| print(f"Loading SANA base reward backbone from {model_id}...") |
|
|
| pretrained_kwargs = _hf_pretrained_kwargs() |
|
|
| precision = os.getenv("ACCELERATE_MIXED_PRECISION", "").strip().lower() |
| module_load_kwargs = dict(pretrained_kwargs) |
| if precision == "bf16": |
| module_load_kwargs["torch_dtype"] = torch.bfloat16 |
| elif precision == "fp16": |
| module_load_kwargs["torch_dtype"] = torch.float16 |
|
|
| self.vae = AutoencoderDC.from_pretrained( |
| model_id, |
| subfolder="vae", |
| **module_load_kwargs, |
| ) |
| self.scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( |
| model_id, |
| subfolder="scheduler", |
| **pretrained_kwargs, |
| ) |
| self.transformer = SanaTransformer2DModel.from_pretrained( |
| model_id, |
| subfolder="transformer", |
| **module_load_kwargs, |
| ) |
|
|
| self.tokenizer = AutoTokenizer.from_pretrained( |
| model_id, |
| subfolder="tokenizer", |
| **pretrained_kwargs, |
| ) |
| self.text_encoder = AutoModel.from_pretrained( |
| model_id, |
| subfolder="text_encoder", |
| **module_load_kwargs, |
| ) |
|
|
| self.tokenizer_2 = None |
| self.text_encoder_2 = None |
| try: |
| self.tokenizer_2 = AutoTokenizer.from_pretrained( |
| model_id, |
| subfolder="tokenizer_2", |
| **pretrained_kwargs, |
| ) |
| self.text_encoder_2 = AutoModel.from_pretrained( |
| model_id, |
| subfolder="text_encoder_2", |
| **module_load_kwargs, |
| ) |
| except Exception: |
| self.tokenizer_2 = None |
| self.text_encoder_2 = None |
|
|
| self.text_pad_token_id = self.tokenizer.pad_token_id if self.tokenizer.pad_token_id is not None else 0 |
| self.text_pad_token_id_2 = ( |
| self.tokenizer_2.pad_token_id |
| if self.tokenizer_2 is not None and self.tokenizer_2.pad_token_id is not None |
| else self.text_pad_token_id |
| ) |
|
|
| text_in_dim = _infer_hidden_size(self.text_encoder) |
| image_in_dim = self.transformer.config.in_channels |
|
|
| self.text_projection = nn.Linear(text_in_dim, projection_dim, bias=False) |
| self.visual_projection = nn.Linear(image_in_dim, projection_dim, bias=False) |
| nn.init.normal_(self.text_projection.weight, std=0.02) |
| nn.init.normal_(self.visual_projection.weight, std=0.02) |
|
|
| self.text_projection_2 = None |
| if self.text_encoder_2 is not None: |
| text_in_dim_2 = _infer_hidden_size(self.text_encoder_2) |
| self.text_projection_2 = nn.Linear(text_in_dim_2, projection_dim, bias=False) |
| nn.init.normal_(self.text_projection_2.weight, std=0.02) |
|
|
| self.logit_scale = nn.Parameter(torch.ones([]) * logit_scale_init_value) |
|
|
| self.to(device) |
| self.eval() |
|
|
| if lrm_model_path: |
| self.load_lrm_weights(lrm_model_path) |
|
|
| print("✓ SANA LRM Reward Model initialized successfully!") |
|
|
| def load_lrm_weights(self, model_path): |
| ckpt_path = Path(model_path) |
| model_file = ckpt_path / "model.safetensors" if ckpt_path.is_dir() else ckpt_path |
|
|
| if not model_file.exists(): |
| raise FileNotFoundError( |
| f"Missing SANA reward checkpoint file: {model_file}. " |
| "Provide checkpoint directory or model.safetensors path." |
| ) |
|
|
| print(f"Loading SANA reward checkpoint from {model_file}...") |
| state = load_file(str(model_file)) |
| missing, unexpected = self.load_state_dict(state, strict=False) |
| self.to(self.device) |
| self.eval() |
|
|
| print(f"✓ Loaded checkpoint keys: {len(state)}") |
| print(f"✓ Missing keys: {len(missing)} | Unexpected keys: {len(unexpected)}") |
|
|
| def encode_prompt(self, prompt): |
| text_input_ids = self.tokenizer( |
| prompt, |
| padding="max_length", |
| max_length=self.max_sequence_length, |
| truncation=True, |
| return_tensors="pt", |
| ).input_ids |
|
|
| if self.tokenizer_2 is not None: |
| text_input_ids_2 = self.tokenizer_2( |
| prompt, |
| padding="max_length", |
| max_length=self.max_sequence_length_2, |
| truncation=True, |
| return_tensors="pt", |
| ).input_ids |
| else: |
| text_input_ids_2 = text_input_ids.clone() |
|
|
| return text_input_ids, text_input_ids_2 |
|
|
| def _encode_prompts(self, prompt, batch_size): |
| if isinstance(prompt, str): |
| prompt = [prompt] |
|
|
| if len(prompt) == 1 and batch_size > 1: |
| prompt = prompt * batch_size |
| elif len(prompt) != batch_size: |
| raise ValueError(f"Prompt batch size mismatch: got {len(prompt)} prompts for batch_size={batch_size}") |
|
|
| text_input_ids, text_input_ids_2 = self.encode_prompt(prompt) |
|
|
| text_input_ids = text_input_ids.to(_module_device(self.text_encoder)) |
| attention_mask = _build_attention_mask(text_input_ids, self.text_pad_token_id) |
|
|
| text_outputs = self.text_encoder(input_ids=text_input_ids, attention_mask=attention_mask) |
| prompt_embeds = _extract_hidden_states(text_outputs) |
| pooled_prompt = _masked_mean(prompt_embeds, attention_mask) |
| text_features = self.text_projection(pooled_prompt) |
|
|
| if self.text_encoder_2 is not None: |
| text_input_ids_2 = text_input_ids_2.to(_module_device(self.text_encoder_2)) |
| attention_mask_2 = _build_attention_mask(text_input_ids_2, self.text_pad_token_id_2) |
| text_outputs_2 = self.text_encoder_2(input_ids=text_input_ids_2, attention_mask=attention_mask_2) |
| prompt_embeds_2 = _extract_hidden_states(text_outputs_2) |
| pooled_prompt_2 = _masked_mean(prompt_embeds_2, attention_mask_2) |
| if self.text_projection_2 is not None: |
| text_features_2 = self.text_projection_2(pooled_prompt_2) |
| text_features = (text_features + text_features_2) / 2.0 |
|
|
| prompt_embeds = prompt_embeds.to(dtype=self.transformer.dtype) |
| attention_mask = attention_mask.to(device=prompt_embeds.device) |
| return prompt_embeds, attention_mask, text_features |
|
|
| @staticmethod |
| def _prepare_timesteps(timesteps, batch_size, device, dtype): |
| if isinstance(timesteps, (int, float)): |
| ts = torch.full((batch_size,), float(timesteps), device=device, dtype=dtype) |
| elif torch.is_tensor(timesteps): |
| ts = timesteps.to(device=device, dtype=dtype).flatten() |
| if ts.numel() == 1 and batch_size > 1: |
| ts = ts.repeat(batch_size) |
| elif ts.numel() != batch_size: |
| ts = ts[:1].repeat(batch_size) |
| else: |
| ts = torch.tensor([float(timesteps)] * batch_size, device=device, dtype=dtype) |
| return ts |
|
|
| def get_reward_score(self, noisy_latents, prompt, timesteps, enable_grad=False, return_logits=False): |
| """Compute reward score for pre-noised SANA latents at timestep t.""" |
|
|
| def _compute(): |
| latents = noisy_latents.to(self.device) |
| batch_size = latents.shape[0] |
|
|
| encoder_hidden_states, encoder_attention_mask, text_features = self._encode_prompts(prompt, batch_size) |
|
|
| transformer_dtype = self.transformer.dtype |
| hidden_states = latents.to(dtype=transformer_dtype) |
| time_cond = self._prepare_timesteps( |
| timesteps, |
| batch_size=batch_size, |
| device=latents.device, |
| dtype=transformer_dtype, |
| ) |
|
|
| transformer_kwargs = { |
| "hidden_states": hidden_states, |
| "encoder_hidden_states": encoder_hidden_states.to(device=latents.device, dtype=transformer_dtype), |
| "encoder_attention_mask": encoder_attention_mask.to(device=latents.device), |
| "timestep": time_cond, |
| "return_dict": False, |
| } |
|
|
| if getattr(self.transformer.config, "guidance_embeds", False): |
| guidance = torch.full( |
| (batch_size,), |
| self.guidance_scale, |
| device=latents.device, |
| dtype=transformer_dtype, |
| ) |
| transformer_kwargs["guidance"] = guidance |
|
|
| model_pred = self.transformer(**transformer_kwargs)[0] |
| if model_pred.shape[1] == 2 * hidden_states.shape[1]: |
| model_pred = model_pred.chunk(2, dim=1)[0] |
|
|
| pooled_tokens = model_pred.flatten(2).mean(dim=-1) |
| image_features = self.visual_projection(pooled_tokens) |
|
|
| text_features = text_features.to(device=image_features.device, dtype=image_features.dtype) |
| image_features = image_features / image_features.norm(dim=-1, keepdim=True).clamp(min=1e-8) |
| text_features = text_features / text_features.norm(dim=-1, keepdim=True).clamp(min=1e-8) |
|
|
| logits = self.logit_scale.exp() * (text_features * image_features).sum(dim=-1) |
| logits = torch.clamp(logits, min=-30.0, max=30.0) |
|
|
| if return_logits: |
| return logits |
|
|
| scores = torch.sigmoid(logits) |
| return scores |
|
|
| if enable_grad: |
| return _compute() |
| with torch.no_grad(): |
| return _compute() |
|
|
| def forward(self, noisy_latents, prompt, timesteps): |
| return self.get_reward_score(noisy_latents, prompt, timesteps) |
|
|
|
|
| |
| LRMRewardModel = LRMRewardModelSana |
|
|