dallred/Mamba-Chat-2.8B
015
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 