aryadomain's picture
Add files using upload-large-folder tool
533920b verified
Raw
History Blame Contribute Delete
12.1 kB
import json
import os
import sys
from glob import glob
from io import BytesIO
from pathlib import Path
os.environ["HF_DATASETS_OFFLINE"] = "1"
os.environ["HF_METRICS_OFFLINE"] = "1"
os.environ["HF_MODULES_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
os.environ["DIFFUSERS_OFFLINE"] = "1"
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ.setdefault("ACCELERATE_MIXED_PRECISION", "bf16")
import torch
from torch.utils.data import DataLoader
from tqdm.auto import tqdm
from datasets import load_dataset
from PIL import Image
from torchvision import transforms
from transformers import AutoTokenizer
from safetensors.torch import load_file
from trainer.models.sana_preference_model import SanaPreferenceModel, SanaPreferenceModelConfig
# -----------------
# Config
# -----------------
PROJECT_ROOT = Path("/g/data/rr81/LPO/lrm/lrm_sana").resolve()
DEFAULT_CKPT_REL = (
"logs/lrm/reward_model/"
"step_sana_sana_sprint_0_6b_1024_variable-t_lr1e-5_step-8000_filter2_time951/"
"checkpoint-gstep100"
)
CKPT_DIR = Path(os.environ.get("SANA_CKPT_DIR", str(PROJECT_ROOT / DEFAULT_CKPT_REL))).resolve()
BASE_SANA_ID = "Efficient-Large-Model/Sana_Sprint_0.6B_1024px_diffusers"
DATASET_NAME = "pickapic-anonymous/pickapic_v1"
SPLIT = "test_unique"
BATCH_SIZE = 1
NUM_WORKERS = 2
MAX_BATCHES = None # Set e.g. 50 for quick checks
MAX_SEQUENCE_LENGTH = 300
MAX_SEQUENCE_LENGTH_2 = 300
IMAGE_SIZE = 1024
os.chdir(PROJECT_ROOT)
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("Project root:", PROJECT_ROOT)
print("Python:", sys.executable)
print("Torch:", torch.__version__)
print("CUDA available:", torch.cuda.is_available())
print("Device:", DEVICE)
print("Checkpoint dir:", CKPT_DIR)
# -----------------
# Load SANA model + local checkpoint
# -----------------
model_file = CKPT_DIR / "model.safetensors"
if not model_file.exists():
raise FileNotFoundError(f"Missing model checkpoint file: {model_file}")
model_cfg = SanaPreferenceModelConfig(
pretrained_model_name_or_path=BASE_SANA_ID,
pretrained_vae_name_or_path="",
model_profile="sana_sprint_0_6b_1024",
max_sequence_length=MAX_SEQUENCE_LENGTH,
max_sequence_length_2=MAX_SEQUENCE_LENGTH_2,
image_size=IMAGE_SIZE,
)
model = SanaPreferenceModel(model_cfg)
state = load_file(str(model_file))
missing, unexpected = model.load_state_dict(state, strict=False)
model.to(DEVICE).eval()
print("Model loaded from checkpoint.")
print("state_dict keys:", len(state))
print("missing keys:", len(missing))
if missing:
print("missing sample:", missing[:10])
print("unexpected keys:", len(unexpected))
if unexpected:
print("unexpected sample:", unexpected[:10])
print("logit_scale(exp):", float(model.logit_scale.exp().detach().cpu().item()))
# -----------------
# Eval helpers
# -----------------
def features2probs(model_obj, text_features, image_0_features, image_1_features):
image_0_scores = model_obj.logit_scale.exp() * torch.diag(
torch.einsum("bd,cd->bc", text_features, image_0_features)
)
image_1_scores = model_obj.logit_scale.exp() * torch.diag(
torch.einsum("bd,cd->bc", text_features, image_1_features)
)
scores = torch.stack([image_0_scores, image_1_scores], dim=-1)
probs = torch.softmax(scores, dim=-1)
return probs[:, 0], probs[:, 1]
def get_features(model_obj, input_ids, input_ids_2, pixels_0_values, pixels_1_values, timesteps):
all_pixel_values = torch.cat([pixels_0_values, pixels_1_values], dim=0)
timesteps = timesteps.reshape(-1, 2)
timesteps = torch.cat([timesteps[:, 0], timesteps[:, 1]], dim=0)
text_features, all_image_features = model_obj(
text_input_ids=input_ids,
text_input_ids_2=input_ids_2,
image_inputs=all_pixel_values,
time_cond=timesteps,
)
all_image_features = all_image_features / all_image_features.norm(dim=-1, keepdim=True)
text_features = text_features / text_features.norm(dim=-1, keepdim=True)
image_0_features, image_1_features = all_image_features.chunk(2, dim=0)
return image_0_features, image_1_features, text_features
def load_dataset_split_like_sana(dataset_name: str, split: str):
offline_mode = os.getenv("HF_HUB_OFFLINE", "0").strip().lower() in {"1", "true", "yes", "on"}
if not offline_mode:
return load_dataset(dataset_name, split=split)
if "/" not in dataset_name:
return load_dataset(dataset_name, split=split)
org, name = dataset_name.split("/", 1)
cache_candidates = []
for p in [
os.getenv("HF_HUB_CACHE"),
os.getenv("HUGGINGFACE_HUB_CACHE"),
(os.path.join(os.getenv("HF_HOME"), "hub") if os.getenv("HF_HOME") else None),
os.path.expanduser("~/.cache/huggingface/hub"),
"/scratch/rr81/ma5430/.cache/huggingface/hub",
]:
if p and p not in cache_candidates:
cache_candidates.append(p)
repo_cache_dirs = [
os.path.join(cache_root, f"datasets--{org}--{name}")
for cache_root in cache_candidates
if os.path.isdir(os.path.join(cache_root, f"datasets--{org}--{name}"))
]
for repo_cache_dir in repo_cache_dirs:
snapshot_dir = None
ref_main = os.path.join(repo_cache_dir, "refs", "main")
if os.path.isfile(ref_main):
revision = open(ref_main, "r", encoding="utf-8").read().strip()
candidate = os.path.join(repo_cache_dir, "snapshots", revision)
if os.path.isdir(candidate):
snapshot_dir = candidate
if snapshot_dir is None:
snapshots = sorted(glob(os.path.join(repo_cache_dir, "snapshots", "*")))
if snapshots:
snapshot_dir = snapshots[-1]
if snapshot_dir is None:
continue
data_dir = os.path.join(snapshot_dir, "data")
if not os.path.isdir(data_dir):
continue
selected_split = split
parquet_files = sorted(glob(os.path.join(data_dir, f"{selected_split}-*.parquet")))
if not parquet_files and split.startswith("validation"):
for alt_split in ("test_unique", "test"):
alt_files = sorted(glob(os.path.join(data_dir, f"{alt_split}-*.parquet")))
if alt_files:
selected_split = alt_split
parquet_files = alt_files
print(f"Offline cache missing split '{split}', falling back to '{selected_split}'")
break
if parquet_files:
print(
f"Loading cached offline split '{selected_split}' from {len(parquet_files)} parquet shards\n"
f"cache={repo_cache_dir}"
)
return load_dataset("parquet", data_files=parquet_files, split="train")
raise RuntimeError(
"Offline mode is enabled and cached parquet dataset was not found. "
f"Searched cache roots: {cache_candidates}. "
"Set HF_HUB_CACHE/HF_HOME to your predownloaded cache root or disable offline mode."
)
image_transform = transforms.Compose(
[
transforms.Resize((IMAGE_SIZE, IMAGE_SIZE), interpolation=transforms.InterpolationMode.BILINEAR),
transforms.CenterCrop(IMAGE_SIZE),
transforms.ToTensor(),
transforms.Normalize([0.5], [0.5]),
]
)
tokenizer = AutoTokenizer.from_pretrained(BASE_SANA_ID, subfolder="tokenizer")
try:
tokenizer_2 = AutoTokenizer.from_pretrained(BASE_SANA_ID, subfolder="tokenizer_2")
except Exception:
tokenizer_2 = None
def resolve_max_len(tok, requested):
m = getattr(tok, "model_max_length", None)
if m is None:
return requested
if m > 100000:
return requested
return min(requested, m)
max_len_1 = resolve_max_len(tokenizer, MAX_SEQUENCE_LENGTH)
max_len_2 = resolve_max_len(tokenizer_2, MAX_SEQUENCE_LENGTH_2) if tokenizer_2 is not None else max_len_1
raw_test = load_dataset_split_like_sana(DATASET_NAME, SPLIT)
raw_test = raw_test.filter(lambda x: x["has_label"])
def to_image(x):
if isinstance(x, dict):
x = x.get("bytes", x)
if isinstance(x, bytes):
x = Image.open(BytesIO(x))
if isinstance(x, str):
x = Image.open(x)
return x.convert("RGB")
def preprocess_example(example):
caption = example["caption"]
input_ids = tokenizer(
caption,
max_length=max_len_1,
padding="max_length",
truncation=True,
add_special_tokens=True,
return_tensors="pt",
).input_ids.squeeze(0)
if tokenizer_2 is not None:
input_ids_2 = tokenizer_2(
caption,
max_length=max_len_2,
padding="max_length",
truncation=True,
add_special_tokens=True,
return_tensors="pt",
).input_ids.squeeze(0)
else:
input_ids_2 = input_ids.clone()
pixel_0 = image_transform(to_image(example["jpg_0"]))
pixel_1 = image_transform(to_image(example["jpg_1"]))
# Non-train split behavior in SANA dataset pipeline.
timestep = torch.tensor([1, 1], dtype=torch.long)
return {
"input_ids": input_ids,
"input_ids_2": input_ids_2,
"pixel_values_0": pixel_0,
"pixel_values_1": pixel_1,
"label_0": torch.tensor(example["label_0"], dtype=torch.long),
"label_1": torch.tensor(example["label_1"], dtype=torch.long),
"timestep": timestep,
}
def collate_fn(batch):
return {
"input_ids": torch.stack([x["input_ids"] for x in batch], dim=0),
"input_ids_2": torch.stack([x["input_ids_2"] for x in batch], dim=0),
"pixel_values_0": torch.stack([x["pixel_values_0"] for x in batch], dim=0),
"pixel_values_1": torch.stack([x["pixel_values_1"] for x in batch], dim=0),
"label_0": torch.stack([x["label_0"] for x in batch], dim=0),
"label_1": torch.stack([x["label_1"] for x in batch], dim=0),
"timestep": torch.stack([x["timestep"] for x in batch], dim=0),
}
class EvalDataset(torch.utils.data.Dataset):
def __init__(self, hf_ds):
self.hf_ds = hf_ds
def __len__(self):
return len(self.hf_ds)
def __getitem__(self, idx):
return preprocess_example(self.hf_ds[idx])
eval_ds = EvalDataset(raw_test)
loader = DataLoader(
eval_ds,
shuffle=False,
batch_size=BATCH_SIZE,
num_workers=NUM_WORKERS,
collate_fn=collate_fn,
)
# -----------------
# Run evaluation
# -----------------
all_correct = []
num_batches = 0
with torch.no_grad():
for batch in tqdm(loader, desc=f"Evaluating {SPLIT}"):
num_batches += 1
for k, v in list(batch.items()):
if torch.is_tensor(v):
batch[k] = v.to(DEVICE)
image_0_features, image_1_features, text_features = get_features(
model,
batch["input_ids"],
batch["input_ids_2"],
batch["pixel_values_0"],
batch["pixel_values_1"],
batch["timestep"],
)
image_0_probs, image_1_probs = features2probs(model, text_features, image_0_features, image_1_features)
agree_on_0 = (image_0_probs > image_1_probs) * batch["label_0"]
agree_on_1 = (image_0_probs < image_1_probs) * batch["label_1"]
is_correct = (agree_on_0 + agree_on_1).detach().cpu()
all_correct.append(is_correct)
if MAX_BATCHES is not None and num_batches >= MAX_BATCHES:
break
correct_tensor = torch.cat(all_correct).float() if all_correct else torch.tensor([], dtype=torch.float32)
accuracy = float(correct_tensor.mean().item()) if correct_tensor.numel() > 0 else float("nan")
num_samples = int(correct_tensor.numel())
metrics = {
"split": SPLIT,
"accuracy": accuracy,
"num_samples": num_samples,
f"{SPLIT}_accuracy": accuracy,
f"{SPLIT}_num_samples": num_samples,
"logit_scale": float(model.logit_scale.exp().detach().cpu().item()),
"evaluated_batches": num_batches,
"checkpoint": str(CKPT_DIR),
}
print(json.dumps(metrics, indent=2))