1from transformers import AutoModelForCausalLM, AutoTokenizer, GPT2Config
2import torch
3from huggingface_hub import hf_hub_download
4
5def load_moe_model(checkpoint_name="best_val_loss_moe_step_9000.bin", model_id="idhant297/moe-5l-total-test"):
6 """
7 Load a MoE model from HuggingFace Hub with a specific checkpoint.
8
9 Args:
10 checkpoint_name (str): The checkpoint filename to load
11 model_id (str): The HuggingFace model repository ID
12
13 Returns:
14 tuple: (model, tokenizer) loaded from the checkpoint
15 """
16 print(f"Loading MoE model from {model_id} checkpoint {checkpoint_name}...")
17
18 tokenizer = AutoTokenizer.from_pretrained(model_id)
19
20 config = GPT2Config.from_pretrained(model_id)
21
22 model = AutoModelForCausalLM.from_config(config)
23
24 checkpoint_path = hf_hub_download(
25 repo_id=model_id,
26 filename=checkpoint_name
27 )
28
29 state_dict = torch.load(checkpoint_path, map_location="cpu")
30 model.load_state_dict(state_dict)
31 model.eval()
32
33 print(f"✅ MoE model loaded successfully from checkpoint {checkpoint_name}")
34 return model, tokenizer
35
36def generate_text_moe(model, tokenizer, prompt, max_length=100, temperature=0.8, top_p=0.95, num_return_sequences=1):
37 """
38 Generate text using the loaded MoE model.
39
40 Args:
41 model: The loaded MoE model
42 tokenizer: The loaded tokenizer
43 prompt (str): Input text prompt
44 max_length (int): Maximum length of generated text
45 temperature (float): Sampling temperature
46 top_p (float): Top-p sampling parameter
47 num_return_sequences (int): Number of sequences to generate
48
49 Returns:
50 list: Generated text sequences
51 """
52 inputs = tokenizer(prompt, return_tensors="pt")
53
54 with torch.no_grad():
55 outputs = model.generate(
56 inputs["input_ids"],
57 max_length=max_length,
58 temperature=temperature,
59 top_p=top_p,
60 do_sample=True,
61 num_return_sequences=num_return_sequences,
62 pad_token_id=tokenizer.eos_token_id
63 )
64
65 generated_texts = []
66 for output in outputs:
67 text = tokenizer.decode(output, skip_special_tokens=True)
68 generated_texts.append(text)
69
70 return generated_texts
71
72# Example usage for MoE model
73checkpoint_name = "best_val_loss_moe_step_8400.bin"
74model, tokenizer = load_moe_model(checkpoint_name)
75
76prompt = "hello?"
77
78generated = generate_text_moe(model, tokenizer, prompt, max_length=50)
79print(generated)