mirror of
https://github.com/hpcaitech/Open-Sora.git
synced 2026-05-21 11:59:01 +02:00
[feat] add score for inference
This commit is contained in:
parent
3a83a40acc
commit
6094f5fe92
|
|
@ -48,3 +48,6 @@ scheduler = dict(
|
|||
num_sampling_steps=30,
|
||||
cfg_scale=7.0,
|
||||
)
|
||||
|
||||
aes = None
|
||||
flow = None
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ======================================================
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
Loading…
Reference in a new issue