Skip to content

Commit a34806c

Browse files
authored
reuse: support cross-worker and prefix-segment reuse (#1428)
- Persist encoder outputs and request metadata in shared disk caches for Wan, Qwen-Image, and InfiniteTalk. - Add InfiniteTalk prefix-segment reuse with cached motion boundaries and previous-result prefix merging. - Publish reuse caches only after successful generation and reject reuse for streaming or tensor-returning requests. - Add reuse capability checks, request schema wiring, cache configuration, and a request example.
1 parent 65fe600 commit a34806c

16 files changed

Lines changed: 414 additions & 99 deletions

File tree

configs/infinitetalk/infinitetalk_480p_multi_reuse.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,5 +32,6 @@
3232
"norm_output_audio": true,
3333
"mxfp8_fuse_enable": false,
3434
"use_timestep_transform": true,
35-
"enable_reuse": true
35+
"enable_reuse": true,
36+
"reuse_cache_path": "save_results/reuse_cache/infinitetalk_480p_multi_reuse"
3637
}

configs/qwen_image/qwen_image_i2i_2511_reuse.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,5 +10,6 @@
1010
"CONDITION_IMAGE_SIZE": 147456,
1111
"USE_IMAGE_ID_IN_PROMPT": true,
1212
"use_compile": true,
13-
"enable_reuse": true
13+
"enable_reuse": true,
14+
"reuse_cache_path": "save_results/reuse_cache/qwen_image_i2i_2511"
1415
}

configs/qwen_image/qwen_image_t2i_2512_reuse.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,5 +8,6 @@
88
"enable_cfg": true,
99
"sample_guide_scale": 4.0,
1010
"use_compile": true,
11-
"enable_reuse": true
11+
"enable_reuse": true,
12+
"reuse_cache_path": "save_results/reuse_cache/qwen_image_t2i_2512"
1213
}

configs/wan/wan_i2v_reuse.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,5 +10,6 @@
1010
"sample_shift": 3,
1111
"enable_cfg": true,
1212
"cpu_offload": false,
13-
"enable_reuse": true
13+
"enable_reuse": true,
14+
"reuse_cache_path": "save_results/reuse_cache/wan_i2v"
1415
}

configs/wan/wan_t2v_reuse.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,5 +11,6 @@
1111
"sample_shift": 8,
1212
"enable_cfg": true,
1313
"cpu_offload": false,
14-
"enable_reuse": true
14+
"enable_reuse": true,
15+
"reuse_cache_path": "save_results/reuse_cache/wan_t2v"
1516
}

configs/wan22/wan_moe_i2v_reuse.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,5 +17,6 @@
1717
"offload_granularity": "model",
1818
"boundary": 0.900,
1919
"use_image_encoder": false,
20-
"enable_reuse": true
20+
"enable_reuse": true,
21+
"reuse_cache_path": "save_results/reuse_cache/wan22_moe_i2v"
2122
}

configs/wan22/wan_moe_t2v_reuse.json

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,5 +18,6 @@
1818
"t5_cpu_offload": false,
1919
"vae_cpu_offload": false,
2020
"boundary": 0.875,
21-
"enable_reuse": true
21+
"enable_reuse": true,
22+
"reuse_cache_path": "save_results/reuse_cache/wan22_moe_t2v"
2223
}

lightx2v/models/runners/base_runner.py

Lines changed: 14 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ def __init__(self, config):
2222
self.input_info = None
2323
self.enable_reuse = config.get("enable_reuse", False)
2424
self.reuse = False
25-
self._reuse_cache = None
25+
self.reuse_prefix_segments = 0
2626
self._gc_frozen = False # one-shot guard for _maybe_freeze_gc()
2727
self._init_modules_depth = 0
2828
self._warmup_done = False
@@ -68,24 +68,25 @@ def warmup(self):
6868
if self.config.get("warmup", False):
6969
raise NotImplementedError(f"Warmup is not supported for {type(self).__name__}")
7070

