Add low memory mode, model unloading between stages, and a CPU thread cap
This commit is contained in:
@@ -82,6 +82,9 @@ GLOBAL_DEFAULTS = {
|
||||
"target_seconds": 60,
|
||||
"clip_seconds": 8,
|
||||
"scene_timing": "narration",
|
||||
"low_memory": True,
|
||||
"free_models_between_stages": True,
|
||||
"cpu_threads": 0,
|
||||
"error_cooldown_minutes": 5,
|
||||
"work_hours_enabled": False,
|
||||
"work_days": "mon,tue,wed,thu,fri,sat,sun",
|
||||
@@ -956,6 +959,21 @@ def ollama_pull(s, model):
|
||||
return True
|
||||
|
||||
|
||||
def unload_ollama(s, model):
|
||||
model = (model or "").strip()
|
||||
if not model or not s.get("free_models_between_stages", True):
|
||||
return
|
||||
try:
|
||||
requests.post(
|
||||
f"{s['ollama_url'].rstrip('/')}/api/generate",
|
||||
json={"model": model, "prompt": "", "keep_alive": 0, "stream": False},
|
||||
timeout=20,
|
||||
)
|
||||
logger.debug("Asked Ollama to release %s", model)
|
||||
except Exception as e:
|
||||
logger.debug("Could not release %s: %s", model, e)
|
||||
|
||||
|
||||
def ensure_local_deps(s):
|
||||
if not ensure_module(s, "torch", "torch") or not ensure_module(s, "diffusers", "diffusers accelerate safetensors transformers".split()[0]):
|
||||
raise RuntimeError("Could not install the image packages. Run ./install-local.sh")
|
||||
@@ -1470,10 +1488,21 @@ def load_pipeline(s):
|
||||
device, dtype = "cuda", torch.float16
|
||||
else:
|
||||
device, dtype = "cpu", torch.float32
|
||||
threads = int(s.get("cpu_threads", 0) or 0)
|
||||
if threads > 0:
|
||||
torch.set_num_threads(threads)
|
||||
logger.info("Limiting image generation to %d CPU threads", threads)
|
||||
logger.info("Loading image model %s on %s (first run downloads several GB)", model, device)
|
||||
pipe = AutoPipelineForText2Image.from_pretrained(model, torch_dtype=dtype, variant="fp16" if dtype == torch.float16 else None)
|
||||
pipe = pipe.to(device)
|
||||
pipe.set_progress_bar_config(disable=True)
|
||||
if s.get("low_memory", True):
|
||||
for enable in ("enable_attention_slicing", "enable_vae_slicing", "enable_vae_tiling"):
|
||||
try:
|
||||
getattr(pipe, enable)()
|
||||
except Exception:
|
||||
pass
|
||||
logger.info("Low memory mode is on for image generation")
|
||||
pipeline["id"] = model
|
||||
pipeline["obj"] = pipe
|
||||
logger.info("Image model ready")
|
||||
@@ -1800,6 +1829,8 @@ def run_job(s):
|
||||
logger.info("Loaded %d existing channel titles to avoid repeats", len(s["_channel_titles"]))
|
||||
script, tags, compliance = produce_script(s, src)
|
||||
remember_story(s["profile_id"], script)
|
||||
if engine == "local":
|
||||
unload_ollama(s, s["ollama_model"])
|
||||
compliance["engine"] = engine
|
||||
update_history(job_id, title=script["title"], notes="; ".join(compliance.get("reviewer_notes", []))[:400])
|
||||
write_json(job_dir / "script.json", {"profile": s["name"], "source": src, "script": script, "tags": tags})
|
||||
@@ -1909,6 +1940,8 @@ def run_job(s):
|
||||
build_scene(i, script["scenes"][i - 1], clips[i - 1])
|
||||
pending = failed
|
||||
|
||||
if kids:
|
||||
unload_ollama(s, s["vision_model"])
|
||||
anatomy_note = f"Possible anatomy glitches kept in scenes {', '.join(str(i) for i in anatomy_flags)}" if anatomy_flags else ""
|
||||
if kids:
|
||||
compliance["vision_checks"] = vision_log
|
||||
|
||||
Reference in New Issue
Block a user