1from transformers import AutoModelForCausalLM, AutoTokenizer, GPT2Config
2import torch
3from huggingface_hub import hf_hub_download
4
5def load_model(checkpoint_name="best_val_loss_dense_step_9000.bin", model_id="idhant297/dense-5l-test"):
6 """
7 Load a model from HuggingFace Hub with a specific checkpoint.
8
9 Args:
10 checkpoint_step (int): The training step checkpoint 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 model from {model_id} at step {checkpoint_step}...")
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_filename
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"Model loaded successfully from checkpoint step {checkpoint_name}")
34 return model, tokenizer
35
36def generate_text(model, tokenizer, prompt, max_length=100, temperature=0.8, top_p=0.95, num_return_sequences=1):
37 """
38 Generate text using the loaded model.
39
40 Args:
41 model: The loaded 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
73checkpoint_name = "best_val_loss_dense_step_9000.bin"
74model, tokenizer = load_model(checkpoint_name)
75
76prompt = "The quick brown fox"
77
78generated = generate_text(model, tokenizer, prompt, max_length=20)
79print(generated)