|
|
| import os |
| from abc import ABC, abstractmethod |
|
|
| import torch |
|
|
| from huggingface_hub import InferenceClient |
| from transformers import pipeline, AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig |
| from peft import PeftModel, PeftConfig |
| from safetensors.torch import load_file, save_file |
| from huggingface_hub import hf_hub_download |
| import tempfile |
| import shutil |
| from pathlib import Path |
|
|
|
|
| HF_TOKEN = os.environ.get("HF_TOKEN") |
|
|
| |
| _model_cache = {} |
| model_configs = { |
| "baseline": { |
| "id": "meta-llama/Llama-3.2-1B-Instruct", |
| "type": "endpoint" |
| }, |
| "fine_tuned": { |
| "base_model_id": "meta-llama/Llama-3.2-1B-Instruct", |
| "adapter_id": "aracape/teaching-assistant-1B-dpo", |
| "type": "pipeline" |
| }, |
| "prompted": { |
| "id": "meta-llama/Llama-3.2-1B-Instruct", |
| "type": "endpoint" |
| }, |
| } |
|
|
| class ModelWrapper(ABC): |
| """Abstract base class for model wrappers""" |
| |
| def __init__(self, model_id: str): |
| self.model_id = model_id |
| |
| @abstractmethod |
| def generate(self, messages: list[dict], max_tokens: int, temperature: float) -> str: |
| """Generate a response given messages""" |
| pass |
|
|
|
|
| class InferenceEndpointModel(ModelWrapper): |
| """Wrapper for models deployed as HF inference endpoints""" |
| |
| def __init__(self, model_id: str): |
| super().__init__(model_id) |
| self.client = InferenceClient(model_id, token=HF_TOKEN) |
| |
| def generate(self, messages: list[dict], max_tokens: int, temperature: float) -> str: |
| """Generate a complete response""" |
| result = self.client.chat_completion( |
| messages, |
| max_tokens=max_tokens, |
| temperature=temperature |
| ) |
| return result.choices[0].message["content"] |
| |
| def generate(self, messages: list[dict], max_tokens: int, temperature: float): |
| """Generate a streaming response""" |
| response = "" |
| for message_chunk in self.client.chat_completion( |
| messages, |
| max_tokens=max_tokens, |
| stream=True, |
| temperature=temperature, |
| ): |
| if message_chunk.choices and message_chunk.choices[0].delta.content: |
| token = message_chunk.choices[0].delta.content |
| response += token |
|
|
| return response |
|
|
|
|
| class PipelineModel(ModelWrapper): |
| """Wrapper for models loaded locally with transformers pipeline""" |
| |
| def __init__(self, base_model_id: str, adapter_id: str): |
| super().__init__(adapter_id) |
| |
| self._pipeline = None |
| self.base_model_id = base_model_id |
| self.adapter_id = adapter_id |
|
|
| def _load_model(self): |
|
|
| if torch.cuda.is_available(): |
| quantization_config = BitsAndBytesConfig( |
| load_in_4bit=True, |
| bnb_4bit_compute_dtype="bfloat16", |
| bnb_4bit_quant_type="nf4", |
| bnb_4bit_use_double_quant=True, |
| ) |
| else: |
| quantization_config = None |
|
|
| print(f"Loading base model: {self.base_model_id}") |
| base_model = AutoModelForCausalLM.from_pretrained( |
| self.base_model_id, |
| dtype="auto", |
| device_map="auto", |
| quantization_config=quantization_config, |
| token=HF_TOKEN |
| ) |
| |
| |
| print(f"Loading adapter: {self.adapter_id}") |
| model = self._load_adapter_robust(base_model) |
| return model |
|
|
| def _load_adapter_robust(self, base_model): |
| """ |
| Robustly load the adapter, handling potential key mismatches |
| (common with TRL/DPO trained models having extra nesting). |
| """ |
| |
| offload_folder = Path("offload") |
| offload_folder.mkdir(exist_ok=True) |
| |
| try: |
| return PeftModel.from_pretrained( |
| base_model, |
| self.adapter_id, |
| offload_folder=str(offload_folder), |
| token=HF_TOKEN |
| ) |
| except KeyError as e: |
| |
| |
| |
| if "base_model.model.model" in str(e) or "embed_tokens" in str(e): |
| print(f"Encountered KeyError loading adapter: {e}. Attempting to fix keys...") |
| return self._load_adapter_with_key_fix(base_model, offload_folder=str(offload_folder)) |
| raise e |
|
|
| def _load_adapter_with_key_fix(self, base_model, offload_folder=None): |
| """ |
| Download adapter, fix keys by removing extra 'model' nesting, and load. |
| """ |
| print("Patching adapter weights to fix key mismatch...") |
| with tempfile.TemporaryDirectory() as temp_dir: |
| temp_path = Path(temp_dir) |
| try: |
| config_path = hf_hub_download(self.adapter_id, "adapter_config.json", token=HF_TOKEN) |
| shutil.copy(config_path, temp_path / "adapter_config.json") |
| except Exception as e: |
| raise RuntimeError(f"Failed to download adapter config: {e}") |
|
|
| |
| is_safetensors = True |
| try: |
| weights_path = hf_hub_download(self.adapter_id, "adapter_model.safetensors", token=HF_TOKEN) |
| weights = load_file(weights_path) |
| except Exception: |
| try: |
| weights_path = hf_hub_download(self.adapter_id, "adapter_model.bin", token=HF_TOKEN) |
| weights = torch.load(weights_path, map_location="cpu") |
| is_safetensors = False |
| except Exception as e: |
| raise RuntimeError(f"Failed to download adapter weights: {e}") |
|
|
| |
| new_weights = {} |
| fixed_count = 0 |
| for k, v in weights.items(): |
| if "base_model.model.model." in k: |
| new_k = k.replace("base_model.model.model.", "base_model.model.") |
| new_weights[new_k] = v |
| fixed_count += 1 |
| else: |
| new_weights[k] = v |
| |
| print(f"Fixed {fixed_count} keys in adapter weights.") |
|
|
| |
| if is_safetensors: |
| save_file(new_weights, temp_path / "adapter_model.safetensors") |
| else: |
| torch.save(new_weights, temp_path / "adapter_model.bin") |
| |
| |
| print("Loading PeftModel with manual state_dict assignment...") |
| config = PeftConfig.from_pretrained(temp_dir) |
| model = PeftModel(base_model, config) |
| |
| |
| |
| try: |
| model.load_state_dict(new_weights, strict=False, assign=True) |
| except TypeError: |
| print("Warning: load_state_dict does not support assign=True. Loading might fail for meta tensors.") |
| model.load_state_dict(new_weights, strict=False) |
| |
| return model |
| |
| @property |
| def pipe(self): |
| if self._pipeline is None: |
| model = self._load_model() |
| tokenizer = AutoTokenizer.from_pretrained(self.adapter_id, token=HF_TOKEN) |
| self._pipeline = pipeline( |
| "text-generation", |
| model=model, |
| tokenizer=tokenizer, |
| ) |
| return self._pipeline |
| |
| def generate(self, messages: list[dict], max_tokens: int, temperature: float) -> str: |
| """Generate a complete response""" |
| result = self.pipe( |
| messages, |
| max_new_tokens=max_tokens, |
| temperature=temperature, |
| do_sample=temperature > 0, |
| ) |
| return result[0]["generated_text"][-1]["content"] |
|
|
|
|
|
|
| def create_model(model_key: str) -> ModelWrapper: |
| """Factory function to create the appropriate model wrapper""" |
| |
| config = model_configs[model_key] |
| |
| if config["type"] == "endpoint": |
| return InferenceEndpointModel(config["id"]) |
| elif config["type"] == "pipeline": |
| return PipelineModel(config["base_model_id"], config["adapter_id"]) |
| else: |
| raise ValueError(f"Unknown model type: {config['type']}") |
|
|
|
|
| def get_model(model_key: str) -> ModelWrapper: |
| """Get or create a model instance""" |
| if model_key not in _model_cache: |
| _model_cache[model_key] = create_model(model_key) |
| return _model_cache[model_key] |
|
|