# Copyright (c) 2025 Perforated AI

import copy
import math
import os
import pdb
import sys
import time
from datetime import datetime

import numpy as np
import torch
import torch.nn as nn
import traceback

from perforatedai import globals_perforatedai as GPA
from perforatedai import utils_perforatedai as UPA

try:
    from perforatedbp import modules_pbp as MPB
except ModuleNotFoundError as e:
    # Only pass if perforatedbp package itself is missing
    if e.name == "perforatedbp":
        pass
    else:
        # perforatedbp exists but is missing a dependency
        raise


# Values for Dendrite training, minimally used in open source version
_DENDRITE_TENSOR_VALUES_BASE = [
    "shape"
]  # Shape is tensor of same shape as total neurons in module
_DENDRITE_SINGLE_VALUES_BASE = []

DENDRITE_INIT_VALUES = ["initialized", "current_d_init"]

_VALUE_TRACKER_ARRAYS_BASE = ["dendrite_outs"]

# Cached values to avoid recomputation (each tracks its own state)
_cached_dendrite_tensor_values = None
_cached_dendrite_tensor_pb_state = None
_cached_dendrite_single_values = None
_cached_dendrite_single_pb_state = None
_cached_value_tracker_arrays = None
_cached_value_tracker_pb_state = None


def get_DENDRITE_TENSOR_VALUES():
    """Get DENDRITE_TENSOR_VALUES, updating from MPB if perforated_backpropagation is enabled.

    Parameters
    ----------
    None

    Returns
    -------
    list[str]
        Names of tensor attributes used for dendrite state handling.
    """
    global _cached_dendrite_tensor_values, _cached_dendrite_tensor_pb_state
    current_pb_state = GPA.pc.get_perforated_backpropagation()

    if (
        _cached_dendrite_tensor_values is None
        or _cached_dendrite_tensor_pb_state != current_pb_state
    ):
        _cached_dendrite_tensor_pb_state = current_pb_state
        if current_pb_state:
            _cached_dendrite_tensor_values = MPB.update_dendrite_tensor_values(
                _DENDRITE_TENSOR_VALUES_BASE.copy()
            )
        else:
            _cached_dendrite_tensor_values = _DENDRITE_TENSOR_VALUES_BASE.copy()
        if current_pb_state:
            _cached_dendrite_tensor_values = _cached_dendrite_tensor_values + MPB._variant_tensor_values

    return _cached_dendrite_tensor_values


def get_DENDRITE_SINGLE_VALUES():
    """Get DENDRITE_SINGLE_VALUES, updating from MPB if perforated_backpropagation is enabled.

    Parameters
    ----------
    None

    Returns
    -------
    list[str]
        Names of scalar attributes used for dendrite state handling.
    """
    global _cached_dendrite_single_values, _cached_dendrite_single_pb_state
    current_pb_state = GPA.pc.get_perforated_backpropagation()

    if (
        _cached_dendrite_single_values is None
        or _cached_dendrite_single_pb_state != current_pb_state
    ):
        _cached_dendrite_single_pb_state = current_pb_state
        if current_pb_state:
            _cached_dendrite_single_values = MPB.update_dendrite_single_values(
                _DENDRITE_SINGLE_VALUES_BASE.copy()
            )
        else:
            _cached_dendrite_single_values = _DENDRITE_SINGLE_VALUES_BASE.copy()
        if current_pb_state:
            _cached_dendrite_single_values = _cached_dendrite_single_values + MPB._variant_single_values

    return _cached_dendrite_single_values


def get_VALUE_TRACKER_ARRAYS():
    """Get VALUE_TRACKER_ARRAYS, updating from MPB if perforated_backpropagation is enabled.

    Parameters
    ----------
    None

    Returns
    -------
    list[str]
        Names of array-valued fields tracked in ``DendriteValueTracker``.
    """
    global _cached_value_tracker_arrays, _cached_value_tracker_pb_state
    current_pb_state = GPA.pc.get_perforated_backpropagation()

    if (
        _cached_value_tracker_arrays is None
        or _cached_value_tracker_pb_state != current_pb_state
    ):
        _cached_value_tracker_pb_state = current_pb_state
        if current_pb_state:
            _cached_value_tracker_arrays = MPB.update_value_tracker_arrays(
                _VALUE_TRACKER_ARRAYS_BASE.copy()
            )
        else:
            _cached_value_tracker_arrays = _VALUE_TRACKER_ARRAYS_BASE.copy()

    return _cached_value_tracker_arrays


def get_DENDRITE_REINIT_VALUES():
    """Get DENDRITE_REINIT_VALUES.

    Parameters
    ----------
    None

    Returns
    -------
    list[str]
        Combined list of attribute names that must be reinitialized.
    """
    return get_DENDRITE_TENSOR_VALUES() + get_DENDRITE_SINGLE_VALUES()


def get_DENDRITE_SAVE_VALUES():
    """Get DENDRITE_SAVE_VALUES.

    Parameters
    ----------
    None

    Returns
    -------
    list[str]
        Combined list of attribute names persisted for save/load.
    """
    return (
        get_DENDRITE_TENSOR_VALUES()
        + get_DENDRITE_SINGLE_VALUES()
        + DENDRITE_INIT_VALUES
    )


