CoolFace
Modelpublic

dallred/Mamba-Chat-2.8B

sourceHugging Faceapache-2.0updated 7mo agoView on Hugging Face
0likes15downloads
handler.py53 linesDownload Raw Back to root
1 2import torch3from transformers import AutoTokenizer4from mamba_ssm.models.mixer_seq_simple import MambaLMHeadModel5 6 7class Pipeline:8    def __init__(self, model_id: str, **kwargs):9        # Load tokenizer10        self.tokenizer = AutoTokenizer.from_pretrained("havenhq/mamba-chat")11        self.tokenizer.eos_token = "<|endoftext|>"12        self.tokenizer.pad_token = self.tokenizer.eos_token13 14        # Load Zephyr chat template15        zephyr_tok = AutoTokenizer.from_pretrained("HuggingFaceH4/zephyr-7b-beta")16        self.tokenizer.chat_template = zephyr_tok.chat_template17 18        # Load model (CUDA as shown in example)19        self.device = "cuda" if torch.cuda.is_available() else "cpu"20        self.model = MambaLMHeadModel.from_pretrained(21            model_id,22            device=self.device,23            dtype=torch.float16 if self.device == "cuda" else torch.float32,24        )25 26    def __call__(self, inputs: str):27        # Build chat message list28        messages = [{"role": "user", "content": inputs}]29 30        # Apply chat template31        input_ids = self.tokenizer.apply_chat_template(32            messages,33            return_tensors="pt",34            add_generation_prompt=True35        ).to(self.device)36 37        # Generate38        output = self.model.generate(39            input_ids=input_ids,40            max_length=2000,41            temperature=0.9,42            top_p=0.7,43            eos_token_id=self.tokenizer.eos_token_id,44        )45 46        decoded = self.tokenizer.batch_decode(output)[0]47 48        # Extract assistant response49        if "<|assistant|>\n" in decoded:50            decoded = decoded.split("<|assistant|>\n")[-1]51 52        return decoded.strip()53