Add low memory mode, model unloading between stages, and a CPU thread cap

This commit is contained in:
Justin Oros
2026-09-26 21:17:39 -07:00
parent fcd9d77a4d
commit df5bcbb185
2 changed files with 38 additions and 1 deletions
+33
View File
@@ -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