def filter_backward(grad_out, values, module=None):
    """Filter backward pass for gradient processing.

    This function processes gradients during the backward pass,
    ensuring correct input dimensions,and applying perforated backpropagation if enabled.

    Parameters
    ----------
    grad_out : torch.Tensor
        The gradient output tensor from the backward pass.
    values : DendriteValueTracker
        A DendriteValueTracker instance containing values associated with the module being processed.
    module : PAINeuronModule, optional
        The owning PAINeuronModule instance. When provided and PBP is disabled,
        the hook deregisters itself after the one-time initialization completes,
        eliminating per-batch Python overhead on all subsequent backward passes.

    Returns
    -------
    None
    """
    if GPA.pc.get_extra_verbose():
        print(f"{values[0].layer_name} calling backward")

    with torch.no_grad():
        val = grad_out.detach()
        # Fast path: Python bool set after first backward in this process.
        # Fallback: current_d_init is a saved buffer — survives checkpoint reload
        # even when _fb_init_done resets to False in a fresh process.
        already_init = (module is not None and module._fb_init_done) or \
                       values[0].current_d_init.item()
        if not already_init:
            # If input dimensions and gradient don't have same shape trigger error and quit
            if len(values[0].this_output_dimensions) != len(grad_out.shape):
                print(
                    "The following module has not properly set this_output_dimensions"
                )
                print(values[0].layer_name)
                print("it is expecting:")
                print(values[0].this_output_dimensions)
                print("but received")
                print(grad_out.shape)
                print(
                    "to check these all at once set GPA.pc.set_debugging_output_dimensions(1)"
                )
                print(
                    f"Call MODEL_VARIABLE{values[0].layer_name}.set_this_output_dimensions([...]) on this module after perforate_model"
                )
                print(
                    "where the ... is replaced with the correct vector as described in section 4 of customization.md"
                )
                if not GPA.pc.get_debugging_output_dimensions():
                    sys.exit(0)
                else:
                    GPA.pc.set_debugging_output_dimensions(2)
                    return
            # Make sure that the input dimensions are correct
            for i in range(len(values[0].this_output_dimensions)):
                if values[0].this_output_dimensions[i] == 0:
                    continue
                # Make sure all input dimensions are either -1 (reduce), 1 (retain), or exact values (old format)
                if (
                    not (grad_out.shape[i] == values[0].this_output_dimensions[i])
                    and not values[0].this_output_dimensions[i] == -1
                    and not values[0].this_output_dimensions[i] == 1
                ):
                    print(
                        "The following module has not properly set this_output_dimensions with this incorrect shape"
                    )
                    print(values[0].layer_name)
                    print("it is expecting:")
                    print(values[0].this_output_dimensions)
                    print("but received")
                    print(grad_out.shape)
                    print(
                        "to check these all at once set GPA.pc.set_debugging_output_dimensions(1)"
                    )
                    if not GPA.pc.get_debugging_output_dimensions():
                        sys.exit(0)
                    else:
                        GPA.pc.set_debugging_output_dimensions(2)
                        return
            # Setup the arrays with the now known shape
            with torch.no_grad():
                if GPA.pc.get_verbose():
                    print("setting d shape for")
                    print(values[0].layer_name)
                    print(val.size())

                values[0].set_out_channels(val.size())
                ndim = len(values[0].this_output_dimensions)
                storage_shape = [1] * ndim
                for _i in range(ndim):
                    if values[0].this_output_dimensions[_i] == 1:
                        storage_shape[_i] = val.shape[_i]
                storage_shape[values[0].this_node_index.item()] = values[0].out_channels
                values[0].setup_arrays(storage_shape)
            # Flag that it has been setup (both the GPU tensor and the fast Python bool)
            values[0].current_d_init[0] = 1
            # If fixed_input_sizes is enabled, populate the tuple caches now
            # that val.shape is known. get_tuples_and_mult will read these on
            # every subsequent call instead of recomputing.
            if GPA.pc.get_perforated_backpropagation() and GPA.pc.get_fixed_input_sizes():
                from perforatedbp import modules_pbp as _MPB
                math_tuple, view_tuple, full_mult = _MPB.get_tuples_and_mult(val, values[0])
                ndim = len(val.shape)
                # math_tuple can be shorter than ndim (excludes this_node_index and
                # retained dims). Pad with -1 sentinel to fill the ndim-length buffer.
                padded_math = math_tuple + [-1] * (ndim - len(math_tuple))
                values[0].math_tuple_cache.copy_(
                    torch.tensor(padded_math, dtype=torch.long, device=val.device)
                )
                values[0].view_tuple_cache.copy_(
                    torch.tensor(view_tuple, dtype=torch.long, device=val.device)
                )
                values[0].full_mult_cache[0] = full_mult
            if module is not None:
                module._fb_init_done = True
                # When PBP is disabled this hook has no further work to do.
                # Deregister it so it never fires again, eliminating per-batch
                # Python overhead on all subsequent backward passes.
                if not GPA.pc.get_perforated_backpropagation() and module._fb_hook_handle is not None:
                    module._fb_hook_handle.remove()
                    module._fb_hook_handle = None
        if GPA.pc.get_perforated_backpropagation():
            MPB.filter_backward_pb(val, values)


def set_wrapped_params(model):
    """Set parameters as wrapped with dendrites.

    Parameters
    ----------
    model : torch.nn.Module
        The model whose parameters are to be marked as wrapped.

    Returns
    -------
    None

    """
    for param in model.parameters():
        param.wrapped = True


def set_tracked_params(model):
    """Set parameters as tracked without dendrites.

    Parameters
    ----------
    model : torch.nn.Module
        The model whose parameters are to be marked as tracked.

    Returns
    -------
    None
    """
    for param in model.parameters():
        param.tracked = True


