File size: 6,118 Bytes
9dce144
38c558b
9dce144
e2eeef7
38c558b
854779e
38c558b
 
 
756cf6c
093ee81
 
 
756cf6c
093ee81
756cf6c
9dce144
e2eeef7
 
38c558b
9dce144
 
 
 
 
 
 
 
 
 
 
 
756cf6c
e2eeef7
9dce144
 
e2eeef7
 
 
 
854779e
9dce144
854779e
38c558b
9dce144
 
 
 
 
 
 
 
 
 
 
38c558b
9dce144
 
38c558b
9dce144
 
 
854779e
9dce144
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
756cf6c
 
9dce144
 
38c558b
 
9dce144
38c558b
9dce144
 
 
 
 
854779e
756cf6c
9dce144
38c558b
9dce144
 
 
 
 
 
 
38c558b
854779e
38c558b
 
9dce144
 
 
 
 
 
 
 
 
 
 
 
 
 
38c558b
9dce144
 
756cf6c
9dce144
 
 
 
 
 
 
38c558b
9dce144
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
# handler.py - PRODUCTION VERSION FOR INFERENCE ENDPOINTS
from __future__ import annotations

import os
from typing import Any, Dict, List, Union

import torch
from transformers import AutoTokenizer, AutoModelForCausalLM

PROMPT_PREFIX = (
    "Ты – модель, которая строго переписывает дореформенный русский текст "
    "в современную орфографию, не меняя смысл и пунктуацию. "
    "Не добавляй комментарии и не переводь текст.\n\nТекст:\n"
)
PROMPT_SUFFIX = "\n\nСовременный орфографический вариант:"


def _as_list(x: Union[str, List[str]]) -> List[str]:
    return [x] if isinstance(x, str) else [str(t) for t in x]


# Load from unsloth's 4-bit quantized version which doesn't use custom files
# OR use the base openai model with trust_remote_code
USE_BASE_MODEL = os.getenv("USE_BASE_MODEL", "true").lower() == "true"

if USE_BASE_MODEL:
    MODEL_ID = "openai/gpt-oss-20b"
    TRUST_REMOTE_CODE = True
else:
    # Alternative: use unsloth's version which may have custom files included
    MODEL_ID = "unsloth/gpt-oss-20b-bnb-4bit"
    TRUST_REMOTE_CODE = False

GEN_KW = {
    "do_sample": False,
    "temperature": 0.0,
    "num_beams": 1,
    "max_new_tokens": int(os.getenv("GEN_MAX_NEW_TOKENS", "512")),
    "repetition_penalty": 1.0,
}


class EndpointHandler:
    def __init__(self, model_dir: str):
        """
        Initialize the endpoint handler.
        
        NOTE: For Inference Endpoints, model_dir points to /repository
        but we're loading from HuggingFace Hub instead since your
        quantized model is missing the custom architecture files.
        """
        print(f"[handler] Model directory provided: {model_dir}")
        print(f"[handler] Loading model from: {MODEL_ID}")
        print(f"[handler] Trust remote code: {TRUST_REMOTE_CODE}")
        
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        
        # Load tokenizer
        self.tokenizer = AutoTokenizer.from_pretrained(
            MODEL_ID,
            use_fast=True,
            trust_remote_code=TRUST_REMOTE_CODE
        )
        
        # Load model
        # The openai/gpt-oss-20b model uses MXFP4 quantization by default
        # which requires specific hardware (H100/A100)
        # For general deployment, we use bfloat16 or float16
        if torch.cuda.is_available():
            # Check if we can use MXFP4 (ideal)
            dtype = "auto"  # Will use MXFP4 if available, otherwise bf16/f16
        else:
            dtype = torch.float32
            
        try:
            self.model = AutoModelForCausalLM.from_pretrained(
                MODEL_ID,
                torch_dtype=dtype,
                device_map="auto" if torch.cuda.is_available() else None,
                trust_remote_code=TRUST_REMOTE_CODE,
                low_cpu_mem_usage=True,
            )
        except Exception as e:
            print(f"[handler] Error loading with MXFP4/auto dtype: {e}")
            print(f"[handler] Falling back to bfloat16...")
            # Fallback to bfloat16 if MXFP4 not supported
            self.model = AutoModelForCausalLM.from_pretrained(
                MODEL_ID,
                torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
                device_map="auto" if torch.cuda.is_available() else None,
                trust_remote_code=TRUST_REMOTE_CODE,
                low_cpu_mem_usage=True,
            )
        
        # Set pad token
        if self.tokenizer.pad_token is None:
            self.tokenizer.pad_token = self.tokenizer.eos_token
            
        # Disable caching on CPU
        if not torch.cuda.is_available():
            self.model.config.use_cache = False
            
        self.model.eval()
        
        print(f"[handler] ✓ Model loaded successfully")
        print(f"[handler] Device: {self.model.device}")
        print(f"[handler] Dtype: {self.model.dtype}")
        print(f"[handler] Model architecture: {self.model.config.architectures}")

    def _encode(self, texts: List[str]) -> Dict[str, Any]:
        """Encode texts with task-specific prompt."""
        prompts = [f"{PROMPT_PREFIX}{t}{PROMPT_SUFFIX}" for t in texts]
        toks = self.tokenizer(
            prompts,
            return_tensors="pt",
            padding=True,
            truncation=True,
            max_length=2048  # Prevent overly long inputs
        )
        return {k: v.to(self.model.device) for k, v in toks.items()}

    @torch.inference_mode()
    def __call__(self, data: Dict[str, Any]) -> List[Dict[str, str]]:
        """
        Process inference request.
        
        Expected input format:
        {
            "inputs": "дореформенный текст" or ["текст1", "текст2"]
        }
        
        Returns:
        [
            {"generated_text": "современный текст"},
            ...
        ]
        """
        if "inputs" not in data:
            return [{"error": "missing 'inputs' field"}]
            
        texts = _as_list(data["inputs"])
        
        if not texts or all(not t.strip() for t in texts):
            return [{"error": "empty input text"}]
        
        try:
            inputs = self._encode(texts)
            outputs = self.model.generate(**inputs, **GEN_KW)

            results: List[Dict[str, str]] = []
            for i, seq in enumerate(outputs):
                # Remove input tokens from output
                in_len = inputs["input_ids"][i].shape[-1]
                gen_only = seq[in_len:]
                text = self.tokenizer.decode(gen_only, skip_special_tokens=True).strip()
                results.append({"generated_text": text})
                
            return results
            
        except Exception as e:
            print(f"[handler] Error during generation: {e}")
            import traceback
            traceback.print_exc()
            return [{"error": f"generation failed: {str(e)}"}]