Files
Athena-Deck/prompt_enhancer.py

155 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Isolated, one-shot Qwen image prompt enhancement owned by Deck."""
import json,os,shutil,signal,subprocess,tempfile,threading,time,uuid
from pathlib import Path
from image_test import probe,cgroup_headroom,decode_references
from image_runtime import BUNDLED
from profiles import PROMPT_ENHANCER_REPOS
WORKER=Path(__file__).with_name('prompt_enhancer_worker.py')
GIB=1024**3
class PromptEnhancer:
def __init__(self,root,catalog,profiles,runtime):
self.root=Path(root);self.catalog=catalog;self.profiles=profiles;self.runtime=runtime
self.lock=threading.RLock();self.cancel=threading.Event();self.process=None;self.job=None;self.preview=None;self.acquire=lambda wait=False:lambda:None
def directory(self,task):return self.root/task
def python(self):return BUNDLED[0] if BUNDLED[0].is_file() else self.runtime.paths()[0]
def installed(self,task):
path=self.directory(task);marker=path/'installed.json'
if not marker.is_file():return None
try:
data=json.loads(marker.read_text());return data if data['repo']==PROMPT_ENHANCER_REPOS[task] and (path/'model.safetensors.index.json').is_file() else None
except (OSError,ValueError,KeyError):return None
def status(self):
with self.lock:
return dict(models={task:dict(repo=repo,installed=bool(self.installed(task)),revision=(self.installed(task) or {}).get('revision')) for task,repo in PROMPT_ENHANCER_REPOS.items()},job=dict(self.job) if self.job else None,preview=dict(self.preview) if self.preview else None)
def install(self,task,revision=None):
if task not in PROMPT_ENHANCER_REPOS:raise ValueError('Unbekannter Prompt-Aufwerter.')
with self.lock:
if self.job and self.job['state']=='running':raise ValueError('Ein Aufwerter-Download läuft bereits.')
if self.preview and self.preview['state']=='running':raise ValueError('Eine Prompt-Vorschau läuft; zuerst abschließen.')
if self.installed(task):raise ValueError('Dieser Aufwerter ist bereits installiert.')
if not self.runtime.status()['installed']:raise ValueError('Zuerst die Bildlaufzeit einrichten.')
if shutil.disk_usage(self.root if self.root.exists() else self.root.parent).free<30*GIB:raise ValueError('Mindestens 30 GiB freier Plattenspeicher erforderlich.')
source=self.catalog.files(PROMPT_ENHANCER_REPOS[task],revision) if revision else self.catalog.files(PROMPT_ENHANCER_REPOS[task]);revision=source['revision']
expected=sum(f['size'] for f in source['files'] if f['name'].endswith('.safetensors'))
if expected<15*GIB or expected>25*GIB:raise ValueError('Unerwartete Größe der offiziellen Aufwerter-Gewichte; Rezept prüfen.')
self.root.mkdir(parents=True,exist_ok=True,mode=0o700);self.cancel.clear()
self.job=dict(id=uuid.uuid4().hex,task=task,state='running',phase='Offizielle Dateien werden heruntergeladen',repo=source['repo'],revision=revision,total_bytes=expected,started_at=time.time())
threading.Thread(target=self._download,args=(task,revision),daemon=True).start();return self.status()
def _download(self,task,revision):
target=self.directory(task);target.mkdir(parents=True,exist_ok=True,mode=0o700)
env=dict(os.environ,HF_TOKEN=self.catalog.hub_auth.token() or '',HF_HUB_DISABLE_TELEMETRY='1',HOME=str(self.root),PYTHONDONTWRITEBYTECODE='1')
data=dict(repo=PROMPT_ENHANCER_REPOS[task],revision=revision,directory=str(target))
state='failed';phase='Aufwerter-Download fehlgeschlagen. Netzwerk, Hub-Zugang und freien Speicher prüfen.'
try:
python=self.python()
with self.lock:
if self.cancel.is_set():raise InterruptedError()
self.process=subprocess.Popen([str(python),str(WORKER),'download'],stdin=subprocess.PIPE,stdout=subprocess.DEVNULL,stderr=subprocess.DEVNULL,text=True,start_new_session=True,env=env)
process=self.process
process.stdin.write(json.dumps(data));process.stdin.close()
deadline=time.monotonic()+21600
while process.poll() is None:
if self.cancel.wait(1):raise InterruptedError()
if time.monotonic()>deadline:raise ValueError('Download-Zeitlimit erreicht.')
if shutil.disk_usage(target).free<10*GIB:raise ValueError('Plattenreserve unter 10 GiB.')
if self.cancel.is_set():raise InterruptedError()
if process.returncode:raise ValueError('Aufwerter-Dateien konnten nicht vollständig heruntergeladen werden.')
marker=target/'installed.tmp';marker.write_text(json.dumps(data));marker.replace(target/'installed.json')
state='complete';phase='Offizieller Aufwerter installiert. Es läuft kein Modell.'
except InterruptedError:state='cancelled';phase='Download abgebrochen.'
except ValueError as exc:phase=str(exc)
except Exception:pass
finally:
self._terminate()
with self.lock:self.job.update(state=state,phase=phase,finished_at=time.time())
def _terminate(self):
with self.lock:process=self.process;self.process=None
if process and process.poll() is None:
try:
os.killpg(process.pid,signal.SIGTERM)
try:process.wait(timeout=5)
except subprocess.TimeoutExpired:os.killpg(process.pid,signal.SIGKILL);process.wait()
except ProcessLookupError:pass
def stop(self):
self.cancel.set();self._terminate();return {'cancellation_requested':True}
def stop_preview(self):
with self.lock:active=bool(self.preview and self.preview['state']=='running')
if active:self.cancel.set();self._terminate()
return {'cancellation_requested':active}
def _devices(self,choice):
if choice=='cpu':return ''
rows=probe();names=['5080','3060'] if choice=='auto' else [choice];selected=[]
for name in names:
row=next((g for g in rows if ('RTX '+name) in g['name']),None)
if not row or row['processes'] or row['free_mib']<('5080'==name and 10000 or 7000):
if choice=='auto':continue
raise ValueError('Gewählte GPU für Prompt-Aufbereitung nicht frei oder mit zu wenig VRAM.')
selected.append(row['uuid'])
if not selected:raise ValueError('Keine geeignete freie GPU für den Prompt-Aufwerter.')
return ','.join(selected)
def rewrite(self,profile,prompt,references=(),cancel=None):
task='i2i' if references else 't2i';config=profile.get('prompt_enhancer') or {};repo=config.get(task)
if not repo:return dict(prompt=prompt,enhanced=False,wh_ratio='')
if repo!=PROMPT_ENHANCER_REPOS[task]:raise ValueError('Unbekannter Prompt-Aufwerter im Profil.')
if not self.installed(task):raise ValueError('Prompt-Aufwerter fehlt. Unter Bildgenerierung → Aufwerter installieren.')
if not isinstance(prompt,str) or not 1<=len(prompt.strip())<=4000:raise ValueError('Prompt mit 1–4000 Zeichen erforderlich.')
headroom=cgroup_headroom()
if headroom is not None and headroom<22*GIB:raise ValueError('Für den Prompt-Aufwerter sind mindestens 22 GiB freier Deck-RAM nötig.')
device=config.get('device','auto');visible=self._devices(device)
rows={g['uuid']:g for g in probe()}
gpu_limits=[min(12 if '5080' in rows[ident]['name'] else 9,max(4,int(rows[ident]['free_mib']/1024)-2)) for ident in visible.split(',') if ident]
self.root.mkdir(parents=True,exist_ok=True,mode=0o700)
with tempfile.TemporaryDirectory(prefix='rewrite-',dir=self.root) as folder:
paths=[]
for i,(raw,extension) in enumerate(references):
path=Path(folder)/f'image-{i}.{extension}';path.write_bytes(raw);paths.append(str(path))
env=dict(os.environ,CUDA_VISIBLE_DEVICES=visible,HF_HUB_OFFLINE='1',TRANSFORMERS_OFFLINE='1',HF_HUB_DISABLE_TELEMETRY='1',OMP_NUM_THREADS='2',HOME=folder,PYTHONDONTWRITEBYTECODE='1')
data=dict(directory=str(self.directory(task)),task=task,images=paths,prompt=prompt,device=device,gpu_limits=gpu_limits)
with self.lock:
if self.job and self.job['state']=='running':raise ValueError('Aufwerter-Download läuft; Bildauftrag danach erneut starten.')
if self.process and self.process.poll() is None:raise ValueError('Ein Prompt-Aufwerter läuft bereits.')
self.process=subprocess.Popen([str(self.python()),str(WORKER),'rewrite'],stdin=subprocess.PIPE,stdout=subprocess.PIPE,stderr=subprocess.DEVNULL,text=True,start_new_session=True,env=env)
process=self.process
try:
payload=json.dumps(data);deadline=time.monotonic()+900
while True:
if cancel and cancel.is_set():raise InterruptedError()
if time.monotonic()>deadline:raise ValueError('Prompt-Aufwerter hat das Zeitlimit erreicht.')
try:output,_=process.communicate(payload,timeout=1);break
except subprocess.TimeoutExpired:payload=None
if cancel and cancel.is_set():raise InterruptedError()
if process.returncode:raise ValueError('Prompt-Aufwerter fehlgeschlagen. GPU/RAM, Modell und Laufzeit prüfen; Bild wurde nicht gestartet.')
result=json.loads(output)
if not isinstance(result.get('prompt'),str) or not 1<=len(result['prompt'])<=4000:raise ValueError('Prompt-Aufwerter lieferte keinen gültigen Prompt.')
return dict(result,enhanced=True,source=repo)
finally:self._terminate()
def preview_start(self,profile_id,prompt,reference_images):
profile=next((p for p in self.profiles.status()['profiles'] if p['id']==profile_id and p['kind']=='image'),None)
if not profile:raise ValueError('Bildprofil nicht gefunden.')
if not isinstance(prompt,str) or not 1<=len(prompt.strip())<=4000:raise ValueError('Prompt mit 1–4000 Zeichen erforderlich.')
references=decode_references(reference_images,profile.get('capabilities',{}).get('reference_images',0))
task='i2i' if references else 't2i'
if not (profile.get('prompt_enhancer') or {}).get(task):raise ValueError('Für diese Bildaufgabe ist im Profil kein Prompt-Aufwerter aktiviert.')
if not self.installed(task):raise ValueError('Prompt-Aufwerter zuerst im Bildprofil installieren.')
release=self.acquire(False)
try:
with self.lock:
if self.preview and self.preview['state']=='running':raise ValueError('Eine Prompt-Vorschau läuft bereits.')
if self.job and self.job['state']=='running':raise ValueError('Aufwerter-Download läuft bereits.')
self.cancel.clear();self.preview=dict(id=uuid.uuid4().hex,state='running',phase='Prompt wird aufbereitet',profile_id=profile_id,task=task,started_at=time.time())
threading.Thread(target=self._preview,args=(profile,prompt,references,release),daemon=True).start();return dict(self.preview)
except Exception:release();raise
def _preview(self,profile,prompt,references,release):
try:
result=self.rewrite(profile,prompt,references,self.cancel)
with self.lock:self.preview.update(state='complete',phase='Prompt-Aufbereitung fertig',result=result)
except InterruptedError:
with self.lock:self.preview.update(state='cancelled',phase='Prompt-Aufbereitung abgebrochen')
except ValueError as exc:
with self.lock:self.preview.update(state='failed',phase=str(exc))
except Exception:
with self.lock:self.preview.update(state='failed',phase='Prompt-Aufwerter fehlgeschlagen. Ressourcen und Laufzeit prüfen.')
finally:release()