networkbot / functions /generate_text.py
Adieva-15's picture
fixed generate.py
774feb5
Raw
History Blame Contribute Delete
3.94 kB
# from config import model_gt
# import torch
# from transformers import GPT2Tokenizer, GPT2LMHeadModel
#
# tokenizer = GPT2Tokenizer.from_pretrained(model_gt)
# model = GPT2LMHeadModel.from_pretrained(model_gt)
#
# # text = "Replace me by any text you'd like."
#
# def generate_text(prompt:str, max_length:int=100)->str:
# """
# Генерирует продолжение текста на основе заданного промпта.
# """
# try:
# tokenizer.pad_token=tokenizer.eos_token
# inputs = tokenizer([prompt],
# return_tensors="pt", # PyTorch тензоры
# truncation=True,
# padding=True,
# add_special_tokens=True,
# max_length=512)
# with torch.no_grad():
# outputs=model.generate(
# input_ids=inputs.input_ids,
# attention_mask=inputs.attention_mask,
# max_length=max_length,
# num_return_sequences=1,
# pad_token_id=tokenizer.eos_token_id,
# do_sample=True, # включаем вероятностную выборку
# temperature=0.6, # креативность (0 – детерминированно, 1 – случайно)
# )
# generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True)
#
# return generated_text
#
# except Exception as e:
# return f'error of generaion: {str(e)}'
import os
import requests
# Загружаем токен из переменных окружения
HF_TOKEN = os.getenv("HF_TOKEN")
if not HF_TOKEN:
raise ValueError("HF_TOKEN не задан в окружении")
MODEL_GT = "ai-forever/rugpt3small_based_on_gpt2" # можно заменить на другую модель
API_URL = f"https://api-inference.huggingface.co/models/{MODEL_GT}"
def generate_text(prompt: str, max_length: int = 80) -> str:
"""
Генерирует продолжение текста через Hugging Face Inference API.
"""
if not prompt.strip():
return "Промпт пуст."
headers = {"Authorization": f"Bearer {HF_TOKEN}"}
payload = {
"inputs": prompt,
"parameters": {
"max_new_tokens": max_length, # не max_length, а max_new_tokens!
"temperature": 0.7,
"top_p": 0.9,
"do_sample": True,
"repetition_penalty": 1.2,
}
}
try:
# Отправляем POST-запрос к API
response = requests.post(API_URL, headers=headers, json=payload, timeout=30)
response.raise_for_status() # выбросит исключение при HTTP-ошибке
result = response.json()
# API возвращает список, в котором первый элемент — словарь с полем generated_text
if isinstance(result, list) and len(result) > 0:
generated = result[0].get('generated_text', '')
# Иногда модель возвращает полный текст вместе с промптом — убираем промпт
if generated.startswith(prompt):
# generated = generated[len(prompt):].strip()
generated = generated.strip()
return generated if generated else "Не удалось сгенерировать текст."
else:
return f"Неожиданный ответ API: {result}"
except requests.exceptions.Timeout:
return "Ошибка: время ожидания API истекло."
except requests.exceptions.RequestException as e:
return f"Ошибка запроса к API: {str(e)}"
except Exception as e:
return f"Неизвестная ошибка: {str(e)}"