Team Ai
Apppublic

zfir/TypeScriptMate

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py569 linesDownload Raw Back to root
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