J-space / math_jlens /model.py
ayh015's picture
Upload folder using huggingface_hub
85b17bd verified
Raw
History Blame Contribute Delete
1.58 kB
"""Minimal Hugging Face adapter for Qwen decoder-only models."""
from __future__ import annotations
from typing import Any
import torch
from torch import nn
class QwenLensModel:
def __init__(self, model: nn.Module, tokenizer: Any) -> None:
self.model = model
self.tokenizer = tokenizer
self.decoder = model.model
self.layers = self.decoder.layers
self.n_layers = len(self.layers)
self.d_model = model.config.hidden_size
model.eval()
for parameter in model.parameters():
parameter.requires_grad_(False)
@property
def input_device(self) -> torch.device:
return self.decoder.embed_tokens.weight.device
def forward(self, input_ids: torch.Tensor) -> Any:
return self.decoder(input_ids=input_ids, use_cache=False)
def encode(self, text: str, max_length: int) -> torch.Tensor:
encoded = self.tokenizer(
text, return_tensors="pt", truncation=True, max_length=max_length,
add_special_tokens=True,
)
return encoded.input_ids.to(self.input_device)
def load_qwen(model_path: str, *, device: str, dtype: torch.dtype) -> QwenLensModel:
"""Load a complete local checkpoint without network requests."""
from transformers import AutoModelForCausalLM, AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(model_path, local_files_only=True)
model = AutoModelForCausalLM.from_pretrained(
model_path, local_files_only=True, torch_dtype=dtype
).to(device)
return QwenLensModel(model, tokenizer)