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, "", [])
class PAILayer(torch.nn.modules.module.Module):
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

PAILayer( layer_array, processor_array, dendrites_to_top, dendrites_to_dendrites, node_index, num_cycles, view_tuple)
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.
layer_array
processor_array
node_index
internal_nonlinearity
def unWrap_params(model):
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 wrapped attribute.
Returns
  • None: This function does not return a value.
def convert_to_pai_layer_block(pretrained_dendrite):
 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.
def get_pretrained_pai_attr(pretrained_dendrite, member):
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 None when the module is None.
def get_pretrained_pai_var(pretrained_dendrite, submodule_id):
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 None when the source module is None.
def optimize_module(net, depth, name_so_far, converted_list):
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.
def blockwise_network(net):
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.