[feat] add score for inference

This commit is contained in:
zhengzangw 2024-06-09 12:01:26 +00:00
parent 3a83a40acc
commit 6094f5fe92
4 changed files with 15 additions and 1 deletions

View file

@ -48,3 +48,6 @@ scheduler = dict(
num_sampling_steps=30,
cfg_scale=7.0,
)
aes = None
flow = None

View file

@ -41,9 +41,10 @@ pretrained_models = {
def reparameter(ckpt, name=None, model=None):
model_name = name
name = os.path.basename(name)
if not dist.is_initialized() or dist.get_rank() == 0:
get_logger().info("loading pretrained model: %s", name)
get_logger().info("loading pretrained model: %s", model_name)
if name in ["DiT-XL-2-512x512.pt", "DiT-XL-2-256x256.pt"]:
ckpt["x_embedder.proj.weight"] = ckpt["x_embedder.proj.weight"].unsqueeze(2)
del ckpt["pos_embed"]

View file

@ -62,6 +62,8 @@ def parse_args(training=False):
parser.add_argument("--condition-frame-length", default=None, type=int, help="condition frame length")
parser.add_argument("--reference-path", default=None, type=str, nargs="+", help="reference path")
parser.add_argument("--mask-strategy", default=None, type=str, nargs="+", help="mask strategy")
parser.add_argument("--aes", default=None, type=float, help="aesthetic score")
parser.add_argument("--flow", default=None, type=float, help="flow score")
# ======================================================
# Training
# ======================================================

View file

@ -110,6 +110,14 @@ def main():
if prompts is None:
assert cfg.get("prompt_path", None) is not None, "Prompt or prompt_path must be provided"
prompts = load_prompts(cfg.prompt_path, start_idx, cfg.get("end_index", None))
score_prompts = []
if cfg.get("aes", None) is not None:
score_prompts.append(f"{prompt} aesthetic score: {cfg.aes:.1f}" for prompt in prompts)
if cfg.get("flow", None) is not None:
score_prompts.append(f"{prompt} motion score: {cfg.flow:.1f}" for prompt in prompts)
if len(score_prompts) > 0:
score_text = ", ".join(score_prompts)
prompts = [f"{prompt} [{score_text}]" for prompt in prompts]
# == prepare reference ==
reference_path = cfg.get("reference_path", [""] * len(prompts))