1+ import torch
2+ import torch .nn as nn
3+ import torch .optim as optim
4+ import torchvision .datasets as datasets
5+ import torchvision .transforms as transforms
6+ import os
7+ import CoOccurringFD
8+ # Define your custom matrix multiplication function
9+ def my_matmul (x , y ):
10+ rows ,cols = x .shape
11+ if cols > 20 :
12+ sketchSize = cols / 10
13+ else :
14+ sketchSize = 10
15+ # Your implementation here
16+ return CoOccurringFD .FDAMM (x ,y ,int (sketchSize ))
17+
18+ # Define a custom Linear layer that uses your custom matrix multiplication function
19+ class CustomLinear (nn .Module ):
20+ def __init__ (self , in_features , out_features , bias = True ):
21+ super ().__init__ ()
22+ self .in_features = in_features
23+ self .out_features = out_features
24+ self .weight = nn .Parameter (torch .Tensor (out_features , in_features ))
25+ if bias :
26+ self .bias = nn .Parameter (torch .Tensor (out_features ))
27+ else :
28+ self .register_parameter ('bias' , None )
29+ self .reset_parameters ()
30+ def mySqrt (self ,a :float ):
31+ y = torch .sqrt (torch .tensor (a , dtype = torch .float32 ))
32+ return y .item ()
33+ def reset_parameters (self ):
34+ nn .init .kaiming_uniform_ (self .weight , a = self .mySqrt (5.0 ))
35+ if self .bias is not None :
36+ fan_in , _ = nn .init ._calculate_fan_in_and_fan_out (self .weight )
37+ bound = 1 / self .mySqrt (fan_in )
38+ nn .init .uniform_ (self .bias , - bound , bound )
39+
40+ def forward (self , input ):
41+ # Use your custom matrix multiplication function instead of torch.matmul
42+ output = my_matmul (input , self .weight .t ())
43+ if self .bias is not None :
44+ output += self .bias
45+ return output
46+
47+ # Define your neural network architecture
48+ class MyNet (nn .Module ):
49+ def __init__ (self ):
50+ super (MyNet , self ).__init__ ()
51+ self .fc1 = nn .Linear (784 , 128 )
52+ self .fc2 = nn .Linear (128 , 128 )
53+ self .fc3 = nn .Linear (128 , 10 )
54+ def forward (self , x ):
55+ x = x .view (- 1 , 784 )
56+ x = nn .functional .relu (self .fc1 (x ))
57+ x = self .fc2 (x )
58+ x = nn .functional .relu (self .fc3 (x ))
59+ return x
60+ def testNN (net ,test_loader ):
61+ #first, load parameters
62+ pretrained_params = torch .load ('pretrained_model.pt' )
63+ custom_params = net .state_dict ()
64+
65+ for name in custom_params :
66+ if name in pretrained_params :
67+ custom_params [name ] = pretrained_params [name ]
68+ net .load_state_dict (custom_params )
69+ correct = 0
70+ total = 0
71+ #then, run test
72+ net2 = net
73+ for data in test_loader :
74+ images , labels = data
75+ outputs = net2 (images )
76+ _ , predicted = torch .max (outputs .data , 1 )
77+ total += labels .size (0 )
78+ correct += (predicted == labels ).sum ().item ()
79+ print (f"Accuracy on test set: { correct / total } " )
80+ print (f"Accuracy on test set: { correct / total } " )
81+ return correct / total
82+ def main ():
83+ device = 'cuda'
84+ # Load the MNIST dataset
85+ train_dataset = datasets .MNIST (root = './data' , train = True , transform = transforms .ToTensor (), download = True )
86+ test_dataset = datasets .MNIST (root = './data' , train = False , transform = transforms .ToTensor ())
87+
88+ # Set up the data loaders
89+ batch_size = 64
90+ train_loader = torch .utils .data .DataLoader (train_dataset , batch_size = batch_size , shuffle = True )
91+ test_loader = torch .utils .data .DataLoader (test_dataset , batch_size = batch_size , shuffle = False )
92+
93+ # Train the neural network using the default Linear layers
94+ net = MyNet ()
95+
96+ if os .path .exists ('pretrained_model.pt' ):
97+ print ('find pretrained model, run test' )
98+ print ('first run default version' )
99+ accuracy0 = testNN (net ,test_loader )
100+ print ('then run coocuuring 1 version' )
101+ # Replace the Linear layers with your custom Linear layers and load the pre-trained weights
102+ net .fc1 = CustomLinear (784 , 128 )
103+ accuracy1 = testNN (net ,test_loader )
104+ print ('next run coocuuring 2 version' )
105+ net .fc1 = nn .Linear (784 ,128 )
106+ net .fc2 = CustomLinear (128 , 128 )
107+ accuracy2 = testNN (net ,test_loader )
108+ print ('finally run coocuuring 3 version' )
109+ net .fc2 = nn .Linear (128 , 128 )
110+ net .fc3 = CustomLinear (128 , 10 )
111+ accuracy3 = testNN (net ,test_loader )
112+ print ('default accuracy=' ,accuracy0 )
113+ print ('co-occuring 1 accuracy=' ,accuracy1 )
114+ print ('co-occuring 2 accuracy=' ,accuracy2 )
115+ print ('co-occuring 3 accuracy=' ,accuracy3 )
116+
117+
118+ else :
119+ print ('build pretrain model first' )
120+ criterion = nn .CrossEntropyLoss ()
121+ optimizer = optim .SGD (net .parameters (), lr = 0.1 )
122+ net = net .to (device )
123+ for epoch in range (10 ):
124+ running_loss = 0.0
125+ for i , data in enumerate (train_loader , 0 ):
126+ inputs , labels = data
127+ inputs = inputs .to (device )
128+ labels = labels .to (device )
129+ optimizer .zero_grad ()
130+ outputs = net (inputs )
131+ loss = criterion (outputs , labels )
132+ loss .backward ()
133+ optimizer .step ()
134+ running_loss += loss .item ()
135+ print (f"Epoch { epoch + 1 } : loss = { running_loss / len (train_loader )} " )
136+ net = net .to ('cpu' )
137+ # Save the pre-trained model
138+ torch .save (net .state_dict (), 'pretrained_model.pt' )
139+
140+
141+
142+
143+
144+ # Evaluate
145+ if __name__ == '__main__' :
146+ main ()
0 commit comments