-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathevaluate.py
More file actions
78 lines (64 loc) · 2.57 KB
/
Copy pathevaluate.py
File metadata and controls
78 lines (64 loc) · 2.57 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
#!/usr/bin/env python3
"""
HiSVD Model Evaluation
Evaluate perplexity for HuggingFace or compressed (.pt) models
on WikiText-2 and C4 benchmarks.
"""
import os
import argparse
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
os.environ.setdefault('HF_DATASETS_OFFLINE', '1')
os.environ.setdefault('HF_HUB_OFFLINE', '1')
os.environ.setdefault('TRANSFORMERS_OFFLINE', '1')
from utils.eval_utils import evaluate_ppl
def main():
parser = argparse.ArgumentParser(description='HiSVD Model Evaluation')
parser.add_argument('--model_path', type=str, required=True,
help='HuggingFace model path or compressed .pt file')
parser.add_argument('--datasets', type=str, default='wikitext2,c4',
help='Comma-separated dataset names')
parser.add_argument('--batch_size', type=int, default=4)
parser.add_argument('--seq_len', type=int, default=2048)
parser.add_argument('--device', type=str, default='cuda')
args = parser.parse_args()
print("=" * 80)
print("HiSVD Evaluation")
print("=" * 80)
print(f" Model: {args.model_path}")
print(f" Datasets: {args.datasets}")
if args.model_path.endswith('.pt'):
print(" Loading compressed model (.pt)...")
data = torch.load(args.model_path, weights_only=False, map_location='cpu')
model = data['model']
tokenizer = data['tokenizer']
model = model.to(args.device)
if 'ppl_results' in data:
print(f" Stored PPL: {data['ppl_results']}")
if 'ranks' in data:
total_sublayers = sum(len(v) for v in data['ranks'].values())
print(f" Compressed sublayers: {total_sublayers}")
else:
print(" Loading HuggingFace model...")
model = AutoModelForCausalLM.from_pretrained(
args.model_path, torch_dtype=torch.float16,
low_cpu_mem_usage=True, local_files_only=True
).to(args.device)
try:
tokenizer = AutoTokenizer.from_pretrained(args.model_path, local_files_only=True)
except Exception:
tokenizer = AutoTokenizer.from_pretrained(args.model_path)
model.seqlen = args.seq_len
ds_list = [d.strip() for d in args.datasets.split(',')]
results = evaluate_ppl(
model, tokenizer, device=args.device,
batch_size=args.batch_size, datasets=ds_list, seq_len=args.seq_len
)
print(f"\n{'=' * 80}")
print("Results")
print(f"{'=' * 80}")
for ds, ppl in results.items():
print(f" {ds}: {ppl:.2f}")
print(f"{'=' * 80}")
if __name__ == '__main__':
main()