CTI-llma3 / handler.py
kls123's picture
Update handler.py
74bce33 verified
Raw
History Blame Contribute Delete
4.96 kB
# handler.py - Hugging Face Inference Endpoints için Custom Handler
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
from peft import LoraConfig
from typing import Dict, List, Any
from huggingface_hub import login
import os
class EndpointHandler():
def __init__(self, path=""):
# Token ile login ol
token = os.getenv("HUGGING_FACE_HUB_TOKEN")
if token:
login(token=token)
print("Initializing CTI model...")
# Model'i yükle
self.model = AutoModelForCausalLM.from_pretrained(
path,
torch_dtype=torch.float16,
trust_remote_code=True
)
# GPU'ya taşı
self.model = self.model.to("cuda")
print("Model moved to CUDA")
# Tokenizer yükle
self.tokenizer = AutoTokenizer.from_pretrained(path, trust_remote_code=True)
self.tokenizer.add_special_tokens({'pad_token': '<PAD>'})
# LoRA adapter ekle
lora_config = LoraConfig(
r=8,
target_modules=["q_proj", "o_proj", "k_proj", "v_proj", "gate_proj", "up_proj", "down_proj"],
bias="none",
task_type="CAUSAL_LM",
)
adapter_name = f"adapter_{hash(str(lora_config))}"
try:
self.model.add_adapter(lora_config, adapter_name=adapter_name)
print(f"LoRA adapter added: {adapter_name}")
except ValueError as e:
if "already exists" in str(e):
print(f"Adapter already exists: {e}")
else:
raise e
print("CTI model initialization completed!")
def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:
"""
Process inference request
Args:
data (Dict): Request data containing:
- inputs (str): The input prompt for analysis
- parameters (dict, optional): Generation parameters
- max_length (int): Maximum length of generated text
- temperature (float): Sampling temperature
- top_p (float): Top-p sampling parameter
- top_k (int): Top-k sampling parameter
- do_sample (bool): Whether to use sampling
Returns:
List[Dict]: Generated text response
"""
try:
# Input'u al
inputs = data.get("inputs", "")
if not inputs:
return [{"error": "No inputs provided"}]
# Parameters'ı al (opsiyonel)
parameters = data.get("parameters", {})
# Default değerler (mevcut kodunuzdaki ayarlar)
max_length = parameters.get("max_length", 2048)
temperature = parameters.get("temperature", 0.7)
top_p = parameters.get("top_p", 0.9)
top_k = parameters.get("top_k", 50)
do_sample = parameters.get("do_sample", True)
num_return_sequences = parameters.get("num_return_sequences", 1)
# Input'u tokenize et
tokenized_inputs = self.tokenizer(inputs, return_tensors="pt")
tokenized_inputs = tokenized_inputs.to("cuda")
# Text generate et (mevcut kodunuzdaki ayarlarla)
with torch.no_grad():
outputs = self.model.generate(
**tokenized_inputs,
max_length=max_length,
num_return_sequences=num_return_sequences,
do_sample=do_sample,
top_p=top_p,
top_k=top_k,
temperature=temperature,
pad_token_id=self.tokenizer.pad_token_id,
eos_token_id=self.tokenizer.eos_token_id
)
# Output'u decode et
generated_text = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
# Response format (HF standardına uygun)
return [{"generated_text": generated_text}]
except Exception as e:
print(f"Error in handler: {e}")
return [{"error": str(e)}]
# Test fonksiyonu (geliştirme amaçlı)
def test_handler():
"""Test the handler locally"""
try:
# Handler'ı initialize et
handler = EndpointHandler(".")
# Test data
test_data = {
"inputs": "What is my name?",
"parameters": {
"max_length": 2048,
"temperature": 0.7,
"top_p": 0.9,
"top_k": 50
}
}
# Test et
result = handler(test_data)
print("Test result:", result)
except Exception as e:
print(f"Test error: {e}")
if __name__ == "__main__":
test_handler()