Skip to content

Commit 625c740

Browse files
committed
feat(ml):debug custom jit
1 parent ab30bf0 commit 625c740

4 files changed

Lines changed: 13 additions & 3 deletions

File tree

models/diffusion_networks.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -289,6 +289,16 @@ def define_G(
289289

290290
elif G_netG == "vit_vid":
291291
variant = getattr(opt, "G_vit_variant", "")
292+
if variant and variant not in JiTVid_VARIANT_CONFIGS:
293+
if variant.startswith("JiT-"):
294+
alias = f"JiTVid-{variant[len('JiT-'):]}"
295+
if alias in JiTVid_VARIANT_CONFIGS:
296+
variant = alias
297+
if variant not in JiTVid_VARIANT_CONFIGS:
298+
raise ValueError(
299+
f"Unknown G_vit_variant '{variant}'. "
300+
f"Valid: {sorted(JiTVid_VARIANT_CONFIGS.keys())}"
301+
)
292302
base = JiTVid_VARIANT_CONFIGS.get(variant, {})
293303
cfg = {
294304
"depth": getattr(opt, "G_vit_depth", base.get("depth", 12)),

options/common_options.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -428,7 +428,7 @@ def initialize(self, parser):
428428
"--G_vit_variant",
429429
type=str,
430430
default="JiT-B/16",
431-
help="Selects the ViT backbone when --G_netG vit",
431+
help="Selects the ViT backbone when --G_netG vit (use JiT-*) or vit_vid (use JiTVid-*)",
432432
)
433433
parser.add_argument(
434434
"--G_vit_disable_bottleneck",

scripts/gen_single_image_diffusion.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -81,7 +81,7 @@ def load_model(
8181
if opt.model_type in ["cm", "cm_gan", "sc", "b2b"]:
8282
opt.alg_palette_sampling_method = sampling_method
8383
opt.alg_diffusion_cond_embed_dim = 256
84-
model = diffusion_networks.define_G(**vars(opt))
84+
model = diffusion_networks.define_G(opt=opt, **vars(opt))
8585
model.eval()
8686

8787
# handle old models

scripts/gen_vid_autoregressive_diffusion_forward_NoCanny_offline.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,7 @@ def load_model(
101101
if opt.model_type in ["cm", "cm_gan", "b2b"]:
102102
opt.alg_palette_sampling_method = sampling_method
103103
opt.alg_diffusion_cond_embed_dim = 256
104-
model = diffusion_networks.define_G(**vars(opt))
104+
model = diffusion_networks.define_G(opt=opt, **vars(opt))
105105
model.eval()
106106

107107
# handle old models

0 commit comments

Comments
 (0)