File size: 1,755 Bytes
ef79214
 
 
 
 
 
 
 
 
059fecb
 
ef79214
 
 
 
 
059fecb
 
ef79214
517d1d5
 
43cd549
ef79214
 
 
 
 
 
5ecca15
d241a52
ef79214
 
d241a52
ef79214
5ecca15
 
 
 
d241a52
ef79214
 
 
 
 
 
b09feb9
ef79214
b09feb9
ef79214
 
 
 
 
 
 
ff58438
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
from fastapi import FastAPI, HTTPException, Depends, Request
from transformers import AutoTokenizer, AutoModelForCausalLM
import torch
import os
import huggingface_hub

app = FastAPI()

EXPECTED_TOKEN = os.environ.get("EXPECTED_TOKEN")
HF_TOKEN = os.environ.get('ACCESS_TOKEN')

REPO_ID = "Day23/coder-personal-use"
MODEL_FOLDER = "model"

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

huggingface_hub.login(token=HF_TOKEN)

model_dir = huggingface_hub.snapshot_download(repo_id=REPO_ID, allow_patterns=["model/*"])
model_dir = os.path.join(model_dir, 'model')

tokenizer = AutoTokenizer.from_pretrained(model_dir, trust_remote_code=True, use_auth_token=True)
model = AutoModelForCausalLM.from_pretrained(
    model_dir,
    trust_remote_code=True,
    device_map=device,
)

@app.post("/generate")
async def generate_text(message: str, token: str):
    """Gera um texto com base na entrada fornecida."""
    
    if not message:
        raise HTTPException(status_code=400, detail="O campo 'message' é obrigatório.")

    if token != EXPECTED_TOKEN:
        raise HTTPException(status_code=401, detail="Token inválido")
        
    messages = [{'role': 'user', 'content': message}]
    
    inputs = tokenizer.apply_chat_template(messages, add_generation_prompt=True, return_tensors="pt").to(device)
    
    with torch.no_grad():
        outputs = model.generate(
            inputs,
            max_new_tokens=350,
            do_sample=True,
            top_k=1,
            top_p=0.95,
            num_return_sequences=1,
            eos_token_id=tokenizer.eos_token_id
        )
    
    generated_text = tokenizer.decode(outputs[0][len(inputs[0]):], skip_special_tokens=True)
    
    return {"response": generated_text}