-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrun_pipeline.py
More file actions
94 lines (77 loc) · 3.38 KB
/
Copy pathrun_pipeline.py
File metadata and controls
94 lines (77 loc) · 3.38 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
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
import sys
import os
sys.path.append(os.path.dirname(os.path.dirname(__file__)))
import torch
from simple_model import SimpleModel
from parser.pytorch_parser import parse_pytorch_model
from passes.lowerer import hir_to_mir
from passes.quantize import quantization_pass
from passes.activation_quantize import activation_quantization_pass
from executor.run import run_mir_naive, run_mir_correct, verify_numerical_equivalence, debug_node_by_node_comparison, run_mir_quantized
from executor.quantized_kernels import run_mir_quantized_fast
from executor.optimized_executor import run_mir_optimized
from passes.per_channel_quantize import per_channel_quantization_pass
from passes.fusion import fusion_pass
from passes.profile import profile
from ir.hir import validate_hir
from ir.mir import validate_mir
def main():
model = SimpleModel()
model.eval()
# parse to HIR
hir = parse_pytorch_model(model, input_shape=(1,4))
hir_dict = hir.to_dict()
print("HIR:")
print(hir.to_json())
validate_hir(hir_dict)
# lower to MIR
mir = hir_to_mir(hir_dict)
from ir.mir import MIR
print("\nMIR (pre-quantize):")
import json
print(json.dumps(mir, indent=2))
validate_mir(mir)
# annotate quantization metadata
mir = quantization_pass(mir, model, num_bits=8)
# Create input tensor for testing
x = torch.randn(1,4)
# add activation quantization
mir = activation_quantization_pass(mir, model, x, num_bits=8)
print("\nMIR (post-weight + activation quantization metadata):")
print(json.dumps(mir.get("metadata", {}), indent=2))
# run (naive)
stats = profile(run_mir_naive, mir, model, x)
print("\nNaive execution stats:")
print(stats)
# run (correct) and verify numerical equivalence
print("\n" + "="*50)
print("NUMERICAL EQUIVALENCE VERIFICATION")
print("="*50)
verification = verify_numerical_equivalence(model, mir, x)
print(f"Original model output: {verification['original_output']}")
print(f"MIR executor output: {verification['mir_output']}")
print(f"Max difference: {verification['max_difference']:.2e}")
print(f"Numerically equivalent: {verification['equivalent']} (tolerance: {verification['tolerance']:.0e})")
# Uncomment the line below to run detailed node-by-node debugging:
# debug_results = debug_node_by_node_comparison(model, mir, x)
# Profile the correct executor
stats_correct = profile(run_mir_correct, mir, model, x)
print(f"\nCorrect execution stats:")
print(f"Runtime: {stats_correct['runtime_sec']:.4f}s")
print(f"Peak memory: {stats_correct['memory_peak_bytes']} bytes")
# run quantized execution
print("\n" + "="*50)
print("QUANTIZED EXECUTION (Weights Only)")
print("="*50)
quantized_stats = profile(run_mir_quantized, mir, model, x)
quantized_result = quantized_stats['result']
print(f"\nQuantized result: {quantized_result}")
# Compare quantized vs float results
original_output = verification['original_output']
max_diff_quantized = torch.max(torch.abs(original_output - quantized_result)).item()
print(f"Max difference (float vs quantized): {max_diff_quantized:.2e}")
print(f"\nQuantized execution stats:")
print(f"Runtime: {quantized_stats['runtime_sec']:.4f}s")
print(f"Peak memory: {quantized_stats['memory_peak_bytes']} bytes")
if __name__ == "__main__":
main()