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            elif type(net.get_submodule(submodule_id)) is PA.TrackedNeuronModule:
229                if GPA.pc.get_extra_verbose():
230                    print(
231                        "Seq sub is PAITracked so unwrapping: %s" % name_so_far
232                        + "["
233                        + str(submodule_id)
234                        + "]"
235                    )
236                setattr(
237                    net,
238                    submodule_id,
239                    net.get_submodule(submodule_id).main_module,
240                )
241            else:
242                if net != net.get_submodule(submodule_id):
243                    # this currently just always returns false, not sure what it was for
244                    converted_list += [name_so_far + "[" + str(submodule_id) + "]"]
245                    setattr(
246                        net,
247                        submodule_id,
248                        optimize_module(
249                            net.get_submodule(submodule_id),
250                            depth + 1,
251                            name_so_far + "[" + str(submodule_id) + "]",
252                            converted_list,
253                        ),
254                    )
255                else:
256                    if GPA.pc.get_extra_verbose():
257                        print(
258                            "%s is a self pointer so skipping"
259                            % (name_so_far + "[" + str(submodule_id) + "]")
260                        )
261    else:
262        for member in all_members:
263            if isinstance(getattr(type(net), member, None), property):
264                continue
265            try:
266                getattr(net, member, None)
267            except:
268                continue
269            sub_name = name_so_far + "." + member
270            if (
271                sub_name in GPA.pc.get_module_names_to_not_save()
272                or sub_name in converted_list
273            ):
274                if GPA.pc.get_extra_verbose():
275                    print("Skipping %s during save" % sub_name)
276                continue
277            if type(getattr(net, member, None)) is PA.PAINeuronModule:
278                if GPA.pc.get_extra_verbose():
279                    print(
280                        "Sub is in conversion list so initiating optimization for: %s"
281                        % name_so_far
282                        + "."
283                        + member
284                    )
285                setattr(net, member, convert_to_pai_layer_block(getattr(net, member)))
286            elif type(getattr(net, member, None)) is PA.TrackedNeuronModule:
287                if GPA.pc.get_extra_verbose():
288                    print(
289                        "Sub is PAITracked so unwrapping: %s"
290                        % name_so_far
291                        + "."
292                        + member
293                    )
294                setattr(net, member, getattr(net, member).main_module)
295            elif issubclass(type(getattr(net, member, None)), nn.Module):
296                if net != getattr(net, member):
297                    converted_list += [sub_name]
298                    setattr(
299                        net,
300                        member,
301                        optimize_module(
302                            getattr(net, member),
303                            depth + 1,
304                            sub_name,
305                            converted_list,
306                        ),
307                    )
308                else:
309                    if GPA.pc.get_extra_verbose():
310                        print("%s is a self pointer so skipping" % (sub_name))
311    return net
312
313def blockwise_network(net):
314    """Convert all eligible modules in a network to blockwise inference form.
315
316    Parameters
317    ----------
318    net : nn.Module
319        Input model.
320
321    Returns
322    -------
323    nn.Module
324        Converted model.
325    """
326    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            elif type(net.get_submodule(submodule_id)) is PA.TrackedNeuronModule:
230                if GPA.pc.get_extra_verbose():
231                    print(
232                        "Seq sub is PAITracked so unwrapping: %s" % name_so_far
233                        + "["
234                        + str(submodule_id)
235                        + "]"
236                    )
237                setattr(
238                    net,
239                    submodule_id,
240                    net.get_submodule(submodule_id).main_module,
241                )
242            else:
243                if net != net.get_submodule(submodule_id):
244                    # this currently just always returns false, not sure what it was for
245                    converted_list += [name_so_far + "[" + str(submodule_id) + "]"]
246                    setattr(
247                        net,
248                        submodule_id,
249                        optimize_module(
250                            net.get_submodule(submodule_id),
251                            depth + 1,
252                            name_so_far + "[" + str(submodule_id) + "]",
253                            converted_list,
254                        ),
255                    )
256                else:
257                    if GPA.pc.get_extra_verbose():
258                        print(
259                            "%s is a self pointer so skipping"
260                            % (name_so_far + "[" + str(submodule_id) + "]")
261                        )
262    else:
263        for member in all_members:
264            if isinstance(getattr(type(net), member, None), property):
265                continue
266            try:
267                getattr(net, member, None)
268            except:
269                continue
270            sub_name = name_so_far + "." + member
271            if (
272                sub_name in GPA.pc.get_module_names_to_not_save()
273                or sub_name in converted_list
274            ):
275                if GPA.pc.get_extra_verbose():
276                    print("Skipping %s during save" % sub_name)
277                continue
278            if type(getattr(net, member, None)) is PA.PAINeuronModule:
279                if GPA.pc.get_extra_verbose():
280                    print(
281                        "Sub is in conversion list so initiating optimization for: %s"
282                        % name_so_far
283                        + "."
284                        + member
285                    )
286                setattr(net, member, convert_to_pai_layer_block(getattr(net, member)))
287            elif type(getattr(net, member, None)) is PA.TrackedNeuronModule:
288                if GPA.pc.get_extra_verbose():
289                    print(
290                        "Sub is PAITracked so unwrapping: %s"
291                        % name_so_far
292                        + "."
293                        + member
294                    )
295                setattr(net, member, getattr(net, member).main_module)
296            elif issubclass(type(getattr(net, member, None)), nn.Module):
297                if net != getattr(net, member):
298                    converted_list += [sub_name]
299                    setattr(
300                        net,
301                        member,
302                        optimize_module(
303                            getattr(net, member),
304                            depth + 1,
305                            sub_name,
306                            converted_list,
307                        ),
308                    )
309                else:
310                    if GPA.pc.get_extra_verbose():
311                        print("%s is a self pointer so skipping" % (sub_name))
312    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):
314def blockwise_network(net):
315    """Convert all eligible modules in a network to blockwise inference form.
316
317    Parameters
318    ----------
319    net : nn.Module
320        Input model.
321
322    Returns
323    -------
324    nn.Module
325        Converted model.
326    """
327    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.