import json import os import glob import shutil import subprocess import sys from typing import Any, Dict, List def env_bool(name: str, default: bool = False) -> bool: value = os.getenv(name) if value is None: return default return value.strip().lower() in {"1", "true", "yes", "on"} def messages_to_prompt(history: List[Dict[str, Any]], system_prompt: str = "") -> str: parts = [] if system_prompt: parts.append(system_prompt.strip()) for message in history or []: role = message.get("role", "user") text = extract_text(message) if text: parts.append(f"{role}: {text}") parts.append("assistant:") return "\n".join(parts) def extract_text(message: Dict[str, Any]) -> str: content = message.get("content") or "" if isinstance(content, str): return content if isinstance(content, list): return "".join( item.get("text", "") for item in content if isinstance(item, dict) and item.get("type") == "text" ) return str(content) class DryRunAdapter: def __init__(self, model_id: str): self.model_id = model_id def generate(self, history, system_prompt, temperature, max_tokens, top_p, top_k): last_user = "" for message in reversed(history or []): if message.get("role") == "user": last_user = extract_text(message) break return ( f"[DRY_RUN:{self.model_id}] Historic Chat Space contract is working. " f"Last user message: {last_user or 'none'}" ) class TransformersAdapter: def __init__(self, model_id: str): import torch from transformers import AutoModelForCausalLM, AutoTokenizer kwargs = parse_json_env("RFAB_HISTORIC_MODEL_KWARGS", {}) if "torch_dtype" in kwargs and kwargs["torch_dtype"] == "auto": kwargs["torch_dtype"] = "auto" elif "torch_dtype" not in kwargs and torch.cuda.is_available(): kwargs["torch_dtype"] = torch.float16 if "device_map" not in kwargs: kwargs["device_map"] = "auto" if torch.cuda.is_available() else None kwargs = {key: value for key, value in kwargs.items() if value is not None} self.model_id = model_id self.tokenizer = AutoTokenizer.from_pretrained(model_id) self.model = AutoModelForCausalLM.from_pretrained(model_id, **kwargs) if not torch.cuda.is_available() and kwargs.get("device_map") is None: self.model.to("cpu") def generate(self, history, system_prompt, temperature, max_tokens, top_p, top_k): import torch prompt = messages_to_prompt(history, system_prompt) inputs = self.tokenizer(prompt, return_tensors="pt", return_token_type_ids=False) device = next(self.model.parameters()).device inputs = {key: value.to(device) for key, value in inputs.items()} with torch.no_grad(): generate_kwargs = { **inputs, "max_new_tokens": int(max_tokens), "temperature": float(temperature), "top_p": float(top_p), "do_sample": float(temperature) > 0, "pad_token_id": self.tokenizer.eos_token_id, } if int(top_k) > 0: generate_kwargs["top_k"] = int(top_k) output = self.model.generate(**generate_kwargs) generated = self.tokenizer.decode(output[0], skip_special_tokens=True) if generated.startswith(prompt): generated = generated[len(prompt):] return generated.strip() class NanochatAdapter: def __init__(self, model_id: str): self.model_id = model_id self.model = None self.tokenizer = None self.device = None def _ensure_loaded(self): if self.model is not None and self.tokenizer is not None: return import torch from huggingface_hub import hf_hub_download, snapshot_download ensure_nanochat_runtime() from nanochat.gpt import GPT, GPTConfig from nanochat.tokenizer import RustBPETokenizer hf_token = os.getenv("HF_TOKEN") or os.getenv("HUGGING_FACE_HUB_TOKEN") local_dir = snapshot_download( repo_id=self.model_id, token=hf_token, ) tokenizer_dir = os.getenv("RFAB_NANOCHAT_TOKENIZER_DIR", "tokenizer") tokenizer_path = os.path.join(local_dir, tokenizer_dir) self.tokenizer = RustBPETokenizer.from_directory(tokenizer_path) meta_file = os.getenv("RFAB_NANOCHAT_META_FILE") or first_match(local_dir, "meta_*.json") model_file = os.getenv("RFAB_NANOCHAT_MODEL_FILE") or first_match(local_dir, "model_*.pt") meta_path = file_in_snapshot(local_dir, meta_file) model_path = file_in_snapshot(local_dir, model_file) if not meta_path: meta_path = hf_hub_download(self.model_id, meta_file, token=hf_token) if not model_path: model_path = hf_hub_download(self.model_id, model_file, token=hf_token) with open(meta_path, "r", encoding="utf-8") as f: meta = json.load(f) self.device = "cuda" if torch.cuda.is_available() else "cpu" config = GPTConfig(**meta["model_config"]) with torch.device("meta"): model = GPT(config) model.to_empty(device=self.device) model.init_weights() state_dict = torch.load(model_path, map_location=self.device) state_dict = {k.removeprefix("_orig_mod."): v for k, v in state_dict.items()} model.load_state_dict(state_dict, strict=True, assign=True) model.eval() self.model = model def generate(self, history, system_prompt, temperature, max_tokens, top_p, top_k): import torch self._ensure_loaded() bos = self.tokenizer.get_bos_token_id() user_start = self.tokenizer.encode_special("<|user_start|>") user_end = self.tokenizer.encode_special("<|user_end|>") assistant_start = self.tokenizer.encode_special("<|assistant_start|>") assistant_end = self.tokenizer.encode_special("<|assistant_end|>") tokens = [bos] if system_prompt: tokens += [user_start] tokens += self.tokenizer.encode(system_prompt.strip()) tokens += [user_end, assistant_start, assistant_end] for message in history or []: role = message.get("role") text = extract_text(message) if not text: continue if role == "user": tokens += [user_start] tokens += self.tokenizer.encode(text) tokens += [user_end] elif role == "assistant": tokens += [assistant_start] tokens += self.tokenizer.encode(text) tokens += [assistant_end] tokens += [assistant_start] generated = [] generate_kwargs = { "max_tokens": int(max_tokens), "temperature": float(temperature), } if int(top_k) > 0: generate_kwargs["top_k"] = int(top_k) if self.device == "cuda": context = torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16) else: context = nullcontext() with torch.no_grad(), context: for token in self.model.generate(tokens, **generate_kwargs): if token in (assistant_end, bos): break generated.append(token) return self.tokenizer.decode(generated).strip() class nullcontext: def __enter__(self): return None def __exit__(self, exc_type, exc, tb): return False def first_match(root, pattern): matches = sorted(glob.glob(os.path.join(root, pattern))) if not matches: raise FileNotFoundError(f"No file matching {pattern} in {root}") return os.path.basename(matches[0]) def file_in_snapshot(root, filename): path = filename if os.path.isabs(filename) else os.path.join(root, filename) if os.path.exists(path): return path basename = os.path.basename(filename) matches = glob.glob(os.path.join(root, "**", basename), recursive=True) if matches: return matches[0] return None def ensure_nanochat_runtime(): try: import nanochat.gpt # noqa: F401 import nanochat.tokenizer # noqa: F401 return except Exception: pass repo_url = os.getenv("RFAB_NANOCHAT_REPO", "https://github.com/karpathy/nanochat.git") ref = os.getenv("RFAB_NANOCHAT_REF", "dc54a1a3077cab11d68fac4c5d1cd5c51f5d8c7a") cache_dir = os.getenv("RFAB_NANOCHAT_CACHE_DIR", "/tmp/rfab_nanochat_runtime") if not os.path.exists(os.path.join(cache_dir, "nanochat")): if os.path.exists(cache_dir): shutil.rmtree(cache_dir) subprocess.check_call([ "git", "clone", "--depth", "1", repo_url, cache_dir, ]) subprocess.check_call(["git", "fetch", "--depth", "1", "origin", ref], cwd=cache_dir) subprocess.check_call(["git", "checkout", ref], cwd=cache_dir) if cache_dir not in sys.path: sys.path.insert(0, cache_dir) import nanochat.gpt # noqa: F401 import nanochat.tokenizer # noqa: F401 def parse_json_env(name: str, default): raw = os.getenv(name) if not raw: return default return json.loads(raw) def create_adapter(): model_id = os.getenv("RFAB_HISTORIC_MODEL_ID", "dry-run-model") if env_bool("DRY_RUN", True): return DryRunAdapter(model_id) adapter = os.getenv("RFAB_HISTORIC_ADAPTER", "transformers").strip().lower() if adapter == "transformers": return TransformersAdapter(model_id) if adapter == "nanochat": return NanochatAdapter(model_id) raise ValueError(f"Unsupported RFAB_HISTORIC_ADAPTER: {adapter}")