Team Ai
Apppublic

Csdfg/ChatGLM-6B-Int4-API-OpenAI-Compatible

sourceHugging Faceapache-2.0updated 4y agoView on Hugging Face
0likes
main.py115 linesDownload Raw Back to root
1import json2from typing import List3 4import torch5from fastapi import FastAPI, Request, status, HTTPException6from pydantic import BaseModel7from torch.cuda import get_device_properties8from transformers import AutoModel, AutoTokenizer9from sse_starlette.sse import EventSourceResponse10from fastapi.middleware.cors import CORSMiddleware11import uvicorn12 13import os14 15os.environ['TRANSFORMERS_CACHE'] = ".cache"16 17bits = 418kernel_path = "models/models--silver--chatglm-6b-int4-slim/quantization_kernels.so"19model_path = "./models/models--silver--chatglm-6b-int4-slim/snapshots/02e096b3805c579caf5741a6d8eddd5ba7a74e0d"20cache_dir = './models'21model_name = 'chatglm-6b-int4'22min_memory = 5.523tokenizer = None24model = None25 26app = FastAPI()27 28app.add_middleware(29    CORSMiddleware,30    allow_origins=["*"],31    allow_credentials=True,32    allow_methods=["*"],33    allow_headers=["*"],34)35 36 37@app.on_event('startup')38def init():39    global tokenizer, model40    tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True, cache_dir=cache_dir)41    model = AutoModel.from_pretrained(model_path, trust_remote_code=True, cache_dir=cache_dir)42 43    if torch.cuda.is_available() and get_device_properties(0).total_memory / 1024 ** 3 > min_memory:44        model = model.half().quantize(bits=bits).cuda()45        print("Using GPU")46    else:47        model = model.float().quantize(bits=bits)48        if torch.cuda.is_available():49            print("Total Memory: ", get_device_properties(0).total_memory / 1024 ** 3)50        else:51            print("No GPU available")52        print("Using CPU")53    model = model.eval()54    if os.environ.get("ngrok_token") is not None:55        ngrok_connect()56 57 58class Message(BaseModel):59    role: str60    content: str61 62 63class Body(BaseModel):64    messages: List[Message]65    model: str66    stream: bool67    max_tokens: int68 69 70@app.get("/")71def read_root():72    return {"Hello": "World!"}73 74 75@app.post("/chat/completions")76async def completions(body: Body, request: Request):77    if not body.stream or body.model != model_name:78        raise HTTPException(status.HTTP_400_BAD_REQUEST, "Not Implemented")79 80    question = body.messages[-1]81    if question.role == 'user':82        question = question.content83    else:84        raise HTTPException(status.HTTP_400_BAD_REQUEST, "No Question Found")85 86    user_question = ''87    history = []88    for message in body.messages:89        if message.role == 'user':90            user_question = message.content91        elif message.role == 'system' or message.role == 'assistant':92            assistant_answer = message.content93            history.append((user_question, assistant_answer))94 95    async def event_generator():96        for response in model.stream_chat(tokenizer, question, history, max_length=max(2048, body.max_tokens)):97            if await request.is_disconnected():98                return99            yield json.dumps({"response": response[0]})100        yield "[DONE]"101 102    return EventSourceResponse(event_generator())103 104 105def ngrok_connect():106    from pyngrok import ngrok, conf107    conf.set_default(conf.PyngrokConfig(ngrok_path="./ngrok"))108    ngrok.set_auth_token(os.environ["ngrok_token"])109    http_tunnel = ngrok.connect(8000)110    print(http_tunnel.public_url)111 112 113if __name__ == "__main__":114    uvicorn.run("main:app", reload=True, app_dir=".")115