class PAINeuronModule(nn.Module):
    """Wrapper to set a module as one that will have dendritic copies."""

    def __init__(self, start_module, name):
        """Initialize PAINeuronModule.

        This function sets up the neuron module to wrap the start_module
        and manage its dendritic connections.

        Parameters
        ----------
        start_module : nn.Module
            The module to wrap.
        name : str
            The name of the neuron module.
        """
        super(PAINeuronModule, self).__init__()

        if isinstance(start_module, nn.Module):
            self.main_module = start_module
        else:
            print("start_module must be nn.Module: %s" % name)
            print(type(start_module))
            print(start_module)
            sys.exit(-1)
        self.name = name
        # Per-module config: loads custom settings from {save_name}_config.json if present.
        # Passes both the instance name (id) and the module type so load_config can
        # fall back to type-level settings when no name-specific entry exists.
        _module_type_name = type(start_module).__name__
        self.module_config = GPA.PAIConfig(
            module_name=self.name, module_type=_module_type_name
        )

        set_wrapped_params(self.main_module)
        if self.module_config.get_verbose():
            print(
                f"initing a module {self.name} with main type {type(self.main_module)}"
            )
            print(start_module)

        # If this main_module is one that requires processing set the processor
        if type(self.main_module) in self.module_config.get_modules_with_processing():
            module_index = self.module_config.get_modules_with_processing().index(
                type(self.main_module)
            )
            self.processor = self.module_config.get_modules_processing_classes()[
                module_index
            ]()
            if self.module_config.get_verbose():
                print("with processor")
                print(self.processor)
        elif (
            type(self.main_module).__name__
            in self.module_config.get_module_names_with_processing()
        ):
            module_index = self.module_config.get_module_names_with_processing().index(
                type(self.main_module).__name__
            )
            self.processor = self.module_config.get_module_by_name_processing_classes()[
                module_index
            ]()
            if self.module_config.get_verbose():
                print("with processor")
                print(self.processor)
        else:
            self.processor = None

        # Field that can be filled in if your activation function requires a parameter
        self.activation_function_value = -1
        self.type = "neuron_module"

        self.register_buffer(
            "this_output_dimensions",
            (torch.tensor(self.module_config.get_output_dimensions())),
        )
        if (self.this_output_dimensions == 0).sum() != 1:
            print(f"5 Need exactly one 0 in the input dimensions: {self.name}")
            print(self.this_output_dimensions)
            sys.exit(-1)
        self.register_buffer(
            "this_node_index",
            torch.tensor(self.module_config.get_output_dimensions().index(0)),
        )
        self.dendrite_modules_added = 0

        # Values for dendrite to neuron weights
        self.dendrites_to_top = nn.ParameterList()
        self.register_parameter("newest_dendrite_to_top", None)
        self.candidate_to_top = nn.ParameterList()
        self.register_parameter("current_candidate_to_top", None)
        # Create the dendrite module
        self.dendrite_module = PAIDendriteModule(
            self.main_module,
            activation_function_value=self.activation_function_value,
            name=self.name,
            output_dimensions=self.this_output_dimensions,
        )
        # If it is linear and default has convolutional dimensions, automatically set to just be batch size and neuron indexes
        if (
            issubclass(type(start_module), nn.Linear)
            or (
                issubclass(type(start_module), GPA.PAISequential)
                and issubclass(type(start_module.model[0]), nn.Linear)
            )
        ) and (
            np.array(self.this_output_dimensions)[2:] == -1
        ).all():  # Everything past 2 is a negative 1
            self.set_this_output_dimensions(self.this_output_dimensions[0:2])
        if (
            issubclass(type(start_module), nn.Conv1d)
            or (
                issubclass(type(start_module), GPA.PAISequential)
                and issubclass(type(start_module.model[0]), nn.Conv1d)
            )
        ) and (
            np.array(self.this_output_dimensions)[3:] == -1
        ).all():  # Everything past 2 is a negative 1
            self.set_this_output_dimensions(self.this_output_dimensions[0:3])
        # Apply per-module output_dimensions override from config if present
        _custom_dims = self.module_config.__dict__.get("_output_dimensions")
        if _custom_dims is not None:
            self.set_this_output_dimensions(torch.tensor(_custom_dims))
        GPA.pai_tracker.add_pai_neuron_module(self)
        if self.module_config.get_perforated_backpropagation():
            MPB.set_neuron_parameters(self.main_module)

        # Track filter_backward initialization state as a
        # plain Python bool (avoids a GPU .item() sync on every backward pass).
        # _fb_hook_handle stores the registered hook handle so it can be
        # removed after initialization completes when PBP is disabled.
        self._fb_init_done = False
        self._fb_hook_handle = None

    def __getattr__(self, name):
        """Get member variables from the main module.

        Parameters
        ----------
        name : str
            The name of the variable to retrieve.
        Returns
        -------
        The requested variable.

        Notes
        -----
        This method first attempts to retrieve the attribute from the PAINeuronModule instance.
        If it fails, it tries to get the attribute from the wrapped main_module.
        This allows seamless access to the main module's attributes without modifying original code.
        """
        try:
            return super().__getattr__(name)
        except AttributeError:
            return getattr(self.main_module, name)

    def __getitem__(self, index):
        """Support indexing operations on the main module.

        Parameters
        ----------
        index : int or slice
            The index or slice to retrieve.

        Returns
        -------
        The indexed item from the main module.
        """
        return self.main_module[index]

    def apply_pb_grads(self):
        """Apply perforated backpropagation gradients if enabled.

        Parameters
        ----------
        None

        Returns
        -------
        None
            This function does not return a value.
        """
        self.dendrite_module.apply_pb_grads()

    def apply_pb_zero(self):
        """Clear leftover saved tensors if there are any.

        Parameters
        ----------
        None

        Returns
        -------
        None
            This function does not return a value.
        """
        self.dendrite_module.apply_pb_zero()

    def clear_processors(self):
        """Clear processors if they save values for DeepCopy and save.

        Parameters
        ----------
        None

        Returns
        -------
        None
        """

        if not self.processor:
            return
        else:
            self.processor.clear_processor()
            self.dendrite_module.clear_processors()

    def clear_dendrites(self):
        """Clear and reset dendrites before loading from a state dict.

        Parameters
        ----------
        None

        Returns
        -------
        None

        """
        # Loading a saved state reconstructs PAIDendriteModule before simulating
        # its saved cycles. Preserve a registered variant factory so candidate
        # creation does not silently fall back to deep-copying the parent module.
        create_dendrite_fn = self.dendrite_module._create_dendrite_fn
        self.dendrite_modules_added = 0
        self.dendrites_to_top = nn.ParameterList()
        self.candidate_to_top = nn.ParameterList()
        self.dendrite_module = PAIDendriteModule(
            self.main_module,
            activation_function_value=self.activation_function_value,
            name=self.name,
            output_dimensions=self.this_output_dimensions,
        )
        if create_dendrite_fn is not None:
            self.dendrite_module.set_create_dendrite(create_dendrite_fn)

    def __str__(self):
        """String representation of the module.

        Parameters
        ----------
        None

        Returns
        -------
        str
            String representation of the module.

        Notes
        -----
        Setting for verbose changes level of details in the string output.
        """
        # If verbose print the whole module otherwise just print the module type as a PAIModule
        if self.module_config.get_verbose():
            total_string = self.main_module.__str__()
            total_string = "PAIModule(" + total_string + ")"
            return total_string + self.dendrite_module.__str__()
        else:
            total_string = self.main_module.__str__()
            total_string = "PAIModule(" + total_string + ")"
            return total_string

    def __repr__(self):
        """Representation of the module."""
        return self.__str__()

    def set_this_output_dimensions(self, new_output_dimensions):
        """Set the input dimensions for the neuron and dendrite blocks.

        Signals to this NeuronModule that its input dimensions are different
        than the global default.

        Parameters
        ----------
        new_output_dimensions : list
            A list or tensor specifying the new input dimensions.
        Returns
        -------
        None

        """
        if type(new_output_dimensions) is list:
            new_output_dimensions = torch.tensor(new_output_dimensions)
        delattr(self, "this_output_dimensions")
        self.register_buffer(
            "this_output_dimensions", new_output_dimensions.detach().clone()
        )
        if (new_output_dimensions == 0).sum() != 1:
            print(f"6 need exactly one 0 in the input dimensions: {self.name}")
            print(new_output_dimensions)
        self.this_node_index.copy_(
            (new_output_dimensions == 0).nonzero(as_tuple=True)[0][0]
        )
        self.dendrite_module.set_this_output_dimensions(new_output_dimensions)

    def set_create_dendrite(self, fn):
        """Set a custom function for creating dendrite modules.

        Override how dendrite copies are created from the parent module.  By default
        the parent module is deep-copied.  Pass any callable with the signature
        ``fn(parent_module) -> nn.Module`` to replace that behaviour.

        Parameters
        ----------
        fn : callable
            A function with signature ``fn(parent_module) -> nn.Module``.

        Returns
        -------
        None
        """
        self.dendrite_module.set_create_dendrite(fn)


    def set_mode(self, mode):
        """Switch between neuron training and dendrite training.

        Parameters
        ----------
        mode : str
            The mode to set. Either "n" for neuron training or "p" for pai-dendrite training.

        Returns
        -------
        bool
            True if mode was set successfully, False otherwise.

        Notes
        -----
        If False is returned, the mode was not changed due to an error.
        This is a problem that should not be ignored, but it can be ignored
        by calling PGA.pc.set_checked_skipped_modules(True)
        """

        if self.module_config.get_verbose():
            print(f"{self.name} calling set mode {mode}")
        # If returning to neuron training
        if mode == "n":
            self.dendrite_module.set_mode(mode)
            # Initialize the dendrite to neuron connections
            if self.dendrite_modules_added > 0:
                if self.module_config.get_learn_dendrites_live():
                    values = torch.cat(
                        (
                            self.dendrites_to_top[self.dendrite_modules_added - 1],
                            nn.Parameter(
                                self.candidate_to_top.detach()
                                .clone()
                                .to(dtype=self.module_config.get_d_type())
                            ),
                        ),
                        0,
                    )
                else:
                    values = torch.cat(
                        (
                            self.dendrites_to_top[self.dendrite_modules_added - 1],
                            nn.Parameter(
                                torch.zeros(
                                    (1, self.out_channels),
                                    device=self.dendrites_to_top[
                                        self.dendrite_modules_added - 1
                                    ].device,
                                    dtype=self.module_config.get_d_type(),
                                )
                            ),
                        ),
                        0,
                    )
                # Freeze the previous dendrites_to_top entry
                # so the optimizer no longer updates stale entries. Only the
                # most recently appended entry is ever used in the forward pass.
                if len(self.dendrites_to_top) > 0:
                    self.dendrites_to_top[-1].requires_grad_(False)
                self.dendrites_to_top.append(
                    nn.Parameter(
                        values.detach()
                        .clone()
                        .to(
                            device=self.module_config.get_device(),
                            dtype=self.module_config.get_d_type(),
                        ),
                        requires_grad=True,
                    )
                )
            else:
                if self.module_config.get_learn_dendrites_live():
                    # Freeze previous entry before appending new one.
                    if len(self.dendrites_to_top) > 0:
                        self.dendrites_to_top[-1].requires_grad_(False)
                    self.dendrites_to_top.append(
                        nn.Parameter(
                            self.candidate_to_top.detach()
                            .clone()
                            .to(dtype=self.module_config.get_d_type()),
                            requires_grad=True,
                        )
                    )
                else:
                    # Freeze previous entry before appending new one.
                    if len(self.dendrites_to_top) > 0:
                        self.dendrites_to_top[-1].requires_grad_(False)
                    self.dendrites_to_top.append(
                        nn.Parameter(
                            torch.zeros(
                                (1, self.out_channels),
                                device=self.module_config.get_device(),
                                dtype=self.module_config.get_d_type(),
                            )
                            .detach()
                            .clone(),
                            requires_grad=True,
                        )
                    )
            self.dendrite_modules_added += 1
            if self.module_config.get_perforated_backpropagation():
                MPB.set_module_n_pb(self)
                MPB.set_neuron_parameters(self.dendrites_to_top)

        # If starting dendrite training
        else:
            try:
                # Save the values that were calculated in filter_backward
                self.out_channels = self.dendrite_module.dendrite_values[0].out_channels
                self.dendrite_module.out_channels = (
                    self.dendrite_module.dendrite_values[0].out_channels
                )
            except Exception as e:
                print(e)
                print(
                    f"this occurred in module: {self.dendrite_module.dendrite_values[0].layer_name}"
                )
                print(
                    "Module should be added to module_names_to_track so it doesn't have dendrites added"
                )
                print("If you are getting here but out_channels has not been set")
                print(
                    "A common reason is that this module never had gradients flow through it."
                )
                print("I have seen this happen because:")
                print("-The weights were frozen (requires_grad = False)")
                print(
                    "-A model is added but not used so it was converted to a perforated module initialized"
                )
                print(
                    "-A module was converted that doesn't have weights that get modified so backward doesn't flow through it"
                )
                print(
                    "If this is normal behavior set GPA.pc.set_checked_skipped_modules(True) in the main to ignore"
                )
                print(
                    "You can also set right now in this pdb terminal to have this not happen more after checking all modules this cycle."
                )
                if not self.module_config.get_checked_skipped_modules():
                    pdb.set_trace()
                return False
            # Only change mode if it makes it past the above exception
            self.dendrite_module.set_mode(mode)
            if self.module_config.get_perforated_backpropagation():
                MPB.set_module_p_pb(self)
        return True

    def create_new_dendrite_module(self):
        """Add an additional dendrite module.

        Parameters
        ----------
        None

        Returns
        -------
        None
        """
        self.dendrite_module.create_new_dendrite_module(self.main_module)

    def forward(self, *args, **kwargs):
        """Forward pass through the neuron module.

        Parameters
        ----------
        *args : tuple
            Positional arguments for the forward pass.
        **kwargs : dict
            Keyword arguments for the forward pass.

        Returns
        -------
        Any
            The output of the module after processing through the neuron and dendrite modules.

        Notes
        -----
            The output of this forward function will have the same format as the output
            of the original module
        """

        # If debugging all input dimensions, quit program on first forward call
        if self.module_config.get_debugging_output_dimensions() == 2:
            print("all input dim problems now printed")
            sys.exit(0)
        if self.module_config.get_extra_verbose():
            print(f"{self.name} calling forward")
        # Call the main modules forward
        out = self.main_module(*args, **kwargs)
        # Filter with the processor if required
        if self.processor is not None:
            try:
                out = self.processor.post_n1(out)
            except Exception as e:
                traceback.print_exc(limit=None, chain=True)
                print(f"Your post_n1 processor for {self.name} caused this error")
                print(
                    f"You must check how this is defined and ensure that it is properly"
                )
                print(f"accepting outputs from the neuron module and returning the")
                print(f"single tensor to be combined with the dendrites output tensor")
                sys.exit()
        # Call the forwards for all of the Dendrites
        (
            dendrite_outs,
            candidate_outs,
            candidate_nonlinear_outs,
            candidate_outs_non_zeroed,
        ) = self.dendrite_module(*args, **kwargs)
        # If there are dendrites add all of their outputs to the neurons output
        if self.dendrite_modules_added > 0:
            for i in range(0, self.dendrite_modules_added):
                to_top = self.dendrites_to_top[self.dendrite_modules_added - 1][i, :]
                for dim in range(len(dendrite_outs[i].shape)):
                    if dim == self.this_node_index:
                        continue
                    to_top = to_top.unsqueeze(dim)
                if self.module_config.get_confirm_correct_sizes():
                    to_top = to_top.expand(
                        list(dendrite_outs[i].size())[0 : self.this_node_index]
                        + [self.out_channels]
                        + list(dendrite_outs[i].size())[self.this_node_index + 1 :]
                    )
                out = out + (dendrite_outs[i].to(out.device) * to_top.to(out.device))

        # If learning live, add the candidate's output to the neuron's output via the live weight
        if self.module_config.get_perforated_backpropagation():
            out = MPB.apply_live_candidate_to_output(
                self, out, candidate_nonlinear_outs
            )

        # Catch if processors are required
        if type(out) is tuple:
            print(self)
            print(
                f"The output of the above module {self.name} is a tuple when it must be a single tensor"
            )
            print(
                "This must be fixed to enable the dendrite and neuron output to be combined"
            )
            print(
                "Look in the API customization.md at section 2.2 regarding processors to fix this."
            )
            pdb.set_trace()

        # Call filter backward to ensure the neuron index is setup correctly.
        # Register only when not yet initialized (PBP=False) or always when PBP
        # is enabled. Store the handle so the hook can deregister itself after
        # the one-time initialization completes (PBP=False path).
        if out.requires_grad and (not self._fb_init_done or GPA.pc.get_perforated_backpropagation()):
            self._fb_hook_handle = out.register_hook(
                lambda grad: filter_backward(grad, self.dendrite_module.dendrite_values, self)
            )

        # If there is a processor apply the second neuron stage
        if self.processor is not None:
            try:
                out = self.processor.post_n2(out)
            except Exception as e:
                traceback.print_exc(limit=None, chain=True)
                print(f"Your post_n2 processor for {self.name} caused this error")
                print(
                    f"You must check how this is defined and ensure that it is properly"
                )
                print(
                    f"accepting the output tensor after combining the neuron's output "
                )
                print(f"with the dendrite's output and returning something that is the")
                print(f"same format as your original module's return")
                sys.exit()
        return out