71-
def set_reuse(self, reuse):
71+
def set_reuse(self, reuse, reuse_prefix_segments=0):
7272
if reuse and not self.enable_reuse:
7373
raise ValueError(f"This {type(self).__name__} service does not enable reuse")
7474
if reuse and self.config.get("disagg_mode"):
7575
raise NotImplementedError(f"{type(self).__name__} reuse does not support disaggregated inference")
76+
if reuse_prefix_segments and not reuse:
77+
raise ValueError("reuse_prefix_segments requires reuse=true")
78+
if reuse:
79+
self.check_reuse_support()
80+
if reuse_prefix_segments:
81+
self.check_segment_reuse_support()
7682
self.reuse = reuse
83+
self.reuse_prefix_segments = reuse_prefix_segments
7784

78-
def _reuse_key(self):
79-
raise NotImplementedError
80-
81-
def _get_reused_inputs(self):
82-
reuse_cache = self._reuse_cache
83-
if reuse_cache is None:
84-
raise RuntimeError("No previous successful request is available for reuse")
85-
if self._reuse_key() != reuse_cache["reuse_key"]:
86-
raise ValueError("Reuse inputs must match the previous successful request")
87-
logger.info("[Reuse] Reusing the previous request's input encoder output")
88-
return reuse_cache["inputs"]
85+
def check_reuse_support(self):
86+
raise NotImplementedError(f"{type(self).__name__} does not support reuse")
87+
88+
def check_segment_reuse_support(self):
89+
raise NotImplementedError(f"{type(self).__name__} does not support segment reuse")
8990

9091
def _maybe_freeze_gc(self):
9192
"""Move the steady-state object graph into the GC's permanent generation once."""

lightx2v/models/runners/default_runner.py

Lines changed: 121 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
import gc
2+
import json
23
import os
4+
import shutil
35

46
import numpy as np
57
import requests
@@ -16,7 +18,7 @@
1618
from lightx2v.utils.generate_task_id import generate_task_id
1719
from lightx2v.utils.global_paras import CALIB
1820
from lightx2v.utils.profiler import *
19-
from lightx2v.utils.utils import fixed_shape_resize, get_optimal_patched_size_with_sp, isotropic_crop_resize, mux_audio_from_video, save_to_image, save_to_video, wan_vae_to_comfy
21+
from lightx2v.utils.utils import fixed_shape_resize, get_optimal_patched_size_with_sp, is_main_process, isotropic_crop_resize, mux_audio_from_video, save_to_image, save_to_video, wan_vae_to_comfy
2022
from lightx2v_platform.base.global_var import AI_DEVICE
2123

2224
torch_device_module = getattr(torch, AI_DEVICE)
@@ -88,6 +90,14 @@ def __init__(self, config):
8890
super().__init__(config)
8991
self.has_prompt_enhancer = False
9092
self.progress_callback = None
93+
self.reuse_cache_path = self.config.get("reuse_cache_path")
94+
if self.enable_reuse and not self.reuse_cache_path:
95+
raise ValueError("enable_reuse requires reuse_cache_path")
96+
self.reuse_cache_dir = None
97+
self.reuse_cache_stage_dir = None
98+
self.final_result_path = None
99+
self.previous_result_path = None
100+
self.work_result_path = None
91101
if self.config["task"] == "t2v" and self.config.get("sub_servers", {}).get("prompt_enhancer") is not None:
92102
self.has_prompt_enhancer = True
93103
if not self.check_sub_servers("prompt_enhancer"):
@@ -98,6 +108,116 @@ def __init__(self, config):
98108
self.set_init_device()
99109
self.init_scheduler()
100110

