zfir/TypeScriptMate
0
1import os2import time3import threading4import csv5import uuid6import json7from typing import Union, List, Optional8from datetime import datetime9 10import torch11from fastapi import FastAPI, BackgroundTasks, HTTPException, Request12from fastapi.templating import Jinja2Templates13from fastapi.responses import HTMLResponse, StreamingResponse14from pydantic import BaseModel, Field15from transformers import GPT2TokenizerFast, AutoModelForCausalLM16from huggingface_hub import snapshot_download17from starlette.concurrency import run_in_threadpool18from supabase import create_client, Client19from peft import PeftConfig, PeftModel20 21MODEL_REPO_ID = os.getenv("MODEL_REPO_ID")22HF_TOKEN = os.getenv("HF_TOKEN")23USE_LORA = bool(int(os.getenv("USE_LORA", "0")))24USE_QUANTIZATION = bool(int(os.getenv("USE_QUANTIZATION", "1")))25SUPABASE_URL = os.getenv("SUPABASE_URL")26SUPABASE_KEY = os.getenv("SUPABASE_SERVICE_ROLE_KEY")27 28if SUPABASE_URL and SUPABASE_KEY:29 supabase: Client = create_client(SUPABASE_URL, SUPABASE_KEY)30 print("Supabase client connected")31else:32 supabase = None33 print("Supabase client not connected")34 35BUCKET = "typescriptmate"36FEEDBACK_LOG = "feedbacks.csv"37MODIFIED_FEEDBACK_LOG = "feedbacks.modified.csv"38COMPLETION_LOG = "completions.csv"39 40for log_name in (COMPLETION_LOG, FEEDBACK_LOG, MODIFIED_FEEDBACK_LOG):41 try:42 res = supabase.storage.from_(BUCKET).download(log_name)43 data = res.content if hasattr(res, "content") else res44 with open(log_name, "wb") as f:45 f.write(data)46 print(f"Loaded existing {log_name} from Supabase")47 except Exception:48 print(f"No existing {log_name} in bucket; starting fresh")49 50torch.set_num_threads(8)51torch.set_num_interop_threads(2)52 53print("Starting app…")54 55app = FastAPI()56 57templates = Jinja2Templates(directory="templates")58 59MODEL_PATH: str = None60tokenizer: GPT2TokenizerFast = None61model: torch.nn.Module = None62 63def write_feedback_log(event: dict):64 file_exists = os.path.isfile(FEEDBACK_LOG)65 with open(FEEDBACK_LOG, "a", newline="", encoding="utf-8") as f:66 writer = csv.DictWriter(67 f,68 fieldnames=[69 # Base fields70 "timestamp", "userId", "userAgent", "selectedProfileId", 71 "eventName", "schema",72 # Autocomplete fields73 "disable", "maxPromptTokens", "debounceDelay", 74 "maxSuffixPercentage", "prefixPercentage", "transform",75 "template", "multilineCompletions", "slidingWindowPrefixPercentage",76 "slidingWindowSize", "useCache", "onlyMyCode", "useRecentlyEdited",77 "useImports", "accepted", "time", "prefix", "suffix",78 "prompt", "completion", "modelProvider", "modelName",79 "cacheHit", "filepath", "gitRepo", "completionId", "uniqueId",80 # Metadata fields81 "event_name", "schema_version", "level", "profile_id"82 ]83 )84 if not file_exists:85 writer.writeheader()86 writer.writerow(event)87 88 if supabase:89 with open(FEEDBACK_LOG, "rb") as file_obj:90 try:91 supabase.storage.from_(BUCKET).upload(92 FEEDBACK_LOG,93 file_obj,94 file_options={"upsert": "true"},95 )96 except Exception as e:97 print("Failed to upload feedback log:", e)98 99def write_completion_log(event: dict):100 file_exists = os.path.isfile(COMPLETION_LOG)101 with open(COMPLETION_LOG, "a", newline="", encoding="utf-8") as csvfile:102 writer = csv.DictWriter(103 csvfile,104 fieldnames=["prompt", "model", "completion", "latency_s", "timestamp"]105 )106 if not file_exists:107 writer.writeheader()108 writer.writerow(event)109 110 if supabase:111 with open(COMPLETION_LOG, "rb") as file_obj:112 try:113 supabase.storage.from_(BUCKET).upload(114 COMPLETION_LOG,115 file_obj,116 file_options={"upsert": "true"},117 )118 except Exception as e:119 print("Failed to upload completion log:", e)120 121 122def load_model():123 global tokenizer, model, MODEL_PATH124 125 if HF_TOKEN and MODEL_REPO_ID:126 MODEL_PATH = snapshot_download(repo_id=MODEL_REPO_ID, token=HF_TOKEN)127 print(f"Model files: {os.listdir(MODEL_PATH)}")128 else:129 MODEL_PATH = "model"130 print("No HF_TOKEN; using local ./model directory")131 132 if USE_LORA:133 print("Loading LoRA model…")134 135 print("Loading adapter config…")136 adapter_config = PeftConfig.from_pretrained(MODEL_PATH)137 base_model_name_or_path = adapter_config.base_model_name_or_path138 139 print("Loading tokenizer…")140 tokenizer = GPT2TokenizerFast.from_pretrained(base_model_name_or_path)141 tokenizer.pad_token = tokenizer.eos_token142 143 print("Loading base model…")144 base_model = AutoModelForCausalLM.from_pretrained(145 base_model_name_or_path,146 torch_dtype=torch.float32 147 )148 model = PeftModel.from_pretrained(149 base_model, 150 MODEL_PATH, 151 torch_dtype=torch.float32, 152 local_files_only=True153 )154 model = model.merge_and_unload()155 156 else:157 print("Loading vanilla model…")158 159 print("Loading tokenizer…")160 tokenizer = GPT2TokenizerFast.from_pretrained(MODEL_PATH)161 tokenizer.pad_token = tokenizer.eos_token162 163 print("Loading base model…")164 model = AutoModelForCausalLM.from_pretrained(MODEL_PATH)165 166 print("Supported quantization engines:", torch.backends.quantized.supported_engines)167 torch.backends.quantized.engine = 'qnnpack'168 169 if USE_QUANTIZATION:170 print("Quantizing model…")171 model = torch.quantization.quantize_dynamic(172 model, 173 {torch.nn.Linear}, 174 dtype=torch.qint8175 )176 177 model.eval()178 179 print("Warming up (1 token)…")180 dummy = tokenizer("console.log", return_tensors="pt")181 with torch.inference_mode():182 _ = model.generate(**dummy, max_new_tokens=1)183 print("Warming up (40 tokens)…")184 with torch.inference_mode():185 _ = model.generate(**dummy, max_new_tokens=40)186 print("Warming up (100 tokens)…")187 with torch.inference_mode():188 _ = model.generate(**dummy, max_new_tokens=100)189 print("Warming up (200 tokens)…")190 with torch.inference_mode():191 _ = model.generate(**dummy, max_new_tokens=200)192 print("Warming up (400 tokens)…")193 with torch.inference_mode():194 _ = model.generate(**dummy, max_new_tokens=400)195 print("Warming up (800 tokens)…")196 with torch.inference_mode():197 _ = model.generate(**dummy, max_new_tokens=800)198 print("Warm-up complete.")199 200 201threading.Thread(target=load_model, daemon=True).start()202 203 204class CompletionRequest(BaseModel):205 prompt: str206 max_tokens: int = 40207 208class OpenAICompletionRequest(BaseModel):209 model: str210 prompt: Union[str, List[str]] = Field(..., description="Either a string or a list of strings")211 max_tokens: int = 40212 temperature: float = 1.0213 top_p: float = 1.0214 n: int = 1215 stream: bool = False216 logprobs: Optional[int] = None217 218class ContinueAutocompleteData(BaseModel):219 timestamp: float220 userId: Optional[str] = None221 userAgent: Optional[str] = None222 selectedProfileId: Optional[str] = None223 eventName: Optional[str] = None224 schema: Optional[str] = None225 disable: Optional[bool] = None226 maxPromptTokens: Optional[int] = None227 debounceDelay: Optional[int] = None228 maxSuffixPercentage: Optional[float] = None229 prefixPercentage: Optional[float] = None230 transform: Optional[Union[bool, str]] = None231 template: Optional[str] = None232 multilineCompletions: Optional[Union[bool, str]] = None233 slidingWindowPrefixPercentage: Optional[float] = None234 slidingWindowSize: Optional[int] = None235 useCache: Optional[bool] = None236 onlyMyCode: Optional[bool] = None237 useRecentlyEdited: Optional[bool] = None238 useImports: Optional[bool] = None239 accepted: Optional[bool] = None240 time: Optional[float] = None241 prefix: Optional[str] = None242 suffix: Optional[str] = None243 prompt: Optional[str] = None244 completion: Optional[str] = None245 modelProvider: Optional[str] = None246 modelName: Optional[str] = None247 cacheHit: Optional[bool] = None248 filepath: Optional[str] = None249 gitRepo: Optional[str] = None250 completionId: Optional[str] = None251 uniqueId: Optional[str] = None252 253class ContinueFeedback(BaseModel):254 name: str255 data: ContinueAutocompleteData256 schema: str257 level: Optional[str] = None258 profileId: Optional[str] = None259 260@app.get("/")261@app.get("/health")262def index_and_health():263 return {264 "status": "healthy",265 "model_loaded": model is not None,266 "model_path": MODEL_PATH,267 "time": time.time()268 }269 270 271@app.get("/logs", response_class=HTMLResponse)272def logs(request: Request):273 def read_last_rows(path: str, max_rows: int = 20):274 try:275 with open(path, newline="", encoding="utf-8") as f:276 rows = list(csv.reader(f))277 except FileNotFoundError:278 return [], []279 280 if not rows:281 return [], []282 283 header, *entries = rows284 last = entries[-max_rows:] if len(entries) > max_rows else entries285 return header, last286 287 def preprocess_rows(header, rows):288 processed = []289 for row in rows:290 row_dict = dict(zip(header, row))291 292 for key in header:293 val = row_dict.get(key, "")294 295 if "timestamp" in key.lower():296 try:297 ts = float(val)298 if ts > 1e12:299 ts /= 1000300 dt = datetime.fromtimestamp(ts)301 row_dict[key] = dt.strftime('%Y-%m-%d %H:%M:%S')302 except:303 try:304 dt = datetime.fromisoformat(val.replace("Z", "+00:00"))305 row_dict[key] = dt.strftime('%Y-%m-%d %H:%M:%S')306 except:307 pass308 309 elif key.lower() == "time":310 try:311 seconds = int(float(val))312 hours, remainder = divmod(seconds, 3600)313 minutes, secs = divmod(remainder, 60)314 row_dict[key] = f"{hours:02}:{minutes:02}:{secs:02}"315 except:316 pass317 318 elif key.lower() in ["prompt", "completion"] and len(val) > 30:319 row_dict[key] = val[:30] + "..."320 321 processed.append([row_dict.get(col, "") for col in header])322 return processed323 324 comp_header, comp_rows = read_last_rows(COMPLETION_LOG)325 fb_header, fb_rows = read_last_rows(MODIFIED_FEEDBACK_LOG)326 327 comp_rows = preprocess_rows(comp_header, comp_rows) if comp_header else []328 fb_rows = preprocess_rows(fb_header, fb_rows) if fb_header else []329 330 return templates.TemplateResponse("logs.html", {331 "request": request,332 "comp": {"header": comp_header, "rows": comp_rows} if comp_header else None,333 "fb": {"header": fb_header, "rows": fb_rows} if fb_header else None334 })335 336def get_max_sequence_length():337 if hasattr(model, 'config') and hasattr(model.config, 'max_position_embeddings'):338 return model.config.max_position_embeddings339 elif hasattr(model, 'config') and hasattr(model.config, 'n_positions'):340 return model.config.n_positions341 else:342 return 1024343 344def truncate_prompt_if_needed(prompt: str, max_tokens: int = 40) -> str:345 max_seq_len = get_max_sequence_length()346 max_prompt_len = max_seq_len - max_tokens - 10347 348 inputs = tokenizer(prompt, return_tensors="pt")349 if inputs["input_ids"].shape[-1] > max_prompt_len:350 input_ids = inputs["input_ids"][0][-max_prompt_len:]351 truncated_prompt = tokenizer.decode(input_ids, skip_special_tokens=True)352 print(f"Warning: Prompt truncated from {inputs['input_ids'].shape[-1]} to {len(input_ids)} tokens")353 return truncated_prompt354 return prompt355 356@app.post("/v1/completions")357async def complete(358 req: OpenAICompletionRequest,359 background_tasks: BackgroundTasks360):361 if model is None:362 raise HTTPException(status_code=503, detail="Model still loading…")363 364 prompts = req.prompt if isinstance(req.prompt, list) else [req.prompt]365 366 truncated_prompts = [truncate_prompt_if_needed(prompt, req.max_tokens) for prompt in prompts]367 368 start_all = time.time()369 370 if req.stream:371 async def generate_stream():372 for idx, single_prompt in enumerate(truncated_prompts):373 inputs = tokenizer(single_prompt, return_tensors="pt")374 prompt_len = inputs["input_ids"].shape[-1]375 376 for choice_idx in range(req.n):377 with torch.inference_mode():378 outputs = await run_in_threadpool(379 lambda: model.generate(380 **inputs,381 max_new_tokens=req.max_tokens,382 pad_token_id=tokenizer.eos_token_id,383 temperature=req.temperature,384 top_p=req.top_p,385 do_sample=(req.temperature != 0.0 or req.top_p < 1.0),386 return_dict_in_generate=True,387 output_scores=True,388 )389 )390 391 generated_ids = outputs.sequences[0][prompt_len:]392 completion_text = tokenizer.decode(generated_ids, skip_special_tokens=True)393 394 elapsed = time.time() - start_all395 396 event = {397 "prompt": single_prompt,398 "model": MODEL_REPO_ID if MODEL_REPO_ID else "local",399 "completion": completion_text,400 "latency_s": elapsed,401 "timestamp": time.time(),402 }403 background_tasks.add_task(write_completion_log, event)404 405 response = {406 "id": str(uuid.uuid4()),407 "object": "text_completion",408 "created": int(time.time()),409 "choices": [{410 "text": completion_text,411 "index": float(idx * req.n + choice_idx),412 "logprobs": None,413 "finish_reason": "length"414 }],415 "model": req.model416 }417 yield f"data: {json.dumps(response)}\n\n"418 yield "data: [DONE]\n\n"419 420 return StreamingResponse(generate_stream(), media_type="text/event-stream")421 422 all_choices = []423 usage_prompt_tokens = 0424 usage_completion_tokens = 0425 for idx, single_prompt in enumerate(truncated_prompts):426 inputs = tokenizer(single_prompt, return_tensors="pt")427 prompt_len = inputs["input_ids"].shape[-1]428 usage_prompt_tokens += prompt_len429 430 for choice_idx in range(req.n):431 with torch.inference_mode():432 outputs = await run_in_threadpool(433 lambda: model.generate(434 **inputs,435 max_new_tokens=req.max_tokens,436 pad_token_id=tokenizer.eos_token_id,437 temperature=req.temperature,438 top_p=req.top_p,439 do_sample=(req.temperature != 0.0 or req.top_p < 1.0),440 )441 )442 443 generated_ids = outputs[0][prompt_len:]444 num_generated = generated_ids.shape[0]445 usage_completion_tokens += num_generated446 447 completion_text = tokenizer.decode(generated_ids, skip_special_tokens=True)448 all_choices.append({449 "text": completion_text,450 "index": float(idx * req.n + choice_idx),451 "logprobs": None,452 "finish_reason": "length"453 })454 455 total_tokens = usage_prompt_tokens + usage_completion_tokens456 elapsed = time.time() - start_all457 458 event = {459 "prompt": truncated_prompts[0],460 "model": MODEL_REPO_ID if MODEL_REPO_ID else "local",461 "completion": all_choices[0]["text"],462 "latency_s": elapsed,463 "timestamp": time.time(),464 }465 background_tasks.add_task(write_completion_log, event)466 467 response = {468 "id": str(uuid.uuid4()),469 "object": "text_completion",470 "created": int(time.time()),471 "model": req.model,472 "choices": all_choices,473 "usage": {474 "prompt_tokens": usage_prompt_tokens,475 "completion_tokens": usage_completion_tokens,476 "total_tokens": total_tokens477 }478 }479 return response480 481@app.post("/complete")482async def legacy_complete(483 req: CompletionRequest,484 background_tasks: BackgroundTasks485):486 if model is None:487 raise HTTPException(status_code=503, detail="Model still loading…")488 489 truncated_prompt = truncate_prompt_if_needed(req.prompt, req.max_tokens)490 491 start = time.time()492 inputs = tokenizer(truncated_prompt, return_tensors="pt")493 with torch.inference_mode():494 outputs = await run_in_threadpool(495 lambda: model.generate(496 **inputs,497 max_new_tokens=req.max_tokens,498 pad_token_id=tokenizer.eos_token_id,499 )500 )501 502 input_len = inputs["input_ids"].shape[-1]503 generated_ids = outputs[0][input_len:]504 completion = tokenizer.decode(generated_ids, skip_special_tokens=True)505 506 latency = time.time() - start507 event = {508 "prompt": req.prompt,509 "model": MODEL_REPO_ID if MODEL_REPO_ID else "local",510 "completion": completion,511 "latency_s": latency,512 "timestamp": time.time(),513 }514 515 background_tasks.add_task(write_completion_log, event)516 517 return {"completion": completion}518 519 520@app.post("/feedback")521async def feedback(request: Request, background_tasks: BackgroundTasks):522 try:523 body = await request.json()524 525 try:526 ev = ContinueFeedback(**body)527 except Exception as e:528 print("Validation error:", str(e))529 raise HTTPException(530 status_code=422,531 detail={532 "error": "Validation error",533 "message": str(e),534 "received_data": body535 }536 )537 538 event = ev.data.dict(exclude_none=True)539 540 event.update({541 "event_name": ev.name,542 "schema_version": ev.schema,543 "level": ev.level,544 "profile_id": ev.profileId,545 "modelName": MODEL_REPO_ID if MODEL_REPO_ID else "local",546 })547 548 background_tasks.add_task(write_feedback_log, event)549 return {"status": "ok"}550 551 except json.JSONDecodeError as e:552 print("JSON decode error:", str(e))553 raise HTTPException(554 status_code=400,555 detail={556 "error": "Invalid JSON",557 "message": str(e)558 }559 )560 except Exception as e:561 print("Unexpected error:", str(e))562 raise HTTPException(563 status_code=500,564 detail={565 "error": "Internal server error",566 "message": str(e)567 }568 )569 