class TrackedNeuronModule(nn.Module):
    """Wrapper for modules you don't want to add dendrites to. Ensures all modules are accounted for."""

    def __init__(self, start_module, name):
        """Initialize TrackedNeuronModule.

        This function sets up the tracked neuron module to wrap the start_module
        without adding dendrites.

        Parameters
        ----------
        start_module : nn.Module
            The module to wrap.
        name : str
            The name of the neuron module.
        """
        super(TrackedNeuronModule, self).__init__()

        if isinstance(start_module, nn.Module):
            self.main_module = start_module
        else:
            print("start_module must be nn.Module: %s" % name)
            print(type(start_module))
            print(start_module)
            sys.exit(-1)
        self.name = name

        self.type = "tracked_module"
        set_tracked_params(self.main_module)
        if GPA.pc.get_verbose():
            print(
                f"tracking a module {self.name} with main type {type(self.main_module)}"
            )
            print(start_module)
        GPA.pai_tracker.add_tracked_neuron_module(self)
        if GPA.pc.get_perforated_backpropagation():
            MPB.set_neuron_parameters(self.main_module)

    def __getattr__(self, name):
        """Get member variables from the main module.

        Parameters
        ----------
        name : str
            The name of the variable to retrieve.
        Returns
        -------
        The requested variable.

        Notes
        -----
        This method first attempts to retrieve the attribute from the PAINeuronModule instance.
        If it fails, it tries to get the attribute from the wrapped main_module.
        This allows seamless access to the main module's attributes without modifying original code.
        """
        try:
            return super().__getattr__(name)
        except AttributeError:
            return getattr(self.main_module, name)

    def __getitem__(self, index):
        """Support indexing operations on the main module.

        Parameters
        ----------
        index : int or slice
            The index or slice to retrieve.

        Returns
        -------
        The indexed item from the main module.
        """
        return self.main_module[index]

    def set_mode(self, mode):
        """Set mode for tracked module.

        Parameters
        ----------
        mode : str
            The mode to set. Either "n" for neuron training or "p" for pai-dendrite training.

        Returns
        -------
        bool
            True.

        Notes
        -----
        This function does not change any behavior since this is a tracked module.
        """

        if GPA.pc.get_verbose():
            print(f"{self.name} calling set mode {mode}")
        return True

    def forward(self, *args, **kwargs):
        """Forward pass for tracked module.

        Parameters
        ----------
        *args : tuple
            Positional arguments for the forward pass.
        **kwargs : dict
            Keyword arguments for the forward pass.

        Returns
        -------
        Any
            The output of the module

        Notes
        -----
            The output of this forward function will have the same format as the output
            of the original module
        """
        return self.main_module(*args, **kwargs)

    def __str__(self):
        """String representation of the module.

        Parameters
        ----------
        None

        Returns
        -------
        str
            String representation of the module.

        Notes
        -----
        Setting for verbose changes level of details in the string output.
        """

        if GPA.pc.get_verbose():
            total_string = self.main_module.__str__()
            total_string = "PAITrackedModule(" + total_string + ")"
            return total_string
        else:
            total_string = self.main_module.__str__()
            total_string = "PAITrackedModule(" + total_string + ")"
            return total_string

    def __repr__(self):
        """Representation of the module."""
        return self.__str__()


