Team Ai
Apppublic

Nephalem/llama-cpp-telegram_bot

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
TelegramBotGenerator.py170 linesDownload Raw Back to root
1from llama_cpp import Llama2from huggingface_hub import hf_hub_download3import logging4import os5import multiprocessing6
7logging.basicConfig(level=logging.INFO)
8logger = logging.getLogger(__name__)
9
10# LLM configuration
11n_ctx = 204812seed = 013n_gpu_layers = 014n_threads = multiprocessing.cpu_count()15
16# Environment variables
17MODEL_REF = os.getenv("MODEL_REF")  # "repo_id:filename" or local path
18MODEL_DIR = os.getenv("MODEL_DIR")  # optional directory for downloaded models
19HF_TOKEN = os.getenv("HF_TOKEN")    # required for gated/private repos
20
21# Optional directory containing model reference files
22MODELS_DIR = os.getenv("MODELS_DIR", "models")
23
24# Ensure directories exist, falling back to /tmp if needed
25try:
26    os.makedirs(MODELS_DIR, exist_ok=True)
27except PermissionError:
28    MODELS_DIR = os.path.join("/tmp", os.path.basename(MODELS_DIR))
29    os.makedirs(MODELS_DIR, exist_ok=True)
30    logger.warning("Using fallback models directory %s", MODELS_DIR)
31
32# Ensure a writable directory for downloaded models
33if not MODEL_DIR:
34    MODEL_DIR = "/tmp/model_cache"
35try:
36    os.makedirs(MODEL_DIR, exist_ok=True)
37except PermissionError:
38    MODEL_DIR = "/tmp/model_cache"
39    os.makedirs(MODEL_DIR, exist_ok=True)
40    logger.warning("Using fallback model cache %s", MODEL_DIR)
41
42
43def _resolve_model_path(model_ref: str) -> str:
44    """Return path to the model file, downloading from HF if necessary."""
45    # Use the path directly if it exists on disk
46    if os.path.exists(model_ref):
47        return model_ref
48    repo, *file_name = model_ref.split(":")
49    kwargs = {
50        "repo_id": repo,
51        "token": HF_TOKEN,
52        "local_dir": MODEL_DIR,
53    }
54    if file_name:
55        kwargs["filename"] = file_name[0]
56    # Download the file using huggingface_hub
57    return hf_hub_download(**kwargs)
58
59
60if not MODEL_REF:
61    raise EnvironmentError(
62        "MODEL_REF environment variable must be set to a local path or 'repo_id:filename'"
63    )
64
65# Resolve and load the model
66logger.info("Resolving model from %s", MODEL_REF)
67model_path = _resolve_model_path(MODEL_REF)
68logger.info("Loading model from %s", model_path)
69llm_generator: Llama = Llama(70    model_path=model_path,71    n_ctx=n_ctx,72    n_threads=n_threads,73    seed=seed,74    n_gpu_layers=n_gpu_layers,75)76 77 78def infer_model_properties(llm: Llama) -> dict:79    """Return basic properties for logging and diagnostics."""80    return {81        "context_window": llm.n_ctx,82        "embedding_size": llm.n_embd,83        "quantization": llm.model_path.split(".")[-2],84    }85 86 87MODEL_PROPERTIES = infer_model_properties(llm_generator)88logger.info(89    "Model loaded: %s | ctx=%s | threads=%s | quant=%s",90    model_path,91    MODEL_PROPERTIES["context_window"],92    n_threads,93    MODEL_PROPERTIES["quantization"],94)95
96
97def get_answer(
98    prompt,
99    generation_params,
100    eos_token,
101    stopping_strings,
102    default_answer: str,
103    turn_template='',
104    **kwargs
105):
106    """Generate a reply using the loaded model."""
107
108    answer = default_answer
109
110    try:
111        answer = llm_generator.create_completion(
112            prompt=prompt,
113            temperature=generation_params["temperature"],
114            top_p=generation_params["top_p"],
115            top_k=generation_params["top_k"],
116            repeat_penalty=generation_params["repetition_penalty"],
117            stop=stopping_strings,
118            max_tokens=generation_params["max_new_tokens"],119            echo=False)120        answer = answer["choices"][0]["text"].replace(prompt, "")
121    except Exception as exception:
122        logger.error("generator_wrapper get answer error %s", exception)
123    return answer
124
125
126def tokens_count(text: str):
127    """Return the token count for a given string."""
128    return len(llm_generator.tokenize(text.encode(encoding="utf-8", errors="strict")))
129
130
131def get_model_list():
132    """Return a list of model files located in the models directory."""
133    bins = []
134    if not os.path.exists(MODELS_DIR):
135        return bins
136    for i in os.listdir(MODELS_DIR):
137
138        if i.endswith((".bin", ".gguf")):
139            bins.append(i)
140    return bins
141
142
143def load_model(model_file: str):
144    """Load and return a new Llama model."""
145    logger.info("Loading model %s", model_file)
146    path_in_models = os.path.join(MODELS_DIR, model_file)
147
148    if os.path.exists(path_in_models):
149        with open(path_in_models, "r", encoding="utf-8") as model:
150            model_ref = model.read().strip()
151    else:
152        model_ref = model_file
153    model_path = _resolve_model_path(model_ref)154    logger.info("Resolved model path %s", model_path)155    llm = Llama(156        model_path=model_path,157        n_ctx=n_ctx,158        n_threads=n_threads,159        seed=seed,160    )161    props = infer_model_properties(llm)162    logger.info(163        "Model loaded: %s | ctx=%s | threads=%s | quant=%s",164        model_path,165        props["context_window"],166        n_threads,167        props["quantization"],168    )169    return llm170