Views
No views yet
fid/(fid+ε) 归一化)num_sampling_steps=1,共 62500 stepmodel,格式与官方 JiT-B_FD-SIM.pth 一致,可直接 --load_fromcheckpoints/
JiT-B_FD-SIM_bs1024_step12500.pth # ema=edm_500
JiT-B_FD-SIM_bs1024_step25000.pth # ema=edm_500
JiT-B_FD-SIM_bs1024_step37500.pth # ema=edm_500
JiT-B_FD-SIM_bs1024_step50000.pth # ema=edm_500
JiT-B_FD-SIM_bs1024_step62500.pth # ema=edm_250 (final)model / step / samples_seen / ema_label。| step | EMA | FID(ADM) ↓ | FDr⁶ ↓ |
|---|---|---|---|
| 12.5k | edm_500 | 3.298 | 7.609 |
| 25k | edm_500 | 1.987 | 6.103 |
| 37.5k | edm_500 | 1.311 | 5.388 |
| 50k | edm_500 | 1.027 | 5.237 |
| 62.5k | edm_250 | 0.944 | 5.063 |
1# 中国大陆建议走镜像
2export HF_ENDPOINT=https://hf-mirror.com
3
4pip install -U huggingface_hub
5hf download shy0423/JiT-B-FD-SIM-bs1024 \
6 --local-dir . \
7 --include "checkpoints/*.pth"FD-Loss)。核心是 JiTDenoiser.generate:从噪声 t=1 欧拉一步到 t=0。1export HF_ENDPOINT=https://hf-mirror.com # 如需下依赖权重
2
3CKPT=checkpoints/JiT-B_FD-SIM_bs1024_step62500.pth
4
5torchrun --nproc_per_node=8 eval_all_fds.py \
6 --model JiT_B \
7 --rope_2d --learned_pe --legacy_time_convention --ema_type edm \
8 --cfg 1.0 --cfg_list 1.0 \
9 --interval_min 0.1 --interval_max 1.0 \
10 --num_sampling_steps 1 \
11 --eval_ema_labels online \
12 --eval_bsz 128 --num_images 50000 \
13 --load_from "$CKPT" \
14 --output_dir work_dirs/eval_simfull1024 \
15 --project eval --exp_name JiT-B-FD-SIM-bs1024-625001PRESET=JiT_B \
2CKPT_PATH=checkpoints/JiT-B_FD-SIM_bs1024_step62500.pth \
3CFG_OVERRIDE=1.0 \
4bash scripts/evaluate_released_ckpt.sh发布权重已是 best-EMA,因此--eval_ema_labels online即可复现上表 FID。本 run 的 1-step 评测 cfg 固定为 1.0。
1import argparse
2import torch
3from torchvision.utils import save_image
4
5from utils.builders import create_generation_model
6from utils.checkpoint_util import ckpt_resume
7from utils.sampling_util import generate_images
8
9args = argparse.Namespace(
10 model="JiT_B",
11 img_size=256,
12 num_classes=1000,
13 label_drop_prob=0.1,
14 attn_dropout=0.0,
15 proj_dropout=0.0,
16 P_mean=-0.8,
17 P_std=0.8,
18 t_eps=0.05,
19 rope_2d=True,
20 learned_pe=True,
21 legacy_time_convention=True,
22 ema_type="edm",
23 ema_rates=None,
24 ema_halflife_kimg=[250, 500, 1000, 2000],
25 load_from="checkpoints/JiT-B_FD-SIM_bs1024_step62500.pth",
26 resume_from=None,
27 auto_resume=False,
28 num_sampling_steps=1, # 一步生成
29 sampling_method="euler",
30 cfg=1.0,
31 interval_min=0.1,
32 interval_max=1.0,
33 same_noise=False,
34 enable_amp=True,
35 amp_dtype=torch.bfloat16,
36 # builders 里其它默认字段按 main_fd / eval_all_fds 的 parser 补齐即可
37)
38
39model, ema_model = create_generation_model(args)
40ckpt_resume(args, model, optimizer=None, model_ema=ema_model)
41
42labels = torch.randint(0, 1000, (8,), device="cuda")
43images = generate_images(args, model, labels=labels, cfg=1.0) # [0,1]
44save_image(images, "samples.png", nrow=4)generate_images → JiTDenoiser.generate:z ~ N(0, I)(noise_scale=1)t: 1 → 0,num_sampling_steps=1 时只做一次 Euler:z ← z + (0-1)·v_θ(z,t=1,y)[-1,1],再线性映射到 [0,1]sample_images_with_grad)把 cfg 固定为 1.0;评测/采样走 generate,本发布建议 cfg=1.0。1@article{yang2026fdloss,
2 title={Representation Fr\'echet Loss for Visual Generation},
3 author={Yang, Jiawei and Geng, Zhengyang and Ju, Xuan and Tian, Yonglong and Wang, Yue},
4 journal={arXiv:2604.28190},
5 year={2026}
6}