pmnet / copy_paste /visualize_copy_paste.py
phasorkinetics's picture
Upload 32 files
b24b632 verified
Raw
History Blame Contribute Delete
4.36 kB
import json
import os
import torch
import numpy as np
import matplotlib.pyplot as plt
from tqdm import tqdm
from transformers import AutoConfig, AutoModelForCausalLM
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
torch.set_default_dtype(torch.float32)
print(f"Using device: {DEVICE}")
PMNET_MODEL_ID = "phasorkinetics/pmnet"
PMNET_CONFIG_PATH = "config_pmnet.json"
PMNET_CKPT = "copy_model_pmnet.safetensors"
PMNET_NO_MEM_CONFIG_PATH = "config_pmnet_no_mem.json"
PMNET_NO_MEM_CKPT = "copy_model_pmnet_no_mem.safetensors"
VOCAB_SIZE = 250
DELIM_TOKEN = 251
SEQ_LENGTHS = [10, 64, 128, 256, 512, 768, 896, 1024, 1088, 1152, 1280]
NUM_SAMPLES = 100
BATCH_SIZE = 10
def load_eval_model(config_path, ckpt_path):
print(f"Loading model from {ckpt_path}...")
with open(config_path, "r", encoding="utf-8") as f:
config_dict = json.load(f)
config = AutoConfig.from_pretrained(PMNET_MODEL_ID, trust_remote_code=True)
config.update(config_dict)
model = AutoModelForCausalLM.from_config(config, trust_remote_code=True)
from safetensors.torch import load_file
state_dict = load_file(ckpt_path)
model.load_state_dict(state_dict, strict=False)
model.float().to(DEVICE)
model.eval()
return model
def evaluate_copy_paste(model, seq_lengths, num_samples, batch_size):
accuracies = []
with torch.no_grad():
for seq_len in seq_lengths:
print(f"Evaluating Sequence Length: {seq_len}...")
total_correct = 0
total_tokens = 0
for _ in range(0, num_samples, batch_size):
current_batch_size = min(batch_size, num_samples - _)
orig_seq = torch.randint(1, VOCAB_SIZE + 1, (current_batch_size, seq_len), device=DEVICE)
delim = torch.full((current_batch_size, 1), DELIM_TOKEN, device=DEVICE)
input_ids = torch.cat([orig_seq, delim, orig_seq], dim=1)
outputs = model(input_ids)
logits = outputs.logits
shifted_logits = logits[:, :-1, :].contiguous()
shifted_labels = input_ids[:, 1:].contiguous()
target_logits = shifted_logits[:, seq_len:, :]
target_labels = shifted_labels[:, seq_len:]
preds = target_logits.argmax(dim=-1)
correct = (preds == target_labels).sum().item()
total_correct += correct
total_tokens += (current_batch_size * seq_len)
acc = (total_correct / total_tokens) * 100.0
accuracies.append(acc)
print(f" -> Accuracy for len {seq_len}: {acc:.2f}%")
return accuracies
def main():
model_pmnet = load_eval_model(PMNET_CONFIG_PATH, PMNET_CKPT)
acc_pmnet = evaluate_copy_paste(model_pmnet, SEQ_LENGTHS, NUM_SAMPLES, BATCH_SIZE)
del model_pmnet
torch.cuda.empty_cache()
model_no_mem = load_eval_model(PMNET_NO_MEM_CONFIG_PATH, PMNET_NO_MEM_CKPT)
acc_no_mem = evaluate_copy_paste(model_no_mem, SEQ_LENGTHS, NUM_SAMPLES, BATCH_SIZE)
del model_no_mem
torch.cuda.empty_cache()
plt.figure(figsize=(8, 6))
plt.plot(SEQ_LENGTHS, acc_pmnet, marker='o', linestyle='-', color='seagreen', linewidth=2.5, label='PMNet (Full Memory)')
plt.plot(SEQ_LENGTHS, acc_no_mem, marker='s', linestyle='--', color='mediumvioletred', linewidth=2.5, label='PMNet (No Memory)')
# plt.axvline(x=128, color='gray', linestyle=':', linewidth=2, label='SWA Window Size (128)')
plt.title('Copy-Paste Task: Accuracy', fontsize=16, fontweight='bold')
plt.xlabel('Sequence Length to Memorize (Bytes)', fontsize=14)
plt.ylabel('Exact Token Accuracy (%)', fontsize=14)
# plt.xscale('log', base=2)
plt.xticks(SEQ_LENGTHS, [str(x) for x in SEQ_LENGTHS], rotation=35, ha='right', fontsize=12)
plt.yticks(np.arange(0, 101, 10), fontsize=12)
plt.ylim(-5, 105)
plt.grid(True, which="both", ls="--", alpha=0.5)
plt.legend(fontsize=12, loc='lower left')
plt.tight_layout()
save_path = "copy_paste.pdf"
plt.savefig(save_path, dpi=300)
print(f"\nEvaluation complete! Graph saved to {save_path}")
if __name__ == "__main__":
main()