"""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) # Backward-compatible alias used by eval scripts. LRMRewardModel = LRMRewardModelSana