111+
def reuse_key(self):
112+
raise NotImplementedError
113+
114+
def reuse_inputs_path(self, cache_dir):
115+
rank = dist.get_rank() if dist.is_initialized() else 0
116+
return os.path.join(cache_dir, f"inputs_rank_{rank:05d}.pt")
117+
118+
def reuse_input_info(self):
119+
return {}
120+
121+
def load_reuse_state(self, map_location=AI_DEVICE):
122+
manifest_path = os.path.join(self.reuse_cache_dir, "manifest.json")
123+
if not os.path.isfile(manifest_path):
124+
raise RuntimeError(f"No previous successful {type(self).__name__} request is available for reuse")
125+
with open(manifest_path, encoding="utf-8") as f:
126+
manifest = json.load(f)
127+
if manifest["reuse_key"] != self.reuse_key():
128+
raise ValueError("Reuse inputs must match the previous successful request")
129+
self.previous_result_path = manifest["result_path"]
130+
return torch.load(self.reuse_inputs_path(self.reuse_cache_dir), map_location=map_location, weights_only=True)
131+
132+
def load_reused_inputs(self):
133+
cached = self.load_reuse_state()
134+
for name, value in cached["input_info"].items():
135+
setattr(self.input_info, name, value)
136+
logger.info("[Reuse] Loaded the previous request's input encoder output from disk")
137+
return cached["inputs"]
138+
139+
def save_reuse_inputs(self):
140+
torch.save(
141+
{"inputs": self.inputs, "input_info": self.reuse_input_info()},
142+
self.reuse_inputs_path(self.reuse_cache_stage_dir),
143+
)
144+
145+
def prepare_reuse_output(self):
146+
self.reuse_cache_dir = None
147+
self.reuse_cache_stage_dir = None
148+
self.final_result_path = None
149+
self.previous_result_path = None
150+
self.work_result_path = None
151+
152+
output_path = self.input_info.save_result_path
153+
local_output = bool(output_path) and not output_path.startswith(("http://", "https://", "rtmp://"))
154+
reuse_cache_enabled = self.enable_reuse and local_output and not self.input_info.return_result_tensor
155+
if self.reuse and not reuse_cache_enabled:
156+
raise ValueError(f"{type(self).__name__} reuse requires a local output and return_result_tensor=false")
157+
if not reuse_cache_enabled:
158+
return
159+
160+
self.final_result_path = os.path.abspath(os.path.expanduser(output_path))
161+
self.reuse_cache_dir = os.path.abspath(os.path.expanduser(self.reuse_cache_path))
162+
self.reuse_cache_stage_dir = f"{self.reuse_cache_dir}.tmp"
163+
164+
def stage_reuse_cache(self):
165+
if self.reuse_cache_dir is None:
166+
return
167+
168+
if is_main_process():
169+
shutil.rmtree(self.reuse_cache_stage_dir, ignore_errors=True)
170+
os.makedirs(self.reuse_cache_stage_dir)
171+
if dist.is_initialized():
172+
dist.barrier()
173+
174+
if self.reuse:
175+
shutil.copy2(
176+
self.reuse_inputs_path(self.reuse_cache_dir),
177+
self.reuse_inputs_path(self.reuse_cache_stage_dir),
178+
)
179+
else:
180+
self.save_reuse_inputs()
181+
182+
if dist.is_initialized():
183+
dist.barrier()
184+
if is_main_process():
185+
with open(os.path.join(self.reuse_cache_stage_dir, "manifest.json"), "w", encoding="utf-8") as f:
186+
json.dump(
187+
{"reuse_key": self.reuse_key(), "result_path": self.final_result_path},
188+
f,
189+
ensure_ascii=False,
190+
sort_keys=True,
191+
)
192+
193+
def commit_reuse_result(self):
194+
if self.reuse_cache_dir is None or not is_main_process():
195+
return
196+
197+
cache_backup_dir = f"{self.reuse_cache_dir}.old"
198+
shutil.rmtree(cache_backup_dir, ignore_errors=True)
199+
cache_backed_up = os.path.isdir(self.reuse_cache_dir)
200+
if cache_backed_up:
201+
os.replace(self.reuse_cache_dir, cache_backup_dir)
202+
try:
203+
os.replace(self.reuse_cache_stage_dir, self.reuse_cache_dir)
204+
if self.work_result_path is not None:
205+
os.replace(self.work_result_path, self.final_result_path)
206+
except Exception:
207+
shutil.rmtree(self.reuse_cache_dir, ignore_errors=True)
208+
if cache_backed_up:
209+
os.replace(cache_backup_dir, self.reuse_cache_dir)
210+
raise
211+
shutil.rmtree(cache_backup_dir, ignore_errors=True)
212+
213+
def discard_reuse_result(self):
214+
if not is_main_process():
215+
return
216+
if self.reuse_cache_stage_dir:
217+
shutil.rmtree(self.reuse_cache_stage_dir, ignore_errors=True)
218+
if self.work_result_path and os.path.exists(self.work_result_path):
219+
os.remove(self.work_result_path)
220+
101221
def warmup(self):
102222
if not self.config.get("warmup", False):
103223
return