def init_params(module, neuron_main_module):
    """Randomize weights after duplicating the main module for the next set of dendrites.

    Parameters
    ----------
    module : nn.Module
        The new dendrite module to initialize.
    neuron_main_module : nn.Module
        The main module of the neuron for potential weight scaling.


    Returns
    -------
    None
        This function does not return a value.
    """
    for param in module.parameters():
        if param.dtype == torch.uint8:
            param.data = torch.randint(0, 256, param.size(), dtype=torch.uint8)
        else:
            # If factoring in the main modules weights multiply the randn()
            #  by the average abs value of the main modules weights
            if GPA.pc.get_candidate_weight_init_by_main():
                main_module_abs = 0
                total_main_params = 0
                for main_param in neuron_main_module.parameters():
                    main_module_abs += main_param.abs().sum().item()
                    total_main_params += main_param.numel()
                if total_main_params > 0:
                    main_module_abs /= total_main_params
                else:
                    main_module_abs = 1.0
                multiplier = main_module_abs
            else:
                multiplier = 1.0
            param.data = (
                torch.randn(param.size(), dtype=param.dtype)
                * GPA.pc.get_candidate_weight_initialization_multiplier()
                * multiplier
            )


class PAIDendriteModule(nn.Module):
    """Module containing all dendrites modules added to the neuron module."""

    def __init__(
        self,
        initial_module,
        activation_function_value=0.3,
        name="no_name_given",
        output_dimensions=None,
    ):
        """Initialize PAINeuronModule.

        This function sets up the dendrite module to create candidate and permanent
        dendrite modules based on the initial_module provided.

        Parameters
        ----------
        initial_module : nn.Module
            The module to copy.
        activation_function_value : float, optional
            A value associated with the activation function, by default 0.3.
        name : str
            The name of the neuron module.
        output_dimensions : vector, optional
            The dimensions of the input vector
        """
        super(PAIDendriteModule, self).__init__()

        if output_dimensions is None:
            output_dimensions = []

        self.layers = nn.ModuleList([])
        self.processors = []
        self.candidate_processors = []
        self.num_dendrites = 0
        self._create_dendrite_fn = None
        # Number of dendrite cycles performed
        self.register_buffer(
            "num_cycles",
            torch.zeros(1, device=GPA.pc.get_device(), dtype=GPA.pc.get_d_type()),
        )
        self.mode = "n"
        self.name = name
        # Create a copy of the parent module so you don't have a pointer to the real one which causes save errors
        self.parent_module = UPA.deep_copy_pai(initial_module)
        if GPA.pc.get_perforated_backpropagation():
            MPB.set_ignored_parameters(self.parent_module)
        # Setup the input dimensions and node index for combining dendrite outputs
        if GPA.pc.get_perforated_backpropagation():
            MPB.create_extra_tensors(self)
        if output_dimensions == []:
            self.register_buffer(
                "this_output_dimensions", torch.tensor(GPA.pc.get_output_dimensions())
            )
        else:
            self.register_buffer(
                "this_output_dimensions", output_dimensions.detach().clone()
            )
        if (self.this_output_dimensions == 0).sum() != 1:
            print(f"1 need exactly one 0 in the input dimensions: {self.name}")
            print(self.this_output_dimensions)
            sys.exit(-1)
        self.register_buffer(
            "this_node_index", torch.tensor(GPA.pc.get_output_dimensions().index(0))
        )

        # Initialize dendrite to dendrite connections
        self.dendrites_to_candidates = nn.ParameterList()
        self.dendrites_to_dendrites = nn.ParameterList()

        # Store an activation function value if required
        self.activation_function_value = activation_function_value
        self.dendrite_values = nn.ModuleList([])
        for j in range(0, GPA.pc.get_global_candidates()):
            if GPA.pc.get_verbose():
                print(f"creating dendrite Values for {self.name}")
            self.dendrite_values.append(
                DendriteValueTracker(
                    False,
                    self.activation_function_value,
                    self.name,
                    self.this_output_dimensions,
                )
            )
        if GPA.pc.get_perforated_backpropagation():
            self.apply_pb_grads = MPB.apply_pb_grads.__get__(self, type(self))
            self.apply_pb_zero = MPB.apply_pb_zero.__get__(self, type(self))

    def __getstate__(self):
        """Tell pickle what to save when this object is serialized (e.g. torch.save).

        apply_pb_grads and apply_pb_zero are bound methods of functions defined
        in modules_pbp and cannot be pickled.  Strip them out; __setstate__ will
        re-attach them after loading.
        """
        import types

        pickle_safe_state = {}
        for attr_name, attr_value in self.__dict__.items():
            if not isinstance(attr_value, types.MethodType):
                pickle_safe_state[attr_name] = attr_value

        return pickle_safe_state

    def __setstate__(self, saved_state):
        """Restore this object from a pickled state (e.g. torch.load).

        Restores all normal attributes, then re-attaches apply_pb_grads and
        apply_pb_zero if perforated backpropagation is enabled.
        """
        self.__dict__.update(saved_state)

        # Re-attach the PBP bound methods that were stripped by __getstate__.
        # dendrite_loss_fn being present on the saved state means PBP was active
        # when the checkpoint was saved.
        if "dendrite_loss_fn" in saved_state:
            import perforatedbp.modules_pbp as MPB
            self.apply_pb_grads = MPB.apply_pb_grads.__get__(self, type(self))
            self.apply_pb_zero = MPB.apply_pb_zero.__get__(self, type(self))

    def create_dendrite(self, parent_module):
        """Create a dendrite module from the parent module.

        Override this function via set_create_dendrite to control how the dendrite
        module is created (e.g. to avoid a deep copy).

        Parameters
        ----------
        parent_module : nn.Module
            The module to create a dendrite from.

        Returns
        -------
        nn.Module
            The new dendrite module.
        """
        if self._create_dendrite_fn is not None:
            return self._create_dendrite_fn(parent_module)
        return UPA.deep_copy_pai(parent_module)

    def set_create_dendrite(self, fn):
        """Set a custom function for creating dendrite modules.

        Call this on a PAIDendriteModule instance to override how dendrites are
        created from the parent module. The function receives the parent module
        and must return a new nn.Module.

        Parameters
        ----------
        fn : callable
            A function with signature ``fn(parent_module) -> nn.Module``.

        Returns
        -------
        None
        """
        self._create_dendrite_fn = fn


    def set_this_output_dimensions(self, new_output_dimensions):
        """Set input dimensions for dendrite module.

        Signals to this DendriteModule that its input dimensions are different
        than the global default.

        Parameters
        ----------
        new_output_dimensions : list
            A list or tensor specifying the new input dimensions.
        Returns
        -------
        None

        """

        if type(new_output_dimensions) is list:
            new_output_dimensions = torch.tensor(new_output_dimensions)
        delattr(self, "this_output_dimensions")
        self.register_buffer(
            "this_output_dimensions", new_output_dimensions.detach().clone()
        )
        if (new_output_dimensions == 0).sum() != 1:
            print(f"2 Need exactly one 0 in the input dimensions: {self.name}")
            print(new_output_dimensions)
            sys.exit(-1)
        self.this_node_index.copy_(
            (new_output_dimensions == 0).nonzero(as_tuple=True)[0][0]
        )
        for j in range(0, GPA.pc.get_global_candidates()):
            self.dendrite_values[j].set_this_output_dimensions(new_output_dimensions)

    def create_new_dendrite_module(self, neuron_main_module):
        """Add a new set of dendrites.

        Parameters
        ----------
        neuron_main_module : Any
            PyTorch module to be used for dendritic learning.
                Typically a copy of the original neuron module.

        Returns
        -------
        None
            This function does not return a value.
        """
        # Candidate module
        self.candidate_module = nn.ModuleList([])
        # Copy that is unused for open source version
        self.best_candidate_module = nn.ModuleList([])
        if GPA.pc.get_verbose():
            print(self.name)
            print("Setting candidate processors")
        self.candidate_processors = []
        with torch.no_grad():
            for i in range(0, GPA.pc.get_global_candidates()):

                new_module = self.create_dendrite(self.parent_module)
                init_params(new_module, neuron_main_module)
                self.candidate_module.append(new_module)
                self.best_candidate_module.append(self.create_dendrite(new_module))
                if type(self.parent_module) in GPA.pc.get_modules_with_processing():
                    module_index = GPA.pc.get_modules_with_processing().index(
                        type(self.parent_module)
                    )
                    self.candidate_processors.append(
                        GPA.pc.get_modules_processing_classes()[module_index]()
                    )
                elif (
                    type(self.parent_module).__name__
                    in GPA.pc.get_module_names_with_processing()
                ):
                    module_index = GPA.pc.get_module_names_with_processing().index(
                        type(self.parent_module).__name__
                    )
                    self.candidate_processors.append(
                        GPA.pc.get_module_by_name_processing_classes()[module_index]()
                    )
                if GPA.pc.get_perforated_backpropagation():
                    MPB.set_candidate_parameters(self.candidate_module[i])
                    MPB.set_ignored_parameters(self.best_candidate_module[i])

        for i in range(0, GPA.pc.get_global_candidates()):
            self.candidate_module[i].to(GPA.pc.get_device())
            self.best_candidate_module[i].to(GPA.pc.get_device())

        # Reset the dendrite_values objects
        for j in range(0, GPA.pc.get_global_candidates()):
            self.dendrite_values[j].reinitialize_for_pai()

        # If there are already dendrites initialize the dendrite to dendrite connections
        if self.num_dendrites > 0:
            self.dendrites_to_candidates = nn.ParameterList()
            for j in range(0, GPA.pc.get_global_candidates()):
                self.dendrites_to_candidates.append(
                    nn.Parameter(
                        torch.zeros(
                            (self.num_dendrites, self.out_channels),
                            device=GPA.pc.get_device(),
                            dtype=GPA.pc.get_d_type(),
                        ),
                        requires_grad=True,
                    )
                )
                if GPA.pc.get_perforated_backpropagation():
                    MPB.init_candidates(self, j)
            if GPA.pc.get_perforated_backpropagation():
                MPB.set_candidate_parameters(self.dendrites_to_candidates)
            # Initialize best_dendrites_to_candidates_saved to snapshot peak-correlation weights at epoch boundaries
            self.best_dendrites_to_candidates_saved = []
            for j in range(0, GPA.pc.get_global_candidates()):
                self.best_dendrites_to_candidates_saved.append(
                    torch.zeros(
                        (self.num_dendrites, self.out_channels),
                        device=GPA.pc.get_device(),
                        dtype=GPA.pc.get_d_type(),
                    )
                )

    def clear_processors(self):
        """Clear processors.

        Parameters
        ----------
        None

        Returns
        -------
        None
            This function does not return a value.
        """
        for processor in self.processors:
            if not processor:
                continue
            else:
                processor.clear_processor()
        for processor in self.candidate_processors:
            if not processor:
                continue
            else:
                processor.clear_processor()

    def set_mode(self, mode):
        """Perform actions when switching between neuron and dendrite training.

        Parameters
        ----------
        mode : str
            The mode to set. Either "n" for neuron training or "p" for pai-dendrite training.

        Returns
        -------
        None
        """

        self.mode = mode
        self.num_cycles += 1
        if GPA.pc.get_verbose():
            print(f"PAI calling set mode {mode} : {self.num_cycles}")
        if not GPA.pc.get_silent():
            print(f"Module {self.name} calling set mode {mode} : {self.num_cycles}")
        # When switching back to neuron training mode convert candidates modules into accepted modules
        if mode == "n":
            if GPA.pc.get_verbose():
                print("So calling all the things to add to modules")
            # Copy weights/bias from correct candidates
            if self.num_dendrites == 1:
                self.dendrites_to_dendrites = nn.ParameterList()
                self.dendrites_to_dendrites.append(torch.tensor([]))
            if self.num_dendrites >= 1:
                self.dendrites_to_dendrites.append(
                    torch.nn.Parameter(
                        torch.zeros(
                            [self.num_dendrites, self.out_channels],
                            device=GPA.pc.get_device(),
                            dtype=GPA.pc.get_d_type(),
                        ),
                        # Grad is true if not pb or if pb and dendrite_update_mode is true
                        requires_grad=(not GPA.pc.get_perforated_backpropagation())
                        or GPA.pc.get_dendrite_update_mode(),
                    )
                )
            with torch.no_grad():
                if GPA.pc.get_global_candidates() > 1:
                    print(
                        "This was a flag that will be needed if using multiple candidates. "
                        "It's not set up yet but nice work finding it."
                    )
                    print(
                        "Note: with multiple candidates, best-score ranking in new_best() uses "
                        "unnormalized covariance (prev_dendrite_candidate_correlation) rather than "
                        "the normalized correlation coefficient. Candidates with larger output "
                        "magnitude will be favored regardless of true correlation quality. "
                        "Fix by tracking running sigma_V and sigma_E and dividing in new_best()."
                    )
                    pdb.set_trace()
                plane_max_index = 0
                self.layers.append(
                    UPA.deep_copy_pai(self.best_candidate_module[plane_max_index])
                )
                self.layers[self.num_dendrites].to(GPA.pc.get_device())
                if self.num_dendrites > 0:
                    self.dendrites_to_dendrites[self.num_dendrites].copy_(
                        self.best_dendrites_to_candidates_saved[plane_max_index]
                    )
                if type(self.parent_module) in GPA.pc.get_modules_with_processing():
                    self.processors.append(self.candidate_processors[plane_max_index])
                if (
                    type(self.parent_module).__name__
                    in GPA.pc.get_module_names_with_processing()
                ):
                    self.processors.append(self.candidate_processors[plane_max_index])
            if GPA.pc.get_perforated_backpropagation():
                MPB.set_pb_mode(self, mode)
            del self.candidate_module, self.best_candidate_module

            self.num_dendrites += 1
            if GPA.pc.get_perforated_backpropagation():
                MPB.set_dendrite_parameters(self.dendrites_to_dendrites)
                MPB.set_dendrite_parameters(self.layers)

    def forward(self, *args, **kwargs):
        """Forward pass for dendrite module.

        Parameters
        ----------
        *args : tuple
            Positional arguments for the forward pass.
        **kwargs : dict
            Keyword arguments for the forward pass.

        Returns
        -------
        Any
            The output of the module after processing through the neuron and dendrite modules.
        Any
            Remaining outputs are only used for Perforated Backpropagation.
        Any
            Remaining outputs are only used for Perforated Backpropagation.
        Any
            Remaining outputs are only used for Perforated Backpropagation.

        Notes
        -----
        If using Perforated Backpropagation, the additional outputs will be moved around in
        this code but left unused and only passed into separate PB functions.
        """

        outs = {}

        # For all modules apply processors, call the modules, then apply post processors
        args2, kwargs2 = args, kwargs
        for c in range(0, self.num_dendrites):
            if GPA.pc.get_perforated_backpropagation():
                args2, kwargs2 = MPB.preprocess_pb(*args, **kwargs)
            if self.processors != []:
                try:
                    args2, kwargs2 = self.processors[c].pre_d(*args2, **kwargs2)
                except Exception as e:
                    traceback.print_exc(limit=None, chain=True)
                    print(f"Your pre_d processor for {self.name} caused this error")
                    print(
                        f"You must check how this is defined and ensure that it is properly"
                    )
                    print(
                        f"accepting inputs to the PAIModule and returning what will then be"
                    )
                    print(f"the input to the dendrite module")
                    sys.exit()
            out_values = self.layers[c](*args2, **kwargs2)
            if self.processors != []:
                try:
                    outs[c] = self.processors[c].post_d(out_values)
                except Exception as e:
                    traceback.print_exc(limit=None, chain=True)
                    print(f"Your post_d processor for {self.name} caused this error")
                    print(
                        f"You must check how this is defined and ensure that it is properly"
                    )
                    print(
                        f"accepting outputs from the dendrite module and returning the"
                    )
                    print(
                        f"single tensor to be combined with the neurons output tensor"
                    )
                    sys.exit()
            else:
                outs[c] = out_values

        # Create dendrite outputs
        # Each dendrite has input from previously created dendrites
        # So activation is added before the nonlinearity is called
        view_tuple = []
        for out_index in range(0, self.num_dendrites):
            current_out = outs[out_index]
            view_tuple = []
            for dim in range(len(current_out.shape)):
                if dim == self.this_node_index:
                    view_tuple.append(-1)
                    continue
                view_tuple.append(1)

            for in_index in range(0, out_index):
                if view_tuple == [
                    1
                ]:  # This is only the case when passing a single datapoint rather than a batch
                    current_out = (
                        current_out
                        + self.dendrites_to_dendrites[out_index][in_index, :].to(
                            current_out.device
                        )
                        * outs[in_index]
                    )
                else:
                    current_out = (
                        current_out
                        + self.dendrites_to_dendrites[out_index][in_index, :]
                        .view(view_tuple)
                        .to(current_out.device)
                        * outs[in_index]
                    )
            outs[out_index] = GPA.pc.get_pai_forward_function()(current_out)
        # Return a dict which has all dendritic outputs after the activation functions were called
        if GPA.pc.get_perforated_backpropagation():
            candidate_outs, candidate_nonlinear_outs, candidate_non_zeroed = (
                MPB.forward_candidates(self, view_tuple, outs, *args2, **kwargs2)
            )
        else:
            candidate_outs, candidate_nonlinear_outs, candidate_non_zeroed = (
                {},
                {},
                {},
            )
        return outs, candidate_outs, candidate_nonlinear_outs, candidate_non_zeroed


