model42M-v1 / examples /generate_example.py
Heet24's picture
Upload folder using huggingface_hub
5c7f987 verified
Raw
History Blame Contribute Delete
5.21 kB
import os
import sys
import torch
import sentencepiece as spm
# Add the parent directory of this script to the python path
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)
from model.config import VaghLMConfig
from model.transformer import VaghLM
def run_generation():
# 1. Initialize configuration and setup device
config = VaghLMConfig()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
# 2. Load the SentencePiece Tokenizer
tokenizer_path = os.path.join(ROOT, "tokenizer", "vocab", "tokenizer.model")
if not os.path.exists(tokenizer_path):
raise FileNotFoundError(f"Tokenizer not found at {tokenizer_path}. Please check your installation.")
sp = spm.SentencePieceProcessor()
sp.load(tokenizer_path)
print(f"[OK] Tokenizer loaded successfully (Vocab size: {sp.vocab_size()})")
# 3. Initialize Model Architecture
model = VaghLM(config)
# 4. Load the Model Weights (checkpoints)
# Check for finetuned_v2_best.pt first, then fallback to best_checkpoint.pt
checkpoint_names = ["finetuned_v2_best.pt", "best_checkpoint.pt"]
checkpoint_path = None
for ckpt_name in checkpoint_names:
p = os.path.join(ROOT, "checkpoints", ckpt_name)
if os.path.exists(p):
checkpoint_path = p
break
if checkpoint_path is None:
print("[WARNING] No pre-trained checkpoint found in checkpoints/. Initializing with random weights.")
else:
print(f"Loading checkpoint: {checkpoint_path}")
ckpt = torch.load(checkpoint_path, map_location=device)
model.load_state_dict(ckpt["model_state"])
print(f"[OK] Model weights loaded (Step: {ckpt.get('step', 'N/A')}, Loss: {ckpt.get('loss', 'N/A'):.4f})")
model.to(device)
model.eval()
print(f"Model parameters: {model.count_parameters():,}")
# 5. Define Special Tokens
BOS_ID = config.bos_id
USER_ID = config.user_id
ASSISTANT_ID = config.assistant_id
EOS_ID = config.eos_id
# 6. Set Generation Parameters
prompt = "Explain what a black hole is."
max_new_tokens = 150
temperature = 0.7
top_k = 40
top_p = 0.9
repetition_penalty = 1.3
print(f"\nPrompt: {prompt}")
print(f"Generation Parameters: temp={temperature}, top_k={top_k}, top_p={top_p}, rep_penalty={repetition_penalty}")
print("-" * 50)
# 7. Format the input with special tokens
user_tokens = sp.encode(prompt)
input_ids = [BOS_ID, USER_ID] + user_tokens + [ASSISTANT_ID]
x = torch.tensor([input_ids], dtype=torch.long, device=device)
generated = []
# 8. Generation Loop
with torch.no_grad():
for _ in range(max_new_tokens):
# Truncate input if it exceeds the model's max context length
if x.size(1) > config.context_length:
x = x[:, -config.context_length:]
# Forward pass to obtain logits for the next token
logits, _ = model(x)
logits = logits[:, -1, :] # Select only the last sequence position
# Apply Repetition Penalty
for token_id in set(x[0].tolist()):
if logits[0, token_id] > 0:
logits[0, token_id] /= repetition_penalty
else:
logits[0, token_id] *= repetition_penalty
# Apply Temperature
if temperature > 0:
logits = logits / temperature
# Apply Top-k Filtering
if top_k is not None and top_k > 0:
values, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < values[:, [-1]]] = -float("inf")
# Apply Top-p (Nucleus) Filtering
if top_p is not None and top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
# Mask tokens that fall outside cumulative probability threshold
sorted_indices_to_remove = (cumulative_probs - torch.softmax(sorted_logits, dim=-1)) > top_p
sorted_logits[sorted_indices_to_remove] = -float("inf")
# Restore original indexing
logits = torch.zeros_like(logits).scatter_(1, sorted_indices, sorted_logits)
# Sample from the probability distribution
if temperature == 0:
next_token = torch.argmax(logits, dim=-1, keepdim=True)
else:
probs = torch.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
x = torch.cat([x, next_token], dim=1)
token_val = next_token.item()
generated.append(token_val)
# Break generation if End Of Sequence token is encountered
if token_val == EOS_ID:
break
# 9. Decode and print output
response = sp.decode(generated)
print(f"\nResponse:\n{response}\n")
if __name__ == "__main__":
run_generation()