perforatedai.blockwise_perforatedai
1# Copyright (c) 2025 Perforated AI 2 3from perforatedai import globals_perforatedai as GPA 4from perforatedai import modules_perforatedai as PA 5from perforatedai import utils_perforatedai as UPA 6import torch.nn as nn 7import torch 8import pdb 9import numpy as np 10import string 11import copy 12 13# This is the cleaner inference version of PAI modules 14class PAILayer(nn.Module): 15 def __init__( 16 self, 17 layer_array, 18 processor_array, 19 dendrites_to_top, 20 dendrites_to_dendrites, 21 node_index, 22 num_cycles, 23 view_tuple, 24 ): 25 """Initialize an inference-time PAI layer block. 26 This functionaly is the same but has dendritic scaffolding cleaned up 27 for faster infrance and smaller memory footprint. 28 29 Parameters 30 ---------- 31 layer_array : nn.Module 32 Ordered module stack used for dendrite and neuron modules. 33 processor_array : list 34 Processor objects aligned with ``layer_array``. 35 dendrites_to_top : nn.ParameterList or None 36 Connection weights from dendrites to the top module. 37 dendrites_to_dendrites : nn.ParameterList or None 38 Connection weights between dendrite modules. 39 node_index : int 40 Index of the feature dimension that corresponds to nodes. 41 num_cycles : torch.Tensor 42 Number of dendrite cycles simulate to simulate in order to set up architecture.. 43 view_tuple : tuple 44 Broadcast shape used for dendrite weights. 45 """ 46 super(PAILayer, self).__init__() 47 self.layer_array = layer_array 48 self.register_buffer("num_cycles", num_cycles) 49 self.register_buffer("view_tuple", torch.tensor(view_tuple)) 50 51 self.processor_array = processor_array 52 if dendrites_to_dendrites: 53 self.skip_weights = dendrites_to_dendrites 54 else: 55 """ 56 This will only be the case if there is less than 2 dendrites, in these cases an empty array 57 should still be added so that dendrites_to_top is included at the correct index 58 """ 59 self.skip_weights = nn.ParameterList() 60 if dendrites_to_top: 61 self.skip_weights.append(dendrites_to_top[len(dendrites_to_top) - 1]) 62 else: 63 self.skip_weights = nn.ParameterList() 64 65 # Delete skip_weights if it's empty (only 1 layer, no skip connections) 66 if len(self.skip_weights) == 0: 67 delattr(self, 'skip_weights') 68 69 self.node_index = node_index 70 self.internal_nonlinearity = GPA.pc.get_pai_forward_function() 71 72def unWrap_params(model): 73 """Remove wrapped parameter attributes from a model in-place. 74 75 Parameters 76 ---------- 77 model : nn.Module 78 Model whose parameters may contain a temporary ``wrapped`` attribute. 79 80 Returns 81 ------- 82 None 83 This function does not return a value. 84 """ 85 for p in model.parameters(): 86 if "wrapped" in p.__dir__(): 87 del p.wrapped 88 89# This converts one training PAI module into an inference PAI module 90def convert_to_pai_layer_block(pretrained_dendrite): 91 """Convert a training-time PAI neuron module to an inference PAI layer. 92 93 Parameters 94 ---------- 95 pretrained_dendrite : PAINeuronModule 96 Trained neuron module containing dendrite layers and processors. 97 98 Returns 99 ------- 100 PAILayer 101 Inference-ready wrapper with converted processors and dendrite connections. 102 """ 103 unWrap_params(pretrained_dendrite) 104 layer_array = [] 105 processor_array = [] 106 for layer_id in range(len(pretrained_dendrite.dendrite_module.layers)): 107 layer_array.append(pretrained_dendrite.dendrite_module.layers[layer_id]) 108 if pretrained_dendrite.dendrite_module.processors == []: 109 processor_array.append(None) 110 else: 111 if not pretrained_dendrite.dendrite_module.processors[layer_id] is None: 112 pretrained_dendrite.dendrite_module.processors[layer_id].pre = ( 113 pretrained_dendrite.dendrite_module.processors[layer_id].pre_d 114 ) 115 pretrained_dendrite.dendrite_module.processors[layer_id].post = ( 116 pretrained_dendrite.dendrite_module.processors[layer_id].post_d 117 ) 118 processor_array.append( 119 pretrained_dendrite.dendrite_module.processors[layer_id] 120 ) 121 layer_array.append(pretrained_dendrite.main_module) 122 if not pretrained_dendrite.processor is None: 123 pretrained_dendrite.processor.pre = pretrained_dendrite.processor.post_n1 124 pretrained_dendrite.processor.post = pretrained_dendrite.processor.post_n2 125 processor_array.append(pretrained_dendrite.processor) 126 127 view_tuple = [] 128 for dim in range( 129 len( 130 pretrained_dendrite.dendrite_module.dendrite_values[0].this_output_dimensions 131 ) 132 ): 133 if ( 134 dim 135 == pretrained_dendrite.dendrite_module.dendrite_values[0].this_node_index 136 ): 137 view_tuple.append(-1) 138 continue 139 view_tuple.append(1) 140 return PAILayer( 141 nn.Sequential(*layer_array), 142 processor_array, 143 pretrained_dendrite.dendrites_to_top, 144 pretrained_dendrite.dendrite_module.dendrites_to_dendrites, 145 pretrained_dendrite.this_node_index, 146 pretrained_dendrite.dendrite_module.num_cycles, 147 view_tuple, 148 ) 149 150 151def get_pretrained_pai_attr(pretrained_dendrite, member): 152 """Safely get an attribute from a possibly missing module. 153 154 Parameters 155 ---------- 156 pretrained_dendrite : nn.Module or None 157 Source module that may be ``None``. 158 member : str 159 Attribute name to retrieve. 160 161 Returns 162 ------- 163 Any 164 Requested attribute value, or ``None`` when the module is ``None``. 165 """ 166 if pretrained_dendrite is None: 167 return None 168 else: 169 return getattr(pretrained_dendrite, member) 170 171 172def get_pretrained_pai_var(pretrained_dendrite, submodule_id): 173 """Safely get a submodule by index/key from a possibly missing module. 174 175 Parameters 176 ---------- 177 pretrained_dendrite : nn.Module or None 178 Source module that may be ``None``. 179 submodule_id : str or int 180 Submodule identifier to retrieve. 181 182 Returns 183 ------- 184 nn.Module or None 185 Retrieved submodule, or ``None`` when the source module is ``None``. 186 """ 187 if pretrained_dendrite is None: 188 return None 189 else: 190 return pretrained_dendrite[submodule_id] 191 192# This optimizes a network recursively from training modules to inference modules 193def optimize_module(net, depth, name_so_far, converted_list): 194 """Recursively replace training PAI modules with inference PAI layers. 195 196 Parameters 197 ---------- 198 net : nn.Module 199 Module tree to traverse and optimize. 200 depth : int 201 Current recursion depth. 202 name_so_far : str 203 Dotted/Indexed path to the current module. 204 converted_list : list 205 Mutable list of module names already converted. 206 207 Returns 208 ------- 209 nn.Module 210 Updated module tree with converted inference modules. 211 """ 212 all_members = net.__dir__() 213 if issubclass(type(net), nn.Sequential) or issubclass(type(net), nn.ModuleList): 214 for submodule_id, layer in net.named_children(): 215 if type(net.get_submodule(submodule_id)) is PA.PAINeuronModule: 216 if GPA.pc.get_extra_verbose(): 217 print( 218 "Seq sub is PAI so optimizing: %s" % name_so_far 219 + "[" 220 + str(submodule_id) 221 + "]" 222 ) 223 setattr( 224 net, 225 submodule_id, 226 convert_to_pai_layer_block(net.get_submodule(submodule_id)), 227 ) 228 else: 229 if net != net.get_submodule(submodule_id): 230 # this currently just always returns false, not sure what it was for 231 converted_list += [name_so_far + "[" + str(submodule_id) + "]"] 232 setattr( 233 net, 234 submodule_id, 235 optimize_module( 236 net.get_submodule(submodule_id), 237 depth + 1, 238 name_so_far + "[" + str(submodule_id) + "]", 239 converted_list, 240 ), 241 ) 242 else: 243 if GPA.pc.get_extra_verbose(): 244 print( 245 "%s is a self pointer so skipping" 246 % (name_so_far + "[" + str(submodule_id) + "]") 247 ) 248 else: 249 for member in all_members: 250 if isinstance(getattr(type(net), member, None), property): 251 continue 252 try: 253 getattr(net, member, None) 254 except: 255 continue 256 sub_name = name_so_far + "." + member 257 if ( 258 sub_name in GPA.pc.get_module_names_to_not_save() 259 or sub_name in converted_list 260 ): 261 if GPA.pc.get_extra_verbose(): 262 print("Skipping %s during save" % sub_name) 263 continue 264 if type(getattr(net, member, None)) is PA.PAINeuronModule: 265 if GPA.pc.get_extra_verbose(): 266 print( 267 "Sub is in conversion list so initiating optimization for: %s" 268 % name_so_far 269 + "." 270 + member 271 ) 272 setattr(net, member, convert_to_pai_layer_block(getattr(net, member))) 273 elif issubclass(type(getattr(net, member, None)), nn.Module): 274 if net != getattr(net, member): 275 converted_list += [sub_name] 276 setattr( 277 net, 278 member, 279 optimize_module( 280 getattr(net, member), 281 depth + 1, 282 sub_name, 283 converted_list, 284 ), 285 ) 286 else: 287 if GPA.pc.get_extra_verbose(): 288 print("%s is a self pointer so skipping" % (sub_name)) 289 return net 290 291def blockwise_network(net): 292 """Convert all eligible modules in a network to blockwise inference form. 293 294 Parameters 295 ---------- 296 net : nn.Module 297 Input model. 298 299 Returns 300 ------- 301 nn.Module 302 Converted model. 303 """ 304 return optimize_module(net, 0, "", [])
15class PAILayer(nn.Module): 16 def __init__( 17 self, 18 layer_array, 19 processor_array, 20 dendrites_to_top, 21 dendrites_to_dendrites, 22 node_index, 23 num_cycles, 24 view_tuple, 25 ): 26 """Initialize an inference-time PAI layer block. 27 This functionaly is the same but has dendritic scaffolding cleaned up 28 for faster infrance and smaller memory footprint. 29 30 Parameters 31 ---------- 32 layer_array : nn.Module 33 Ordered module stack used for dendrite and neuron modules. 34 processor_array : list 35 Processor objects aligned with ``layer_array``. 36 dendrites_to_top : nn.ParameterList or None 37 Connection weights from dendrites to the top module. 38 dendrites_to_dendrites : nn.ParameterList or None 39 Connection weights between dendrite modules. 40 node_index : int 41 Index of the feature dimension that corresponds to nodes. 42 num_cycles : torch.Tensor 43 Number of dendrite cycles simulate to simulate in order to set up architecture.. 44 view_tuple : tuple 45 Broadcast shape used for dendrite weights. 46 """ 47 super(PAILayer, self).__init__() 48 self.layer_array = layer_array 49 self.register_buffer("num_cycles", num_cycles) 50 self.register_buffer("view_tuple", torch.tensor(view_tuple)) 51 52 self.processor_array = processor_array 53 if dendrites_to_dendrites: 54 self.skip_weights = dendrites_to_dendrites 55 else: 56 """ 57 This will only be the case if there is less than 2 dendrites, in these cases an empty array 58 should still be added so that dendrites_to_top is included at the correct index 59 """ 60 self.skip_weights = nn.ParameterList() 61 if dendrites_to_top: 62 self.skip_weights.append(dendrites_to_top[len(dendrites_to_top) - 1]) 63 else: 64 self.skip_weights = nn.ParameterList() 65 66 # Delete skip_weights if it's empty (only 1 layer, no skip connections) 67 if len(self.skip_weights) == 0: 68 delattr(self, 'skip_weights') 69 70 self.node_index = node_index 71 self.internal_nonlinearity = GPA.pc.get_pai_forward_function()
Base class for all neural network modules.
Your models should also subclass this class.
Modules can also contain other Modules, allowing them to be nested in a tree structure. You can assign the submodules as regular attributes::
import torch.nn as nn
import torch.nn.functional as F
class Model(nn.Module):
def __init__(self) -> None:
super().__init__()
self.conv1 = nn.Conv2d(1, 20, 5)
self.conv2 = nn.Conv2d(20, 20, 5)
def forward(self, x):
x = F.relu(self.conv1(x))
return F.relu(self.conv2(x))
Submodules assigned in this way will be registered, and will also have their
parameters converted when you call to(), etc.
As per the example above, an __init__() call to the parent class
must be made before assignment on the child.
:ivar training: Boolean represents whether this module is in training or evaluation mode. :vartype training: bool
16 def __init__( 17 self, 18 layer_array, 19 processor_array, 20 dendrites_to_top, 21 dendrites_to_dendrites, 22 node_index, 23 num_cycles, 24 view_tuple, 25 ): 26 """Initialize an inference-time PAI layer block. 27 This functionaly is the same but has dendritic scaffolding cleaned up 28 for faster infrance and smaller memory footprint. 29 30 Parameters 31 ---------- 32 layer_array : nn.Module 33 Ordered module stack used for dendrite and neuron modules. 34 processor_array : list 35 Processor objects aligned with ``layer_array``. 36 dendrites_to_top : nn.ParameterList or None 37 Connection weights from dendrites to the top module. 38 dendrites_to_dendrites : nn.ParameterList or None 39 Connection weights between dendrite modules. 40 node_index : int 41 Index of the feature dimension that corresponds to nodes. 42 num_cycles : torch.Tensor 43 Number of dendrite cycles simulate to simulate in order to set up architecture.. 44 view_tuple : tuple 45 Broadcast shape used for dendrite weights. 46 """ 47 super(PAILayer, self).__init__() 48 self.layer_array = layer_array 49 self.register_buffer("num_cycles", num_cycles) 50 self.register_buffer("view_tuple", torch.tensor(view_tuple)) 51 52 self.processor_array = processor_array 53 if dendrites_to_dendrites: 54 self.skip_weights = dendrites_to_dendrites 55 else: 56 """ 57 This will only be the case if there is less than 2 dendrites, in these cases an empty array 58 should still be added so that dendrites_to_top is included at the correct index 59 """ 60 self.skip_weights = nn.ParameterList() 61 if dendrites_to_top: 62 self.skip_weights.append(dendrites_to_top[len(dendrites_to_top) - 1]) 63 else: 64 self.skip_weights = nn.ParameterList() 65 66 # Delete skip_weights if it's empty (only 1 layer, no skip connections) 67 if len(self.skip_weights) == 0: 68 delattr(self, 'skip_weights') 69 70 self.node_index = node_index 71 self.internal_nonlinearity = GPA.pc.get_pai_forward_function()
Initialize an inference-time PAI layer block. This functionaly is the same but has dendritic scaffolding cleaned up for faster infrance and smaller memory footprint.
Parameters
- layer_array (nn.Module): Ordered module stack used for dendrite and neuron modules.
- processor_array (list):
Processor objects aligned with
layer_array. - dendrites_to_top (nn.ParameterList or None): Connection weights from dendrites to the top module.
- dendrites_to_dendrites (nn.ParameterList or None): Connection weights between dendrite modules.
- node_index (int): Index of the feature dimension that corresponds to nodes.
- num_cycles (torch.Tensor): Number of dendrite cycles simulate to simulate in order to set up architecture..
- view_tuple (tuple): Broadcast shape used for dendrite weights.
73def unWrap_params(model): 74 """Remove wrapped parameter attributes from a model in-place. 75 76 Parameters 77 ---------- 78 model : nn.Module 79 Model whose parameters may contain a temporary ``wrapped`` attribute. 80 81 Returns 82 ------- 83 None 84 This function does not return a value. 85 """ 86 for p in model.parameters(): 87 if "wrapped" in p.__dir__(): 88 del p.wrapped
Remove wrapped parameter attributes from a model in-place.
Parameters
- model (nn.Module):
Model whose parameters may contain a temporary
wrappedattribute.
Returns
- None: This function does not return a value.
91def convert_to_pai_layer_block(pretrained_dendrite): 92 """Convert a training-time PAI neuron module to an inference PAI layer. 93 94 Parameters 95 ---------- 96 pretrained_dendrite : PAINeuronModule 97 Trained neuron module containing dendrite layers and processors. 98 99 Returns 100 ------- 101 PAILayer 102 Inference-ready wrapper with converted processors and dendrite connections. 103 """ 104 unWrap_params(pretrained_dendrite) 105 layer_array = [] 106 processor_array = [] 107 for layer_id in range(len(pretrained_dendrite.dendrite_module.layers)): 108 layer_array.append(pretrained_dendrite.dendrite_module.layers[layer_id]) 109 if pretrained_dendrite.dendrite_module.processors == []: 110 processor_array.append(None) 111 else: 112 if not pretrained_dendrite.dendrite_module.processors[layer_id] is None: 113 pretrained_dendrite.dendrite_module.processors[layer_id].pre = ( 114 pretrained_dendrite.dendrite_module.processors[layer_id].pre_d 115 ) 116 pretrained_dendrite.dendrite_module.processors[layer_id].post = ( 117 pretrained_dendrite.dendrite_module.processors[layer_id].post_d 118 ) 119 processor_array.append( 120 pretrained_dendrite.dendrite_module.processors[layer_id] 121 ) 122 layer_array.append(pretrained_dendrite.main_module) 123 if not pretrained_dendrite.processor is None: 124 pretrained_dendrite.processor.pre = pretrained_dendrite.processor.post_n1 125 pretrained_dendrite.processor.post = pretrained_dendrite.processor.post_n2 126 processor_array.append(pretrained_dendrite.processor) 127 128 view_tuple = [] 129 for dim in range( 130 len( 131 pretrained_dendrite.dendrite_module.dendrite_values[0].this_output_dimensions 132 ) 133 ): 134 if ( 135 dim 136 == pretrained_dendrite.dendrite_module.dendrite_values[0].this_node_index 137 ): 138 view_tuple.append(-1) 139 continue 140 view_tuple.append(1) 141 return PAILayer( 142 nn.Sequential(*layer_array), 143 processor_array, 144 pretrained_dendrite.dendrites_to_top, 145 pretrained_dendrite.dendrite_module.dendrites_to_dendrites, 146 pretrained_dendrite.this_node_index, 147 pretrained_dendrite.dendrite_module.num_cycles, 148 view_tuple, 149 )
Convert a training-time PAI neuron module to an inference PAI layer.
Parameters
- pretrained_dendrite (PAINeuronModule): Trained neuron module containing dendrite layers and processors.
Returns
- PAILayer: Inference-ready wrapper with converted processors and dendrite connections.
152def get_pretrained_pai_attr(pretrained_dendrite, member): 153 """Safely get an attribute from a possibly missing module. 154 155 Parameters 156 ---------- 157 pretrained_dendrite : nn.Module or None 158 Source module that may be ``None``. 159 member : str 160 Attribute name to retrieve. 161 162 Returns 163 ------- 164 Any 165 Requested attribute value, or ``None`` when the module is ``None``. 166 """ 167 if pretrained_dendrite is None: 168 return None 169 else: 170 return getattr(pretrained_dendrite, member)
Safely get an attribute from a possibly missing module.
Parameters
- pretrained_dendrite (nn.Module or None):
Source module that may be
None. - member (str): Attribute name to retrieve.
Returns
- Any: Requested attribute value, or
Nonewhen the module isNone.
173def get_pretrained_pai_var(pretrained_dendrite, submodule_id): 174 """Safely get a submodule by index/key from a possibly missing module. 175 176 Parameters 177 ---------- 178 pretrained_dendrite : nn.Module or None 179 Source module that may be ``None``. 180 submodule_id : str or int 181 Submodule identifier to retrieve. 182 183 Returns 184 ------- 185 nn.Module or None 186 Retrieved submodule, or ``None`` when the source module is ``None``. 187 """ 188 if pretrained_dendrite is None: 189 return None 190 else: 191 return pretrained_dendrite[submodule_id]
Safely get a submodule by index/key from a possibly missing module.
Parameters
- pretrained_dendrite (nn.Module or None):
Source module that may be
None. - submodule_id (str or int): Submodule identifier to retrieve.
Returns
- nn.Module or None: Retrieved submodule, or
Nonewhen the source module isNone.
194def optimize_module(net, depth, name_so_far, converted_list): 195 """Recursively replace training PAI modules with inference PAI layers. 196 197 Parameters 198 ---------- 199 net : nn.Module 200 Module tree to traverse and optimize. 201 depth : int 202 Current recursion depth. 203 name_so_far : str 204 Dotted/Indexed path to the current module. 205 converted_list : list 206 Mutable list of module names already converted. 207 208 Returns 209 ------- 210 nn.Module 211 Updated module tree with converted inference modules. 212 """ 213 all_members = net.__dir__() 214 if issubclass(type(net), nn.Sequential) or issubclass(type(net), nn.ModuleList): 215 for submodule_id, layer in net.named_children(): 216 if type(net.get_submodule(submodule_id)) is PA.PAINeuronModule: 217 if GPA.pc.get_extra_verbose(): 218 print( 219 "Seq sub is PAI so optimizing: %s" % name_so_far 220 + "[" 221 + str(submodule_id) 222 + "]" 223 ) 224 setattr( 225 net, 226 submodule_id, 227 convert_to_pai_layer_block(net.get_submodule(submodule_id)), 228 ) 229 else: 230 if net != net.get_submodule(submodule_id): 231 # this currently just always returns false, not sure what it was for 232 converted_list += [name_so_far + "[" + str(submodule_id) + "]"] 233 setattr( 234 net, 235 submodule_id, 236 optimize_module( 237 net.get_submodule(submodule_id), 238 depth + 1, 239 name_so_far + "[" + str(submodule_id) + "]", 240 converted_list, 241 ), 242 ) 243 else: 244 if GPA.pc.get_extra_verbose(): 245 print( 246 "%s is a self pointer so skipping" 247 % (name_so_far + "[" + str(submodule_id) + "]") 248 ) 249 else: 250 for member in all_members: 251 if isinstance(getattr(type(net), member, None), property): 252 continue 253 try: 254 getattr(net, member, None) 255 except: 256 continue 257 sub_name = name_so_far + "." + member 258 if ( 259 sub_name in GPA.pc.get_module_names_to_not_save() 260 or sub_name in converted_list 261 ): 262 if GPA.pc.get_extra_verbose(): 263 print("Skipping %s during save" % sub_name) 264 continue 265 if type(getattr(net, member, None)) is PA.PAINeuronModule: 266 if GPA.pc.get_extra_verbose(): 267 print( 268 "Sub is in conversion list so initiating optimization for: %s" 269 % name_so_far 270 + "." 271 + member 272 ) 273 setattr(net, member, convert_to_pai_layer_block(getattr(net, member))) 274 elif issubclass(type(getattr(net, member, None)), nn.Module): 275 if net != getattr(net, member): 276 converted_list += [sub_name] 277 setattr( 278 net, 279 member, 280 optimize_module( 281 getattr(net, member), 282 depth + 1, 283 sub_name, 284 converted_list, 285 ), 286 ) 287 else: 288 if GPA.pc.get_extra_verbose(): 289 print("%s is a self pointer so skipping" % (sub_name)) 290 return net
Recursively replace training PAI modules with inference PAI layers.
Parameters
- net (nn.Module): Module tree to traverse and optimize.
- depth (int): Current recursion depth.
- name_so_far (str): Dotted/Indexed path to the current module.
- converted_list (list): Mutable list of module names already converted.
Returns
- nn.Module: Updated module tree with converted inference modules.
292def blockwise_network(net): 293 """Convert all eligible modules in a network to blockwise inference form. 294 295 Parameters 296 ---------- 297 net : nn.Module 298 Input model. 299 300 Returns 301 ------- 302 nn.Module 303 Converted model. 304 """ 305 return optimize_module(net, 0, "", [])
Convert all eligible modules in a network to blockwise inference form.
Parameters
- net (nn.Module): Input model.
Returns
- nn.Module: Converted model.