File size: 2,501 Bytes
338f0c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
57
58
59
60
61
62
63
64
65
66
67
68
69
# handler.py

import os
import torch
from typing import Dict, Any, Optional
from transformers import (
    AutoTokenizer,
    AutoModelForCausalLM,
    GenerationConfig,
)

class EndpointHandler:
    def __init__(self, model_path: str = None):
        # Load tokenizer & model
        self.tokenizer = AutoTokenizer.from_pretrained(model_path, use_fast=True)
        self.model     = AutoModelForCausalLM.from_pretrained(model_path)
        # Ensure pad token
        if self.tokenizer.pad_token_id is None:
            self.tokenizer.pad_token = self.tokenizer.eos_token
        
        # Put model on GPU if available
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        self.model.to(self.device)

    def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
        # 1) Extract user input & parameters
        user_input = data.get("inputs", "")
        params     = data.get("parameters", {})

        # 2) Build prompt
        prompt = f"User: {user_input}\nAssistant:"
        
        # 3) Tokenize & prepare tensors
        encoded = self.tokenizer(
            prompt,
            return_tensors="pt",
            padding=False,
        ).to(self.device)
        input_ids      = encoded.input_ids
        attention_mask = encoded.attention_mask
        prompt_len     = input_ids.shape[1]

        # 4) Merge default gen args with overrides
        gen_args = {
            "max_new_tokens":   params.get("max_new_tokens", 128),
            "temperature":      params.get("temperature", 1.0),
            "top_p":            params.get("top_p", 1.0),
            "top_k":            params.get("top_k", 50),
            "do_sample":        params.get("do_sample", True),
            "repetition_penalty": params.get("repetition_penalty", 1.0),
            "pad_token_id":     self.tokenizer.pad_token_id,
            "eos_token_id":     self.tokenizer.eos_token_id,
        }
        # Build a GenerationConfig for cleaner API
        gen_config = GenerationConfig(**gen_args)

        # 5) Call generate
        output = self.model.generate(
            inputs=input_ids,
            attention_mask=attention_mask,
            generation_config=gen_config,
        )
        
        # 6) Decode only the new tokens
        #    output is shape [1, prompt_len + gen_len]

        full = self.tokenizer.decode(output[0, prompt_len:], skip_special_tokens=True)
        first_reply = full.split("\nUser:")[0].strip()
        return {"generated_text": first_reply}