mirror of
https://github.com/storytold/cloud-worker.git
synced 2026-10-09 00:09:43 +00:00
bench: --ref-downscale option — pre-shrink refs before the node (24GB OOM mitigation)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
+24
-11
@@ -85,7 +85,8 @@ def snap_length(seconds: float) -> int:
|
||||
|
||||
def build_graph(task, model_file, prompt, width, height, length, steps, seed,
|
||||
sampler="res_multistep", scheduler="simple", image=BENCH_IMAGE,
|
||||
ref_image_size="match", ref_count=1, filename_prefix="bench/run"):
|
||||
ref_image_size="match", ref_count=1, ref_downscale=0.0,
|
||||
filename_prefix="bench/run"):
|
||||
g = {
|
||||
"1": {"class_type": "UNETLoader",
|
||||
"inputs": {"unet_name": model_file, "weight_dtype": "default"}},
|
||||
@@ -122,15 +123,25 @@ def build_graph(task, model_file, prompt, width, height, length, steps, seed,
|
||||
**common}}
|
||||
elif task == "ref2v":
|
||||
ref_inputs = {}
|
||||
|
||||
def ref_source(nid_load, fname):
|
||||
g[nid_load] = {"class_type": "LoadImage", "inputs": {"image": fname}}
|
||||
if not ref_downscale:
|
||||
return [nid_load, 0]
|
||||
nid_scale = str(int(nid_load) + 40)
|
||||
g[nid_scale] = {"class_type": "ImageScaleToTotalPixels",
|
||||
"inputs": {"image": [nid_load, 0],
|
||||
"upscale_method": "lanczos",
|
||||
"megapixels": ref_downscale,
|
||||
"resolution_steps": 1}}
|
||||
return [nid_scale, 0]
|
||||
|
||||
if ref_count <= 1:
|
||||
g["10"] = {"class_type": "LoadImage", "inputs": {"image": image}}
|
||||
ref_inputs["ref_images.ref_image_0"] = ["10", 0]
|
||||
ref_inputs["ref_images.ref_image_0"] = ref_source("10", image)
|
||||
else:
|
||||
for i in range(ref_count):
|
||||
nid = str(20 + i)
|
||||
g[nid] = {"class_type": "LoadImage",
|
||||
"inputs": {"image": BENCH_REF_SET[i % len(BENCH_REF_SET)]}}
|
||||
ref_inputs[f"ref_images.ref_image_{i}"] = [nid, 0]
|
||||
ref_inputs[f"ref_images.ref_image_{i}"] = ref_source(
|
||||
str(20 + i), BENCH_REF_SET[i % len(BENCH_REF_SET)])
|
||||
g["5"] = {"class_type": "MiniMaxH3ReferenceToVideo",
|
||||
"inputs": {"clip": ["2", 0], "vae": ["3", 0], "audio_vae": ["4", 0],
|
||||
"ref_image_size": ref_image_size,
|
||||
@@ -173,7 +184,7 @@ def run_one(api, cfg, poll=2.0, timeout=3600):
|
||||
graph = build_graph(**{k: v for k, v in cfg.items() if k in (
|
||||
"task", "model_file", "prompt", "width", "height", "length", "steps", "seed",
|
||||
"sampler", "scheduler", "image", "ref_image_size", "ref_count",
|
||||
"filename_prefix")})
|
||||
"ref_downscale", "filename_prefix")})
|
||||
peak = {"vram": 0}
|
||||
stop = threading.Event()
|
||||
|
||||
@@ -263,7 +274,7 @@ def default_prompt(task):
|
||||
|
||||
def make_cfg(task, model_file, width, height, seconds, steps, seed, label, warm,
|
||||
sampler="res_multistep", scheduler="simple", ref_image_size="match",
|
||||
ref_count=1):
|
||||
ref_count=1, ref_downscale=0.0):
|
||||
prompt = ref2v_prompt(ref_count) if task == "ref2v" else default_prompt(task)
|
||||
return {
|
||||
"label": label, "task": task, "model_file": model_file,
|
||||
@@ -271,7 +282,7 @@ def make_cfg(task, model_file, width, height, seconds, steps, seed, label, warm,
|
||||
"seconds": seconds, "length": snap_length(seconds), "steps": steps,
|
||||
"seed": seed, "sampler": sampler, "scheduler": scheduler,
|
||||
"image": BENCH_IMAGE, "ref_image_size": ref_image_size,
|
||||
"ref_count": ref_count, "warm": warm,
|
||||
"ref_count": ref_count, "ref_downscale": ref_downscale, "warm": warm,
|
||||
"filename_prefix": f"bench/{label}",
|
||||
}
|
||||
|
||||
@@ -395,6 +406,8 @@ def main():
|
||||
ap.add_argument("--ref-image-size", default="match", choices=["match", "max"])
|
||||
ap.add_argument("--ref-count", type=int, default=1,
|
||||
help="number of reference images for ref2v (1-8)")
|
||||
ap.add_argument("--ref-downscale", type=float, default=0.0,
|
||||
help="pre-downscale refs to this many megapixels before the node (0=off)")
|
||||
ap.add_argument("--label", default=None)
|
||||
ap.add_argument("--prefix", default="", help="label prefix for suites, e.g. prunedint8_")
|
||||
ap.add_argument("--warm", action="store_true", help="mark run as warm (model preloaded)")
|
||||
@@ -423,7 +436,7 @@ def main():
|
||||
cfgs = [make_cfg(args.task, model, args.width, args.height, args.seconds,
|
||||
args.steps, args.seed, label, args.warm,
|
||||
args.sampler, args.scheduler, args.ref_image_size,
|
||||
args.ref_count)]
|
||||
args.ref_count, args.ref_downscale)]
|
||||
else:
|
||||
ap.error("need --suite or --task")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user