Mettle / verify_against_flat.py
kykatro's picture
Add Mettle RC2 weights, loader, and model card
0e83a2b verified
Raw
History Blame Contribute Delete
3.63 kB
"""Numerically compare the release model with a flat Mettle checkpoint."""
import argparse
import gc
from pathlib import Path
import timm
import torch
from configuration_mettle import MettleConfig
from modeling_mettle import MettleModel, MettleRefineHead
from package_model import convert_state_dict, validate_shapes
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("checkpoint", help="flat Mettle checkpoint")
parser.add_argument(
"--device",
default="cuda" if torch.cuda.is_available() else "cpu",
)
parser.add_argument("--atol", type=float, default=1e-5)
return parser.parse_args()
def load_config():
return MettleConfig.from_json_file(
str(Path(__file__).with_name("config.json"))
)
def reference_outputs(config, state, pixel_values):
backbone = timm.create_model(
config.backbone_name,
pretrained=False,
num_classes=0,
init_values=1e-5,
dynamic_img_size=False,
img_size=config.image_size,
)
backbone_state = {
key: value for key, value in state.items()
if not key.startswith("head.")
}
backbone.load_state_dict(backbone_state, strict=True)
head = MettleRefineHead(
dim=config.hidden_size,
num_atoms=config.head_num_atoms,
rank=config.head_rank,
hidden_size=config.head_hidden_size,
)
head_state = {
key.removeprefix("head."): value for key, value in state.items()
if key.startswith("head.")
}
head.load_state_dict(head_state, strict=True)
backbone.to(pixel_values.device).eval()
head.to(pixel_values.device).eval()
with torch.inference_mode():
tokens = backbone.forward_features(pixel_values)
cls = head(tokens[:, 0].to(torch.float32))
mean_patch = tokens[:, int(backbone.num_prefix_tokens):].float().mean(1)
cls_mean = torch.cat((cls, mean_patch), dim=-1)
return cls.cpu(), cls_mean.cpu()
def release_outputs(config, state, pixel_values):
model = MettleModel(config)
model.load_state_dict(convert_state_dict(state), strict=True)
model.to(pixel_values.device).eval()
with torch.inference_mode():
cls = model.encode(pixel_values, "cls")
cls_mean = model.encode(pixel_values, "cls_mean")
return cls.cpu(), cls_mean.cpu()
def main():
args = parse_args()
config = load_config()
state = torch.load(args.checkpoint, map_location="cpu", weights_only=True)
validate_shapes(state)
generator = torch.Generator().manual_seed(20260730)
pixel_values = torch.randn(
1,
3,
config.image_size,
config.image_size,
generator=generator,
).to(args.device)
reference_cls, reference_cls_mean = reference_outputs(
config,
state,
pixel_values,
)
if args.device.startswith("cuda"):
torch.cuda.empty_cache()
gc.collect()
release_cls, release_cls_mean = release_outputs(
config,
state,
pixel_values,
)
cls_delta = (reference_cls - release_cls).abs().max().item()
cls_mean_delta = (
reference_cls_mean - release_cls_mean
).abs().max().item()
print(f"CLS max absolute delta: {cls_delta:.3e}")
print(f"CLS+mean max absolute delta: {cls_mean_delta:.3e}")
if cls_delta > args.atol or cls_mean_delta > args.atol:
raise SystemExit(
"FAILED: release outputs differ from the flat-checkpoint reference"
)
print("PASS: both feature views reproduce the flat-checkpoint reference")
if __name__ == "__main__":
main()