aryadomain's picture
Add files using upload-large-folder tool
b871133 verified
Raw
History Blame Contribute Delete
13.3 kB
"""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