lightx2v/models/runners/qwen_image/qwen_image_runner.py

Lines changed: 37 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,9 @@ def __init__(self, config):
7373
if self.text_encoder_type in ["lightllm_service", "lightllm_kernel"]:
7474
logger.info(f"Using LightLLM text encoder: {self.text_encoder_type}")
7575

76+
def check_reuse_support(self):
77+
pass
78+
7679
@ProfilingContext4DebugL1("Warmup")
7780
def run_warmup(self):
7881
task = self.config.get("task")
@@ -567,29 +570,44 @@ def _finalize_pipeline_outputs(self, input_info, images, latents=None, generator
567570
elif input_info.save_result_path is not None:
568571
return {"images": None}
569572

570-
def _reuse_key(self):
571-
reuse_key = (self.input_info.prompt, self.input_info.negative_prompt)
573+
def reuse_key(self):
574+
reuse_key = {
575+
"prompt": self.input_info.prompt,
576+
"negative_prompt": self.input_info.negative_prompt,
577+
}
572578
if self.config["task"] == "i2i":
573-
reuse_key += (tuple(self.input_info.image_path.split(",")),)
579+
reuse_key["image_path"] = self.input_info.image_path.split(",")
574580
return reuse_key
575581

582+
def reuse_input_info(self):
583+
input_info = {"txt_seq_lens": list(self.input_info.txt_seq_lens)}
584+
if self.config["task"] == "i2i":
585+
input_info["original_size"] = list(self.input_info.original_size)
586+
return input_info
587+
576588
def _run_pipeline_local(self, input_info):
577-
if self.reuse:
578-
self.inputs = self._get_reused_inputs()
579-
else:
580-
self.inputs = self.run_input_encoder()
581-
if self.enable_reuse:
582-
self._reuse_cache = {"reuse_key": self._reuse_key(), "inputs": self.inputs}
583-
if self.config["task"] == "i2i" and "image_encoder_output" in self.inputs:
584-
self.input_info.image_encoder_output = self.inputs["image_encoder_output"]
585-
self.set_target_shape()
586-
self.set_img_shapes()
587-
logger.info(f"input_info: {self.input_info}")
588-
latents, generator = self.run_dit()
589-
images = self.run_vae_decoder(latents)
590-
self.end_run()
591-
self._save_images(images, input_info, log_prefix="Image saved")
592-
return self._finalize_pipeline_outputs(input_info, images, latents=latents, generator=generator)
589+
self.prepare_reuse_output()
590+
try:
591+
self.inputs = self.load_reused_inputs() if self.reuse else self.run_input_encoder()
592+
self.stage_reuse_cache()
593+
if self.config["task"] == "i2i" and "image_encoder_output" in self.inputs:
594+
self.input_info.image_encoder_output = self.inputs["image_encoder_output"]
595+
self.set_target_shape()
596+
self.set_img_shapes()
597+
logger.info(f"input_info: {self.input_info}")
598+
latents, generator = self.run_dit()
599+
images = self.run_vae_decoder(latents)
600+
self.end_run()
601+
self._save_images(images, input_info, log_prefix="Image saved")
602+
result = self._finalize_pipeline_outputs(input_info, images, latents=latents, generator=generator)
603+
self.commit_reuse_result()
604+
return result
605+
except Exception:
606+
self.discard_reuse_result()
607+
raise
608+
finally:
609+
if self.input_info is not None:
610+
self.end_run()
593611

594612
def _run_pipeline_disagg_encoder(self):
595613
self.inputs = self.run_input_encoder()

0 commit comments

Comments
 (0)