class DendriteValueTracker(nn.Module):
    """Tracker object that maintains certain values for each set of dendrites."""

    def __init__(
        self,
        initialized,
        activation_function_value,
        name,
        output_dimensions,
        out_channels=-1,
    ):
        """Initialize DendriteValueTracker.

        This function sets up the value tracker to maintain statistics and values
        for each set of dendrites.

        Parameters
        ----------
        initialized : int
            Whether the dendrite has been initialized (1) or not (0).
        activation_function_value : float
            A value associated with the activation function.
        name : str
            The name of the associated neuron module.
        output_dimensions : vector
            The dimensions of the input vector.
        out_channels : int
            The number of output channels
        """
        super(DendriteValueTracker, self).__init__()

        self.layer_name = name
        for val_name in DENDRITE_INIT_VALUES:
            self.register_buffer(
                val_name,
                torch.zeros(1, device=GPA.pc.get_device(), dtype=GPA.pc.get_d_type()),
            )
        self.initialized[0] = initialized
        self.activation_function_value = activation_function_value
        self.register_buffer(
            "this_output_dimensions", output_dimensions.clone().detach()
        )
        if (self.this_output_dimensions == 0).sum() != 1:
            print(f"3 need exactly one 0 in the input dimensions: {self.layer_name}")
            print(self.this_output_dimensions)
            sys.exit(-1)
        self.register_buffer(
            "this_node_index", (output_dimensions == 0).nonzero(as_tuple=True)[0]
        )
        if out_channels != -1:
            ndim = len(output_dimensions)
            init_shape = [1] * ndim
            init_shape[(output_dimensions == 0).nonzero(as_tuple=True)[0].item()] = out_channels
            self.setup_arrays(init_shape)
        else:
            self.out_channels = -1

    def print(self):
        """Print value tracker information.

        Parameters
        ----------
        None

        Returns
        -------
        None
            This function does not return a value.
        """
        total_string = "Value Tracker:"
        for val_name in DENDRITE_INIT_VALUES:
            total_string += f"\t{val_name}:\n\t\t"
            total_string += getattr(self, val_name).__repr__()
            total_string += "\n"
        for val_name in get_DENDRITE_TENSOR_VALUES():
            if getattr(self, val_name, None) is not None:
                total_string += f"\t{val_name}:\n\t\t"
                total_string += getattr(self, val_name).__repr__()
                total_string += "\n"
        print(total_string)

    def set_this_output_dimensions(self, new_output_dimensions):
        """Set input dimensions for value tracker

        Signals to this DendriteValueTracker that its input dimensions are different
        than the global default.

        Parameters
        ----------
        new_output_dimensions : list
            A list or tensor specifying the new input dimensions.
        Returns
        -------
        None

        """
        if type(new_output_dimensions) is list:
            new_output_dimensions = torch.tensor(new_output_dimensions)
        delattr(self, "this_output_dimensions")
        self.register_buffer(
            "this_output_dimensions", new_output_dimensions.detach().clone()
        )
        if (new_output_dimensions == 0).sum() != 1:
            print(f"4 need exactly one 0 in the input dimensions: {self.layer_name}")
            print(new_output_dimensions)
            sys.exit(-1)
        self.this_node_index.copy_(
            (new_output_dimensions == 0).nonzero(as_tuple=True)[0][0]
        )

    def set_out_channels(self, shape_values):
        """Set output channels based on shape values and saved node index

        Parameters
        ----------
        shape_values : list or torch.Size
            A list or tensor specifying the shape values.

        Returns
        -------
        None
        """
        if type(shape_values) == torch.Size:
            self.out_channels = int(shape_values[self.this_node_index])
        else:
            self.out_channels = int(shape_values[self.this_node_index].item())

    def setup_arrays(self, storage_shape):
        """Setup arrays for value tracker.

        Parameters
        ----------
        storage_shape : list
            Shape for the tracking tensors: 1 at every dim except the channel
            dim (this_node_index), which holds out_channels.  E.g. [1, 5] for
            a linear layer with 5 outputs.
        Returns
        -------
        None

        """
        # storage_shape is a list with 1 at every dim except the channel dim.
        # Passed in directly from filter_backward so it is always derived from
        # the live gradient — not from out_channels, which is not saved/loaded.
        self.out_channels = storage_shape[self.this_node_index.item()]
        self.register_buffer(
            "dendrite_storage_shape",
            torch.tensor(storage_shape, dtype=torch.long, device=GPA.pc.get_device()),
        )
        for val_name in get_DENDRITE_TENSOR_VALUES():
            self.register_buffer(
                val_name,
                torch.zeros(
                    storage_shape, device=GPA.pc.get_device(), dtype=GPA.pc.get_d_type()
                ),
            )

        for name in get_VALUE_TRACKER_ARRAYS():
            setattr(self, name, {})
            count = 1
            if torch.cuda.device_count() > count:
                count = torch.cuda.device_count()
            for i in range(count):
                getattr(self, name)[i] = []
        for val_name in get_DENDRITE_SINGLE_VALUES():
            self.register_buffer(
                val_name,
                torch.zeros(1, device=GPA.pc.get_device(), dtype=GPA.pc.get_d_type()),
            )

        # If fixed_input_sizes is enabled, register cache buffers now that the
        # final ndim is known (output_dimensions may have been corrected after
        # __init__ for Linear layers). Mirrors the pattern of this_output_dimensions
        # — registered as buffers so they are saved and loaded automatically.
        if GPA.pc.get_perforated_backpropagation() and GPA.pc.get_fixed_input_sizes():
            ndim = len(storage_shape)
            if not hasattr(self, 'math_tuple_cache'):
                self.register_buffer(
                    "math_tuple_cache",
                    torch.zeros(ndim, dtype=torch.long, device=GPA.pc.get_device()),
                )
                self.register_buffer(
                    "view_tuple_cache",
                    torch.zeros(ndim, dtype=torch.long, device=GPA.pc.get_device()),
                )
                self.register_buffer(
                    "full_mult_cache",
                    torch.zeros(1, dtype=torch.long, device=GPA.pc.get_device()),
            )

    def reinitialize_for_pai(self):
        """Reinitialize value tracker to add the next set of dendrites

        Parameters
        ----------
        None

        Returns
        -------
        None
            This function does not return a value.
        """

        if self.out_channels == -1:
            print("You have a perforated module that was never initialized")
            print("This likely means it is not being added to the autograd graph")
            print("Check your forward function that it is actually being used")
            print("If its not you should really delete it, but you can also add")
            print(self.layer_name)
            print("with:")
            print("GPA.pc.append_module_ids_to_track(['" + self.layer_name + "'])")
            print("This can also happen while testing_dendrite_capacity if you")
            print(
                "run a validation cycle and try to add Dendrites before doing any training.\n"
            )
            pdb.set_trace()

        self.initialized[0] = 0
        if GPA.pc.get_perforated_backpropagation():
            MPB.reinitialize_for_pb(self)
        else:
            for val_name in get_DENDRITE_REINIT_VALUES():
                setattr(self, val_name, getattr(self, val_name) * 0)
