File size: 6,862 Bytes
b4b0f75
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""Small HTTP transport and owned single-GPU native-server lifecycle."""
from contextlib import contextmanager
from datetime import datetime, timezone
import json
import os
from pathlib import Path
import signal
import socket
import subprocess
import threading
import time
import urllib.request


class Client:
    def __init__(self, base='http://127.0.0.1:5991', adapter_id=None, timeout=120):
        self.base=base.rstrip('/').removesuffix('/v1')
        self.adapter_id=adapter_id
        self.timeout=timeout
        self.calls=[]
        self.lock=threading.RLock()

    def health(self):
        with urllib.request.urlopen(self.base+'/health',timeout=5) as stream:return json.load(stream)

    def verify_profile(self):
        with urllib.request.urlopen(self.base+'/lora-adapters',timeout=5) as stream:loaded=json.load(stream)
        if self.adapter_id is None and loaded:
            raise RuntimeError('Base-only profile requires a backend with no adapters loaded')
        if self.adapter_id is not None and self.adapter_id not in [r['id'] for r in loaded]:
            raise RuntimeError('Requested adapter ID is not loaded by the backend')
        return loaded

    def post(self, path, payload):
        body=dict(payload)
        if self.adapter_id is None:
            body.pop('lora',None)
        else:
            body['lora']=[dict(id=self.adapter_id,scale=1)]
        body.update(cache_prompt=False,chat_template_kwargs=dict(enable_thinking=False))
        req=urllib.request.Request(self.base+path,data=json.dumps(body).encode(),headers={'Content-Type':'application/json'})
        tick=time.perf_counter()
        with self.lock, urllib.request.urlopen(req,timeout=self.timeout) as stream:
            result=json.load(stream)
        usage=result.get('usage') or {}
        self.calls.append(dict(seconds=time.perf_counter()-tick,prompt_tokens=usage.get('prompt_tokens'),
                               completion_tokens=usage.get('completion_tokens'),
                               cached_tokens=(usage.get('prompt_tokens_details') or {}).get('cached_tokens',0),
                               finish_reason=(result.get('choices') or [{}])[0].get('finish_reason'),adapter_id=self.adapter_id))
        return result

    def generate(self, prompt, max_tokens=1536, json_output=False, schema=None):
        body=dict(model='qwen',messages=[dict(role='user',content=prompt)],max_tokens=max_tokens,temperature=0)
        if json_output: body['response_format']={'type':'json_object'}
        if schema is not None:body['response_format']={'type':'json_object','schema':schema}
        result=self.post('/v1/chat/completions',body)
        choice=result['choices'][0]
        if choice['finish_reason']=='length':
            raise ValueError('Structured response was truncated; no actions may execute')
        return choice['message'].get('content') or ''


class NativeServer:
    def __init__(self, root, gpu=0, port=5991, model='models/Ternary-Bonsai-2-27B-PQ2_0.gguf', adapter=None,
                 binary='llama.cpp-b2/build/bin/llama-server', context=4096):
        self.root=Path(root); self.root.mkdir(parents=True,exist_ok=True)
        self.gpu,self.port,self.model,self.adapter,self.binary,self.context=gpu,port,str(model),adapter,str(binary),context
        self.proc=None; self.samples=[]; self.stop=threading.Event()

    def event(self,kind,**values):
        row=dict(time_utc=datetime.now(timezone.utc).isoformat(),kind=kind,**values)
        with (self.root/'runtime-events.jsonl').open('a') as f:f.write(json.dumps(row)+'\n')

    def memory(self):
        return int(subprocess.check_output(['nvidia-smi',f'--id={self.gpu}','--query-gpu=memory.used','--format=csv,noheader,nounits'],text=True))

    def __enter__(self):
        with socket.socket() as sock:
            if sock.connect_ex(('127.0.0.1',self.port))==0:raise RuntimeError('Port occupied; no existing server will be stopped')
        if self.memory()>=100:raise RuntimeError(f'GPU {self.gpu} is occupied')
        env=dict(os.environ,CUDA_VISIBLE_DEVICES=str(self.gpu),PYTHONDONTWRITEBYTECODE='1')
        env['LD_LIBRARY_PATH']=str(Path(self.binary).resolve().parent)+((':'+env['LD_LIBRARY_PATH']) if env.get('LD_LIBRARY_PATH') else '')
        command=[self.binary,'-m',self.model,'--host','0.0.0.0','--port',str(self.port),'-c',str(self.context),'-ngl','99','-fa','on',
                 '--jinja','--reasoning','off','-np','1','--no-context-shift','--cache-ram','0','--no-cache-idle-slots']
        if self.adapter:command+=['--lora',str(self.adapter)]
        self.log=(self.root/f'server-{time.time_ns()}.log').open('x')
        self.proc=subprocess.Popen(command,env=env,stdout=self.log,stderr=subprocess.STDOUT)
        self.event('server_started',pid=self.proc.pid,gpu=self.gpu,command=command)
        try:
            deadline=time.monotonic()+180
            while time.monotonic()<deadline:
                if self.proc.poll() is not None:raise RuntimeError('Native server exited during startup')
                try:
                    with urllib.request.urlopen(f'http://127.0.0.1:{self.port}/health',timeout=1) as r:
                        if r.status==200:break
                except OSError:pass
                time.sleep(.25)
            else:raise TimeoutError('Native server readiness timed out')
            with urllib.request.urlopen(f'http://127.0.0.1:{self.port}/lora-adapters',timeout=5) as r:loaded=json.load(r)
            if [x['path'] for x in loaded] != ([str(self.adapter)] if self.adapter else []):
                raise RuntimeError('Loaded adapters do not match requested profile')
            def sample():
                while not self.stop.is_set():
                    self.samples.append(dict(time=time.time(),mib=self.memory()))
                    self.stop.wait(.5)
            self.thread=threading.Thread(target=sample,daemon=True);self.thread.start()
            self.event('server_ready',memory_mib=self.memory(),loaded_adapters=loaded)
            return Client(f'http://127.0.0.1:{self.port}',0 if self.adapter else None)
        except BaseException:
            self.__exit__(None,None,None)
            raise

    def __exit__(self,*args):
        self.stop.set()
        if hasattr(self,'thread'):self.thread.join(timeout=5)
        if self.proc:
            for sig in [signal.SIGTERM,signal.SIGINT,signal.SIGKILL]:
                if self.proc.poll() is not None:break
                self.proc.send_signal(sig)
                try:self.proc.wait(timeout=10)
                except subprocess.TimeoutExpired:pass
            self.log.close()
            self.event('server_stopped',pid=self.proc.pid,exit_code=self.proc.poll(),peak_mib=max((s['mib'] for s in self.samples),default=None))
            with (self.root/f'memory-{self.proc.pid}.json').open('x') as f:json.dump(self.samples,f)