perforatedai.tracker_perforatedai

   1# Copyright (c) 2025 Perforated AI
   2
   3import io
   4import math
   5import os
   6import shutil
   7import sys
   8import time
   9from datetime import datetime
  10from pydoc import locate
  11
  12import matplotlib as mpl
  13import matplotlib.pyplot as plt
  14import numpy as np
  15import pandas as pd
  16import torch
  17import torch.nn as nn
  18import torch.nn.functional as F
  19import torch.nn.init as init
  20import pdb
  21
  22from perforatedai import globals_perforatedai as GPA
  23from perforatedai import modules_perforatedai as PA
  24from perforatedai import utils_perforatedai as UPA
  25
  26try:
  27    from dashboard_utils.event_emitter import emitter as _dashboard_emitter
  28except ImportError:
  29    _dashboard_emitter = None
  30
  31
  32def _pai_log(level, message):
  33    """Emit a tracker log message to dashboard or stdout.
  34
  35    Parameters
  36    ----------
  37    level : str
  38        Log level such as ``info``, ``warning``, or ``error``.
  39    message : str
  40        Message text to emit.
  41    """
  42    if _dashboard_emitter is not None:
  43        _dashboard_emitter.log(GPA.pc, level, message)
  44    else:
  45        if level in ("warning", "error") or not GPA.pc.get_silent():
  46            print(message)
  47
  48
  49try:
  50    from perforatedbp import tracker_pbp as TPB
  51except ModuleNotFoundError:
  52    pass  # Module not found, pass silently
  53except ImportError as e:
  54    print(f"Import error occurred: {e}")
  55
  56mpl.use("Agg")
  57
  58# Status constants for restructuring during add_validation_score
  59NO_MODEL_UPDATE = 0
  60NETWORK_RESTRUCTURED = 1
  61TRAINING_COMPLETE = 2
  62
  63# Status constant for each batch
  64STEP_CLEARED = 0
  65STEP_CALLED = 1
  66
  67
  68def update_restructuring_status(old_status, new_status):
  69    """Update restructured variable during add_validation_score
  70
  71    Update the restructuring status based on the new status.
  72    If the new status is that there was not an update,
  73    dont overwrite the old status which may show there was an update.
  74
  75    Parameters
  76    ----------
  77    old_status : int
  78        The old restructuring status.
  79    new_status : int
  80        The new restructuring status.
  81
  82    Returns
  83    -------
  84    int
  85        The updated restructuring status.
  86
  87    """
  88    if new_status == NETWORK_RESTRUCTURED or new_status == TRAINING_COMPLETE:
  89        return NETWORK_RESTRUCTURED
  90    else:
  91        return old_status
  92
  93
  94def update_learning_rate():
  95    """Update the learning rate in the tracker.
  96
  97    Parameters
  98    ----------
  99    None
 100
 101    Returns
 102    -------
 103    None
 104        This function does not return a value.
 105    """
 106    for param_group in GPA.pai_tracker.member_vars["optimizer_instance"].param_groups:
 107        learning_rate = param_group["lr"]
 108    GPA.pai_tracker.add_learning_rate(learning_rate)
 109
 110
 111def update_param_count(net):
 112    """Update the parameter count in the tracker if not already set.
 113
 114    Parameters
 115    ----------
 116    net : torch.nn.Module
 117        The neural network model to count parameters for.
 118    Returns
 119    -------
 120    None
 121    """
 122    if len(GPA.pai_tracker.member_vars["param_counts"]) == 0:
 123        GPA.pai_tracker.member_vars["param_counts"].append(UPA.count_params(net))
 124
 125
 126def check_input_problems(net, accuracy):
 127    """Check for potential input problems in add_validation_score.
 128
 129    Parameters
 130    ----------
 131    net : torch.nn.Module
 132        The neural network model to check.
 133    accuracy : float, int, or torch.Tensor
 134        The accuracy score to validate.
 135
 136    Returns
 137    -------
 138    float
 139        The validated accuracy score.
 140
 141    """
 142
 143    # Make sure you are passing in the model and not the dataparallel wrapper
 144    if issubclass(type(net), nn.DataParallel):
 145        _pai_log("error", "Need to call .module when using add validation score")
 146        pdb.set_trace()
 147        sys.exit(-1)
 148
 149    if "module" in net.__dir__():
 150        _pai_log("error", "Need to call .module when using add validation score")
 151        pdb.set_trace()
 152        sys.exit(-1)
 153
 154    if not isinstance(accuracy, (float, int)):
 155        try:
 156            accuracy = accuracy.item()
 157        except:
 158            _pai_log(
 159                "error",
 160                f"Scores added for add_validation_score should be float, int, or tensor, yours is a: {type(accuracy)}",
 161            )
 162            pdb.set_trace()
 163            sys.exit(-1)
 164    return accuracy
 165
 166
 167def update_running_accuracy(accuracy, epochs_since_cycle_switch):
 168    """Add the new accuracy to the tracker.
 169
 170    Parameters
 171    ----------
 172    accuracy : float, int, or torch.Tensor
 173        The accuracy score to add.
 174    epochs_since_cycle_switch : int
 175        The number of epochs since the last cycle switch.
 176
 177    Returns
 178    -------
 179    None
 180
 181    """
 182    # Only update running_accuracy when neurons are being updated
 183    if GPA.pai_tracker.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
 184        if epochs_since_cycle_switch < GPA.pc.get_initial_history_after_switches():
 185            if epochs_since_cycle_switch <= 0:
 186                GPA.pai_tracker.member_vars["running_accuracy"] = accuracy
 187            else:
 188                GPA.pai_tracker.member_vars[
 189                    "running_accuracy"
 190                ] = GPA.pai_tracker.member_vars["running_accuracy"] * (
 191                    1 - (1.0 / (epochs_since_cycle_switch + 1))
 192                ) + accuracy * (
 193                    1.0 / (epochs_since_cycle_switch + 1)
 194                )
 195        else:
 196            GPA.pai_tracker.member_vars[
 197                "running_accuracy"
 198            ] = GPA.pai_tracker.member_vars["running_accuracy"] * (
 199                1.0 - 1.0 / GPA.pc.get_history_lookback()
 200            ) + accuracy * (
 201                1.0 / GPA.pc.get_history_lookback()
 202            )
 203
 204    GPA.pai_tracker.member_vars["accuracies"].append(accuracy)
 205    if GPA.pai_tracker.member_vars["mode"] == "n":
 206        GPA.pai_tracker.member_vars["n_accuracies"].append(accuracy)
 207
 208    if (
 209        GPA.pc.get_drawing_pai()
 210        or GPA.pai_tracker.member_vars["mode"] == "n"
 211        or GPA.pc.get_learn_dendrites_live()
 212    ):
 213        GPA.pai_tracker.member_vars["running_accuracies"].append(
 214            GPA.pai_tracker.member_vars["running_accuracy"]
 215        )
 216
 217
 218def score_beats_current_best(new_score, old_score):
 219    """Check if the new score beats the current best score.
 220
 221    Parameters
 222    ----------
 223    new_score : float
 224        The new score to compare.
 225    old_score : float
 226        The old score to compare against.
 227
 228    Returns
 229    -------
 230    bool
 231        True if the new score beats the old score, False otherwise.
 232
 233    Notes
 234    -----
 235    Must beat the old score by the margins set in globals for improvement thresholds.
 236
 237    """
 238    return (
 239        GPA.pai_tracker.member_vars["maximizing_score"]
 240        and (new_score * (1.0 - GPA.pc.get_improvement_threshold()) > old_score)
 241        and new_score - GPA.pc.get_improvement_threshold_raw() > old_score
 242    ) or (
 243        (not GPA.pai_tracker.member_vars["maximizing_score"])
 244        and (new_score * (1.0 + GPA.pc.get_improvement_threshold()) < old_score)
 245        and (new_score + GPA.pc.get_improvement_threshold_raw()) < old_score
 246    )
 247
 248
 249def check_new_best(net, accuracy, epochs_since_cycle_switch):
 250    """Check if the new accuracy is a new best.
 251
 252    Performs saves if new best score is found.
 253
 254    Parameters
 255    ----------
 256    net : torch.nn.Module
 257        The neural network model being trained.
 258    accuracy : float
 259        The accuracy score to check.
 260    epochs_since_cycle_switch : int
 261        The number of epochs since the last cycle switch.
 262
 263    Returns
 264    -------
 265    None
 266
 267    """
 268    score_improved = score_beats_current_best(
 269        GPA.pai_tracker.member_vars["running_accuracy"],
 270        GPA.pai_tracker.member_vars["current_best_validation_score"],
 271    )
 272
 273    enough_time = (
 274        epochs_since_cycle_switch > GPA.pc.get_initial_history_after_switches()
 275    ) or (GPA.pai_tracker.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME)
 276
 277    if (
 278        score_improved
 279        or GPA.pai_tracker.member_vars["current_best_validation_score"] == 0
 280    ) and enough_time:
 281
 282        if GPA.pai_tracker.member_vars["maximizing_score"]:
 283            if GPA.pc.get_verbose():
 284                print(
 285                    f"\n\nGot score of {accuracy:.10f} "
 286                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
 287                    f"*{1-GPA.pc.get_improvement_threshold()}="
 288                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 - GPA.pc.get_improvement_threshold())}) '
 289                    f'which is higher than {GPA.pai_tracker.member_vars["current_best_validation_score"]:.10f} '
 290                    f"by {GPA.pc.get_improvement_threshold_raw()} so setting epoch to "
 291                    f'{GPA.pai_tracker.member_vars["num_epochs_run"]}\n\n'
 292                )
 293        else:
 294            if GPA.pc.get_verbose():
 295                print(
 296                    f"\n\nGot score of {accuracy:.10f} "
 297                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
 298                    f"*{1+GPA.pc.get_improvement_threshold()}="
 299                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 + GPA.pc.get_improvement_threshold())}) '
 300                    f'which is lower than {GPA.pai_tracker.member_vars["current_best_validation_score"]:.10f} '
 301                    f'so setting epoch to {GPA.pai_tracker.member_vars["num_epochs_run"]}\n\n'
 302                )
 303
 304        # Set the new best score
 305        GPA.pai_tracker.member_vars["current_best_validation_score"] = (
 306            GPA.pai_tracker.member_vars["running_accuracy"]
 307        )
 308        GPA.pai_tracker.member_vars["epoch_last_improved"] = (
 309            GPA.pai_tracker.member_vars["num_epochs_run"]
 310        )
 311        if GPA.pc.get_verbose():
 312            print(
 313                f'2 epoch improved is {GPA.pai_tracker.member_vars["epoch_last_improved"]}'
 314            )
 315        # Immediately update this list before saving so loading will have it correctly
 316        GPA.pai_tracker.member_vars["last_improved_accuracies"].append(
 317            GPA.pai_tracker.member_vars["epoch_last_improved"]
 318        )
 319        # Check if global best
 320        is_global_best = score_beats_current_best(
 321            GPA.pai_tracker.member_vars["current_best_validation_score"],
 322            GPA.pai_tracker.member_vars["global_best_validation_score"],
 323        )
 324
 325        if (
 326            is_global_best
 327            or GPA.pai_tracker.member_vars["global_best_validation_score"] == 0
 328        ):
 329            if GPA.pc.get_verbose():
 330                print(
 331                    f"This also beats global best of "
 332                    f'{GPA.pai_tracker.member_vars["global_best_validation_score"]} so saving'
 333                )
 334            GPA.pai_tracker.member_vars["global_best_validation_score"] = (
 335                GPA.pai_tracker.member_vars["current_best_validation_score"]
 336            )
 337            GPA.pai_tracker.member_vars["current_n_set_global_best"] = True
 338            UPA.save_system(net, GPA.pc.get_save_name(), "best_model")
 339            if GPA.pc.get_pai_saves():
 340                UPA.pai_save_system(net, GPA.pc.get_save_name(), "best_model")
 341    else:
 342        if GPA.pc.get_verbose():
 343            print("Not saving new best because:")
 344            if epochs_since_cycle_switch <= GPA.pc.get_initial_history_after_switches():
 345                print(
 346                    f"Not enough history since switch {epochs_since_cycle_switch} <= "
 347                    f"{GPA.pc.get_initial_history_after_switches()}"
 348                )
 349            elif GPA.pai_tracker.member_vars["maximizing_score"]:
 350                print(
 351                    f"Got score of {accuracy} "
 352                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
 353                    f"*{1-GPA.pc.get_improvement_threshold()}="
 354                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 - GPA.pc.get_improvement_threshold())}) '
 355                    f"which is not higher than "
 356                    f'{GPA.pai_tracker.member_vars["current_best_validation_score"]}'
 357                )
 358            else:
 359                print(
 360                    f"Got score of {accuracy} "
 361                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
 362                    f"*{1+GPA.pc.get_improvement_threshold()}="
 363                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 + GPA.pc.get_improvement_threshold())}) '
 364                    f"which is not lower than "
 365                    f'{GPA.pai_tracker.member_vars["current_best_validation_score"]}'
 366                )
 367        GPA.pai_tracker.member_vars["last_improved_accuracies"].append(
 368            GPA.pai_tracker.member_vars["epoch_last_improved"]
 369        )
 370        # If it's the first epoch, save as best anyway
 371        if len(GPA.pai_tracker.member_vars["accuracies"]) == 1:
 372            if GPA.pc.get_verbose():
 373                print("Saving first model or all models")
 374            UPA.save_system(net, GPA.pc.get_save_name(), "best_model")
 375            if GPA.pc.get_pai_saves():
 376                UPA.pai_save_system(net, GPA.pc.get_save_name(), "best_model")
 377
 378
 379def process_no_improvement(net):
 380    """Handle the case where no improvement is observed.
 381
 382    If the new dendrite did not improve scores, but its time to switch modes
 383    either trigger the end of learning or reset to the previous dendrite
 384    to try again.
 385
 386    Parameters
 387    ----------
 388    net : torch.nn.Module
 389        The neural network model being trained.
 390
 391    Returns
 392    -------
 393    int
 394        The status of restructuring or training completion.
 395    torch.nn.Module
 396        The potentially modified neural network model.
 397
 398    """
 399    if GPA.pc.get_verbose():
 400        print(
 401            f"Planning to switch to p mode but best beat last: "
 402            f'{GPA.pai_tracker.member_vars["current_n_set_global_best"]} '
 403            f"current start lr steps: "
 404            f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
 405            f"and last maximum lr steps: "
 406            f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
 407            f'for rate: {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]:.8f}'
 408        )
 409
 410    now = datetime.now()
 411    dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
 412
 413    if GPA.pc.get_verbose():
 414        print(
 415            f'1 saving break {dt_string}_noImprove_lr_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
 416        )
 417
 418    GPA.pai_tracker.save_graphs(
 419        f'{dt_string}_noImprove_lr_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
 420    )
 421
 422    if (
 423        GPA.pai_tracker.member_vars["num_dendrite_tries"]
 424        < GPA.pc.get_max_dendrite_tries() -1
 425    ):
 426        _pai_log(
 427            "info",
 428            f"The newest added dendrites did not improve but current tries "
 429            f'{GPA.pai_tracker.member_vars["num_dendrite_tries"] + 1} '
 430            f"is less than max tries {GPA.pc.get_max_dendrite_tries()} "
 431            f"so loading last switch and trying new Dendrites.",
 432        )
 433        old_tries = GPA.pai_tracker.member_vars["num_dendrite_tries"]
 434        # Load best model from previous n mode
 435        net = UPA.change_learning_modes(
 436            net,
 437            GPA.pc.get_save_name(),
 438            "best_model",
 439            GPA.pai_tracker.member_vars["doing_pai"],
 440        )
 441        GPA.pai_tracker.member_vars["num_dendrite_tries"] = old_tries + 1
 442        return NETWORK_RESTRUCTURED, net
 443    else:
 444        _pai_log(
 445            "info",
 446            f"The newest added dendrites did not improve system and "
 447            f'{GPA.pai_tracker.member_vars["num_dendrite_tries"] + 1} > '
 448            f"{GPA.pc.get_max_dendrite_tries()} so returning training_complete.",
 449        )
 450        _pai_log("info", "You should now exit your training loop and best_model will be your final model for inference")
 451        if not GPA.pc.get_perforated_backpropagation() and GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
 452            _pai_log("info", "For improved results, try perforated backpropagation next time!")
 453        old_silent = GPA.pc.get_silent()
 454        GPA.pc.set_silent(True)
 455        UPA.load_system(net, GPA.pc.get_save_name(), "best_model", switch_call=True)
 456        GPA.pc.set_silent(old_silent)
 457        GPA.pai_tracker.save_graphs()
 458        UPA.pai_save_system(net, GPA.pc.get_save_name(), "final_clean")
 459        return TRAINING_COMPLETE, net
 460
 461
 462def process_final_network(net):
 463    """When the max number of dendrites has been hit load the best_model and return
 464
 465    Parameters
 466    ----------
 467    net : torch.nn.Module
 468        The neural network model being trained.
 469
 470    Returns
 471    -------
 472    torch.nn.Module
 473        The final neural network model.
 474    """
 475
 476    _pai_log("info", f"Last Dendrites were good and this hit the max of {GPA.pc.get_max_dendrites()}")
 477    if not GPA.pc.get_perforated_backpropagation() and GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
 478        _pai_log("info", "For improved results, try perforated backpropagation next time!")
 479    GPA.pai_tracker.save_graphs("before_final")
 480    UPA.load_system(net, GPA.pc.get_save_name(), "best_model", switch_call=True)
 481    GPA.pai_tracker.save_graphs()
 482    UPA.pai_save_system(net, GPA.pc.get_save_name(), "final_clean")
 483    return net
 484
 485
 486def process_scheduler_update(net, accuracy, epochs_since_cycle_switch):
 487    """Updates the scheduler
 488
 489    This increments the scheduler, but if we are automatically sweeping
 490    to find the best initial learning rate for a new set of dendrites
 491    this function also triggers the network at addition time to
 492    try the next value.
 493
 494    Process for finding best initial learning rate for dendrites:
 495    1. Start at default rate
 496    2. Learn at that rate until scheduler increments twice
 497    3. Save that version, start dendrites at LR current increment - 1
 498    4. Repeat 2 and 3 until version has worse final score at set LR
 499    5. Load previous model with best accuracy at that LR as initial rate
 500
 501    Parameters
 502    ----------
 503    net : torch.nn.Module
 504        The neural network model being trained.
 505    accuracy : float
 506        The accuracy of the model at the current learning rate.
 507    epochs_since_cycle_switch : int
 508        The number of epochs since the last cycle switch.
 509
 510    Returns
 511    -------
 512    int
 513        The status of restructuring or training completion.
 514    torch.nn.Module
 515        The potentially modified neural network model.
 516    """
 517
 518    restructured = False
 519    for param_group in GPA.pai_tracker.member_vars["optimizer_instance"].param_groups:
 520        learning_rate1 = param_group["lr"]
 521
 522    if (
 523        type(GPA.pai_tracker.member_vars["scheduler_instance"])
 524        is torch.optim.lr_scheduler.ReduceLROnPlateau
 525    ):
 526        if (
 527            epochs_since_cycle_switch > GPA.pc.get_initial_history_after_switches()
 528            or GPA.pai_tracker.member_vars["mode"] == "p"
 529        ):
 530            if GPA.pc.get_verbose():
 531                print(
 532                    f"Updating scheduler with last improved "
 533                    f'{GPA.pai_tracker.member_vars["epoch_last_improved"]} '
 534                    f'from current {GPA.pai_tracker.member_vars["num_epochs_run"]}'
 535                )
 536            if GPA.pai_tracker.member_vars["scheduler"] is not None:
 537                GPA.pai_tracker.member_vars["scheduler_instance"].step(metrics=accuracy)
 538                if (
 539                    GPA.pai_tracker.member_vars["scheduler"]
 540                    is torch.optim.lr_scheduler.ReduceLROnPlateau
 541                ):
 542                    if GPA.pc.get_verbose():
 543                        print(
 544                            f"Scheduler is now at "
 545                            f'{GPA.pai_tracker.member_vars["scheduler_instance"].num_bad_epochs} bad epochs'
 546                        )
 547        else:
 548            if GPA.pc.get_verbose():
 549                print("Not stepping optimizer since hasnt initialized")
 550
 551    elif GPA.pai_tracker.member_vars["scheduler"] is not None:
 552        if (
 553            epochs_since_cycle_switch > GPA.pc.get_initial_history_after_switches()
 554            or GPA.pai_tracker.member_vars["mode"] == "p"
 555        ):
 556            if GPA.pc.get_verbose():
 557                if hasattr(GPA.pai_tracker.member_vars["scheduler_instance"], '_step_count'):
 558                    count = GPA.pai_tracker.member_vars["scheduler_instance"]._step_count
 559                else:
 560                    count = GPA.pai_tracker.member_vars["scheduler_instance"].last_epoch
 561
 562                print(
 563                    f"Incrementing scheduler to count "
 564                    f'{count}'
 565                )
 566            GPA.pai_tracker.member_vars["scheduler_instance"].step()
 567            if (
 568                GPA.pai_tracker.member_vars["scheduler"]
 569                is torch.optim.lr_scheduler.ReduceLROnPlateau
 570            ):
 571                if GPA.pc.get_verbose():
 572                    print(
 573                        f"Scheduler is now at "
 574                        f'{GPA.pai_tracker.member_vars["scheduler_instance"].num_bad_epochs} bad epochs'
 575                    )
 576
 577    if (
 578        epochs_since_cycle_switch <= GPA.pc.get_initial_history_after_switches()
 579        and GPA.pai_tracker.member_vars["mode"] == "n"
 580    ):
 581        if GPA.pc.get_verbose():
 582            print(
 583                f"Not stepping with history {GPA.pc.get_initial_history_after_switches()} "
 584                f"and current {epochs_since_cycle_switch}"
 585            )
 586
 587    for param_group in GPA.pai_tracker.member_vars["optimizer_instance"].param_groups:
 588        learning_rate2 = param_group["lr"]
 589
 590    stepped = False
 591    at_last_count = False
 592
 593    if GPA.pc.get_verbose():
 594        print(
 595            f"Checking if at last with scores "
 596            f'{len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"])}, '
 597            f"count since switch {epochs_since_cycle_switch} "
 598            f"and last total lr step count "
 599            f'{GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]}'
 600        )
 601
 602    # Check if at double or exactly the test count
 603    if (
 604        len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 0
 605        and epochs_since_cycle_switch
 606        == GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"] * 2
 607    ) or (
 608        len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 1
 609        and epochs_since_cycle_switch
 610        == GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]
 611    ):
 612        at_last_count = True
 613
 614    if GPA.pc.get_verbose():
 615        print(
 616            f"At last count {at_last_count} with count {epochs_since_cycle_switch} "
 617            f'and last LR count {GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]}'
 618        )
 619
 620    if learning_rate1 != learning_rate2:
 621        stepped = True
 622        GPA.pai_tracker.member_vars["current_step_count"] += 1
 623
 624        if GPA.pc.get_verbose():
 625            print(
 626                f"Learning rate just stepped to {learning_rate2:.10e} "
 627                f'with {GPA.pai_tracker.member_vars["current_step_count"]} total steps'
 628            )
 629
 630        if (
 631            GPA.pai_tracker.member_vars["current_step_count"]
 632            == GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]
 633        ):
 634            if GPA.pc.get_verbose():
 635                print(
 636                    f'{GPA.pai_tracker.member_vars["current_step_count"]} '
 637                    f"steps is the max of the last switch mode"
 638                )
 639            # Set it when 1->2 gets to 2, not when 0->1 hits 2 as stopping point
 640            if (
 641                GPA.pai_tracker.member_vars["current_step_count"]
 642                - GPA.pai_tracker.member_vars[
 643                    "current_n_learning_rate_initial_skip_steps"
 644                ]
 645                == 1
 646            ):
 647                GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"] = (
 648                    epochs_since_cycle_switch
 649                )
 650
 651    if GPA.pc.get_verbose():
 652        print(
 653            f"Learning rates were {learning_rate1:.8e} and {learning_rate2:.8e} "
 654            f'started with {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}, '
 655            f'and is now at {GPA.pai_tracker.member_vars["current_step_count"]} '
 656            f'committed {GPA.pai_tracker.member_vars["committed_to_initial_rate"]} '
 657            f"then either this (non zero) or eventually comparing to "
 658            f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
 659            f'steps or rate {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]:.8f}'
 660        )
 661
 662    # If learning rate just stepped, check restart at lower rate
 663    if (
 664        (GPA.pai_tracker.member_vars["scheduler"] is not None)
 665        and
 666        # If potentially might have higher accuracy
 667        (
 668            (GPA.pai_tracker.member_vars["mode"] == "n")
 669            or GPA.pc.get_learn_dendrites_live()
 670        )
 671        and
 672        # And learning rate just stepped
 673        (stepped or at_last_count)
 674    ):
 675
 676        # If this is the first dendrite addition (last_max_learning_rate_steps == 0),
 677        # immediately commit to the initial rate without searching
 678        if GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] == 0:
 679            if GPA.pc.get_verbose():
 680                print(
 681                    f"First dendrite addition detected (last_max_learning_rate_steps == 0), "
 682                    f"immediately committing to initial rate without search"
 683                )
 684            GPA.pai_tracker.member_vars["committed_to_initial_rate"] = True
 685            GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] = (
 686                GPA.pai_tracker.member_vars["current_step_count"]
 687            )
 688            GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = (
 689                learning_rate2
 690            )
 691
 692        # If hasn't committed to a learning rate for this cycle yet
 693        if not GPA.pai_tracker.member_vars["committed_to_initial_rate"]:
 694            best_score_so_far = GPA.pai_tracker.member_vars[
 695                "global_best_validation_score"
 696            ]
 697
 698            if GPA.pc.get_verbose():
 699                print(
 700                    f"In statements to check next learning rate with "
 701                    f"stepped {stepped} and max count {at_last_count}"
 702                )
 703
 704            # If no scores saved for this dendrite and initial LR test did second step
 705            if len(
 706                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]
 707            ) == 0 and (
 708                GPA.pai_tracker.member_vars["current_step_count"]
 709                - GPA.pai_tracker.member_vars[
 710                    "current_n_learning_rate_initial_skip_steps"
 711                ]
 712                == 2
 713                or at_last_count
 714            ):
 715
 716                restructured = True
 717                GPA.pai_tracker.clear_optimizer_and_scheduler()
 718
 719                # Save system for this initial condition
 720                old_global = GPA.pai_tracker.member_vars["global_best_validation_score"]
 721                old_accuracy = GPA.pai_tracker.member_vars[
 722                    "current_best_validation_score"
 723                ]
 724                old_counts = GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]
 725                skip1 = GPA.pai_tracker.member_vars[
 726                    "current_n_learning_rate_initial_skip_steps"
 727                ]
 728
 729                now = datetime.now()
 730                dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
 731
 732                GPA.pai_tracker.save_graphs(
 733                    f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
 734                )
 735
 736                if GPA.pc.get_test_saves():
 737                    UPA.save_system(
 738                        net,
 739                        GPA.pc.get_save_name(),
 740                        f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
 741                    )
 742
 743                if GPA.pc.get_verbose():
 744                    print(
 745                        f"Saving with initial steps: {dt_string}_PBCount_"
 746                        f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
 747                        f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
 748                        f"with current best {old_accuracy}"
 749                    )
 750
 751                # Load back at start and try with lower initial learning rate
 752                net = UPA.load_system(
 753                    net,
 754                    GPA.pc.get_save_name(),
 755                    f'switch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
 756                    switch_call=True,
 757                )
 758                GPA.pai_tracker.member_vars[
 759                    "current_n_learning_rate_initial_skip_steps"
 760                ] = (skip1 + 1)
 761                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"].append(
 762                    old_accuracy
 763                )
 764                GPA.pai_tracker.member_vars["global_best_validation_score"] = old_global
 765                GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"] = old_counts
 766
 767            # If there is one score already, this is first step at next score
 768            elif len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 1:
 769                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"].append(
 770                    GPA.pai_tracker.member_vars["current_best_validation_score"]
 771                )
 772
 773                # If this LR's score was worse than last LR's score
 774                lr_score_worse = False
 775                if GPA.pai_tracker.member_vars["maximizing_score"]:
 776                    lr_score_worse = (
 777                        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]
 778                        > GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]
 779                    )
 780                else:
 781                    lr_score_worse = (
 782                        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]
 783                        < GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]
 784                    )
 785
 786                if lr_score_worse:
 787                    restructured = True
 788                    GPA.pai_tracker.clear_optimizer_and_scheduler()
 789
 790                    if GPA.pc.get_verbose():
 791                        print(
 792                            f'Got initial {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]-1} '
 793                            f'step score {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]} '
 794                            f'and {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
 795                            f'score at step {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]} '
 796                            f"so loading old score"
 797                        )
 798
 799                    prior_best = GPA.pai_tracker.member_vars[
 800                        "current_cycle_lr_max_scores"
 801                    ][0]
 802
 803                    now = datetime.now()
 804                    dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
 805
 806                    GPA.pai_tracker.save_graphs(
 807                        f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
 808                    )
 809
 810                    if GPA.pc.get_test_saves():
 811                        UPA.save_system(
 812                            net,
 813                            GPA.pc.get_save_name(),
 814                            f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
 815                        )
 816
 817                    if GPA.pc.get_verbose():
 818                        print(
 819                            f"Saving with initial steps: {dt_string}_PBCount_"
 820                            f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
 821                            f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
 822                        )
 823
 824                    if GPA.pc.get_test_saves():
 825                        net = UPA.load_system(
 826                            net,
 827                            GPA.pc.get_save_name(),
 828                            f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]-1}',
 829                            switch_call=True,
 830                        )
 831
 832                    # Save graphs for chosen one
 833                    now = datetime.now()
 834                    dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
 835
 836                    GPA.pai_tracker.save_graphs(
 837                        f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}PICKED'
 838                    )
 839
 840                    if GPA.pc.get_test_saves():
 841                        UPA.save_system(
 842                            net,
 843                            GPA.pc.get_save_name(),
 844                            f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
 845                        )
 846
 847                    if GPA.pc.get_verbose():
 848                        print(
 849                            f"Saving with initial steps: {dt_string}_PBCount_"
 850                            f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
 851                            f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
 852                        )
 853
 854                    GPA.pai_tracker.member_vars["committed_to_initial_rate"] = True
 855                    GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] = (
 856                        GPA.pai_tracker.member_vars["current_step_count"]
 857                    )
 858                    GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = (
 859                        learning_rate2
 860                    )
 861                    GPA.pai_tracker.member_vars["current_best_validation_score"] = (
 862                        prior_best
 863                    )
 864
 865                    if GPA.pc.get_verbose():
 866                        print(
 867                            f"Setting last max steps to "
 868                            f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
 869                            f'and lr {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]}'
 870                        )
 871
 872                else:  # Current LR score is better
 873                    if GPA.pc.get_verbose():
 874                        print(
 875                            f'Got initial {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]-1} '
 876                            f'step score {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]} '
 877                            f'and {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
 878                            f'score at step {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]} '
 879                            f"so NOT loading old score and continuing with this score"
 880                        )
 881
 882                    if at_last_count:  # If this is the last one, set it to be picked
 883                        restructured = True
 884                        GPA.pai_tracker.clear_optimizer_and_scheduler()
 885
 886                        now = datetime.now()
 887                        dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
 888
 889                        GPA.pai_tracker.save_graphs(
 890                            f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}PICKED'
 891                        )
 892
 893                        if GPA.pc.get_test_saves():
 894                            UPA.save_system(
 895                                net,
 896                                GPA.pc.get_save_name(),
 897                                f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
 898                            )
 899
 900                        if GPA.pc.get_verbose():
 901                            print(
 902                                f"Saving with initial steps: {dt_string}_PBCount_"
 903                                f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
 904                                f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
 905                            )
 906
 907                        GPA.pai_tracker.member_vars["committed_to_initial_rate"] = True
 908                        GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] = (
 909                            GPA.pai_tracker.member_vars["current_step_count"]
 910                        )
 911                        GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = (
 912                            learning_rate2
 913                        )
 914
 915                        if GPA.pc.get_verbose():
 916                            print(
 917                                f"Setting last max steps to "
 918                                f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
 919                                f'and lr {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]}'
 920                            )
 921
 922                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
 923
 924            elif len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 2:
 925                print(
 926                    "Should never be here. Please let Perforated AI know if this happened."
 927                )
 928                pdb.set_trace()
 929
 930            GPA.pai_tracker.member_vars["global_best_validation_score"] = (
 931                best_score_so_far
 932            )
 933
 934        else:
 935            if GPA.pc.get_verbose():
 936                print(
 937                    f"Setting last max steps to "
 938                    f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
 939                    f'and lr {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]}'
 940                )
 941            GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] += 1
 942            GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = learning_rate2
 943    if restructured:
 944        return NETWORK_RESTRUCTURED, net
 945    else:
 946        return NO_MODEL_UPDATE, net
 947
 948
 949class PAINeuronModuleTracker:
 950    """
 951    Manager class that tracks all neuron layers and dendrite layers,
 952    controls when new dendrites are added, and communicates signals to modules.
 953    """
 954
 955    def __init__(
 956        self,
 957        doing_pai,
 958        save_name,
 959        making_graphs=True,
 960        param_vals_setting=-1,
 961        values_per_train_epoch=-1,
 962        values_per_val_epoch=-1,
 963    ):
 964        """Initialize the tracker
 965
 966        Parameters
 967        ----------
 968        doing_pai : bool
 969            Whether or not dendrites should be used.
 970        save_name : str
 971            The base name for saving models and graphs.
 972        making_graphs : bool, optional
 973            Whether or not to generate graphs, by default True.
 974        param_vals_setting : int, optional
 975            Parameter values setting, by default -1.
 976        values_per_train_epoch : int, optional
 977            The number of values to look back for graphing
 978            during training, by default -1 (all values).
 979        values_per_val_epoch : int, optional
 980            The number of values to look back for graphing
 981            during validation, by default -1 (all values).
 982        Returns
 983        -------
 984        None
 985        """
 986
 987        # Dict of member vars and their types for saving
 988        self.member_vars = {}
 989        self.member_var_types = {}
 990
 991        # Whether or not PAI will be running
 992        self.member_vars["doing_pai"] = doing_pai
 993        self.member_var_types["doing_pai"] = "bool"
 994
 995        # How many Dendrites have been added
 996        self.member_vars["num_dendrites_added"] = 0
 997        self.member_var_types["num_dendrites_added"] = "int"
 998
 999        # How many Dendrites have been successfully integrated, does not count currently training dendrites
1000        self.member_vars["num_dendrites_integrated"] = 0
1001        self.member_var_types["num_dendrites_integrated"] = "int"
1002
1003        # How many cycles have been run, *2 or *2+1 of the above
1004        self.member_vars["num_cycles"] = 0
1005        self.member_var_types["num_cycles"] = "int"
1006
1007        # Pointers to all neuron wrapped modules
1008        self.neuron_module_vector = []
1009
1010        # Pointers to all non neuron modules for tracking
1011        self.tracked_neuron_module_vector = []
1012
1013        # Neuron training or dendrite training mode
1014        self.member_vars["mode"] = "n"
1015        self.member_var_types["mode"] = "string"
1016
1017        # Number of epochs run excluding overwritten epochs
1018        self.member_vars["num_epochs_run"] = -1
1019        self.member_var_types["num_epochs_run"] = "int"
1020
1021        # Number including overwritten epochs
1022        self.member_vars["total_epochs_run"] = -1
1023        self.member_var_types["total_epochs_run"] = "int"
1024
1025        # Last epoch that validation/correlation score was improved
1026        self.member_vars["epoch_last_improved"] = 0
1027        self.member_var_types["epoch_last_improved"] = "int"
1028
1029        # Running validation accuracy
1030        self.member_vars["running_accuracy"] = 0
1031        self.member_var_types["running_accuracy"] = "float"
1032
1033        # True if maxing validation, False if minimizing Loss
1034        self.member_vars["maximizing_score"] = True
1035        self.member_var_types["maximizing_score"] = "bool"
1036
1037        # Mode for switching back and forth between learning modes
1038        self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
1039        self.member_var_types["switch_mode"] = "int"
1040
1041        # Epoch of the last switch
1042        self.member_vars["last_switch"] = 0
1043        self.member_var_types["last_switch"] = "int"
1044
1045        # Highest validation score from current cycle
1046        self.member_vars["current_best_validation_score"] = 0
1047        self.member_var_types["current_best_validation_score"] = "float"
1048
1049        # Last epoch where the learning rate was updated
1050        self.member_vars["initial_lr_test_epoch_count"] = -1
1051        self.member_var_types["initial_lr_test_epoch_count"] = "int"
1052
1053        # Highest validation score of full run
1054        self.member_vars["global_best_validation_score"] = 0
1055        self.member_var_types["global_best_validation_score"] = "float"
1056
1057        # List of switch epochs
1058        self.member_vars["switch_epochs"] = []
1059        self.member_var_types["switch_epochs"] = "int array"
1060
1061        # Parameter counts at each network structure
1062        self.member_vars["param_counts"] = []
1063        self.member_var_types["param_counts"] = "int array"
1064
1065        # List of epochs where switch was made to neuron training
1066        self.member_vars["n_switch_epochs"] = []
1067        self.member_var_types["n_switch_epochs"] = "int array"
1068
1069        # List of epochs where switch was made to dendrite training
1070        self.member_vars["p_switch_epochs"] = []
1071        self.member_var_types["p_switch_epochs"] = "int array"
1072
1073        # List of validation accuracies
1074        self.member_vars["accuracies"] = []
1075        self.member_var_types["accuracies"] = "float array"
1076
1077        # List of epochs where score improved for scheduler updates
1078        self.member_vars["last_improved_accuracies"] = []
1079        self.member_var_types["last_improved_accuracies"] = "int array"
1080
1081        # List of test accuracy scores registered
1082        self.member_vars["test_accuracies"] = []
1083        self.member_var_types["test_accuracies"] = "float array"
1084
1085        # List of accuracies registered during neuron training
1086        self.member_vars["n_accuracies"] = []
1087        self.member_var_types["n_accuracies"] = "float array"
1088
1089        # List of accuracies registered during dendrite training
1090        self.member_vars["p_accuracies"] = []
1091        self.member_var_types["p_accuracies"] = "float array"
1092
1093        # Running average accuracies from recent epochs
1094        self.member_vars["running_accuracies"] = []
1095        self.member_var_types["running_accuracies"] = "float array"
1096
1097        # List of additional scores recorded
1098        self.member_vars["extra_scores"] = {}
1099        self.member_var_types["extra_scores"] = "float array dictionary"
1100
1101        # Extra scores not set to be graphed
1102        self.member_vars["extra_scores_without_graphing"] = {}
1103        self.member_var_types["extra_scores_without_graphing"] = (
1104            "float array dictionary"
1105        )
1106
1107        # List of test scores
1108        self.member_vars["test_scores"] = []
1109        self.member_var_types["test_scores"] = "float array"
1110
1111        # Extra scores calculated during neuron training
1112        self.member_vars["n_extra_scores"] = {}
1113        self.member_var_types["n_extra_scores"] = "float array dictionary"
1114
1115        # List of training losses calculated
1116        self.member_vars["training_loss"] = []
1117        self.member_var_types["training_loss"] = "float array"
1118
1119        # List of learning rates at each epoch
1120        self.member_vars["training_learning_rates"] = []
1121        self.member_var_types["training_learning_rates"] = "float array"
1122
1123        # Best dendrite scores
1124        self.member_vars["best_scores"] = []
1125        self.member_var_types["best_scores"] = "float array array"
1126
1127        # Current dendrite scores
1128        self.member_vars["current_scores"] = []
1129        self.member_var_types["current_scores"] = "float array array"
1130
1131        # Times for neuron training epochs
1132        self.member_vars["n_epoch_times"] = []
1133        self.member_var_types["n_epoch_times"] = "float array"
1134
1135        # Timing values
1136        self.member_vars["p_epoch_times"] = []
1137        self.member_var_types["p_epoch_times"] = "float array"
1138        self.member_vars["n_train_times"] = []
1139        self.member_var_types["n_train_times"] = "float array"
1140        self.member_vars["p_train_times"] = []
1141        self.member_var_types["p_train_times"] = "float array"
1142        self.member_vars["n_val_times"] = []
1143        self.member_var_types["n_val_times"] = "float array"
1144        self.member_vars["p_val_times"] = []
1145        self.member_var_types["p_val_times"] = "float array"
1146
1147        # Setting for tracking timing
1148        self.member_vars["manual_train_switch"] = False
1149        self.member_var_types["manual_train_switch"] = "bool"
1150
1151        # Tracking scores overwritten when reloading best model
1152        self.member_vars["overwritten_extras"] = []
1153        self.member_var_types["overwritten_extras"] = "float array dictionary array"
1154        self.member_vars["overwritten_vals"] = []
1155        self.member_var_types["overwritten_vals"] = "float array array"
1156        self.member_vars["overwritten_epochs"] = 0
1157        self.member_var_types["overwritten_epochs"] = "int"
1158
1159        # Setting for determining scores
1160        self.member_vars["param_vals_setting"] = GPA.pc.get_param_vals_setting()
1161        self.member_var_types["param_vals_setting"] = "int"
1162
1163        # Optimizer and scheduler types and instances
1164        self.member_vars["optimizer"] = None
1165        self.member_var_types["optimizer"] = "type"
1166        self.member_vars["scheduler"] = None
1167        self.member_var_types["scheduler"] = "type"
1168        self.member_vars["optimizer_instance"] = None
1169        self.member_var_types["optimizer_instance"] = "empty array"
1170        self.member_vars["scheduler_instance"] = None
1171        self.member_var_types["scheduler_instance"] = "empty array"
1172
1173        # Flag for if the tracker was loaded
1174        self.loaded = False
1175
1176        # flag for 
1177        self.member_vars["step_status"] = STEP_CLEARED
1178        self.member_var_types["step_status"] = "int"
1179
1180
1181        # Settings for tracking learning rates
1182        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
1183        self.member_var_types["current_n_learning_rate_initial_skip_steps"] = "int"
1184        self.member_vars["last_max_learning_rate_steps"] = 0
1185        self.member_var_types["last_max_learning_rate_steps"] = "int"
1186        self.member_vars["last_max_learning_rate_value"] = -1
1187        self.member_var_types["last_max_learning_rate_value"] = "float"
1188        self.member_vars["current_cycle_lr_max_scores"] = []
1189        self.member_var_types["current_cycle_lr_max_scores"] = "float array"
1190        self.member_vars["current_step_count"] = 0
1191        self.member_var_types["current_step_count"] = "int"
1192        self.member_vars["committed_to_initial_rate"] = True
1193        self.member_var_types["committed_to_initial_rate"] = "bool"
1194        self.member_vars["best_mean_score_improved_this_epoch"] = 0
1195        self.member_var_types["best_mean_score_improved_this_epoch"] = "int"
1196
1197        # Flag for if current dendrite achieved highest global score
1198        self.member_vars["current_n_set_global_best"] = True
1199        self.member_var_types["current_n_set_global_best"] = "bool"
1200
1201        # Number of tries adding this dendrite count
1202        self.member_vars["num_dendrite_tries"] = 0
1203        self.member_var_types["num_dendrite_tries"] = "int"
1204
1205        # Count of batches per epoch
1206        self.values_per_train_epoch = values_per_train_epoch
1207        self.values_per_val_epoch = values_per_val_epoch
1208
1209        self.save_name = save_name
1210        self.making_graphs = making_graphs
1211
1212        self.start_time = time.time()
1213        self.saved_time = 0
1214        self.start_epoch(internal_call=True)
1215
1216        if GPA.pc.get_verbose():
1217            print(f'Initializing with switch_mode {self.member_vars["switch_mode"]}')
1218
1219    def to_string(self):
1220        """Convert tracker values to string for saving with safetensors.
1221
1222        Parameters
1223        ----------
1224        None
1225
1226        Returns
1227        -------
1228        str
1229            Serialized tracker state suitable for storage in a safetensors field.
1230        """
1231
1232        full_string = ""
1233        for var in self.member_vars:
1234            full_string += var + ","
1235            if self.member_vars[var] is None:
1236                full_string += "None"
1237                full_string += "\n"
1238            elif self.member_var_types[var] == "bool":
1239                full_string += str(self.member_vars[var])
1240                full_string += "\n"
1241            elif self.member_var_types[var] in ("int", "float", "string"):
1242                full_string += str(self.member_vars[var])
1243                full_string += "\n"
1244            elif self.member_var_types[var] == "type":
1245                name = (
1246                    self.member_vars[var].__module__
1247                    + "."
1248                    + self.member_vars[var].__name__
1249                )
1250                full_string += str(self.member_vars[var])
1251                full_string += "\n"
1252            elif self.member_var_types[var] == "empty array":
1253                full_string += "[]"
1254                full_string += "\n"
1255            elif self.member_var_types[var] in ("int array", "float array"):
1256                full_string += "\n"
1257                string = ""
1258                for val in self.member_vars[var]:
1259                    string += str(val) + ","
1260                # Remove the last comma
1261                string = string[:-1]
1262                full_string += string
1263                full_string += "\n"
1264            elif self.member_var_types[var] == "float array dictionary array":
1265                full_string += "\n"
1266                for array in self.member_vars[var]:
1267                    for key in array:
1268                        string = key + ","
1269                        for val in array[key]:
1270                            string += str(val) + ","
1271                        # Remove the last comma
1272                        string = string[:-1]
1273                        full_string += string
1274                        full_string += "\n"
1275                    full_string += "endkey"
1276                    full_string += "\n"
1277                full_string += "endarray"
1278                full_string += "\n"
1279            elif self.member_var_types[var] == "float array dictionary":
1280                full_string += "\n"
1281                for key in self.member_vars[var]:
1282                    string = key + ","
1283                    for val in self.member_vars[var][key]:
1284                        string += str(val) + ","
1285                    # Remove the last comma
1286                    string = string[:-1]
1287                    full_string += string
1288                    full_string += "\n"
1289                full_string += "end"
1290                full_string += "\n"
1291            elif self.member_var_types[var] == "float array array":
1292                full_string += "\n"
1293                for array in self.member_vars[var]:
1294                    string = ""
1295                    for val in array:
1296                        string += str(val) + ","
1297                    # Remove the last comma
1298                    string = string[:-1]
1299                    full_string += string
1300                    full_string += "\n"
1301                full_string += "end"
1302                full_string += "\n"
1303            else:
1304                print("Did not find a member variable")
1305                pdb.set_trace()
1306        return full_string
1307
1308    def from_string(self, string):
1309        """Load tracker values from string.
1310
1311        Parameters
1312        ----------
1313        string : str
1314            The string to load from.
1315
1316        Returns
1317        -------
1318        None
1319            This function does not return a value.
1320        """
1321        f = io.StringIO(string)
1322        while True:
1323            line = f.readline()
1324            if not line:
1325                break
1326            vals = line.split(",")
1327            var = vals[0]
1328
1329            if self.member_var_types[var] == "bool":
1330                val = vals[1][:-1]
1331                if val == "True":
1332                    self.member_vars[var] = True
1333                elif val == "False":
1334                    self.member_vars[var] = False
1335                elif val == "1":
1336                    self.member_vars[var] = 1
1337                elif val == "0":
1338                    self.member_vars[var] = 0
1339                else:
1340                    print("Something went wrong with loading")
1341                    pdb.set_trace()
1342            elif self.member_var_types[var] == "int":
1343                val = vals[1]
1344                self.member_vars[var] = int(val)
1345            elif self.member_var_types[var] == "float":
1346                val = vals[1]
1347                self.member_vars[var] = float(val)
1348            elif self.member_var_types[var] == "string":
1349                val = vals[1][:-1]
1350                self.member_vars[var] = val
1351            elif self.member_var_types[var] == "type":
1352                # Ignore loading types, tracker should have them set up
1353                continue
1354            elif self.member_var_types[var] == "empty array":
1355                val = vals[1]
1356                self.member_vars[var] = []
1357            elif self.member_var_types[var] == "int array":
1358                vals = f.readline()[:-1].split(",")
1359                self.member_vars[var] = []
1360                if vals[0] == "":
1361                    continue
1362                for val in vals:
1363                    self.member_vars[var].append(int(val))
1364            elif self.member_var_types[var] == "float array":
1365                vals = f.readline()[:-1].split(",")
1366                self.member_vars[var] = []
1367                if vals[0] == "":
1368                    continue
1369                for val in vals:
1370                    self.member_vars[var].append(float(val))
1371            elif self.member_var_types[var] == "float array dictionary array":
1372                self.member_vars[var] = []
1373                line2 = f.readline()[:-1]
1374                while line2 != "endarray":
1375                    temp = {}
1376                    while line2 != "endkey":
1377                        vals = line2.split(",")
1378                        name = vals[0]
1379                        temp[name] = []
1380                        vals = vals[1:]
1381                        for val in vals:
1382                            temp[name].append(float(val))
1383                        line2 = f.readline()[:-1]
1384                    self.member_vars[var].append(temp)
1385                    line2 = f.readline()[:-1]
1386            elif self.member_var_types[var] == "float array dictionary":
1387                self.member_vars[var] = {}
1388                line2 = f.readline()[:-1]
1389                while line2 != "end":
1390                    vals = line2.split(",")
1391                    name = vals[0]
1392                    self.member_vars[var][name] = []
1393                    vals = vals[1:]
1394                    for val in vals:
1395                        self.member_vars[var][name].append(float(val))
1396                    line2 = f.readline()[:-1]
1397            elif self.member_var_types[var] == "float array array":
1398                self.member_vars[var] = []
1399                line2 = f.readline()[:-1]
1400                while line2 != "end":
1401                    vals = line2.split(",")
1402                    self.member_vars[var].append([])
1403                    if line2:
1404                        for val in vals:
1405                            self.member_vars[var][-1].append(float(val))
1406                    line2 = f.readline()[:-1]
1407            else:
1408                print("Did not find a member variable")
1409
1410                pdb.set_trace()
1411
1412    def from_string_debug(self, string):
1413        """Debug function to print tracker values from string without loading them.
1414
1415        Parameters
1416        ----------
1417        string : str
1418            The string to debug load from.
1419
1420        Returns
1421        -------
1422        None
1423            This function does not return a value.
1424        """
1425        f = io.StringIO(string)
1426        print("=== DEBUGGING TRACKER VARIABLES ===")
1427
1428        while True:
1429            line = f.readline()
1430            if not line:
1431                break
1432            vals = line.split(",")
1433            var = vals[0]
1434
1435            print(f"\nVariable: {var}")
1436            print(f"Type: {self.member_var_types.get(var, 'UNKNOWN TYPE')}")
1437            print(f"Current value: {self.member_vars.get(var, 'NOT SET')}")
1438
1439            if self.member_var_types.get(var) == "bool":
1440                val = vals[1][:-1]
1441                print(f"Would set to: {val} -> {val == 'True'}")
1442
1443            elif self.member_var_types.get(var) == "int":
1444                val = vals[1]
1445                print(f"Would set to: {int(val)}")
1446
1447            elif self.member_var_types.get(var) == "float":
1448                val = vals[1]
1449                print(f"Would set to: {float(val)}")
1450
1451            elif self.member_var_types.get(var) == "string":
1452                val = vals[1][:-1]
1453                print(f"Would set to: '{val}'")
1454
1455            elif self.member_var_types.get(var) == "type":
1456                print("Would skip (type loading)")
1457
1458            elif self.member_var_types.get(var) == "empty array":
1459                val = vals[1]
1460                print(f"Would set to: [] (empty array)")
1461
1462            elif self.member_var_types.get(var) == "int array":
1463                vals_line = f.readline()[:-1].split(",")
1464                print(f"Would set to int array with {len(vals_line)} elements:")
1465                if vals_line[0] != "":
1466                    print(
1467                        f"  Elements: {vals_line[:5]}{'...' if len(vals_line) > 5 else ''}"
1468                    )
1469                else:
1470                    print("  Empty array")
1471
1472            elif self.member_var_types.get(var) == "float array":
1473                vals_line = f.readline()[:-1].split(",")
1474                print(f"Would set to float array with {len(vals_line)} elements:")
1475                if vals_line[0] != "":
1476                    print(
1477                        f"  Elements: {vals_line[:5]}{'...' if len(vals_line) > 5 else ''}"
1478                    )
1479                else:
1480                    print("  Empty array")
1481
1482            elif self.member_var_types.get(var) == "float array dictionary array":
1483                print("Would process float array dictionary array:")
1484                array_count = 0
1485                line2 = f.readline()[:-1]
1486                while line2 != "endarray":
1487                    key_count = 0
1488                    while line2 != "endkey":
1489                        vals_dict = line2.split(",")
1490                        name = vals_dict[0]
1491                        print(
1492                            f"  Array {array_count}, Key '{name}': {len(vals_dict)-1} elements"
1493                        )
1494                        key_count += 1
1495                        line2 = f.readline()[:-1]
1496                    print(f"  Array {array_count} has {key_count} keys")
1497                    array_count += 1
1498                    line2 = f.readline()[:-1]
1499                print(f"  Total arrays: {array_count}")
1500
1501            elif self.member_var_types.get(var) == "float array dictionary":
1502                print("Would process float array dictionary:")
1503                line2 = f.readline()[:-1]
1504                key_count = 0
1505                while line2 != "end":
1506                    vals_dict = line2.split(",")
1507                    name = vals_dict[0]
1508                    print(f"  Key '{name}': {len(vals_dict)-1} elements")
1509                    key_count += 1
1510                    line2 = f.readline()[:-1]
1511                print(f"  Total keys: {key_count}")
1512
1513            elif self.member_var_types.get(var) == "float array array":
1514                print("Would process float array array:")
1515                line2 = f.readline()[:-1]
1516                array_count = 0
1517                while line2 != "end":
1518                    if line2:
1519                        vals_array = line2.split(",")
1520                        print(f"  Array {array_count}: {len(vals_array)} elements")
1521                    else:
1522                        print(f"  Array {array_count}: empty")
1523                    array_count += 1
1524                    line2 = f.readline()[:-1]
1525                print(f"  Total arrays: {array_count}")
1526
1527            else:
1528                print(f"UNKNOWN TYPE: {self.member_var_types.get(var, 'NOT FOUND')}")
1529
1530        print("\n=== END DEBUG ===")
1531
1532    def save_tracker_settings(self):
1533        """Save tracker settings for DistributedDataParallel use.
1534
1535        Saves settings in save_name/array_dims.csv
1536
1537        Parameters
1538        ----------
1539        None
1540        Returns
1541        -------
1542        None
1543
1544        -----
1545        Instructions for use are in API customization.md
1546        """
1547        if not os.path.isdir(self.save_name):
1548            os.makedirs(self.save_name)
1549        f = open(self.save_name + "/array_dims.csv", "w")
1550        for layer in self.neuron_module_vector:
1551            dv = layer.dendrite_module.dendrite_values[0]
1552            shape_str = ",".join(str(s) for s in dv.dendrite_storage_shape.tolist())
1553            f.write(f"{layer.name},{shape_str}\n")
1554        f.close()
1555        if not GPA.pc.get_silent():
1556            print("Tracker settings saved.")
1557            print("You may now delete save_tracker_settings")
1558
1559    def initialize_tracker_settings(self):
1560        """Initialize tracker settings from saved file.
1561
1562        This function loads tracker settings from a CSV file and applies them
1563        to the layers the tracker is managing.
1564
1565        Parameters
1566        ----------
1567        None
1568
1569        Returns
1570        -------
1571        None
1572
1573        """
1574
1575        channels = {}
1576        if not os.path.exists(self.save_name + "/array_dims.csv"):
1577            print(
1578                "You must call save_tracker_settings before "
1579                "initialize_tracker_settings"
1580            )
1581            print("Follow instructions in customization.md")
1582            pdb.set_trace()
1583        f = open(self.save_name + "/array_dims.csv", "r")
1584        for line in f:
1585            parts = line.strip().split(",")
1586            channels[parts[0]] = [int(s) for s in parts[1:]]
1587        for layer in self.neuron_module_vector:
1588            dv = layer.dendrite_module.dendrite_values[0]
1589            dv.setup_arrays(channels[layer.name])
1590
1591    def set_optimizer_instance(self, optimizer_instance, additional_optimizers=[]):
1592        """Set optimizer instance directly.
1593
1594        Parameters
1595        ----------
1596        optimizer_instance : object
1597            The optimizer instance to set.
1598
1599        Returns
1600        -------
1601        None
1602
1603        """
1604        # This call must be first before the parameters get filtered.
1605        optimizer_instance.zero_grad()
1606        try:
1607            for param_group in optimizer_instance.param_groups:
1608                if (
1609                    param_group["weight_decay"] > 0
1610                    and GPA.pc.get_weight_decay_accepted() is False
1611                ):
1612                    _pai_log(
1613                        "warning",
1614                        "For PAI training it is recommended to not use weight decay in your optimizer",
1615                    )
1616
1617        except:
1618            pass
1619        self.member_vars["optimizer_instance"] = optimizer_instance
1620        if GPA.pc.get_perforated_backpropagation():
1621            TPB.setup_optimizer_pb(self.member_vars["optimizer_instance"])
1622            for optimizer in additional_optimizers:
1623                TPB.filter_params(optimizer)
1624
1625    def set_optimizer(self, optimizer):
1626        """Set optimizer type to be initialized later
1627
1628        Parameters
1629        ----------
1630        optimizer : object
1631            The optimizer type to set.
1632
1633        Returns
1634        -------
1635        None
1636
1637        """
1638        self.member_vars["optimizer"] = optimizer
1639
1640    def set_scheduler(self, scheduler):
1641        """Set scheduler type to be initialized later
1642
1643        Parameters
1644        ----------
1645        scheduler : object
1646            The scheduler type to set.
1647
1648        Returns
1649        -------
1650        None
1651
1652        """
1653        if scheduler is not torch.optim.lr_scheduler.ReduceLROnPlateau:
1654            if GPA.pc.get_verbose():
1655                print("Not using ReduceLROnPlateau, this is not recommended")
1656        self.member_vars["scheduler"] = scheduler
1657
1658    def increment_scheduler(self, num_ticks, mode):
1659        """Increment the scheduler a set number of times.
1660
1661        Used for finding best initial learning rate when adding dendrites.
1662
1663        Parameters
1664        ----------
1665        num_ticks : int
1666            The number of scheduler steps to take.
1667        mode : str
1668            The mode for stepping the scheduler. Options are:
1669            - "step_learning_rate": Step based on improved accuracy epochs
1670            - "increment_epoch_count": Step based on total epoch count
1671
1672        Returns
1673        -------
1674        current_steps : int
1675            The number of learning rate changes that occurred.
1676        learning_rate1 : float
1677            The final learning rate after stepping.
1678
1679        """
1680
1681        current_steps = 0
1682        current_ticker = 0
1683
1684        for param_group in GPA.pai_tracker.member_vars[
1685            "optimizer_instance"
1686        ].param_groups:
1687            learning_rate1 = param_group["lr"]
1688
1689        if GPA.pc.get_verbose():
1690            print("Using scheduler:")
1691            print(type(self.member_vars["scheduler_instance"]))
1692
1693        while current_ticker < num_ticks:
1694            if GPA.pc.get_verbose():
1695                print(
1696                    f"Lower start rate initial {learning_rate1} "
1697                    f'stepping {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} times'
1698                )
1699
1700            if (
1701                type(self.member_vars["scheduler_instance"])
1702                is torch.optim.lr_scheduler.ReduceLROnPlateau
1703            ):
1704                if mode == "step_learning_rate":
1705                    # Step with counter as last improved accuracy
1706                    self.member_vars["scheduler_instance"].step(
1707                        metrics=self.member_vars["last_improved_accuracies"][
1708                            GPA.pai_tracker.steps_after_switch() - 1
1709                        ]
1710                    )
1711                elif mode == "increment_epoch_count":
1712                    # Step with improved epoch counts up to current location
1713                    self.member_vars["scheduler_instance"].step(
1714                        metrics=self.member_vars["last_improved_accuracies"][
1715                            -((num_ticks - 1) - current_ticker) - 1
1716                        ]
1717                    )
1718            else:
1719                self.member_vars["scheduler_instance"].step()
1720
1721            for param_group in GPA.pai_tracker.member_vars[
1722                "optimizer_instance"
1723            ].param_groups:
1724                learning_rate2 = param_group["lr"]
1725
1726            if learning_rate2 != learning_rate1:
1727                current_steps += 1
1728                learning_rate1 = learning_rate2
1729                if mode == "step_learning_rate":
1730                    current_ticker += 1
1731                if GPA.pc.get_verbose():
1732                    print(f"1 step {current_steps} to {learning_rate2}")
1733
1734            if mode == "increment_epoch_count":
1735                current_ticker += 1
1736
1737        return current_steps, learning_rate1
1738
1739    def setup_optimizer(self, net, opt_args, sched_args=None, parameters=None):
1740        """Initialize the optimizer and scheduler when added.
1741
1742        Parameters
1743        ----------
1744        net : object
1745            The neural network model.
1746        opt_args : dict
1747            The arguments for the optimizer.
1748        sched_args : dict, optional
1749            The arguments for the scheduler, by default None.
1750
1751        Returns
1752        -------
1753        optimizer : object
1754            The initialized optimizer instance.
1755        scheduler : object or None
1756            The initialized scheduler instance, or None if no scheduler was set.
1757
1758        """
1759        if "weight_decay" in opt_args and not GPA.pc.get_weight_decay_accepted():
1760            _pai_log(
1761                "warning",
1762                "For PAI training it is recommended to not use weight decay in your optimizer",
1763            )
1764
1765        if ("model" not in opt_args.keys()) and "params" not in opt_args.keys():
1766            print("In setup_optimizer it will be depreciated to not pass in params yourself in the future")
1767            print("please change the settings to include params")
1768            if self.member_vars["mode"] == "n":
1769                if parameters is not None:
1770                    opt_args["params"] = parameters
1771                else:
1772                    opt_args["params"] = filter(lambda p: p.requires_grad, net.parameters())
1773            else:
1774                params = UPA.get_pai_network_params(net)
1775                if parameters is not None:
1776                    # Filter parameters to only those in params, preserving weight_decay
1777                    params_set = set(params)
1778                    filtered_params = []
1779                    for param_group in parameters:
1780                        filtered_group_params = [p for p in param_group["params"] if p in params_set]
1781                        if filtered_group_params:
1782                            filtered_params.append({
1783                                "params": filtered_group_params,
1784                                "weight_decay": param_group["weight_decay"]
1785                            })
1786                    opt_args["params"] = filtered_params
1787                else:
1788                    opt_args["params"] = params
1789        elif "params" in opt_args.keys():
1790            # Check if params is a list of param groups (dicts) or a single param group
1791            params_value = opt_args["params"]
1792            if isinstance(params_value, list) and len(params_value) > 0:
1793                # Check if it's a list of dicts (multiple param groups) or list of tensors (single group)
1794                if isinstance(params_value[0], dict):
1795                    # Multiple param groups format: [{"params": [...], "lr": ...}, ...]
1796                    # Filter each param group for requires_grad
1797                    filtered_param_groups = []
1798                    for param_group in params_value:
1799                        filtered_group_params = [p for p in param_group["params"] if p.requires_grad]
1800                        if filtered_group_params:
1801                            new_group = param_group.copy()
1802                            new_group["params"] = filtered_group_params
1803                            filtered_param_groups.append(new_group)
1804                    opt_args["params"] = filtered_param_groups
1805                else:
1806                    # Single param group format: [tensor1, tensor2, ...] or generator
1807                    # Filter for requires_grad
1808                    opt_args["params"] = [p for p in params_value if p.requires_grad]
1809            elif hasattr(params_value, '__iter__'):
1810                # Handle generators or other iterables
1811                opt_args["params"] = [p for p in params_value if p.requires_grad]
1812
1813        optimizer = self.member_vars["optimizer"](**opt_args)
1814        self.set_optimizer_instance(optimizer)
1815
1816        if self.member_vars["scheduler"] is not None:
1817            # Handle SequentialLR specially
1818            if self.member_vars["scheduler"] is torch.optim.lr_scheduler.SequentialLR:
1819                """
1820                sched_args should be a dict with "schedulers" (list of tuples) and "milestones"
1821                For example:
1822                sequential_schedArgs = {
1823                    "schedulers": [
1824                        (warmup_scheduler_class, warmup_schedArgs),
1825                        (main_scheduler_class, main_schedArgs)
1826                    ],
1827                    "milestones": [switch_epoch]
1828                }
1829                """
1830                schedulers = []
1831                milestones = sched_args.get("milestones", [])
1832                scheduler_configs = sched_args.get("schedulers", [])
1833                
1834                for scheduler_class, scheduler_args in scheduler_configs:
1835                    schedulers.append(scheduler_class(optimizer, **scheduler_args))
1836                
1837                self.member_vars["scheduler_instance"] = torch.optim.lr_scheduler.SequentialLR(
1838                    optimizer, schedulers=schedulers, milestones=milestones
1839                )
1840            else:
1841                self.member_vars["scheduler_instance"] = self.member_vars["scheduler"](
1842                    optimizer, **sched_args
1843                )
1844            current_steps = 0
1845
1846            for param_group in GPA.pai_tracker.member_vars[
1847                "optimizer_instance"
1848            ].param_groups:
1849                learning_rate1 = param_group["lr"]
1850
1851            if GPA.pc.get_verbose():
1852                print(
1853                    f"Resetting scheduler with {GPA.pai_tracker.steps_after_switch()} "
1854                    f"steps and {GPA.pc.get_initial_history_after_switches()} initial ticks to skip"
1855                )
1856
1857            # Find setting of previously used learning rate before adding dendrites
1858            if (
1859                GPA.pai_tracker.member_vars[
1860                    "current_n_learning_rate_initial_skip_steps"
1861                ]
1862                != 0
1863            ):
1864                additional_steps, learning_rate1 = self.increment_scheduler(
1865                    GPA.pai_tracker.member_vars[
1866                        "current_n_learning_rate_initial_skip_steps"
1867                    ],
1868                    "step_learning_rate",
1869                )
1870                current_steps += additional_steps
1871
1872            if self.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
1873                initial = GPA.pc.get_initial_history_after_switches()
1874            else:
1875                initial = 0
1876
1877            if GPA.pai_tracker.steps_after_switch() > initial:
1878                # Minus extra 1 because this gets called after start epoch
1879                additional_steps, learning_rate1 = self.increment_scheduler(
1880                    (GPA.pai_tracker.steps_after_switch() - initial) - 1,
1881                    "increment_epoch_count",
1882                )
1883                current_steps += additional_steps
1884
1885            if GPA.pc.get_verbose():
1886                print(
1887                    f"Scheduler update loop with {current_steps} "
1888                    f"ended with {learning_rate1}"
1889                )
1890                print(
1891                    f"Scheduler ended with {current_steps} steps "
1892                    f"and lr of {learning_rate1}"
1893                )
1894
1895            self.member_vars["current_step_count"] = current_steps
1896            return optimizer, self.member_vars["scheduler_instance"]
1897        else:
1898            return optimizer, None
1899
1900    def clear_optimizer_and_scheduler(self):
1901        """Clear the instances for saving.
1902
1903        Parameters
1904        ----------
1905        None
1906
1907        Returns
1908        -------
1909        None
1910            This function does not return a value.
1911        """
1912        self.member_vars["optimizer_instance"] = None
1913        self.member_vars["scheduler_instance"] = None
1914
1915    def switch_time(self):
1916        """Determine if it's time to switch between neuron and dendrite training.
1917
1918        Parameters
1919        ----------
1920        None
1921
1922        Returns
1923        -------
1924        bool
1925            True if it's time to switch, False otherwise.
1926
1927        Notes
1928        -----
1929        Based on current settings and history of scores.
1930        """
1931
1932        switch_phrase = "No mode, this should never be the case."
1933        switch_number = GPA.pc.get_n_epochs_to_switch()
1934        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1935            switch_phrase = "DOING_SWITCH_EVERY_TIME"
1936        elif self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY:
1937            switch_phrase = "DOING_HISTORY"
1938        elif self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH:
1939            switch_phrase = "DOING_FIXED_SWITCH"
1940            switch_number = GPA.pc.get_fixed_switch_num()
1941        elif self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1942            switch_phrase = "DOING_NO_SWITCH"
1943        else:
1944            print(
1945                "A switch mode must be set.  Check your settings for GPA.pc.set_switch_mode()."
1946            )
1947            pdb.set_trace()
1948        if not GPA.pc.get_silent():
1949            if(GPA.pc.get_perforated_backpropagation()):
1950                print(
1951                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1952                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1953                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1954                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1955                    f'n: {switch_number}, p: {GPA.pc.get_p_epochs_to_switch()}, '
1956                    f'num_cycles: {self.member_vars["num_cycles"]}'
1957                )
1958            else:
1959                print(
1960                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1961                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1962                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1963                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1964                    f'n: {switch_number}, num_cycles: {self.member_vars["num_cycles"]}'
1965                )
1966            print(
1967                f'  Score tracking: current_n_set_global_best={self.member_vars["current_n_set_global_best"]}, '
1968                f'global_best={self.member_vars["global_best_validation_score"]:.4f}, '
1969                f'current_best={self.member_vars["current_best_validation_score"]:.4f}'
1970            )
1971        if GPA.pc.get_perforated_backpropagation():
1972            # this will fill in epoch last improved
1973            TPB.best_pai_score_improved_this_epoch(self)  ## CLOSED ONLY
1974        if self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1975            if not GPA.pc.get_silent():
1976                print("Returning False - doing no switch mode")
1977            return False
1978
1979        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1980            if not GPA.pc.get_silent():
1981                print("Returning True - switching every time")
1982            return True
1983
1984        # Check if we're in the middle of learning rate optimization
1985        # If so, block ALL switch triggers until committed
1986        if GPA.pc.get_verbose():
1987            print("=== LR Optimization Check ===")
1988            print(f'  mode == "n": {self.member_vars["mode"] == "n"}')
1989            print(f"  get_learn_dendrites_live(): {GPA.pc.get_learn_dendrites_live()}")
1990            print(f'  committed_to_initial_rate: {GPA.pai_tracker.member_vars["committed_to_initial_rate"]}')
1991            print(f"  get_dont_give_up_unless_learning_rate_lowered(): {GPA.pc.get_dont_give_up_unless_learning_rate_lowered()}")
1992            print(f'  current_n_learning_rate_initial_skip_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"]}')
1993            print(f'  last_max_learning_rate_steps: {self.member_vars["last_max_learning_rate_steps"]}')
1994            print(f'  skip_steps < max_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"] < self.member_vars["last_max_learning_rate_steps"]}')
1995            print(f'  scheduler is not None: {self.member_vars["scheduler"] is not None}')
1996            print("=============================")
1997        
1998        if (
1999            ((self.member_vars["mode"] == "n") or GPA.pc.get_learn_dendrites_live())
2000            and (GPA.pai_tracker.member_vars["committed_to_initial_rate"] is False)
2001            and (GPA.pc.get_dont_give_up_unless_learning_rate_lowered())
2002            and (
2003                self.member_vars["current_n_learning_rate_initial_skip_steps"]
2004                <= self.member_vars["last_max_learning_rate_steps"]
2005            )
2006            and self.member_vars["scheduler"] is not None
2007        ):
2008            if not GPA.pc.get_silent():
2009                print(
2010                    f"Returning False - learning rate optimization in progress. "
2011                    f"Not committed yet. Comparing "
2012                    f'initial {self.member_vars["current_n_learning_rate_initial_skip_steps"]} '
2013                    f'to last max {self.member_vars["last_max_learning_rate_steps"]}'
2014                )
2015            return False
2016
2017        if len(self.member_vars["switch_epochs"]) == 0:
2018            this_count = self.member_vars["num_epochs_run"]
2019        else:
2020            this_count = (
2021                self.member_vars["num_epochs_run"]
2022                - self.member_vars["switch_epochs"][-1]
2023            )
2024        cap_switch = False
2025        if GPA.pc.get_perforated_backpropagation():
2026            cap_switch = TPB.check_cap_switch(self, this_count)
2027
2028        if self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY and (
2029            (
2030                (self.member_vars["mode"] == "n")
2031                and (
2032                    self.member_vars["num_epochs_run"]
2033                    - self.member_vars["epoch_last_improved"]
2034                    >= GPA.pc.get_n_epochs_to_switch()
2035                )
2036                and this_count
2037                >= GPA.pc.get_initial_history_after_switches()
2038                + GPA.pc.get_n_epochs_to_switch()
2039            )
2040            or (GPA.pc.get_perforated_backpropagation() and TPB.history_switch(self))
2041            or cap_switch
2042        ):
2043            if not GPA.pc.get_silent():
2044                print("Returning True - History and last improved is hit")
2045            return True
2046
2047        if self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH and (
2048            (
2049                self.member_vars["total_epochs_run"] % GPA.pc.get_fixed_switch_num()
2050                == GPA.pc.get_fixed_switch_num() - 1
2051            )
2052            and self.member_vars["num_epochs_run"]
2053            >= GPA.pc.get_first_fixed_switch_num() - 1
2054        ):
2055            if not GPA.pc.get_silent():
2056                print("Returning True - Fixed switch number is hit")
2057            return True
2058
2059        if not GPA.pc.get_silent():
2060            print("Returning False - no triggers to switch have been hit")
2061        return False
2062
2063    def steps_after_switch(self):
2064        """Based on settings, return value for steps since a switch.
2065
2066        Different options for param vals setting determine what is returned.
2067
2068        Parameters
2069        ----------
2070        None
2071
2072        Returns
2073        -------
2074        int
2075            The number of epochs since the last switch, or total epochs run,
2076            depending on settings.
2077
2078        """
2079        if self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_TOTAL_EPOCH:
2080            return self.member_vars["num_epochs_run"]
2081        elif (
2082            self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_UPDATE_EPOCH
2083        ):
2084            return self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2085        elif (
2086            self.member_vars["param_vals_setting"]
2087            == GPA.pc.PARAM_VALS_BY_NEURON_EPOCH_START
2088        ):
2089            if self.member_vars["mode"] == "p":
2090                return (
2091                    self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2092                )
2093            else:
2094                return self.member_vars["num_epochs_run"]
2095        else:
2096            print(
2097                f'{self.member_vars["param_vals_setting"]} is not a valid param vals option'
2098            )
2099            pdb.set_trace()
2100
2101    def add_pai_neuron_module(self, new_module, initial_add=True):
2102        """Add neuron modules to internal vectors.
2103
2104        Parameters
2105        ----------
2106        new_module : object
2107            The new module to add.
2108        initial_add : bool, optional
2109            Whether this is the initial addition rather than loading from file
2110
2111        Returns
2112        -------
2113        None
2114
2115        """
2116
2117        # If it's a duplicate, ignore the second addition
2118        if new_module in self.neuron_module_vector:
2119            return
2120        self.neuron_module_vector.append(new_module)
2121        if self.member_vars["doing_pai"]:
2122            PA.set_wrapped_params(new_module)
2123        if initial_add:
2124            self.member_vars["best_scores"].append([])
2125            self.member_vars["current_scores"].append([])
2126
2127    def add_tracked_neuron_module(self, new_module, initial_add=True):
2128        """Add tracked modules to internal vectors
2129
2130        Parameters
2131        ----------
2132        new_module : object
2133            The new module to add.
2134        initial_add : bool, optional
2135            Whether this is the initial addition rather than loading from file
2136
2137        Returns
2138        -------
2139        None
2140
2141        """
2142        # If it's a duplicate, ignore the second addition
2143        if new_module in self.tracked_neuron_module_vector:
2144            return
2145        self.tracked_neuron_module_vector.append(new_module)
2146        if self.member_vars["doing_pai"]:
2147            PA.set_tracked_params(new_module)
2148
2149    def reset_module_vector(self, net, load_from_restart):
2150        """Clear internal vectors and reset from network.
2151
2152        Parameters
2153        ----------
2154        net : object
2155            The neural network model.
2156        load_from_restart : bool
2157            Whether loading from a restart file.
2158
2159        Returns
2160        -------
2161        None
2162
2163        """
2164        self.neuron_module_vector = []
2165        self.tracked_neuron_module_vector = []
2166        this_list = UPA.get_pai_modules(net, 0)
2167        for module in this_list:
2168            self.add_pai_neuron_module(module, initial_add=load_from_restart)
2169        this_list = UPA.get_tracked_modules(net, 0)
2170        for module in this_list:
2171            self.add_tracked_neuron_module(module, initial_add=load_from_restart)
2172
2173    def reset_vals_for_score_reset(self):
2174        """Reset cycle scores for new cycle.
2175
2176        Parameters
2177        ----------
2178        None
2179
2180        Returns
2181        -------
2182        None
2183            This function does not return a value.
2184        """
2185
2186        if GPA.pc.get_find_best_lr():
2187            self.member_vars["committed_to_initial_rate"] = False
2188            print("Resetting committed to initial rate to False")
2189        # If retaining all dendrties always say that the current dendrites set global best for saving and loading
2190        if GPA.pc.get_retain_all_dendrites():
2191            self.member_vars["current_n_set_global_best"] = True
2192            self.member_vars["global_best_validation_score"] = 0
2193        else:
2194            self.member_vars["current_n_set_global_best"] = False
2195
2196        # Don't reset global best, but do reset current best
2197        self.member_vars["current_best_validation_score"] = 0
2198        self.member_vars["initial_lr_test_epoch_count"] = -1
2199
2200    def set_dendrite_training(self):
2201        """Signal all layers to start dendrite training.
2202
2203        Parameters
2204        ----------
2205        None
2206
2207        Returns
2208        -------
2209        None
2210            This function does not return a value.
2211        """
2212        if GPA.pc.get_verbose():
2213            print("Calling set_dendrite_training")
2214
2215        for layer in self.neuron_module_vector[:]:
2216            worked = layer.set_mode("p")
2217            """
2218            worked is False when a layer was added to the neuron module vector
2219            but then it's never actually been used. This can happen when
2220            you have set a layer to have requires_grad = False or when
2221            you have a module as a member variable but it's not actually
2222            part of the network. Should be moved to be a tracked layer
2223            rather than a neuron layer.
2224            """
2225            if not worked:
2226                self.neuron_module_vector.remove(layer)
2227
2228        for layer in self.tracked_neuron_module_vector[:]:
2229            worked = layer.set_mode("p")
2230
2231        self.create_new_dendrite_module()
2232        self.member_vars["mode"] = "p"
2233        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2234
2235        if GPA.pc.get_learn_dendrites_live():
2236            self.reset_vals_for_score_reset()
2237
2238        self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2239            "current_step_count"
2240        ]
2241
2242        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
2243        GPA.pai_tracker.member_vars["num_cycles"] += 1
2244
2245
2246    def set_neuron_training(self):
2247        """Signal all layers to start neuron training.
2248
2249        Parameters
2250        ----------
2251        None
2252
2253        Returns
2254        -------
2255        None
2256            This function does not return a value.
2257        """
2258        for module in self.neuron_module_vector:
2259            module.set_mode("n")
2260        for module in self.tracked_neuron_module_vector[:]:
2261            module.set_mode("n")
2262
2263        self.member_vars["mode"] = "n"
2264        self.member_vars["num_dendrites_added"] += 1
2265        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2266        self.reset_vals_for_score_reset()
2267
2268        self.member_vars["current_cycle_lr_max_scores"] = []
2269        if GPA.pc.get_learn_dendrites_live():
2270            self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2271                "current_step_count"
2272            ]
2273        GPA.pai_tracker.member_vars["num_cycles"] += 1
2274
2275        if GPA.pc.get_reset_best_score_on_switch():
2276            GPA.pai_tracker.member_vars["current_best_validation_score"] = 0
2277            GPA.pai_tracker.member_vars["running_accuracy"] = 0
2278
2279    def start_epoch(self, internal_call=False):
2280        """Perform steps for when a new training epoch is about to begin.
2281
2282        Parameters
2283        ----------
2284        internal_call : bool, optional
2285            Whether this is an internal call or manual call
2286
2287        Returns
2288        -------
2289        None
2290
2291        Notes
2292        -----
2293        If you ever need to call this manually, set internal_call to False.
2294
2295        """
2296        if self.member_vars["manual_train_switch"] and internal_call:
2297            return
2298
2299        if not internal_call and not self.member_vars["manual_train_switch"]:
2300            self.member_vars["manual_train_switch"] = True
2301            self.saved_time = 0
2302            self.member_vars["num_epochs_run"] = -1
2303            self.member_vars["total_epochs_run"] = -1
2304
2305        end = time.time()
2306        if self.member_vars["manual_train_switch"]:
2307            if self.saved_time != 0:
2308                if self.member_vars["mode"] == "p":
2309                    self.member_vars["p_val_times"].append(end - self.saved_time)
2310                else:
2311                    self.member_vars["n_val_times"].append(end - self.saved_time)
2312
2313        if self.member_vars["mode"] == "p":
2314            for layer in self.neuron_module_vector:
2315                for m in range(0, GPA.pc.get_global_candidates()):
2316                    with torch.no_grad():
2317                        if GPA.pc.get_verbose():
2318                            print(f"Resetting score for {layer.name}")
2319                        # Snapshot best_score before reset so we can compute per-epoch improvement
2320                        layer.dendrite_module.dendrite_values[
2321                            m
2322                        ].epoch_start_best_score.copy_(
2323                            layer.dendrite_module.dendrite_values[
2324                                m
2325                            ].best_score.detach()
2326                        )
2327                        layer.dendrite_module.dendrite_values[
2328                            m
2329                        ].best_score_improved_this_epoch = (
2330                            layer.dendrite_module.dendrite_values[
2331                                m
2332                            ].best_score_improved_this_epoch
2333                            * 0
2334                        )
2335                        layer.dendrite_module.dendrite_values[
2336                            m
2337                        ].nodes_best_improved_this_epoch = (
2338                            layer.dendrite_module.dendrite_values[
2339                                m
2340                            ].nodes_best_improved_this_epoch
2341                            * 0
2342                        )
2343                        layer.dendrite_module.dendrite_values[
2344                            m
2345                        ].nodes_improved_any = (
2346                            layer.dendrite_module.dendrite_values[
2347                                m
2348                            ].nodes_improved_any
2349                            * 0
2350                        )
2351            if GPA.pc.get_perforated_backpropagation():
2352                self.member_vars["best_mean_score_improved_this_epoch"] = 0
2353        self.member_vars["num_epochs_run"] += 1
2354        self.member_vars["total_epochs_run"] = (
2355            self.member_vars["num_epochs_run"] + self.member_vars["overwritten_epochs"]
2356        )
2357        self.saved_time = end
2358
2359    def stop_epoch(self, internal_call=False):
2360        """Perform steps when a training epoch has completed.
2361
2362        Parameters
2363        ----------
2364        internal_call : bool, optional
2365            Whether this is an internal call or manual call
2366
2367        Returns
2368        -------
2369        None
2370
2371        Notes
2372        -----
2373        If you ever need to call this manually, set internal_call to False.
2374
2375        """
2376        end = time.time()
2377        if self.member_vars["manual_train_switch"] and internal_call:
2378            return
2379
2380        if self.member_vars["manual_train_switch"]:
2381            if self.member_vars["mode"] == "p":
2382                self.member_vars["p_train_times"].append(end - self.saved_time)
2383            else:
2384                self.member_vars["n_train_times"].append(end - self.saved_time)
2385        else:
2386            if self.member_vars["mode"] == "p":
2387                self.member_vars["p_epoch_times"].append(end - self.saved_time)
2388            else:
2389                self.member_vars["n_epoch_times"].append(end - self.saved_time)
2390
2391        self.saved_time = end
2392
2393    def initialize(
2394        self,
2395        model,
2396        doing_pai=True,
2397        save_name="PAI",
2398        making_graphs=True,
2399        maximizing_score=True,
2400        num_classes=10000,
2401        values_per_train_epoch=-1,
2402        values_per_val_epoch=-1,
2403        zooming_graph=True,
2404    ):
2405        """Setup the tracker with initial settings.
2406
2407
2408        Parameters
2409        ----------
2410        model : object
2411            The neural network model.
2412        doing_pai : bool, optional
2413            Whether to add dendrites, by default True.
2414        save_name : str, optional
2415            The name under which to save the model.
2416        making_graphs : bool, optional
2417            Whether to make graphs, by default True.
2418        maximizing_score : bool, optional
2419            Whether to maximize the score, by default True.
2420        num_classes : int, optional
2421            The number of classes in the dataset, unused
2422        values_per_train_epoch : int, optional
2423            The number of values to look back for graphing
2424            during training, by default -1 (all values).
2425        values_per_val_epoch : int, optional
2426            The number of values to look back for graphing
2427            during validation, by default -1 (all values).
2428        zooming_graph : bool, optional
2429            Whether to zoom on graphs, by default True.
2430
2431
2432        Returns
2433        -------
2434        nn.Module
2435            Converted model instance configured for the tracker settings.
2436        """
2437        model = UPA.convert_network(model)
2438        self.member_vars["doing_pai"] = doing_pai
2439        self.member_vars["maximizing_score"] = maximizing_score
2440        self.save_name = save_name
2441        self.zooming_graph = zooming_graph
2442        self.making_graphs = making_graphs
2443
2444        if not self.loaded:
2445            self.member_vars["running_accuracy"] = (1.0 / num_classes) * 100
2446
2447        self.values_per_train_epoch = values_per_train_epoch
2448        self.values_per_val_epoch = values_per_val_epoch
2449
2450        if GPA.pc.get_testing_dendrite_capacity():
2451            if not GPA.pc.get_silent():
2452                print("Running a test of Dendrite Capacity.")
2453            GPA.pc.set_switch_mode(GPA.pc.DOING_SWITCH_EVERY_TIME)
2454            self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
2455            GPA.pc.set_retain_all_dendrites(True)
2456            GPA.pc.set_max_dendrite_tries(1000)
2457            GPA.pc.set_max_dendrites(1000)
2458            if GPA.pc.get_perforated_backpropagation():
2459                GPA.pc.set_initial_correlation_batches(1)
2460        else:
2461            if not GPA.pc.get_silent():
2462                print("Running Dendrite Experiment")
2463        return model
2464
2465    def generate_accuracy_plots(self, ax, save_folder, extra_string):
2466        """
2467        Generate plots and csvs for accuracy
2468
2469        Parameters
2470        ----------
2471        ax : object
2472            The matplotlib axis to plot on.
2473        save_folder : str
2474            The folder to save the plots and csvs in.
2475        extra_string : str
2476            An extra string to append to the filenames.
2477
2478        Returns
2479        -------
2480        None
2481
2482        """
2483
2484        # If scores are being saved for epochs that get overwritten, plot them
2485        for list_id in range(len(self.member_vars["overwritten_extras"])):
2486            for extra_id in self.member_vars["overwritten_extras"][list_id]:
2487                ax.plot(
2488                    np.arange(
2489                        len(self.member_vars["overwritten_extras"][list_id][extra_id])
2490                    ),
2491                    self.member_vars["overwritten_extras"][list_id][extra_id],
2492                    "r",
2493                )
2494            ax.plot(
2495                np.arange(len(self.member_vars["overwritten_vals"][list_id])),
2496                self.member_vars["overwritten_vals"][list_id],
2497                "b",
2498            )
2499
2500        # Determine which accuracy vector to use
2501        if GPA.pc.get_drawing_pai():
2502            accuracies = self.member_vars["accuracies"]
2503        else:
2504            accuracies = self.member_vars["n_accuracies"]
2505
2506        # Get pointer to additional scores being saved
2507        extra_scores = self.member_vars["extra_scores"]
2508
2509        # Plot the main accuracy scores
2510        ax.plot(np.arange(len(accuracies)), accuracies, label="Validation Scores")
2511        ax.plot(
2512            np.arange(len(self.member_vars["running_accuracies"])),
2513            self.member_vars["running_accuracies"],
2514            label="Validation Running Scores",
2515        )
2516
2517        # Plot additional scores
2518        for extra_score in extra_scores:
2519            ax.plot(
2520                np.arange(len(extra_scores[extra_score])),
2521                extra_scores[extra_score],
2522                label=extra_score,
2523            )
2524
2525        plt.title(save_folder + "/" + self.save_name + "Scores")
2526        plt.xlabel("Epochs")
2527        plt.ylabel("Score")
2528
2529        # Add point at epoch last improved and best validation score
2530        if GPA.pc.get_drawing_pai():
2531            ax.plot(
2532                self.member_vars["epoch_last_improved"],
2533                self.member_vars["global_best_validation_score"],
2534                "bo",
2535                label="Global best (y)",
2536            )
2537            ax.plot(
2538                self.member_vars["epoch_last_improved"],
2539                accuracies[self.member_vars["epoch_last_improved"]],
2540                "go",
2541                label="Epoch Last Improved",
2542            )
2543        else:
2544            if self.member_vars["mode"] == "n":
2545                missed_time = (
2546                    self.member_vars["num_epochs_run"]
2547                    - self.member_vars["epoch_last_improved"]
2548                )
2549                ax.plot(
2550                    (len(self.member_vars["n_accuracies"]) - 1) - missed_time,
2551                    self.member_vars["n_accuracies"][-(missed_time + 1)],
2552                    "go",
2553                    label="Epoch Last Improved",
2554                )
2555
2556        # Generate csv file for the values graphed
2557        pd1 = pd.DataFrame(
2558            {"Epochs": np.arange(len(accuracies)), "Validation Scores": accuracies}
2559        )
2560        pd2 = pd.DataFrame(
2561            {
2562                "Epochs": np.arange(len(self.member_vars["running_accuracies"])),
2563                "Validation Running Scores": self.member_vars["running_accuracies"],
2564            }
2565        )
2566        pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2567        for extra_score in extra_scores:
2568            pd2 = pd.DataFrame(
2569                {
2570                    "Epochs": np.arange(len(extra_scores[extra_score])),
2571                    extra_score: extra_scores[extra_score],
2572                }
2573            )
2574            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2575        extra_scores_without_graphing = self.member_vars[
2576            "extra_scores_without_graphing"
2577        ]
2578        for extra_score in extra_scores_without_graphing:
2579            pd2 = pd.DataFrame(
2580                {
2581                    "Epochs": np.arange(
2582                        len(extra_scores_without_graphing[extra_score])
2583                    ),
2584                    extra_score: extra_scores_without_graphing[extra_score],
2585                }
2586            )
2587            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2588        pd1.to_csv(
2589            save_folder + "/" + self.save_name + extra_string + "Scores.csv",
2590            index=False,
2591        )
2592        del pd1, pd2
2593
2594        # Set y min and max to zoom in on important part of axis
2595        if (
2596            len(self.member_vars["switch_epochs"]) > 0
2597            and self.member_vars["switch_epochs"][0] > 0
2598            and self.zooming_graph
2599        ):
2600            if GPA.pai_tracker.member_vars["maximizing_score"]:
2601                min_val = np.array(
2602                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2603                ).mean()
2604                for extra_score in extra_scores:
2605                    min_pot = np.array(
2606                        extra_scores[extra_score][
2607                            0 : self.member_vars["switch_epochs"][0]
2608                        ]
2609                    ).mean()
2610                    if min_pot < min_val:
2611                        min_val = min_pot
2612                ax.set_ylim(ymin=min_val)
2613            else:
2614                max_val = np.array(
2615                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2616                ).mean()
2617                for extra_score in extra_scores:
2618                    max_pot = np.array(
2619                        extra_scores[extra_score][
2620                            0 : self.member_vars["switch_epochs"][0]
2621                        ]
2622                    ).mean()
2623                    if max_pot > max_val:
2624                        max_val = max_pot
2625                ax.set_ylim(ymax=max_val)
2626
2627        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2628
2629        # Draw vertical lines for epochs where a dendrite switch occurred
2630        if GPA.pc.get_drawing_pai() and self.member_vars["doing_pai"]:
2631            color = "r"
2632            for switcher in self.member_vars["switch_epochs"]:
2633                plt.axvline(x=switcher, ymin=0, ymax=1, color=color)
2634                if color == "r":
2635                    color = "b"
2636                else:
2637                    color = "r"
2638        else:
2639            for switcher in self.member_vars["n_switch_epochs"]:
2640                plt.axvline(x=switcher, ymin=0, ymax=1, color="b")
2641
2642    def generate_time_plots(self, ax, save_folder, extra_string):
2643        """
2644        Generate plots and csvs for timing
2645
2646        Parameters
2647        ----------
2648        ax : object
2649            The matplotlib axis to plot on.
2650        save_folder : str
2651            The folder to save the plots and csvs in.
2652        extra_string : str
2653            An extra string to append to the filenames.
2654
2655        Returns
2656        -------
2657        None
2658
2659        """
2660        if self.member_vars["manual_train_switch"]:
2661            ax.plot(
2662                np.arange(len(self.member_vars["n_train_times"])),
2663                self.member_vars["n_train_times"],
2664                label="Normal Epoch Train Times",
2665            )
2666            ax.plot(
2667                np.arange(len(self.member_vars["p_train_times"])),
2668                self.member_vars["p_train_times"],
2669                label="PAI Epoch Train Times",
2670            )
2671            ax.plot(
2672                np.arange(len(self.member_vars["n_val_times"])),
2673                self.member_vars["n_val_times"],
2674                label="Normal Epoch Val Times",
2675            )
2676            ax.plot(
2677                np.arange(len(self.member_vars["p_val_times"])),
2678                self.member_vars["p_val_times"],
2679                label="PAI Epoch Val Times",
2680            )
2681
2682            plt.title(
2683                save_folder + "/" + self.save_name + "times (by train() and eval())"
2684            )
2685            plt.xlabel("Iteration")
2686            plt.ylabel("Epoch Time in Seconds ")
2687            ax.set_ylim(ymin=0)
2688            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2689
2690            pd1 = pd.DataFrame(
2691                {
2692                    "Epochs": np.arange(len(self.member_vars["n_train_times"])),
2693                    "Normal Epoch Train Times": self.member_vars["n_train_times"],
2694                }
2695            )
2696            pd2 = pd.DataFrame(
2697                {
2698                    "Epochs": np.arange(len(self.member_vars["p_train_times"])),
2699                    "PAI Epoch Train Times": self.member_vars["p_train_times"],
2700                }
2701            )
2702            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2703
2704            pd2 = pd.DataFrame(
2705                {
2706                    "Epochs": np.arange(len(self.member_vars["n_val_times"])),
2707                    "Normal Epoch Val Times": self.member_vars["n_val_times"],
2708                }
2709            )
2710            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2711
2712            pd2 = pd.DataFrame(
2713                {
2714                    "Epochs": np.arange(len(self.member_vars["p_val_times"])),
2715                    "PAI Epoch Val Times": self.member_vars["p_val_times"],
2716                }
2717            )
2718            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2719
2720            pd1.to_csv(
2721                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2722                index=False,
2723            )
2724            del pd1, pd2
2725        else:
2726            ax.plot(
2727                np.arange(len(self.member_vars["n_epoch_times"])),
2728                self.member_vars["n_epoch_times"],
2729                label="Normal Epoch Times",
2730            )
2731            ax.plot(
2732                np.arange(len(self.member_vars["p_epoch_times"])),
2733                self.member_vars["p_epoch_times"],
2734                label="PAI Epoch Times",
2735            )
2736
2737            plt.title(
2738                save_folder + "/" + self.save_name + "times (by train() and eval())"
2739            )
2740            plt.xlabel("Iteration")
2741            plt.ylabel("Epoch Time in Seconds ")
2742            ax.set_ylim(ymin=0)
2743            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2744
2745            pd1 = pd.DataFrame(
2746                {
2747                    "Epochs": np.arange(len(self.member_vars["n_epoch_times"])),
2748                    "Normal Epoch Times": self.member_vars["n_epoch_times"],
2749                }
2750            )
2751            pd2 = pd.DataFrame(
2752                {
2753                    "Epochs": np.arange(len(self.member_vars["p_epoch_times"])),
2754                    "PAI Epoch Times": self.member_vars["p_epoch_times"],
2755                }
2756            )
2757            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2758
2759            pd1.to_csv(
2760                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2761                index=False,
2762            )
2763            del pd1, pd2
2764
2765        if self.values_per_train_epoch != -1 and self.values_per_val_epoch != -1:
2766            ax2 = ax.twinx()  # Second axes sharing same x-axis
2767            ax2.set_ylabel("Single Datapoint Time in Seconds")
2768
2769            ax2.plot(
2770                np.arange(len(self.member_vars["n_train_times"])),
2771                np.array(self.member_vars["n_train_times"])
2772                / self.values_per_train_epoch,
2773                linestyle="dashed",
2774                label="Normal Train Item Times",
2775            )
2776            ax2.plot(
2777                np.arange(len(self.member_vars["p_train_times"])),
2778                np.array(self.member_vars["p_train_times"])
2779                / self.values_per_train_epoch,
2780                linestyle="dashed",
2781                label="PAI Train Item Times",
2782            )
2783            ax2.plot(
2784                np.arange(len(self.member_vars["n_val_times"])),
2785                np.array(self.member_vars["n_val_times"]) / self.values_per_val_epoch,
2786                linestyle="dashed",
2787                label="Normal Val Item Times",
2788            )
2789            ax2.plot(
2790                np.arange(len(self.member_vars["p_val_times"])),
2791                np.array(self.member_vars["p_val_times"]) / self.values_per_val_epoch,
2792                linestyle="dashed",
2793                label="PAI Val Item Times",
2794            )
2795            ax2.tick_params(axis="y")
2796            ax2.set_ylim(ymin=0)
2797            ax2.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2798
2799    def generate_learning_rate_plots(self, ax, save_folder, extra_string):
2800        """
2801        Generate plots and csvs for learning rate
2802
2803        Parameters
2804        ----------
2805        ax : object
2806            The matplotlib axis to plot on.
2807        save_folder : str
2808            The folder to save the plots and csvs in.
2809        extra_string : str
2810            An extra string to append to the filenames.
2811
2812        Returns
2813        -------
2814        None
2815
2816        """
2817        ax.plot(
2818            np.arange(len(self.member_vars["training_learning_rates"])),
2819            self.member_vars["training_learning_rates"],
2820            label="learning_rate",
2821        )
2822        plt.title(save_folder + "/" + self.save_name + "learning_rate")
2823        plt.xlabel("Epochs")
2824        plt.ylabel("learning_rate")
2825        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2826
2827        pd1 = pd.DataFrame(
2828            {
2829                "Epochs": np.arange(len(self.member_vars["training_learning_rates"])),
2830                "learning_rate": self.member_vars["training_learning_rates"],
2831            }
2832        )
2833        pd1.to_csv(
2834            save_folder + "/" + self.save_name + extra_string + "learning_rate.csv",
2835            index=False,
2836        )
2837        del pd1
2838
2839    def get_current_pb_scores(self):
2840        """
2841        Get the latest best PBScore of each dendrite layer, the same numbers
2842        written to the Best PBScores csv.
2843
2844        Returns
2845        -------
2846        dict[str, Any]
2847            Layer name to score.  Empty outside of dendrite scoring phases,
2848            when no candidate dendrites are being scored.
2849
2850
2851        Parameters
2852        ----------
2853        None
2854
2855        """
2856        if not self.member_vars["doing_pai"]:
2857            return {}
2858        if not GPA.pc.get_perforated_backpropagation():
2859            return {}
2860        # Scores only advance while candidate dendrites are being trained
2861        if (
2862            self.member_vars["mode"] != "p"
2863            and not GPA.pc.get_learn_dendrites_live()
2864        ):
2865            return {}
2866
2867        scores = {}
2868        for layer_id in range(len(self.neuron_module_vector)):
2869            if layer_id >= len(self.member_vars["best_scores"]):
2870                continue
2871            layer_scores = self.member_vars["best_scores"][layer_id]
2872            if len(layer_scores) == 0:
2873                continue
2874            score = layer_scores[-1]
2875            if hasattr(score, "item"):
2876                score = score.item()
2877            score = float(score)
2878            if math.isnan(score) or math.isinf(score):
2879                continue
2880            scores[self.neuron_module_vector[layer_id].name] = score
2881        return scores
2882
2883    def generate_dendrite_learning_plots(self, ax, save_folder, extra_string):
2884        """
2885        Generate dendrite score plots for the tracker.
2886        Also saves csv files associated with the plots.
2887
2888        Parameters
2889        ----------
2890        ax : matplotlib.axes.Axes
2891            Axis used for plotting dendrite-learning curves.
2892        save_folder : str
2893            Directory where plot images and CSV summaries are written.
2894        extra_string : str
2895            Filename suffix used to distinguish this output set.
2896
2897        Returns
2898        -------
2899        None
2900            Saves plots and score CSV files to disk.
2901        """
2902        if self.member_vars["doing_pai"]:
2903            pd1 = None
2904            pd2 = None
2905            num_colors = len(self.neuron_module_vector)
2906
2907            cm = plt.get_cmap("gist_rainbow")
2908            layer_colors = [cm(1.0 * i / max(num_colors, 1)) for i in range(num_colors)]
2909
2910            for layer_id in range(len(self.neuron_module_vector)):
2911                color = layer_colors[layer_id]
2912                ax.plot(
2913                    np.arange(len(self.member_vars["best_scores"][layer_id])),
2914                    self.member_vars["best_scores"][layer_id],
2915                    label=self.neuron_module_vector[layer_id].name,
2916                    color=color,
2917                )
2918
2919                pd2 = pd.DataFrame(
2920                    {
2921                        "Epochs": np.arange(
2922                            len(self.member_vars["best_scores"][layer_id])
2923                        ),
2924                        f"Best ever for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2925                            "best_scores"
2926                        ][
2927                            layer_id
2928                        ],
2929                    }
2930                )
2931
2932                if pd1 is None:
2933                    pd1 = pd2
2934                else:
2935                    pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2936
2937                if len(self.member_vars["current_scores"][layer_id]) != 0:
2938                    ax.plot(
2939                        np.arange(len(self.member_vars["current_scores"][layer_id])),
2940                        self.member_vars["current_scores"][layer_id],
2941                        label=f"Current:{self.neuron_module_vector[layer_id].name}",
2942                        color=color,
2943                        linestyle="--",
2944                    )
2945
2946                pd2 = pd.DataFrame(
2947                    {
2948                        "Epochs": np.arange(
2949                            len(self.member_vars["current_scores"][layer_id])
2950                        ),
2951                        f"Best current for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2952                            "current_scores"
2953                        ][
2954                            layer_id
2955                        ],
2956                    }
2957                )
2958                pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2959
2960            plt.title(save_folder + "/" + self.save_name + " Best PBScores")
2961            plt.xlabel("Epochs")
2962            plt.ylabel("Best PBScore")
2963            ax.legend(
2964                bbox_to_anchor=(1.05, 1),
2965                loc="upper left",
2966                ncol=max(1, math.ceil(len(self.neuron_module_vector) / 30)),
2967            )
2968            for switcher in self.member_vars["p_switch_epochs"]:
2969                plt.axvline(x=switcher, ymin=0, ymax=1, color="r")
2970
2971            if self.member_vars["mode"] == "p":
2972                missed_time = (
2973                    self.member_vars["num_epochs_run"]
2974                    - self.member_vars["epoch_last_improved"]
2975                )
2976                plt.axvline(
2977                    x=(len(self.member_vars["best_scores"][0]) - (missed_time + 1)),
2978                    ymin=0,
2979                    ymax=1,
2980                    color="g",
2981                )
2982
2983            # pd1 here will be none if no PB layers are created
2984            if pd1 is not None:
2985                pd1.to_csv(
2986                    save_folder
2987                    + "/"
2988                    + self.save_name
2989                    + extra_string
2990                    + "Best PBScores.csv",
2991                    index=False,
2992                )
2993            del pd1, pd2
2994
2995    def generate_extra_csv_files(self, save_folder, extra_string):
2996        """
2997        Generate additional csvs
2998
2999        Parameters
3000        ----------
3001        save_folder : str
3002            The folder to save the plots and csvs in.
3003        extra_string : str
3004            An extra string to append to the filenames.
3005
3006        Returns
3007        -------
3008        None
3009
3010        """
3011        pd1 = pd.DataFrame(
3012            {
3013                "Switch Number": np.arange(len(self.member_vars["switch_epochs"])),
3014                "Switch Epoch": self.member_vars["switch_epochs"],
3015            }
3016        )
3017        pd1.to_csv(
3018            save_folder + "/" + self.save_name + extra_string + "switch_epochs.csv",
3019            index=False,
3020        )
3021        del pd1
3022
3023        pd1 = pd.DataFrame(
3024            {
3025                "Switch Number": np.arange(len(self.member_vars["param_counts"])),
3026                "Param Count": self.member_vars["param_counts"],
3027            }
3028        )
3029        pd1.to_csv(
3030            save_folder + "/" + self.save_name + extra_string + "param_counts.csv",
3031            index=False,
3032        )
3033        del pd1
3034
3035        """
3036        Create best_arch_scores.csv file
3037        When working with dendrites there is a tradeoff between additional param count and score improvement.
3038        This file will help track that tradeoff by recording the best scores for all extra_scores
3039        and extra_scores_without_graphing for each architecture version.
3040        The scores recorded here are from the epoch when the best validation score was found
3041        within each switch_epoch boundary.
3042        """
3043        switch_counts = len(self.member_vars["switch_epochs"])
3044        best_valid = []
3045        associated_params = []
3046        
3047        # Initialize dictionaries to store best scores for each extra score type
3048        best_extra_scores = {}
3049        for score_name in self.member_vars["extra_scores"]:
3050            best_extra_scores[score_name] = []
3051        for score_name in self.member_vars["extra_scores_without_graphing"]:
3052            best_extra_scores[score_name] = []
3053
3054        for switch in range(0, switch_counts, 2):
3055            start_index = 0
3056            if switch != 0:
3057                start_index = self.member_vars["switch_epochs"][switch - 1] + 1
3058            end_index = self.member_vars["switch_epochs"][switch] + 1
3059
3060            if GPA.pai_tracker.member_vars["maximizing_score"]:
3061                best_valid_index = start_index + np.argmax(
3062                    self.member_vars["accuracies"][start_index:end_index]
3063                )
3064            else:
3065                best_valid_index = start_index + np.argmin(
3066                    self.member_vars["accuracies"][start_index:end_index]
3067                )
3068
3069            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3070            best_valid.append(best_valid_score)
3071            
3072            # Get corresponding scores from all extra_scores
3073            for score_name in self.member_vars["extra_scores"]:
3074                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3075                    best_extra_scores[score_name].append(
3076                        self.member_vars["extra_scores"][score_name][best_valid_index]
3077                    )
3078                else:
3079                    best_extra_scores[score_name].append(None)
3080            
3081            # Get corresponding scores from all extra_scores_without_graphing
3082            for score_name in self.member_vars["extra_scores_without_graphing"]:
3083                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3084                    best_extra_scores[score_name].append(
3085                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3086                    )
3087                else:
3088                    best_extra_scores[score_name].append(None)
3089            
3090            if self.member_vars["doing_pai"]:
3091                associated_params.append(self.member_vars["param_counts"][switch])
3092            else:
3093                associated_params.append(self.member_vars["param_counts"][-1])
3094
3095        # If in neuron training mode but not the very first epoch
3096        if self.member_vars["mode"] == "n" and (
3097            (len(self.member_vars["switch_epochs"]) == 0)
3098            or (
3099                self.member_vars["switch_epochs"][-1] + 1
3100                != len(self.member_vars["accuracies"])
3101            )
3102        ):
3103            start_index = 0
3104            if len(self.member_vars["switch_epochs"]) != 0:
3105                start_index = self.member_vars["switch_epochs"][-1] + 1
3106
3107            if GPA.pai_tracker.member_vars["maximizing_score"]:
3108                best_valid_index = start_index + np.argmax(
3109                    self.member_vars["accuracies"][start_index:]
3110                )
3111            else:
3112                best_valid_index = start_index + np.argmin(
3113                    self.member_vars["accuracies"][start_index:]
3114                )
3115
3116            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3117            best_valid.append(best_valid_score)
3118            
3119            # Get corresponding scores from all extra_scores
3120            for score_name in self.member_vars["extra_scores"]:
3121                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3122                    best_extra_scores[score_name].append(
3123                        self.member_vars["extra_scores"][score_name][best_valid_index]
3124                    )
3125                else:
3126                    best_extra_scores[score_name].append(None)
3127            
3128            # Get corresponding scores from all extra_scores_without_graphing
3129            for score_name in self.member_vars["extra_scores_without_graphing"]:
3130                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3131                    best_extra_scores[score_name].append(
3132                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3133                    )
3134                else:
3135                    best_extra_scores[score_name].append(None)
3136            
3137            associated_params.append(self.member_vars["param_counts"][-1])
3138
3139        # Build dataframe with all columns
3140        csv_data = {
3141            "Param Counts": associated_params,
3142            "Max Valid Scores": best_valid,
3143        }
3144        
3145        # Add columns for each extra score
3146        for score_name in best_extra_scores:
3147            csv_data[score_name] = best_extra_scores[score_name]
3148        
3149        pd1 = pd.DataFrame(csv_data)
3150        pd1.to_csv(
3151            save_folder + "/" + self.save_name + extra_string + "_best_arch_scores.csv",
3152            index=False,
3153        )
3154        del pd1
3155
3156    def save_graphs(self, extra_string=""):
3157        """
3158        Save graphs and csvs for all the values the tracker records
3159
3160        Parameters
3161        ----------
3162        extra_string : str
3163            An extra string to append to the filenames.
3164
3165        Returns
3166        -------
3167        None
3168
3169        """
3170        # If running DDP only save with rank 0
3171        if "RANK" in os.environ:
3172            if int(os.environ["RANK"]) != 0:
3173                return
3174        if not self.making_graphs:
3175            return
3176
3177        save_folder = "./" + self.save_name + "/"
3178
3179        plt.ioff()
3180        fig = plt.figure(figsize=(28, 14))
3181
3182        # Plot with accuracy scores
3183        ax = plt.subplot(221)
3184        self.generate_accuracy_plots(ax, save_folder, extra_string)
3185
3186        # Plot dendrite learning scores
3187        ax = plt.subplot(222)
3188        self.generate_dendrite_learning_plots(ax, save_folder, extra_string)
3189
3190        if GPA.pc.get_drawing_extra_graphs():
3191            # Plot learning rates for each training epoch
3192            ax = plt.subplot(223)
3193            self.generate_learning_rate_plots(ax, save_folder, extra_string)
3194
3195            # Plot the times for each training epoch
3196            ax = plt.subplot(224)
3197            self.generate_time_plots(ax, save_folder, extra_string)
3198
3199        # Generate extra CSV files
3200        self.generate_extra_csv_files(save_folder, extra_string)
3201
3202        fig.tight_layout()
3203        plt.savefig(save_folder + "/" + self.save_name + extra_string + ".png")
3204        plt.close("all")
3205
3206    def add_loss(self, loss):
3207        """Add loss to tracking vectors.
3208
3209        Parameters
3210        ----------
3211        loss : float or int
3212            The loss value to add.
3213
3214        Returns
3215        -------
3216        None
3217
3218        """
3219        if not isinstance(loss, (float, int)):
3220            loss = loss.item()
3221        self.member_vars["training_loss"].append(loss)
3222
3223    def add_learning_rate(self, learning_rate):
3224        """Add learning rate to tracking vectors.
3225
3226        Parameters
3227        ----------
3228        learning_rate : float or int
3229            The learning rate value to add.
3230
3231        Returns
3232        -------
3233        None
3234
3235        """
3236        if not isinstance(learning_rate, (float, int)):
3237            learning_rate = learning_rate.item()
3238        self.member_vars["training_learning_rates"].append(learning_rate)
3239
3240    def add_extra_score(self, score, extra_score_name):
3241        """Add extra score to tracking vectors.
3242
3243        Parameters
3244        ----------
3245        score : float or int
3246            The score value to add.
3247
3248        extra_score_name : str
3249            The name of the extra score.
3250
3251        Returns
3252        -------
3253        None
3254
3255        """
3256        if not isinstance(score, (float, int)):
3257            try:
3258                score = score.item()
3259            except:
3260                print(
3261                    "Scores added for Perforated Backpropagation should be "
3262                    "float, int, or tensor, yours is a:"
3263                )
3264                print(type(score))
3265                pdb.set_trace()
3266
3267        if GPA.pc.get_verbose():
3268            print(f"Adding extra score {extra_score_name} of {float(score)}")
3269
3270        if extra_score_name not in self.member_vars["extra_scores"]:
3271            self.member_vars["extra_scores"][extra_score_name] = []
3272        self.member_vars["extra_scores"][extra_score_name].append(score)
3273
3274        if self.member_vars["mode"] == "n":
3275            if extra_score_name not in self.member_vars["n_extra_scores"]:
3276                self.member_vars["n_extra_scores"][extra_score_name] = []
3277            self.member_vars["n_extra_scores"][extra_score_name].append(score)
3278
3279    def add_extra_score_without_graphing(self, score, extra_score_name):
3280        """Add extra score without graphing to tracking vectors.
3281
3282        Parameters
3283        ----------
3284        score : float or int
3285            The score value to add.
3286
3287        extra_score_name : str
3288            The name of the extra score.
3289
3290        Returns
3291        -------
3292        None
3293
3294        """
3295        if not isinstance(score, (float, int)):
3296            try:
3297                score = score.item()
3298            except:
3299                print(
3300                    "Scores added for Perforated Backpropagation should be "
3301                    "float, int, or tensor, yours is a:"
3302                )
3303                print(type(score))
3304                print("in add_extra_score_without_graphing")
3305                pdb.set_trace()
3306
3307        if GPA.pc.get_verbose():
3308            print(f"Adding extra score {extra_score_name} of {float(score)}")
3309
3310        if extra_score_name not in self.member_vars["extra_scores_without_graphing"]:
3311            self.member_vars["extra_scores_without_graphing"][extra_score_name] = []
3312        self.member_vars["extra_scores_without_graphing"][extra_score_name].append(
3313            score
3314        )
3315
3316    def add_test_score(self, score, extra_score_name):
3317        """Add test score to tracking vectors.
3318
3319        Parameters
3320        ----------
3321        score : float or int
3322            The score value to add.
3323
3324        extra_score_name : str
3325            The name of the extra score.
3326
3327        Returns
3328        -------
3329        None
3330
3331        Notes
3332        -----
3333        This function is a wrapper around `add_extra_score` that separates
3334        test score for adding to best_arch_scores.csv.
3335
3336        """
3337        self.add_extra_score(score, extra_score_name)
3338
3339        if not isinstance(score, (float, int)):
3340            try:
3341                score = score.item()
3342            except:
3343                print(
3344                    "Scores added for Perforated Backpropagation should be "
3345                    "float, int, or tensor, yours is a:"
3346                )
3347                print(type(score))
3348                print("in add_test_score")
3349                pdb.set_trace()
3350
3351        if GPA.pc.get_verbose():
3352            print(f"Adding test score {extra_score_name} of {float(score)}")
3353        self.member_vars["test_scores"].append(score)
3354
3355    def add_validation_score(self, accuracy, net, force_switch=False):
3356        """Function to add the validation score.
3357
3358        This is complex because it determines neuron and dendrite switching.
3359
3360        Parameters
3361        ----------
3362        accuracy : float or int
3363            The accuracy or loss value to add.
3364        net : object
3365            The neural network model.
3366        force_switch : bool, optional
3367            Whether to force a switch, by default False.
3368
3369        Returns
3370        -------
3371        net : object
3372            The potentially modified neural network model.
3373        training_complete : bool
3374            Whether training is complete.
3375        restructured : bool
3376            Whether the model has been restructured.
3377
3378        Notes
3379        -----
3380        WARNING: Do not call self anywhere in this function. When systems
3381        get loaded the actual tracker you are working with can change.
3382        """
3383
3384        _pai_log("info", f"Adding validation score {accuracy:.8f}")
3385
3386        update_learning_rate()
3387        update_param_count(net)
3388
3389        accuracy = check_input_problems(net, accuracy)
3390
3391        if len(GPA.pai_tracker.member_vars["switch_epochs"]) == 0:
3392            epochs_since_cycle_switch = GPA.pai_tracker.member_vars["num_epochs_run"]
3393        else:
3394            epochs_since_cycle_switch = (
3395                GPA.pai_tracker.member_vars["num_epochs_run"]
3396                - GPA.pai_tracker.member_vars["switch_epochs"][-1]
3397            )
3398
3399        update_running_accuracy(accuracy, epochs_since_cycle_switch)
3400        if GPA.pc.get_perforated_backpropagation():
3401            TPB.update_pb_scores(self)
3402
3403        # Captured before any switch below flips the mode and reloads scores
3404        epoch_pb_scores = self.get_current_pb_scores()
3405
3406        GPA.pai_tracker.stop_epoch(internal_call=True)
3407
3408        # If it is neuron training mode
3409        if (
3410            GPA.pai_tracker.member_vars["mode"] == "n"
3411            or GPA.pc.get_learn_dendrites_live()
3412        ):
3413            check_new_best(net, accuracy, epochs_since_cycle_switch)
3414        elif GPA.pc.get_perforated_backpropagation():
3415            TPB.check_best_pai_score_improvement()
3416
3417        # Save the latest model
3418        if GPA.pc.get_test_saves():
3419            UPA.save_system(net, GPA.pc.get_save_name(), "latest")
3420        if GPA.pc.get_pai_saves():
3421            UPA.pai_save_system(net, GPA.pc.get_save_name(), "latest")
3422
3423        restructuring_status_value = NO_MODEL_UPDATE
3424        # If it is time to switch based on scores and counter or a manual switch
3425        if GPA.pai_tracker.switch_time() or force_switch:
3426            # If testing dendrite capacity switch after enough dendrites added
3427            if (
3428                (GPA.pai_tracker.member_vars["mode"] == "n")
3429                and (GPA.pai_tracker.member_vars["num_dendrites_added"] > 2)
3430                and GPA.pc.get_testing_dendrite_capacity()
3431            ):
3432                GPA.pai_tracker.save_graphs()
3433                _pai_log(
3434                    "info",
3435                    "Successfully added 3 dendrites with GPA.pc.set_testing_dendrite_capacity(True) (default). "
3436                    "You may now set that to False and run a real experiment.",
3437                )
3438                return net, False, True
3439
3440            # If doing neuron training but this dendrite count didn't improve
3441            if (
3442                (GPA.pai_tracker.member_vars["mode"] == "n")
3443                or GPA.pc.get_learn_dendrites_live()
3444            ) and (GPA.pai_tracker.member_vars["current_n_set_global_best"] is False):
3445                new_restructuring_status_value, net = process_no_improvement(net)
3446                # if this was the final try return that training is complete
3447                if new_restructuring_status_value == TRAINING_COMPLETE:
3448                    if _dashboard_emitter is not None:
3449                        _dashboard_emitter.emit_run_end(GPA.pc)
3450                    return net, True, True
3451                else:
3452                    restructuring_status_value = update_restructuring_status(
3453                        restructuring_status_value, new_restructuring_status_value
3454                    )
3455            # Else if did improve, do a normal switch process
3456            else:
3457                if GPA.pc.get_verbose():
3458                    print(
3459                        f"Calling switch_mode with "
3460                        f'{GPA.pai_tracker.member_vars["current_n_set_global_best"]}, '
3461                        f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}, '
3462                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]}, '
3463                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_value"]},'
3464                        f'{GPA.pc.get_max_dendrites()},'
3465                        f'{GPA.pai_tracker.member_vars["num_dendrites_added"]},'
3466                        f'{GPA.pai_tracker.member_vars["num_dendrite_tries"]},'
3467                    )
3468                import pdb; pdb.set_trace
3469                # If the max number of dendrites has been hit or not doing pai and adding dendtites
3470                # then return rather than adding more
3471                if (
3472                    (GPA.pai_tracker.member_vars["mode"] == "n")
3473                    and (
3474                        GPA.pc.get_max_dendrites()
3475                        == GPA.pai_tracker.member_vars["num_dendrites_added"]
3476                    )
3477                ) or (GPA.pai_tracker.member_vars["doing_pai"] is False):
3478                    if GPA.pc.get_verbose():
3479                        print(
3480                            "Max dendrites reached or not doing PAI, finishing training"
3481                        )
3482                    net = process_final_network(net)
3483                    # Increment integrated if we have dendrites (means they're integrated)
3484                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3485                        GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3486                        _pai_log("info", f"Final dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3487                        if _dashboard_emitter is not None:
3488                            _dashboard_emitter.emit_dendrite_added(
3489                                GPA.pc,
3490                                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3491                                num_dendrites_integrated=GPA.pai_tracker.member_vars[
3492                                    "num_dendrites_integrated"
3493                                ],
3494                            )
3495                    if _dashboard_emitter is not None:
3496                        _dashboard_emitter.emit_run_end(GPA.pc)
3497                    return net, True, True
3498
3499                # Otherwise if its neuron training mode reset the counter of failed dendrites
3500                # Check if we should increment integrated count BEFORE change_learning_modes loads old state
3501                should_increment_integrated = False
3502                if GPA.pai_tracker.member_vars["mode"] == "n":
3503                    GPA.pai_tracker.member_vars["num_dendrite_tries"] = 0
3504                    if GPA.pc.get_verbose():
3505                        print(
3506                            "Adding new dendrites without resetting which means "
3507                            "the last ones improved. Resetting num_dendrite_tries"
3508                        )
3509                    # Remember to increment after change_learning_modes (which loads old tracker state)
3510                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3511                        should_increment_integrated = True
3512
3513                GPA.pai_tracker.save_graphs(
3514                    f'_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}'
3515                )
3516
3517                if GPA.pc.get_test_saves():
3518                    UPA.save_system(
3519                        net,
3520                        GPA.pc.get_save_name(),
3521                        f'beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3522                    )
3523                    # Copy current best model from this set of dendrites
3524                    # If running DDP only copy with rank 0
3525                    if "RANK" not in os.environ or int(os.environ["RANK"]) == 0:
3526                        shutil.copyfile(
3527                            f"{GPA.pc.get_save_name()}/best_model.pt",
3528                            f'{GPA.pc.get_save_name()}/best_model_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}.pt',
3529                        )
3530
3531                net = UPA.change_learning_modes(
3532                    net,
3533                    GPA.pc.get_save_name(),
3534                    "best_model",
3535                    GPA.pai_tracker.member_vars["doing_pai"],
3536                )
3537                restructuring_status_value = NETWORK_RESTRUCTURED
3538                
3539                # Now increment after change_learning_modes has loaded the best model
3540                # This ensures the increment persists and doesn't get overwritten
3541                if should_increment_integrated:
3542                    GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3543                    _pai_log("info", f"Dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3544                    if _dashboard_emitter is not None:
3545                        _dashboard_emitter.emit_dendrite_added(
3546                            GPA.pc,
3547                            epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3548                            num_dendrites_integrated=GPA.pai_tracker.member_vars[
3549                                "num_dendrites_integrated"
3550                            ],
3551                        )
3552
3553            # If restructured is true, clear scheduler/optimizer before saving
3554            if restructuring_status_value != NETWORK_RESTRUCTURED:
3555                print(
3556                    "Restructured should always be triggered here, let us know if you encounter this situation"
3557                )
3558                pdb.set_trace()
3559
3560            # Since there is a restructuring optimizer and scheduler must be reinitialized after return
3561            GPA.pai_tracker.clear_optimizer_and_scheduler()
3562
3563            # Save the model from after the switch
3564            UPA.save_system(
3565                net,
3566                GPA.pc.get_save_name(),
3567                f'switch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3568            )
3569
3570        # If not time to switch and you have a scheduler, perform the update step
3571        elif GPA.pai_tracker.member_vars["scheduler"] is not None:
3572            new_restructuring_status_value, net = process_scheduler_update(
3573                net, accuracy, epochs_since_cycle_switch
3574            )
3575            restructuring_status_value = update_restructuring_status(
3576                restructuring_status_value, new_restructuring_status_value
3577            )
3578
3579        GPA.pai_tracker.start_epoch(internal_call=True)
3580        if _dashboard_emitter is not None:
3581            _mv = GPA.pai_tracker.member_vars
3582            _lr = _mv["training_learning_rates"][-1] if _mv["training_learning_rates"] else None
3583            _train_score = _mv["extra_scores"].get("train", [None])[-1]
3584            _n_times = _mv["n_epoch_times"] or [(_mv["n_train_times"][-1] + _mv["n_val_times"][-1]) if (_mv["n_train_times"] and _mv["n_val_times"]) else None]
3585            _p_times = _mv["p_epoch_times"] or [(_mv["p_train_times"][-1] + _mv["p_val_times"][-1]) if (_mv["p_train_times"] and _mv["p_val_times"]) else None]
3586            _dashboard_emitter.emit_epoch(
3587                GPA.pc,
3588                epoch=_mv["num_epochs_run"],
3589                validation_score=accuracy,
3590                learning_rate=_lr,
3591                train_score=_train_score,
3592                normal_time=_n_times[-1],
3593                pai_time=_p_times[-1],
3594                pb_scores=epoch_pb_scores,
3595            )
3596        GPA.pai_tracker.save_graphs()
3597
3598        if restructuring_status_value == NETWORK_RESTRUCTURED:
3599            GPA.pai_tracker.member_vars["epoch_last_improved"] = (
3600                GPA.pai_tracker.member_vars["num_epochs_run"]
3601            )
3602            if GPA.pc.get_verbose():
3603                print(
3604                    f"Setting epoch last improved to "
3605                    f'{GPA.pai_tracker.member_vars["epoch_last_improved"]}'
3606                )
3607
3608            now = datetime.now()
3609            dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
3610
3611            if GPA.pc.get_verbose():
3612                print("Not saving restructure right now")
3613
3614            """
3615            This block of code helped with a save issue with safetensors and huggingface, but it breaks DDP.  
3616            Temporarily removing it to avoid DDP issues, but if you encounter save issues try adding it back in.
3617            for param in net.parameters():
3618                param.data = param.data.contiguous()
3619            """
3620        if GPA.pc.get_verbose():
3621            print(
3622                f"Completed adding score. Restructured is {restructuring_status_value}, "
3623                f"\ncurrent switch list is:"
3624            )
3625            print(GPA.pai_tracker.member_vars["switch_epochs"])
3626
3627        if _dashboard_emitter is not None and restructuring_status_value == NETWORK_RESTRUCTURED:
3628            _param_count = UPA.count_params(net)
3629            _dashboard_emitter.emit_switch(
3630                GPA.pc,
3631                switch_number=GPA.pai_tracker.member_vars["num_dendrites_added"],
3632                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3633                param_count=_param_count,
3634                switch_type=GPA.pai_tracker.member_vars["mode"],
3635            )
3636
3637        # Always False for training complete if nothing triggered that training is over
3638        return net, restructuring_status_value, False
3639
3640    def clear_all_processors(self):
3641        """Clear all processors from modules.
3642
3643        Parameters
3644        ----------
3645        None
3646
3647        Returns
3648        -------
3649        None
3650            This function does not return a value.
3651        """
3652        for module in self.neuron_module_vector:
3653            module.clear_processors()
3654
3655    def create_new_dendrite_module(self):
3656        """Add dendrite module to all neuron modules.
3657
3658        Parameters
3659        ----------
3660        None
3661
3662        Returns
3663        -------
3664        None
3665            This function does not return a value.
3666        """
3667        for module in self.neuron_module_vector:
3668            module.create_new_dendrite_module()
3669
3670    def set_create_dendrite_global(self, fn):
3671        """Call set_create_dendrite(fn) on every tracked PAINeuronModule."""
3672        for module in self.neuron_module_vector:
3673            module.set_create_dendrite(fn)
3674
3675    def set_dendrite_loss_fn_global(self, fn):
3676        """Set the global dendrite loss function used by all dendrite modules."""
3677        from perforatedbp import modules_pbp as MPB
3678        MPB.dendrite_loss_fn = fn
3679
3680    def apply_pb_grads(self):
3681        """Apply perforated backpropagation gradients to all modules.
3682
3683        Parameters
3684        ----------
3685        None
3686
3687        Returns
3688        -------
3689        None
3690            This function does not return a value.
3691        """
3692        if self.member_vars["mode"] == "p":
3693            for module in self.neuron_module_vector:
3694                module.apply_pb_grads()
3695
3696    def apply_pb_zero(self):
3697        """Apply perforated backpropagation zero gradients to all modules.
3698
3699        Parameters
3700        ----------
3701        None
3702
3703        Returns
3704        -------
3705        None
3706            This function does not return a value.
3707        """
3708        if self.member_vars["mode"] == "p":
3709            for module in self.neuron_module_vector:
3710                module.apply_pb_zero()
NO_MODEL_UPDATE = 0
NETWORK_RESTRUCTURED = 1
TRAINING_COMPLETE = 2
STEP_CLEARED = 0
STEP_CALLED = 1
def update_restructuring_status(old_status, new_status):
69def update_restructuring_status(old_status, new_status):
70    """Update restructured variable during add_validation_score
71
72    Update the restructuring status based on the new status.
73    If the new status is that there was not an update,
74    dont overwrite the old status which may show there was an update.
75
76    Parameters
77    ----------
78    old_status : int
79        The old restructuring status.
80    new_status : int
81        The new restructuring status.
82
83    Returns
84    -------
85    int
86        The updated restructuring status.
87
88    """
89    if new_status == NETWORK_RESTRUCTURED or new_status == TRAINING_COMPLETE:
90        return NETWORK_RESTRUCTURED
91    else:
92        return old_status

Update restructured variable during add_validation_score

Update the restructuring status based on the new status. If the new status is that there was not an update, dont overwrite the old status which may show there was an update.

Parameters
  • old_status (int): The old restructuring status.
  • new_status (int): The new restructuring status.
Returns
  • int: The updated restructuring status.
def update_learning_rate():
 95def update_learning_rate():
 96    """Update the learning rate in the tracker.
 97
 98    Parameters
 99    ----------
100    None
101
102    Returns
103    -------
104    None
105        This function does not return a value.
106    """
107    for param_group in GPA.pai_tracker.member_vars["optimizer_instance"].param_groups:
108        learning_rate = param_group["lr"]
109    GPA.pai_tracker.add_learning_rate(learning_rate)

Update the learning rate in the tracker.

Parameters
  • None
Returns
  • None: This function does not return a value.
def update_param_count(net):
112def update_param_count(net):
113    """Update the parameter count in the tracker if not already set.
114
115    Parameters
116    ----------
117    net : torch.nn.Module
118        The neural network model to count parameters for.
119    Returns
120    -------
121    None
122    """
123    if len(GPA.pai_tracker.member_vars["param_counts"]) == 0:
124        GPA.pai_tracker.member_vars["param_counts"].append(UPA.count_params(net))

Update the parameter count in the tracker if not already set.

Parameters
  • net (torch.nn.Module): The neural network model to count parameters for.
Returns
  • None
def check_input_problems(net, accuracy):
127def check_input_problems(net, accuracy):
128    """Check for potential input problems in add_validation_score.
129
130    Parameters
131    ----------
132    net : torch.nn.Module
133        The neural network model to check.
134    accuracy : float, int, or torch.Tensor
135        The accuracy score to validate.
136
137    Returns
138    -------
139    float
140        The validated accuracy score.
141
142    """
143
144    # Make sure you are passing in the model and not the dataparallel wrapper
145    if issubclass(type(net), nn.DataParallel):
146        _pai_log("error", "Need to call .module when using add validation score")
147        pdb.set_trace()
148        sys.exit(-1)
149
150    if "module" in net.__dir__():
151        _pai_log("error", "Need to call .module when using add validation score")
152        pdb.set_trace()
153        sys.exit(-1)
154
155    if not isinstance(accuracy, (float, int)):
156        try:
157            accuracy = accuracy.item()
158        except:
159            _pai_log(
160                "error",
161                f"Scores added for add_validation_score should be float, int, or tensor, yours is a: {type(accuracy)}",
162            )
163            pdb.set_trace()
164            sys.exit(-1)
165    return accuracy

Check for potential input problems in add_validation_score.

Parameters
  • net (torch.nn.Module): The neural network model to check.
  • accuracy (float, int, or torch.Tensor): The accuracy score to validate.
Returns
  • float: The validated accuracy score.
def update_running_accuracy(accuracy, epochs_since_cycle_switch):
168def update_running_accuracy(accuracy, epochs_since_cycle_switch):
169    """Add the new accuracy to the tracker.
170
171    Parameters
172    ----------
173    accuracy : float, int, or torch.Tensor
174        The accuracy score to add.
175    epochs_since_cycle_switch : int
176        The number of epochs since the last cycle switch.
177
178    Returns
179    -------
180    None
181
182    """
183    # Only update running_accuracy when neurons are being updated
184    if GPA.pai_tracker.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
185        if epochs_since_cycle_switch < GPA.pc.get_initial_history_after_switches():
186            if epochs_since_cycle_switch <= 0:
187                GPA.pai_tracker.member_vars["running_accuracy"] = accuracy
188            else:
189                GPA.pai_tracker.member_vars[
190                    "running_accuracy"
191                ] = GPA.pai_tracker.member_vars["running_accuracy"] * (
192                    1 - (1.0 / (epochs_since_cycle_switch + 1))
193                ) + accuracy * (
194                    1.0 / (epochs_since_cycle_switch + 1)
195                )
196        else:
197            GPA.pai_tracker.member_vars[
198                "running_accuracy"
199            ] = GPA.pai_tracker.member_vars["running_accuracy"] * (
200                1.0 - 1.0 / GPA.pc.get_history_lookback()
201            ) + accuracy * (
202                1.0 / GPA.pc.get_history_lookback()
203            )
204
205    GPA.pai_tracker.member_vars["accuracies"].append(accuracy)
206    if GPA.pai_tracker.member_vars["mode"] == "n":
207        GPA.pai_tracker.member_vars["n_accuracies"].append(accuracy)
208
209    if (
210        GPA.pc.get_drawing_pai()
211        or GPA.pai_tracker.member_vars["mode"] == "n"
212        or GPA.pc.get_learn_dendrites_live()
213    ):
214        GPA.pai_tracker.member_vars["running_accuracies"].append(
215            GPA.pai_tracker.member_vars["running_accuracy"]
216        )

Add the new accuracy to the tracker.

Parameters
  • accuracy (float, int, or torch.Tensor): The accuracy score to add.
  • epochs_since_cycle_switch (int): The number of epochs since the last cycle switch.
Returns
  • None
def score_beats_current_best(new_score, old_score):
219def score_beats_current_best(new_score, old_score):
220    """Check if the new score beats the current best score.
221
222    Parameters
223    ----------
224    new_score : float
225        The new score to compare.
226    old_score : float
227        The old score to compare against.
228
229    Returns
230    -------
231    bool
232        True if the new score beats the old score, False otherwise.
233
234    Notes
235    -----
236    Must beat the old score by the margins set in globals for improvement thresholds.
237
238    """
239    return (
240        GPA.pai_tracker.member_vars["maximizing_score"]
241        and (new_score * (1.0 - GPA.pc.get_improvement_threshold()) > old_score)
242        and new_score - GPA.pc.get_improvement_threshold_raw() > old_score
243    ) or (
244        (not GPA.pai_tracker.member_vars["maximizing_score"])
245        and (new_score * (1.0 + GPA.pc.get_improvement_threshold()) < old_score)
246        and (new_score + GPA.pc.get_improvement_threshold_raw()) < old_score
247    )

Check if the new score beats the current best score.

Parameters
  • new_score (float): The new score to compare.
  • old_score (float): The old score to compare against.
Returns
  • bool: True if the new score beats the old score, False otherwise.
Notes

Must beat the old score by the margins set in globals for improvement thresholds.

def check_new_best(net, accuracy, epochs_since_cycle_switch):
250def check_new_best(net, accuracy, epochs_since_cycle_switch):
251    """Check if the new accuracy is a new best.
252
253    Performs saves if new best score is found.
254
255    Parameters
256    ----------
257    net : torch.nn.Module
258        The neural network model being trained.
259    accuracy : float
260        The accuracy score to check.
261    epochs_since_cycle_switch : int
262        The number of epochs since the last cycle switch.
263
264    Returns
265    -------
266    None
267
268    """
269    score_improved = score_beats_current_best(
270        GPA.pai_tracker.member_vars["running_accuracy"],
271        GPA.pai_tracker.member_vars["current_best_validation_score"],
272    )
273
274    enough_time = (
275        epochs_since_cycle_switch > GPA.pc.get_initial_history_after_switches()
276    ) or (GPA.pai_tracker.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME)
277
278    if (
279        score_improved
280        or GPA.pai_tracker.member_vars["current_best_validation_score"] == 0
281    ) and enough_time:
282
283        if GPA.pai_tracker.member_vars["maximizing_score"]:
284            if GPA.pc.get_verbose():
285                print(
286                    f"\n\nGot score of {accuracy:.10f} "
287                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
288                    f"*{1-GPA.pc.get_improvement_threshold()}="
289                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 - GPA.pc.get_improvement_threshold())}) '
290                    f'which is higher than {GPA.pai_tracker.member_vars["current_best_validation_score"]:.10f} '
291                    f"by {GPA.pc.get_improvement_threshold_raw()} so setting epoch to "
292                    f'{GPA.pai_tracker.member_vars["num_epochs_run"]}\n\n'
293                )
294        else:
295            if GPA.pc.get_verbose():
296                print(
297                    f"\n\nGot score of {accuracy:.10f} "
298                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
299                    f"*{1+GPA.pc.get_improvement_threshold()}="
300                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 + GPA.pc.get_improvement_threshold())}) '
301                    f'which is lower than {GPA.pai_tracker.member_vars["current_best_validation_score"]:.10f} '
302                    f'so setting epoch to {GPA.pai_tracker.member_vars["num_epochs_run"]}\n\n'
303                )
304
305        # Set the new best score
306        GPA.pai_tracker.member_vars["current_best_validation_score"] = (
307            GPA.pai_tracker.member_vars["running_accuracy"]
308        )
309        GPA.pai_tracker.member_vars["epoch_last_improved"] = (
310            GPA.pai_tracker.member_vars["num_epochs_run"]
311        )
312        if GPA.pc.get_verbose():
313            print(
314                f'2 epoch improved is {GPA.pai_tracker.member_vars["epoch_last_improved"]}'
315            )
316        # Immediately update this list before saving so loading will have it correctly
317        GPA.pai_tracker.member_vars["last_improved_accuracies"].append(
318            GPA.pai_tracker.member_vars["epoch_last_improved"]
319        )
320        # Check if global best
321        is_global_best = score_beats_current_best(
322            GPA.pai_tracker.member_vars["current_best_validation_score"],
323            GPA.pai_tracker.member_vars["global_best_validation_score"],
324        )
325
326        if (
327            is_global_best
328            or GPA.pai_tracker.member_vars["global_best_validation_score"] == 0
329        ):
330            if GPA.pc.get_verbose():
331                print(
332                    f"This also beats global best of "
333                    f'{GPA.pai_tracker.member_vars["global_best_validation_score"]} so saving'
334                )
335            GPA.pai_tracker.member_vars["global_best_validation_score"] = (
336                GPA.pai_tracker.member_vars["current_best_validation_score"]
337            )
338            GPA.pai_tracker.member_vars["current_n_set_global_best"] = True
339            UPA.save_system(net, GPA.pc.get_save_name(), "best_model")
340            if GPA.pc.get_pai_saves():
341                UPA.pai_save_system(net, GPA.pc.get_save_name(), "best_model")
342    else:
343        if GPA.pc.get_verbose():
344            print("Not saving new best because:")
345            if epochs_since_cycle_switch <= GPA.pc.get_initial_history_after_switches():
346                print(
347                    f"Not enough history since switch {epochs_since_cycle_switch} <= "
348                    f"{GPA.pc.get_initial_history_after_switches()}"
349                )
350            elif GPA.pai_tracker.member_vars["maximizing_score"]:
351                print(
352                    f"Got score of {accuracy} "
353                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
354                    f"*{1-GPA.pc.get_improvement_threshold()}="
355                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 - GPA.pc.get_improvement_threshold())}) '
356                    f"which is not higher than "
357                    f'{GPA.pai_tracker.member_vars["current_best_validation_score"]}'
358                )
359            else:
360                print(
361                    f"Got score of {accuracy} "
362                    f'(average {GPA.pai_tracker.member_vars["running_accuracy"]}, '
363                    f"*{1+GPA.pc.get_improvement_threshold()}="
364                    f'{GPA.pai_tracker.member_vars["running_accuracy"]*(1.0 + GPA.pc.get_improvement_threshold())}) '
365                    f"which is not lower than "
366                    f'{GPA.pai_tracker.member_vars["current_best_validation_score"]}'
367                )
368        GPA.pai_tracker.member_vars["last_improved_accuracies"].append(
369            GPA.pai_tracker.member_vars["epoch_last_improved"]
370        )
371        # If it's the first epoch, save as best anyway
372        if len(GPA.pai_tracker.member_vars["accuracies"]) == 1:
373            if GPA.pc.get_verbose():
374                print("Saving first model or all models")
375            UPA.save_system(net, GPA.pc.get_save_name(), "best_model")
376            if GPA.pc.get_pai_saves():
377                UPA.pai_save_system(net, GPA.pc.get_save_name(), "best_model")

Check if the new accuracy is a new best.

Performs saves if new best score is found.

Parameters
  • net (torch.nn.Module): The neural network model being trained.
  • accuracy (float): The accuracy score to check.
  • epochs_since_cycle_switch (int): The number of epochs since the last cycle switch.
Returns
  • None
def process_no_improvement(net):
380def process_no_improvement(net):
381    """Handle the case where no improvement is observed.
382
383    If the new dendrite did not improve scores, but its time to switch modes
384    either trigger the end of learning or reset to the previous dendrite
385    to try again.
386
387    Parameters
388    ----------
389    net : torch.nn.Module
390        The neural network model being trained.
391
392    Returns
393    -------
394    int
395        The status of restructuring or training completion.
396    torch.nn.Module
397        The potentially modified neural network model.
398
399    """
400    if GPA.pc.get_verbose():
401        print(
402            f"Planning to switch to p mode but best beat last: "
403            f'{GPA.pai_tracker.member_vars["current_n_set_global_best"]} '
404            f"current start lr steps: "
405            f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
406            f"and last maximum lr steps: "
407            f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
408            f'for rate: {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]:.8f}'
409        )
410
411    now = datetime.now()
412    dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
413
414    if GPA.pc.get_verbose():
415        print(
416            f'1 saving break {dt_string}_noImprove_lr_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
417        )
418
419    GPA.pai_tracker.save_graphs(
420        f'{dt_string}_noImprove_lr_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
421    )
422
423    if (
424        GPA.pai_tracker.member_vars["num_dendrite_tries"]
425        < GPA.pc.get_max_dendrite_tries() -1
426    ):
427        _pai_log(
428            "info",
429            f"The newest added dendrites did not improve but current tries "
430            f'{GPA.pai_tracker.member_vars["num_dendrite_tries"] + 1} '
431            f"is less than max tries {GPA.pc.get_max_dendrite_tries()} "
432            f"so loading last switch and trying new Dendrites.",
433        )
434        old_tries = GPA.pai_tracker.member_vars["num_dendrite_tries"]
435        # Load best model from previous n mode
436        net = UPA.change_learning_modes(
437            net,
438            GPA.pc.get_save_name(),
439            "best_model",
440            GPA.pai_tracker.member_vars["doing_pai"],
441        )
442        GPA.pai_tracker.member_vars["num_dendrite_tries"] = old_tries + 1
443        return NETWORK_RESTRUCTURED, net
444    else:
445        _pai_log(
446            "info",
447            f"The newest added dendrites did not improve system and "
448            f'{GPA.pai_tracker.member_vars["num_dendrite_tries"] + 1} > '
449            f"{GPA.pc.get_max_dendrite_tries()} so returning training_complete.",
450        )
451        _pai_log("info", "You should now exit your training loop and best_model will be your final model for inference")
452        if not GPA.pc.get_perforated_backpropagation() and GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
453            _pai_log("info", "For improved results, try perforated backpropagation next time!")
454        old_silent = GPA.pc.get_silent()
455        GPA.pc.set_silent(True)
456        UPA.load_system(net, GPA.pc.get_save_name(), "best_model", switch_call=True)
457        GPA.pc.set_silent(old_silent)
458        GPA.pai_tracker.save_graphs()
459        UPA.pai_save_system(net, GPA.pc.get_save_name(), "final_clean")
460        return TRAINING_COMPLETE, net

Handle the case where no improvement is observed.

If the new dendrite did not improve scores, but its time to switch modes either trigger the end of learning or reset to the previous dendrite to try again.

Parameters
  • net (torch.nn.Module): The neural network model being trained.
Returns
  • int: The status of restructuring or training completion.
  • torch.nn.Module: The potentially modified neural network model.
def process_final_network(net):
463def process_final_network(net):
464    """When the max number of dendrites has been hit load the best_model and return
465
466    Parameters
467    ----------
468    net : torch.nn.Module
469        The neural network model being trained.
470
471    Returns
472    -------
473    torch.nn.Module
474        The final neural network model.
475    """
476
477    _pai_log("info", f"Last Dendrites were good and this hit the max of {GPA.pc.get_max_dendrites()}")
478    if not GPA.pc.get_perforated_backpropagation() and GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
479        _pai_log("info", "For improved results, try perforated backpropagation next time!")
480    GPA.pai_tracker.save_graphs("before_final")
481    UPA.load_system(net, GPA.pc.get_save_name(), "best_model", switch_call=True)
482    GPA.pai_tracker.save_graphs()
483    UPA.pai_save_system(net, GPA.pc.get_save_name(), "final_clean")
484    return net

When the max number of dendrites has been hit load the best_model and return

Parameters
  • net (torch.nn.Module): The neural network model being trained.
Returns
  • torch.nn.Module: The final neural network model.
def process_scheduler_update(net, accuracy, epochs_since_cycle_switch):
487def process_scheduler_update(net, accuracy, epochs_since_cycle_switch):
488    """Updates the scheduler
489
490    This increments the scheduler, but if we are automatically sweeping
491    to find the best initial learning rate for a new set of dendrites
492    this function also triggers the network at addition time to
493    try the next value.
494
495    Process for finding best initial learning rate for dendrites:
496    1. Start at default rate
497    2. Learn at that rate until scheduler increments twice
498    3. Save that version, start dendrites at LR current increment - 1
499    4. Repeat 2 and 3 until version has worse final score at set LR
500    5. Load previous model with best accuracy at that LR as initial rate
501
502    Parameters
503    ----------
504    net : torch.nn.Module
505        The neural network model being trained.
506    accuracy : float
507        The accuracy of the model at the current learning rate.
508    epochs_since_cycle_switch : int
509        The number of epochs since the last cycle switch.
510
511    Returns
512    -------
513    int
514        The status of restructuring or training completion.
515    torch.nn.Module
516        The potentially modified neural network model.
517    """
518
519    restructured = False
520    for param_group in GPA.pai_tracker.member_vars["optimizer_instance"].param_groups:
521        learning_rate1 = param_group["lr"]
522
523    if (
524        type(GPA.pai_tracker.member_vars["scheduler_instance"])
525        is torch.optim.lr_scheduler.ReduceLROnPlateau
526    ):
527        if (
528            epochs_since_cycle_switch > GPA.pc.get_initial_history_after_switches()
529            or GPA.pai_tracker.member_vars["mode"] == "p"
530        ):
531            if GPA.pc.get_verbose():
532                print(
533                    f"Updating scheduler with last improved "
534                    f'{GPA.pai_tracker.member_vars["epoch_last_improved"]} '
535                    f'from current {GPA.pai_tracker.member_vars["num_epochs_run"]}'
536                )
537            if GPA.pai_tracker.member_vars["scheduler"] is not None:
538                GPA.pai_tracker.member_vars["scheduler_instance"].step(metrics=accuracy)
539                if (
540                    GPA.pai_tracker.member_vars["scheduler"]
541                    is torch.optim.lr_scheduler.ReduceLROnPlateau
542                ):
543                    if GPA.pc.get_verbose():
544                        print(
545                            f"Scheduler is now at "
546                            f'{GPA.pai_tracker.member_vars["scheduler_instance"].num_bad_epochs} bad epochs'
547                        )
548        else:
549            if GPA.pc.get_verbose():
550                print("Not stepping optimizer since hasnt initialized")
551
552    elif GPA.pai_tracker.member_vars["scheduler"] is not None:
553        if (
554            epochs_since_cycle_switch > GPA.pc.get_initial_history_after_switches()
555            or GPA.pai_tracker.member_vars["mode"] == "p"
556        ):
557            if GPA.pc.get_verbose():
558                if hasattr(GPA.pai_tracker.member_vars["scheduler_instance"], '_step_count'):
559                    count = GPA.pai_tracker.member_vars["scheduler_instance"]._step_count
560                else:
561                    count = GPA.pai_tracker.member_vars["scheduler_instance"].last_epoch
562
563                print(
564                    f"Incrementing scheduler to count "
565                    f'{count}'
566                )
567            GPA.pai_tracker.member_vars["scheduler_instance"].step()
568            if (
569                GPA.pai_tracker.member_vars["scheduler"]
570                is torch.optim.lr_scheduler.ReduceLROnPlateau
571            ):
572                if GPA.pc.get_verbose():
573                    print(
574                        f"Scheduler is now at "
575                        f'{GPA.pai_tracker.member_vars["scheduler_instance"].num_bad_epochs} bad epochs'
576                    )
577
578    if (
579        epochs_since_cycle_switch <= GPA.pc.get_initial_history_after_switches()
580        and GPA.pai_tracker.member_vars["mode"] == "n"
581    ):
582        if GPA.pc.get_verbose():
583            print(
584                f"Not stepping with history {GPA.pc.get_initial_history_after_switches()} "
585                f"and current {epochs_since_cycle_switch}"
586            )
587
588    for param_group in GPA.pai_tracker.member_vars["optimizer_instance"].param_groups:
589        learning_rate2 = param_group["lr"]
590
591    stepped = False
592    at_last_count = False
593
594    if GPA.pc.get_verbose():
595        print(
596            f"Checking if at last with scores "
597            f'{len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"])}, '
598            f"count since switch {epochs_since_cycle_switch} "
599            f"and last total lr step count "
600            f'{GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]}'
601        )
602
603    # Check if at double or exactly the test count
604    if (
605        len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 0
606        and epochs_since_cycle_switch
607        == GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"] * 2
608    ) or (
609        len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 1
610        and epochs_since_cycle_switch
611        == GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]
612    ):
613        at_last_count = True
614
615    if GPA.pc.get_verbose():
616        print(
617            f"At last count {at_last_count} with count {epochs_since_cycle_switch} "
618            f'and last LR count {GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]}'
619        )
620
621    if learning_rate1 != learning_rate2:
622        stepped = True
623        GPA.pai_tracker.member_vars["current_step_count"] += 1
624
625        if GPA.pc.get_verbose():
626            print(
627                f"Learning rate just stepped to {learning_rate2:.10e} "
628                f'with {GPA.pai_tracker.member_vars["current_step_count"]} total steps'
629            )
630
631        if (
632            GPA.pai_tracker.member_vars["current_step_count"]
633            == GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]
634        ):
635            if GPA.pc.get_verbose():
636                print(
637                    f'{GPA.pai_tracker.member_vars["current_step_count"]} '
638                    f"steps is the max of the last switch mode"
639                )
640            # Set it when 1->2 gets to 2, not when 0->1 hits 2 as stopping point
641            if (
642                GPA.pai_tracker.member_vars["current_step_count"]
643                - GPA.pai_tracker.member_vars[
644                    "current_n_learning_rate_initial_skip_steps"
645                ]
646                == 1
647            ):
648                GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"] = (
649                    epochs_since_cycle_switch
650                )
651
652    if GPA.pc.get_verbose():
653        print(
654            f"Learning rates were {learning_rate1:.8e} and {learning_rate2:.8e} "
655            f'started with {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}, '
656            f'and is now at {GPA.pai_tracker.member_vars["current_step_count"]} '
657            f'committed {GPA.pai_tracker.member_vars["committed_to_initial_rate"]} '
658            f"then either this (non zero) or eventually comparing to "
659            f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
660            f'steps or rate {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]:.8f}'
661        )
662
663    # If learning rate just stepped, check restart at lower rate
664    if (
665        (GPA.pai_tracker.member_vars["scheduler"] is not None)
666        and
667        # If potentially might have higher accuracy
668        (
669            (GPA.pai_tracker.member_vars["mode"] == "n")
670            or GPA.pc.get_learn_dendrites_live()
671        )
672        and
673        # And learning rate just stepped
674        (stepped or at_last_count)
675    ):
676
677        # If this is the first dendrite addition (last_max_learning_rate_steps == 0),
678        # immediately commit to the initial rate without searching
679        if GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] == 0:
680            if GPA.pc.get_verbose():
681                print(
682                    f"First dendrite addition detected (last_max_learning_rate_steps == 0), "
683                    f"immediately committing to initial rate without search"
684                )
685            GPA.pai_tracker.member_vars["committed_to_initial_rate"] = True
686            GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] = (
687                GPA.pai_tracker.member_vars["current_step_count"]
688            )
689            GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = (
690                learning_rate2
691            )
692
693        # If hasn't committed to a learning rate for this cycle yet
694        if not GPA.pai_tracker.member_vars["committed_to_initial_rate"]:
695            best_score_so_far = GPA.pai_tracker.member_vars[
696                "global_best_validation_score"
697            ]
698
699            if GPA.pc.get_verbose():
700                print(
701                    f"In statements to check next learning rate with "
702                    f"stepped {stepped} and max count {at_last_count}"
703                )
704
705            # If no scores saved for this dendrite and initial LR test did second step
706            if len(
707                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]
708            ) == 0 and (
709                GPA.pai_tracker.member_vars["current_step_count"]
710                - GPA.pai_tracker.member_vars[
711                    "current_n_learning_rate_initial_skip_steps"
712                ]
713                == 2
714                or at_last_count
715            ):
716
717                restructured = True
718                GPA.pai_tracker.clear_optimizer_and_scheduler()
719
720                # Save system for this initial condition
721                old_global = GPA.pai_tracker.member_vars["global_best_validation_score"]
722                old_accuracy = GPA.pai_tracker.member_vars[
723                    "current_best_validation_score"
724                ]
725                old_counts = GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"]
726                skip1 = GPA.pai_tracker.member_vars[
727                    "current_n_learning_rate_initial_skip_steps"
728                ]
729
730                now = datetime.now()
731                dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
732
733                GPA.pai_tracker.save_graphs(
734                    f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
735                )
736
737                if GPA.pc.get_test_saves():
738                    UPA.save_system(
739                        net,
740                        GPA.pc.get_save_name(),
741                        f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
742                    )
743
744                if GPA.pc.get_verbose():
745                    print(
746                        f"Saving with initial steps: {dt_string}_PBCount_"
747                        f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
748                        f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
749                        f"with current best {old_accuracy}"
750                    )
751
752                # Load back at start and try with lower initial learning rate
753                net = UPA.load_system(
754                    net,
755                    GPA.pc.get_save_name(),
756                    f'switch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
757                    switch_call=True,
758                )
759                GPA.pai_tracker.member_vars[
760                    "current_n_learning_rate_initial_skip_steps"
761                ] = (skip1 + 1)
762                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"].append(
763                    old_accuracy
764                )
765                GPA.pai_tracker.member_vars["global_best_validation_score"] = old_global
766                GPA.pai_tracker.member_vars["initial_lr_test_epoch_count"] = old_counts
767
768            # If there is one score already, this is first step at next score
769            elif len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 1:
770                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"].append(
771                    GPA.pai_tracker.member_vars["current_best_validation_score"]
772                )
773
774                # If this LR's score was worse than last LR's score
775                lr_score_worse = False
776                if GPA.pai_tracker.member_vars["maximizing_score"]:
777                    lr_score_worse = (
778                        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]
779                        > GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]
780                    )
781                else:
782                    lr_score_worse = (
783                        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]
784                        < GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]
785                    )
786
787                if lr_score_worse:
788                    restructured = True
789                    GPA.pai_tracker.clear_optimizer_and_scheduler()
790
791                    if GPA.pc.get_verbose():
792                        print(
793                            f'Got initial {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]-1} '
794                            f'step score {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]} '
795                            f'and {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
796                            f'score at step {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]} '
797                            f"so loading old score"
798                        )
799
800                    prior_best = GPA.pai_tracker.member_vars[
801                        "current_cycle_lr_max_scores"
802                    ][0]
803
804                    now = datetime.now()
805                    dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
806
807                    GPA.pai_tracker.save_graphs(
808                        f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
809                    )
810
811                    if GPA.pc.get_test_saves():
812                        UPA.save_system(
813                            net,
814                            GPA.pc.get_save_name(),
815                            f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
816                        )
817
818                    if GPA.pc.get_verbose():
819                        print(
820                            f"Saving with initial steps: {dt_string}_PBCount_"
821                            f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
822                            f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
823                        )
824
825                    if GPA.pc.get_test_saves():
826                        net = UPA.load_system(
827                            net,
828                            GPA.pc.get_save_name(),
829                            f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]-1}',
830                            switch_call=True,
831                        )
832
833                    # Save graphs for chosen one
834                    now = datetime.now()
835                    dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
836
837                    GPA.pai_tracker.save_graphs(
838                        f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}PICKED'
839                    )
840
841                    if GPA.pc.get_test_saves():
842                        UPA.save_system(
843                            net,
844                            GPA.pc.get_save_name(),
845                            f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
846                        )
847
848                    if GPA.pc.get_verbose():
849                        print(
850                            f"Saving with initial steps: {dt_string}_PBCount_"
851                            f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
852                            f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
853                        )
854
855                    GPA.pai_tracker.member_vars["committed_to_initial_rate"] = True
856                    GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] = (
857                        GPA.pai_tracker.member_vars["current_step_count"]
858                    )
859                    GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = (
860                        learning_rate2
861                    )
862                    GPA.pai_tracker.member_vars["current_best_validation_score"] = (
863                        prior_best
864                    )
865
866                    if GPA.pc.get_verbose():
867                        print(
868                            f"Setting last max steps to "
869                            f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
870                            f'and lr {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]}'
871                        )
872
873                else:  # Current LR score is better
874                    if GPA.pc.get_verbose():
875                        print(
876                            f'Got initial {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]-1} '
877                            f'step score {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][0]} '
878                            f'and {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} '
879                            f'score at step {GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"][1]} '
880                            f"so NOT loading old score and continuing with this score"
881                        )
882
883                    if at_last_count:  # If this is the last one, set it to be picked
884                        restructured = True
885                        GPA.pai_tracker.clear_optimizer_and_scheduler()
886
887                        now = datetime.now()
888                        dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
889
890                        GPA.pai_tracker.save_graphs(
891                            f'{dt_string}_PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}PICKED'
892                        )
893
894                        if GPA.pc.get_test_saves():
895                            UPA.save_system(
896                                net,
897                                GPA.pc.get_save_name(),
898                                f'PBCount_{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}',
899                            )
900
901                        if GPA.pc.get_verbose():
902                            print(
903                                f"Saving with initial steps: {dt_string}_PBCount_"
904                                f'{GPA.pai_tracker.member_vars["num_dendrites_added"]}_startSteps_'
905                                f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}'
906                            )
907
908                        GPA.pai_tracker.member_vars["committed_to_initial_rate"] = True
909                        GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] = (
910                            GPA.pai_tracker.member_vars["current_step_count"]
911                        )
912                        GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = (
913                            learning_rate2
914                        )
915
916                        if GPA.pc.get_verbose():
917                            print(
918                                f"Setting last max steps to "
919                                f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
920                                f'and lr {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]}'
921                            )
922
923                GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
924
925            elif len(GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"]) == 2:
926                print(
927                    "Should never be here. Please let Perforated AI know if this happened."
928                )
929                pdb.set_trace()
930
931            GPA.pai_tracker.member_vars["global_best_validation_score"] = (
932                best_score_so_far
933            )
934
935        else:
936            if GPA.pc.get_verbose():
937                print(
938                    f"Setting last max steps to "
939                    f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]} '
940                    f'and lr {GPA.pai_tracker.member_vars["last_max_learning_rate_value"]}'
941                )
942            GPA.pai_tracker.member_vars["last_max_learning_rate_steps"] += 1
943            GPA.pai_tracker.member_vars["last_max_learning_rate_value"] = learning_rate2
944    if restructured:
945        return NETWORK_RESTRUCTURED, net
946    else:
947        return NO_MODEL_UPDATE, net

Updates the scheduler

This increments the scheduler, but if we are automatically sweeping to find the best initial learning rate for a new set of dendrites this function also triggers the network at addition time to try the next value.

Process for finding best initial learning rate for dendrites:

  1. Start at default rate
  2. Learn at that rate until scheduler increments twice
  3. Save that version, start dendrites at LR current increment - 1
  4. Repeat 2 and 3 until version has worse final score at set LR
  5. Load previous model with best accuracy at that LR as initial rate
Parameters
  • net (torch.nn.Module): The neural network model being trained.
  • accuracy (float): The accuracy of the model at the current learning rate.
  • epochs_since_cycle_switch (int): The number of epochs since the last cycle switch.
Returns
  • int: The status of restructuring or training completion.
  • torch.nn.Module: The potentially modified neural network model.
class PAINeuronModuleTracker:
 950class PAINeuronModuleTracker:
 951    """
 952    Manager class that tracks all neuron layers and dendrite layers,
 953    controls when new dendrites are added, and communicates signals to modules.
 954    """
 955
 956    def __init__(
 957        self,
 958        doing_pai,
 959        save_name,
 960        making_graphs=True,
 961        param_vals_setting=-1,
 962        values_per_train_epoch=-1,
 963        values_per_val_epoch=-1,
 964    ):
 965        """Initialize the tracker
 966
 967        Parameters
 968        ----------
 969        doing_pai : bool
 970            Whether or not dendrites should be used.
 971        save_name : str
 972            The base name for saving models and graphs.
 973        making_graphs : bool, optional
 974            Whether or not to generate graphs, by default True.
 975        param_vals_setting : int, optional
 976            Parameter values setting, by default -1.
 977        values_per_train_epoch : int, optional
 978            The number of values to look back for graphing
 979            during training, by default -1 (all values).
 980        values_per_val_epoch : int, optional
 981            The number of values to look back for graphing
 982            during validation, by default -1 (all values).
 983        Returns
 984        -------
 985        None
 986        """
 987
 988        # Dict of member vars and their types for saving
 989        self.member_vars = {}
 990        self.member_var_types = {}
 991
 992        # Whether or not PAI will be running
 993        self.member_vars["doing_pai"] = doing_pai
 994        self.member_var_types["doing_pai"] = "bool"
 995
 996        # How many Dendrites have been added
 997        self.member_vars["num_dendrites_added"] = 0
 998        self.member_var_types["num_dendrites_added"] = "int"
 999
1000        # How many Dendrites have been successfully integrated, does not count currently training dendrites
1001        self.member_vars["num_dendrites_integrated"] = 0
1002        self.member_var_types["num_dendrites_integrated"] = "int"
1003
1004        # How many cycles have been run, *2 or *2+1 of the above
1005        self.member_vars["num_cycles"] = 0
1006        self.member_var_types["num_cycles"] = "int"
1007
1008        # Pointers to all neuron wrapped modules
1009        self.neuron_module_vector = []
1010
1011        # Pointers to all non neuron modules for tracking
1012        self.tracked_neuron_module_vector = []
1013
1014        # Neuron training or dendrite training mode
1015        self.member_vars["mode"] = "n"
1016        self.member_var_types["mode"] = "string"
1017
1018        # Number of epochs run excluding overwritten epochs
1019        self.member_vars["num_epochs_run"] = -1
1020        self.member_var_types["num_epochs_run"] = "int"
1021
1022        # Number including overwritten epochs
1023        self.member_vars["total_epochs_run"] = -1
1024        self.member_var_types["total_epochs_run"] = "int"
1025
1026        # Last epoch that validation/correlation score was improved
1027        self.member_vars["epoch_last_improved"] = 0
1028        self.member_var_types["epoch_last_improved"] = "int"
1029
1030        # Running validation accuracy
1031        self.member_vars["running_accuracy"] = 0
1032        self.member_var_types["running_accuracy"] = "float"
1033
1034        # True if maxing validation, False if minimizing Loss
1035        self.member_vars["maximizing_score"] = True
1036        self.member_var_types["maximizing_score"] = "bool"
1037
1038        # Mode for switching back and forth between learning modes
1039        self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
1040        self.member_var_types["switch_mode"] = "int"
1041
1042        # Epoch of the last switch
1043        self.member_vars["last_switch"] = 0
1044        self.member_var_types["last_switch"] = "int"
1045
1046        # Highest validation score from current cycle
1047        self.member_vars["current_best_validation_score"] = 0
1048        self.member_var_types["current_best_validation_score"] = "float"
1049
1050        # Last epoch where the learning rate was updated
1051        self.member_vars["initial_lr_test_epoch_count"] = -1
1052        self.member_var_types["initial_lr_test_epoch_count"] = "int"
1053
1054        # Highest validation score of full run
1055        self.member_vars["global_best_validation_score"] = 0
1056        self.member_var_types["global_best_validation_score"] = "float"
1057
1058        # List of switch epochs
1059        self.member_vars["switch_epochs"] = []
1060        self.member_var_types["switch_epochs"] = "int array"
1061
1062        # Parameter counts at each network structure
1063        self.member_vars["param_counts"] = []
1064        self.member_var_types["param_counts"] = "int array"
1065
1066        # List of epochs where switch was made to neuron training
1067        self.member_vars["n_switch_epochs"] = []
1068        self.member_var_types["n_switch_epochs"] = "int array"
1069
1070        # List of epochs where switch was made to dendrite training
1071        self.member_vars["p_switch_epochs"] = []
1072        self.member_var_types["p_switch_epochs"] = "int array"
1073
1074        # List of validation accuracies
1075        self.member_vars["accuracies"] = []
1076        self.member_var_types["accuracies"] = "float array"
1077
1078        # List of epochs where score improved for scheduler updates
1079        self.member_vars["last_improved_accuracies"] = []
1080        self.member_var_types["last_improved_accuracies"] = "int array"
1081
1082        # List of test accuracy scores registered
1083        self.member_vars["test_accuracies"] = []
1084        self.member_var_types["test_accuracies"] = "float array"
1085
1086        # List of accuracies registered during neuron training
1087        self.member_vars["n_accuracies"] = []
1088        self.member_var_types["n_accuracies"] = "float array"
1089
1090        # List of accuracies registered during dendrite training
1091        self.member_vars["p_accuracies"] = []
1092        self.member_var_types["p_accuracies"] = "float array"
1093
1094        # Running average accuracies from recent epochs
1095        self.member_vars["running_accuracies"] = []
1096        self.member_var_types["running_accuracies"] = "float array"
1097
1098        # List of additional scores recorded
1099        self.member_vars["extra_scores"] = {}
1100        self.member_var_types["extra_scores"] = "float array dictionary"
1101
1102        # Extra scores not set to be graphed
1103        self.member_vars["extra_scores_without_graphing"] = {}
1104        self.member_var_types["extra_scores_without_graphing"] = (
1105            "float array dictionary"
1106        )
1107
1108        # List of test scores
1109        self.member_vars["test_scores"] = []
1110        self.member_var_types["test_scores"] = "float array"
1111
1112        # Extra scores calculated during neuron training
1113        self.member_vars["n_extra_scores"] = {}
1114        self.member_var_types["n_extra_scores"] = "float array dictionary"
1115
1116        # List of training losses calculated
1117        self.member_vars["training_loss"] = []
1118        self.member_var_types["training_loss"] = "float array"
1119
1120        # List of learning rates at each epoch
1121        self.member_vars["training_learning_rates"] = []
1122        self.member_var_types["training_learning_rates"] = "float array"
1123
1124        # Best dendrite scores
1125        self.member_vars["best_scores"] = []
1126        self.member_var_types["best_scores"] = "float array array"
1127
1128        # Current dendrite scores
1129        self.member_vars["current_scores"] = []
1130        self.member_var_types["current_scores"] = "float array array"
1131
1132        # Times for neuron training epochs
1133        self.member_vars["n_epoch_times"] = []
1134        self.member_var_types["n_epoch_times"] = "float array"
1135
1136        # Timing values
1137        self.member_vars["p_epoch_times"] = []
1138        self.member_var_types["p_epoch_times"] = "float array"
1139        self.member_vars["n_train_times"] = []
1140        self.member_var_types["n_train_times"] = "float array"
1141        self.member_vars["p_train_times"] = []
1142        self.member_var_types["p_train_times"] = "float array"
1143        self.member_vars["n_val_times"] = []
1144        self.member_var_types["n_val_times"] = "float array"
1145        self.member_vars["p_val_times"] = []
1146        self.member_var_types["p_val_times"] = "float array"
1147
1148        # Setting for tracking timing
1149        self.member_vars["manual_train_switch"] = False
1150        self.member_var_types["manual_train_switch"] = "bool"
1151
1152        # Tracking scores overwritten when reloading best model
1153        self.member_vars["overwritten_extras"] = []
1154        self.member_var_types["overwritten_extras"] = "float array dictionary array"
1155        self.member_vars["overwritten_vals"] = []
1156        self.member_var_types["overwritten_vals"] = "float array array"
1157        self.member_vars["overwritten_epochs"] = 0
1158        self.member_var_types["overwritten_epochs"] = "int"
1159
1160        # Setting for determining scores
1161        self.member_vars["param_vals_setting"] = GPA.pc.get_param_vals_setting()
1162        self.member_var_types["param_vals_setting"] = "int"
1163
1164        # Optimizer and scheduler types and instances
1165        self.member_vars["optimizer"] = None
1166        self.member_var_types["optimizer"] = "type"
1167        self.member_vars["scheduler"] = None
1168        self.member_var_types["scheduler"] = "type"
1169        self.member_vars["optimizer_instance"] = None
1170        self.member_var_types["optimizer_instance"] = "empty array"
1171        self.member_vars["scheduler_instance"] = None
1172        self.member_var_types["scheduler_instance"] = "empty array"
1173
1174        # Flag for if the tracker was loaded
1175        self.loaded = False
1176
1177        # flag for 
1178        self.member_vars["step_status"] = STEP_CLEARED
1179        self.member_var_types["step_status"] = "int"
1180
1181
1182        # Settings for tracking learning rates
1183        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
1184        self.member_var_types["current_n_learning_rate_initial_skip_steps"] = "int"
1185        self.member_vars["last_max_learning_rate_steps"] = 0
1186        self.member_var_types["last_max_learning_rate_steps"] = "int"
1187        self.member_vars["last_max_learning_rate_value"] = -1
1188        self.member_var_types["last_max_learning_rate_value"] = "float"
1189        self.member_vars["current_cycle_lr_max_scores"] = []
1190        self.member_var_types["current_cycle_lr_max_scores"] = "float array"
1191        self.member_vars["current_step_count"] = 0
1192        self.member_var_types["current_step_count"] = "int"
1193        self.member_vars["committed_to_initial_rate"] = True
1194        self.member_var_types["committed_to_initial_rate"] = "bool"
1195        self.member_vars["best_mean_score_improved_this_epoch"] = 0
1196        self.member_var_types["best_mean_score_improved_this_epoch"] = "int"
1197
1198        # Flag for if current dendrite achieved highest global score
1199        self.member_vars["current_n_set_global_best"] = True
1200        self.member_var_types["current_n_set_global_best"] = "bool"
1201
1202        # Number of tries adding this dendrite count
1203        self.member_vars["num_dendrite_tries"] = 0
1204        self.member_var_types["num_dendrite_tries"] = "int"
1205
1206        # Count of batches per epoch
1207        self.values_per_train_epoch = values_per_train_epoch
1208        self.values_per_val_epoch = values_per_val_epoch
1209
1210        self.save_name = save_name
1211        self.making_graphs = making_graphs
1212
1213        self.start_time = time.time()
1214        self.saved_time = 0
1215        self.start_epoch(internal_call=True)
1216
1217        if GPA.pc.get_verbose():
1218            print(f'Initializing with switch_mode {self.member_vars["switch_mode"]}')
1219
1220    def to_string(self):
1221        """Convert tracker values to string for saving with safetensors.
1222
1223        Parameters
1224        ----------
1225        None
1226
1227        Returns
1228        -------
1229        str
1230            Serialized tracker state suitable for storage in a safetensors field.
1231        """
1232
1233        full_string = ""
1234        for var in self.member_vars:
1235            full_string += var + ","
1236            if self.member_vars[var] is None:
1237                full_string += "None"
1238                full_string += "\n"
1239            elif self.member_var_types[var] == "bool":
1240                full_string += str(self.member_vars[var])
1241                full_string += "\n"
1242            elif self.member_var_types[var] in ("int", "float", "string"):
1243                full_string += str(self.member_vars[var])
1244                full_string += "\n"
1245            elif self.member_var_types[var] == "type":
1246                name = (
1247                    self.member_vars[var].__module__
1248                    + "."
1249                    + self.member_vars[var].__name__
1250                )
1251                full_string += str(self.member_vars[var])
1252                full_string += "\n"
1253            elif self.member_var_types[var] == "empty array":
1254                full_string += "[]"
1255                full_string += "\n"
1256            elif self.member_var_types[var] in ("int array", "float array"):
1257                full_string += "\n"
1258                string = ""
1259                for val in self.member_vars[var]:
1260                    string += str(val) + ","
1261                # Remove the last comma
1262                string = string[:-1]
1263                full_string += string
1264                full_string += "\n"
1265            elif self.member_var_types[var] == "float array dictionary array":
1266                full_string += "\n"
1267                for array in self.member_vars[var]:
1268                    for key in array:
1269                        string = key + ","
1270                        for val in array[key]:
1271                            string += str(val) + ","
1272                        # Remove the last comma
1273                        string = string[:-1]
1274                        full_string += string
1275                        full_string += "\n"
1276                    full_string += "endkey"
1277                    full_string += "\n"
1278                full_string += "endarray"
1279                full_string += "\n"
1280            elif self.member_var_types[var] == "float array dictionary":
1281                full_string += "\n"
1282                for key in self.member_vars[var]:
1283                    string = key + ","
1284                    for val in self.member_vars[var][key]:
1285                        string += str(val) + ","
1286                    # Remove the last comma
1287                    string = string[:-1]
1288                    full_string += string
1289                    full_string += "\n"
1290                full_string += "end"
1291                full_string += "\n"
1292            elif self.member_var_types[var] == "float array array":
1293                full_string += "\n"
1294                for array in self.member_vars[var]:
1295                    string = ""
1296                    for val in array:
1297                        string += str(val) + ","
1298                    # Remove the last comma
1299                    string = string[:-1]
1300                    full_string += string
1301                    full_string += "\n"
1302                full_string += "end"
1303                full_string += "\n"
1304            else:
1305                print("Did not find a member variable")
1306                pdb.set_trace()
1307        return full_string
1308
1309    def from_string(self, string):
1310        """Load tracker values from string.
1311
1312        Parameters
1313        ----------
1314        string : str
1315            The string to load from.
1316
1317        Returns
1318        -------
1319        None
1320            This function does not return a value.
1321        """
1322        f = io.StringIO(string)
1323        while True:
1324            line = f.readline()
1325            if not line:
1326                break
1327            vals = line.split(",")
1328            var = vals[0]
1329
1330            if self.member_var_types[var] == "bool":
1331                val = vals[1][:-1]
1332                if val == "True":
1333                    self.member_vars[var] = True
1334                elif val == "False":
1335                    self.member_vars[var] = False
1336                elif val == "1":
1337                    self.member_vars[var] = 1
1338                elif val == "0":
1339                    self.member_vars[var] = 0
1340                else:
1341                    print("Something went wrong with loading")
1342                    pdb.set_trace()
1343            elif self.member_var_types[var] == "int":
1344                val = vals[1]
1345                self.member_vars[var] = int(val)
1346            elif self.member_var_types[var] == "float":
1347                val = vals[1]
1348                self.member_vars[var] = float(val)
1349            elif self.member_var_types[var] == "string":
1350                val = vals[1][:-1]
1351                self.member_vars[var] = val
1352            elif self.member_var_types[var] == "type":
1353                # Ignore loading types, tracker should have them set up
1354                continue
1355            elif self.member_var_types[var] == "empty array":
1356                val = vals[1]
1357                self.member_vars[var] = []
1358            elif self.member_var_types[var] == "int array":
1359                vals = f.readline()[:-1].split(",")
1360                self.member_vars[var] = []
1361                if vals[0] == "":
1362                    continue
1363                for val in vals:
1364                    self.member_vars[var].append(int(val))
1365            elif self.member_var_types[var] == "float array":
1366                vals = f.readline()[:-1].split(",")
1367                self.member_vars[var] = []
1368                if vals[0] == "":
1369                    continue
1370                for val in vals:
1371                    self.member_vars[var].append(float(val))
1372            elif self.member_var_types[var] == "float array dictionary array":
1373                self.member_vars[var] = []
1374                line2 = f.readline()[:-1]
1375                while line2 != "endarray":
1376                    temp = {}
1377                    while line2 != "endkey":
1378                        vals = line2.split(",")
1379                        name = vals[0]
1380                        temp[name] = []
1381                        vals = vals[1:]
1382                        for val in vals:
1383                            temp[name].append(float(val))
1384                        line2 = f.readline()[:-1]
1385                    self.member_vars[var].append(temp)
1386                    line2 = f.readline()[:-1]
1387            elif self.member_var_types[var] == "float array dictionary":
1388                self.member_vars[var] = {}
1389                line2 = f.readline()[:-1]
1390                while line2 != "end":
1391                    vals = line2.split(",")
1392                    name = vals[0]
1393                    self.member_vars[var][name] = []
1394                    vals = vals[1:]
1395                    for val in vals:
1396                        self.member_vars[var][name].append(float(val))
1397                    line2 = f.readline()[:-1]
1398            elif self.member_var_types[var] == "float array array":
1399                self.member_vars[var] = []
1400                line2 = f.readline()[:-1]
1401                while line2 != "end":
1402                    vals = line2.split(",")
1403                    self.member_vars[var].append([])
1404                    if line2:
1405                        for val in vals:
1406                            self.member_vars[var][-1].append(float(val))
1407                    line2 = f.readline()[:-1]
1408            else:
1409                print("Did not find a member variable")
1410
1411                pdb.set_trace()
1412
1413    def from_string_debug(self, string):
1414        """Debug function to print tracker values from string without loading them.
1415
1416        Parameters
1417        ----------
1418        string : str
1419            The string to debug load from.
1420
1421        Returns
1422        -------
1423        None
1424            This function does not return a value.
1425        """
1426        f = io.StringIO(string)
1427        print("=== DEBUGGING TRACKER VARIABLES ===")
1428
1429        while True:
1430            line = f.readline()
1431            if not line:
1432                break
1433            vals = line.split(",")
1434            var = vals[0]
1435
1436            print(f"\nVariable: {var}")
1437            print(f"Type: {self.member_var_types.get(var, 'UNKNOWN TYPE')}")
1438            print(f"Current value: {self.member_vars.get(var, 'NOT SET')}")
1439
1440            if self.member_var_types.get(var) == "bool":
1441                val = vals[1][:-1]
1442                print(f"Would set to: {val} -> {val == 'True'}")
1443
1444            elif self.member_var_types.get(var) == "int":
1445                val = vals[1]
1446                print(f"Would set to: {int(val)}")
1447
1448            elif self.member_var_types.get(var) == "float":
1449                val = vals[1]
1450                print(f"Would set to: {float(val)}")
1451
1452            elif self.member_var_types.get(var) == "string":
1453                val = vals[1][:-1]
1454                print(f"Would set to: '{val}'")
1455
1456            elif self.member_var_types.get(var) == "type":
1457                print("Would skip (type loading)")
1458
1459            elif self.member_var_types.get(var) == "empty array":
1460                val = vals[1]
1461                print(f"Would set to: [] (empty array)")
1462
1463            elif self.member_var_types.get(var) == "int array":
1464                vals_line = f.readline()[:-1].split(",")
1465                print(f"Would set to int array with {len(vals_line)} elements:")
1466                if vals_line[0] != "":
1467                    print(
1468                        f"  Elements: {vals_line[:5]}{'...' if len(vals_line) > 5 else ''}"
1469                    )
1470                else:
1471                    print("  Empty array")
1472
1473            elif self.member_var_types.get(var) == "float array":
1474                vals_line = f.readline()[:-1].split(",")
1475                print(f"Would set to float array with {len(vals_line)} elements:")
1476                if vals_line[0] != "":
1477                    print(
1478                        f"  Elements: {vals_line[:5]}{'...' if len(vals_line) > 5 else ''}"
1479                    )
1480                else:
1481                    print("  Empty array")
1482
1483            elif self.member_var_types.get(var) == "float array dictionary array":
1484                print("Would process float array dictionary array:")
1485                array_count = 0
1486                line2 = f.readline()[:-1]
1487                while line2 != "endarray":
1488                    key_count = 0
1489                    while line2 != "endkey":
1490                        vals_dict = line2.split(",")
1491                        name = vals_dict[0]
1492                        print(
1493                            f"  Array {array_count}, Key '{name}': {len(vals_dict)-1} elements"
1494                        )
1495                        key_count += 1
1496                        line2 = f.readline()[:-1]
1497                    print(f"  Array {array_count} has {key_count} keys")
1498                    array_count += 1
1499                    line2 = f.readline()[:-1]
1500                print(f"  Total arrays: {array_count}")
1501
1502            elif self.member_var_types.get(var) == "float array dictionary":
1503                print("Would process float array dictionary:")
1504                line2 = f.readline()[:-1]
1505                key_count = 0
1506                while line2 != "end":
1507                    vals_dict = line2.split(",")
1508                    name = vals_dict[0]
1509                    print(f"  Key '{name}': {len(vals_dict)-1} elements")
1510                    key_count += 1
1511                    line2 = f.readline()[:-1]
1512                print(f"  Total keys: {key_count}")
1513
1514            elif self.member_var_types.get(var) == "float array array":
1515                print("Would process float array array:")
1516                line2 = f.readline()[:-1]
1517                array_count = 0
1518                while line2 != "end":
1519                    if line2:
1520                        vals_array = line2.split(",")
1521                        print(f"  Array {array_count}: {len(vals_array)} elements")
1522                    else:
1523                        print(f"  Array {array_count}: empty")
1524                    array_count += 1
1525                    line2 = f.readline()[:-1]
1526                print(f"  Total arrays: {array_count}")
1527
1528            else:
1529                print(f"UNKNOWN TYPE: {self.member_var_types.get(var, 'NOT FOUND')}")
1530
1531        print("\n=== END DEBUG ===")
1532
1533    def save_tracker_settings(self):
1534        """Save tracker settings for DistributedDataParallel use.
1535
1536        Saves settings in save_name/array_dims.csv
1537
1538        Parameters
1539        ----------
1540        None
1541        Returns
1542        -------
1543        None
1544
1545        -----
1546        Instructions for use are in API customization.md
1547        """
1548        if not os.path.isdir(self.save_name):
1549            os.makedirs(self.save_name)
1550        f = open(self.save_name + "/array_dims.csv", "w")
1551        for layer in self.neuron_module_vector:
1552            dv = layer.dendrite_module.dendrite_values[0]
1553            shape_str = ",".join(str(s) for s in dv.dendrite_storage_shape.tolist())
1554            f.write(f"{layer.name},{shape_str}\n")
1555        f.close()
1556        if not GPA.pc.get_silent():
1557            print("Tracker settings saved.")
1558            print("You may now delete save_tracker_settings")
1559
1560    def initialize_tracker_settings(self):
1561        """Initialize tracker settings from saved file.
1562
1563        This function loads tracker settings from a CSV file and applies them
1564        to the layers the tracker is managing.
1565
1566        Parameters
1567        ----------
1568        None
1569
1570        Returns
1571        -------
1572        None
1573
1574        """
1575
1576        channels = {}
1577        if not os.path.exists(self.save_name + "/array_dims.csv"):
1578            print(
1579                "You must call save_tracker_settings before "
1580                "initialize_tracker_settings"
1581            )
1582            print("Follow instructions in customization.md")
1583            pdb.set_trace()
1584        f = open(self.save_name + "/array_dims.csv", "r")
1585        for line in f:
1586            parts = line.strip().split(",")
1587            channels[parts[0]] = [int(s) for s in parts[1:]]
1588        for layer in self.neuron_module_vector:
1589            dv = layer.dendrite_module.dendrite_values[0]
1590            dv.setup_arrays(channels[layer.name])
1591
1592    def set_optimizer_instance(self, optimizer_instance, additional_optimizers=[]):
1593        """Set optimizer instance directly.
1594
1595        Parameters
1596        ----------
1597        optimizer_instance : object
1598            The optimizer instance to set.
1599
1600        Returns
1601        -------
1602        None
1603
1604        """
1605        # This call must be first before the parameters get filtered.
1606        optimizer_instance.zero_grad()
1607        try:
1608            for param_group in optimizer_instance.param_groups:
1609                if (
1610                    param_group["weight_decay"] > 0
1611                    and GPA.pc.get_weight_decay_accepted() is False
1612                ):
1613                    _pai_log(
1614                        "warning",
1615                        "For PAI training it is recommended to not use weight decay in your optimizer",
1616                    )
1617
1618        except:
1619            pass
1620        self.member_vars["optimizer_instance"] = optimizer_instance
1621        if GPA.pc.get_perforated_backpropagation():
1622            TPB.setup_optimizer_pb(self.member_vars["optimizer_instance"])
1623            for optimizer in additional_optimizers:
1624                TPB.filter_params(optimizer)
1625
1626    def set_optimizer(self, optimizer):
1627        """Set optimizer type to be initialized later
1628
1629        Parameters
1630        ----------
1631        optimizer : object
1632            The optimizer type to set.
1633
1634        Returns
1635        -------
1636        None
1637
1638        """
1639        self.member_vars["optimizer"] = optimizer
1640
1641    def set_scheduler(self, scheduler):
1642        """Set scheduler type to be initialized later
1643
1644        Parameters
1645        ----------
1646        scheduler : object
1647            The scheduler type to set.
1648
1649        Returns
1650        -------
1651        None
1652
1653        """
1654        if scheduler is not torch.optim.lr_scheduler.ReduceLROnPlateau:
1655            if GPA.pc.get_verbose():
1656                print("Not using ReduceLROnPlateau, this is not recommended")
1657        self.member_vars["scheduler"] = scheduler
1658
1659    def increment_scheduler(self, num_ticks, mode):
1660        """Increment the scheduler a set number of times.
1661
1662        Used for finding best initial learning rate when adding dendrites.
1663
1664        Parameters
1665        ----------
1666        num_ticks : int
1667            The number of scheduler steps to take.
1668        mode : str
1669            The mode for stepping the scheduler. Options are:
1670            - "step_learning_rate": Step based on improved accuracy epochs
1671            - "increment_epoch_count": Step based on total epoch count
1672
1673        Returns
1674        -------
1675        current_steps : int
1676            The number of learning rate changes that occurred.
1677        learning_rate1 : float
1678            The final learning rate after stepping.
1679
1680        """
1681
1682        current_steps = 0
1683        current_ticker = 0
1684
1685        for param_group in GPA.pai_tracker.member_vars[
1686            "optimizer_instance"
1687        ].param_groups:
1688            learning_rate1 = param_group["lr"]
1689
1690        if GPA.pc.get_verbose():
1691            print("Using scheduler:")
1692            print(type(self.member_vars["scheduler_instance"]))
1693
1694        while current_ticker < num_ticks:
1695            if GPA.pc.get_verbose():
1696                print(
1697                    f"Lower start rate initial {learning_rate1} "
1698                    f'stepping {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} times'
1699                )
1700
1701            if (
1702                type(self.member_vars["scheduler_instance"])
1703                is torch.optim.lr_scheduler.ReduceLROnPlateau
1704            ):
1705                if mode == "step_learning_rate":
1706                    # Step with counter as last improved accuracy
1707                    self.member_vars["scheduler_instance"].step(
1708                        metrics=self.member_vars["last_improved_accuracies"][
1709                            GPA.pai_tracker.steps_after_switch() - 1
1710                        ]
1711                    )
1712                elif mode == "increment_epoch_count":
1713                    # Step with improved epoch counts up to current location
1714                    self.member_vars["scheduler_instance"].step(
1715                        metrics=self.member_vars["last_improved_accuracies"][
1716                            -((num_ticks - 1) - current_ticker) - 1
1717                        ]
1718                    )
1719            else:
1720                self.member_vars["scheduler_instance"].step()
1721
1722            for param_group in GPA.pai_tracker.member_vars[
1723                "optimizer_instance"
1724            ].param_groups:
1725                learning_rate2 = param_group["lr"]
1726
1727            if learning_rate2 != learning_rate1:
1728                current_steps += 1
1729                learning_rate1 = learning_rate2
1730                if mode == "step_learning_rate":
1731                    current_ticker += 1
1732                if GPA.pc.get_verbose():
1733                    print(f"1 step {current_steps} to {learning_rate2}")
1734
1735            if mode == "increment_epoch_count":
1736                current_ticker += 1
1737
1738        return current_steps, learning_rate1
1739
1740    def setup_optimizer(self, net, opt_args, sched_args=None, parameters=None):
1741        """Initialize the optimizer and scheduler when added.
1742
1743        Parameters
1744        ----------
1745        net : object
1746            The neural network model.
1747        opt_args : dict
1748            The arguments for the optimizer.
1749        sched_args : dict, optional
1750            The arguments for the scheduler, by default None.
1751
1752        Returns
1753        -------
1754        optimizer : object
1755            The initialized optimizer instance.
1756        scheduler : object or None
1757            The initialized scheduler instance, or None if no scheduler was set.
1758
1759        """
1760        if "weight_decay" in opt_args and not GPA.pc.get_weight_decay_accepted():
1761            _pai_log(
1762                "warning",
1763                "For PAI training it is recommended to not use weight decay in your optimizer",
1764            )
1765
1766        if ("model" not in opt_args.keys()) and "params" not in opt_args.keys():
1767            print("In setup_optimizer it will be depreciated to not pass in params yourself in the future")
1768            print("please change the settings to include params")
1769            if self.member_vars["mode"] == "n":
1770                if parameters is not None:
1771                    opt_args["params"] = parameters
1772                else:
1773                    opt_args["params"] = filter(lambda p: p.requires_grad, net.parameters())
1774            else:
1775                params = UPA.get_pai_network_params(net)
1776                if parameters is not None:
1777                    # Filter parameters to only those in params, preserving weight_decay
1778                    params_set = set(params)
1779                    filtered_params = []
1780                    for param_group in parameters:
1781                        filtered_group_params = [p for p in param_group["params"] if p in params_set]
1782                        if filtered_group_params:
1783                            filtered_params.append({
1784                                "params": filtered_group_params,
1785                                "weight_decay": param_group["weight_decay"]
1786                            })
1787                    opt_args["params"] = filtered_params
1788                else:
1789                    opt_args["params"] = params
1790        elif "params" in opt_args.keys():
1791            # Check if params is a list of param groups (dicts) or a single param group
1792            params_value = opt_args["params"]
1793            if isinstance(params_value, list) and len(params_value) > 0:
1794                # Check if it's a list of dicts (multiple param groups) or list of tensors (single group)
1795                if isinstance(params_value[0], dict):
1796                    # Multiple param groups format: [{"params": [...], "lr": ...}, ...]
1797                    # Filter each param group for requires_grad
1798                    filtered_param_groups = []
1799                    for param_group in params_value:
1800                        filtered_group_params = [p for p in param_group["params"] if p.requires_grad]
1801                        if filtered_group_params:
1802                            new_group = param_group.copy()
1803                            new_group["params"] = filtered_group_params
1804                            filtered_param_groups.append(new_group)
1805                    opt_args["params"] = filtered_param_groups
1806                else:
1807                    # Single param group format: [tensor1, tensor2, ...] or generator
1808                    # Filter for requires_grad
1809                    opt_args["params"] = [p for p in params_value if p.requires_grad]
1810            elif hasattr(params_value, '__iter__'):
1811                # Handle generators or other iterables
1812                opt_args["params"] = [p for p in params_value if p.requires_grad]
1813
1814        optimizer = self.member_vars["optimizer"](**opt_args)
1815        self.set_optimizer_instance(optimizer)
1816
1817        if self.member_vars["scheduler"] is not None:
1818            # Handle SequentialLR specially
1819            if self.member_vars["scheduler"] is torch.optim.lr_scheduler.SequentialLR:
1820                """
1821                sched_args should be a dict with "schedulers" (list of tuples) and "milestones"
1822                For example:
1823                sequential_schedArgs = {
1824                    "schedulers": [
1825                        (warmup_scheduler_class, warmup_schedArgs),
1826                        (main_scheduler_class, main_schedArgs)
1827                    ],
1828                    "milestones": [switch_epoch]
1829                }
1830                """
1831                schedulers = []
1832                milestones = sched_args.get("milestones", [])
1833                scheduler_configs = sched_args.get("schedulers", [])
1834                
1835                for scheduler_class, scheduler_args in scheduler_configs:
1836                    schedulers.append(scheduler_class(optimizer, **scheduler_args))
1837                
1838                self.member_vars["scheduler_instance"] = torch.optim.lr_scheduler.SequentialLR(
1839                    optimizer, schedulers=schedulers, milestones=milestones
1840                )
1841            else:
1842                self.member_vars["scheduler_instance"] = self.member_vars["scheduler"](
1843                    optimizer, **sched_args
1844                )
1845            current_steps = 0
1846
1847            for param_group in GPA.pai_tracker.member_vars[
1848                "optimizer_instance"
1849            ].param_groups:
1850                learning_rate1 = param_group["lr"]
1851
1852            if GPA.pc.get_verbose():
1853                print(
1854                    f"Resetting scheduler with {GPA.pai_tracker.steps_after_switch()} "
1855                    f"steps and {GPA.pc.get_initial_history_after_switches()} initial ticks to skip"
1856                )
1857
1858            # Find setting of previously used learning rate before adding dendrites
1859            if (
1860                GPA.pai_tracker.member_vars[
1861                    "current_n_learning_rate_initial_skip_steps"
1862                ]
1863                != 0
1864            ):
1865                additional_steps, learning_rate1 = self.increment_scheduler(
1866                    GPA.pai_tracker.member_vars[
1867                        "current_n_learning_rate_initial_skip_steps"
1868                    ],
1869                    "step_learning_rate",
1870                )
1871                current_steps += additional_steps
1872
1873            if self.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
1874                initial = GPA.pc.get_initial_history_after_switches()
1875            else:
1876                initial = 0
1877
1878            if GPA.pai_tracker.steps_after_switch() > initial:
1879                # Minus extra 1 because this gets called after start epoch
1880                additional_steps, learning_rate1 = self.increment_scheduler(
1881                    (GPA.pai_tracker.steps_after_switch() - initial) - 1,
1882                    "increment_epoch_count",
1883                )
1884                current_steps += additional_steps
1885
1886            if GPA.pc.get_verbose():
1887                print(
1888                    f"Scheduler update loop with {current_steps} "
1889                    f"ended with {learning_rate1}"
1890                )
1891                print(
1892                    f"Scheduler ended with {current_steps} steps "
1893                    f"and lr of {learning_rate1}"
1894                )
1895
1896            self.member_vars["current_step_count"] = current_steps
1897            return optimizer, self.member_vars["scheduler_instance"]
1898        else:
1899            return optimizer, None
1900
1901    def clear_optimizer_and_scheduler(self):
1902        """Clear the instances for saving.
1903
1904        Parameters
1905        ----------
1906        None
1907
1908        Returns
1909        -------
1910        None
1911            This function does not return a value.
1912        """
1913        self.member_vars["optimizer_instance"] = None
1914        self.member_vars["scheduler_instance"] = None
1915
1916    def switch_time(self):
1917        """Determine if it's time to switch between neuron and dendrite training.
1918
1919        Parameters
1920        ----------
1921        None
1922
1923        Returns
1924        -------
1925        bool
1926            True if it's time to switch, False otherwise.
1927
1928        Notes
1929        -----
1930        Based on current settings and history of scores.
1931        """
1932
1933        switch_phrase = "No mode, this should never be the case."
1934        switch_number = GPA.pc.get_n_epochs_to_switch()
1935        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1936            switch_phrase = "DOING_SWITCH_EVERY_TIME"
1937        elif self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY:
1938            switch_phrase = "DOING_HISTORY"
1939        elif self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH:
1940            switch_phrase = "DOING_FIXED_SWITCH"
1941            switch_number = GPA.pc.get_fixed_switch_num()
1942        elif self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1943            switch_phrase = "DOING_NO_SWITCH"
1944        else:
1945            print(
1946                "A switch mode must be set.  Check your settings for GPA.pc.set_switch_mode()."
1947            )
1948            pdb.set_trace()
1949        if not GPA.pc.get_silent():
1950            if(GPA.pc.get_perforated_backpropagation()):
1951                print(
1952                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1953                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1954                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1955                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1956                    f'n: {switch_number}, p: {GPA.pc.get_p_epochs_to_switch()}, '
1957                    f'num_cycles: {self.member_vars["num_cycles"]}'
1958                )
1959            else:
1960                print(
1961                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1962                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1963                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1964                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1965                    f'n: {switch_number}, num_cycles: {self.member_vars["num_cycles"]}'
1966                )
1967            print(
1968                f'  Score tracking: current_n_set_global_best={self.member_vars["current_n_set_global_best"]}, '
1969                f'global_best={self.member_vars["global_best_validation_score"]:.4f}, '
1970                f'current_best={self.member_vars["current_best_validation_score"]:.4f}'
1971            )
1972        if GPA.pc.get_perforated_backpropagation():
1973            # this will fill in epoch last improved
1974            TPB.best_pai_score_improved_this_epoch(self)  ## CLOSED ONLY
1975        if self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1976            if not GPA.pc.get_silent():
1977                print("Returning False - doing no switch mode")
1978            return False
1979
1980        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1981            if not GPA.pc.get_silent():
1982                print("Returning True - switching every time")
1983            return True
1984
1985        # Check if we're in the middle of learning rate optimization
1986        # If so, block ALL switch triggers until committed
1987        if GPA.pc.get_verbose():
1988            print("=== LR Optimization Check ===")
1989            print(f'  mode == "n": {self.member_vars["mode"] == "n"}')
1990            print(f"  get_learn_dendrites_live(): {GPA.pc.get_learn_dendrites_live()}")
1991            print(f'  committed_to_initial_rate: {GPA.pai_tracker.member_vars["committed_to_initial_rate"]}')
1992            print(f"  get_dont_give_up_unless_learning_rate_lowered(): {GPA.pc.get_dont_give_up_unless_learning_rate_lowered()}")
1993            print(f'  current_n_learning_rate_initial_skip_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"]}')
1994            print(f'  last_max_learning_rate_steps: {self.member_vars["last_max_learning_rate_steps"]}')
1995            print(f'  skip_steps < max_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"] < self.member_vars["last_max_learning_rate_steps"]}')
1996            print(f'  scheduler is not None: {self.member_vars["scheduler"] is not None}')
1997            print("=============================")
1998        
1999        if (
2000            ((self.member_vars["mode"] == "n") or GPA.pc.get_learn_dendrites_live())
2001            and (GPA.pai_tracker.member_vars["committed_to_initial_rate"] is False)
2002            and (GPA.pc.get_dont_give_up_unless_learning_rate_lowered())
2003            and (
2004                self.member_vars["current_n_learning_rate_initial_skip_steps"]
2005                <= self.member_vars["last_max_learning_rate_steps"]
2006            )
2007            and self.member_vars["scheduler"] is not None
2008        ):
2009            if not GPA.pc.get_silent():
2010                print(
2011                    f"Returning False - learning rate optimization in progress. "
2012                    f"Not committed yet. Comparing "
2013                    f'initial {self.member_vars["current_n_learning_rate_initial_skip_steps"]} '
2014                    f'to last max {self.member_vars["last_max_learning_rate_steps"]}'
2015                )
2016            return False
2017
2018        if len(self.member_vars["switch_epochs"]) == 0:
2019            this_count = self.member_vars["num_epochs_run"]
2020        else:
2021            this_count = (
2022                self.member_vars["num_epochs_run"]
2023                - self.member_vars["switch_epochs"][-1]
2024            )
2025        cap_switch = False
2026        if GPA.pc.get_perforated_backpropagation():
2027            cap_switch = TPB.check_cap_switch(self, this_count)
2028
2029        if self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY and (
2030            (
2031                (self.member_vars["mode"] == "n")
2032                and (
2033                    self.member_vars["num_epochs_run"]
2034                    - self.member_vars["epoch_last_improved"]
2035                    >= GPA.pc.get_n_epochs_to_switch()
2036                )
2037                and this_count
2038                >= GPA.pc.get_initial_history_after_switches()
2039                + GPA.pc.get_n_epochs_to_switch()
2040            )
2041            or (GPA.pc.get_perforated_backpropagation() and TPB.history_switch(self))
2042            or cap_switch
2043        ):
2044            if not GPA.pc.get_silent():
2045                print("Returning True - History and last improved is hit")
2046            return True
2047
2048        if self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH and (
2049            (
2050                self.member_vars["total_epochs_run"] % GPA.pc.get_fixed_switch_num()
2051                == GPA.pc.get_fixed_switch_num() - 1
2052            )
2053            and self.member_vars["num_epochs_run"]
2054            >= GPA.pc.get_first_fixed_switch_num() - 1
2055        ):
2056            if not GPA.pc.get_silent():
2057                print("Returning True - Fixed switch number is hit")
2058            return True
2059
2060        if not GPA.pc.get_silent():
2061            print("Returning False - no triggers to switch have been hit")
2062        return False
2063
2064    def steps_after_switch(self):
2065        """Based on settings, return value for steps since a switch.
2066
2067        Different options for param vals setting determine what is returned.
2068
2069        Parameters
2070        ----------
2071        None
2072
2073        Returns
2074        -------
2075        int
2076            The number of epochs since the last switch, or total epochs run,
2077            depending on settings.
2078
2079        """
2080        if self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_TOTAL_EPOCH:
2081            return self.member_vars["num_epochs_run"]
2082        elif (
2083            self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_UPDATE_EPOCH
2084        ):
2085            return self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2086        elif (
2087            self.member_vars["param_vals_setting"]
2088            == GPA.pc.PARAM_VALS_BY_NEURON_EPOCH_START
2089        ):
2090            if self.member_vars["mode"] == "p":
2091                return (
2092                    self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2093                )
2094            else:
2095                return self.member_vars["num_epochs_run"]
2096        else:
2097            print(
2098                f'{self.member_vars["param_vals_setting"]} is not a valid param vals option'
2099            )
2100            pdb.set_trace()
2101
2102    def add_pai_neuron_module(self, new_module, initial_add=True):
2103        """Add neuron modules to internal vectors.
2104
2105        Parameters
2106        ----------
2107        new_module : object
2108            The new module to add.
2109        initial_add : bool, optional
2110            Whether this is the initial addition rather than loading from file
2111
2112        Returns
2113        -------
2114        None
2115
2116        """
2117
2118        # If it's a duplicate, ignore the second addition
2119        if new_module in self.neuron_module_vector:
2120            return
2121        self.neuron_module_vector.append(new_module)
2122        if self.member_vars["doing_pai"]:
2123            PA.set_wrapped_params(new_module)
2124        if initial_add:
2125            self.member_vars["best_scores"].append([])
2126            self.member_vars["current_scores"].append([])
2127
2128    def add_tracked_neuron_module(self, new_module, initial_add=True):
2129        """Add tracked modules to internal vectors
2130
2131        Parameters
2132        ----------
2133        new_module : object
2134            The new module to add.
2135        initial_add : bool, optional
2136            Whether this is the initial addition rather than loading from file
2137
2138        Returns
2139        -------
2140        None
2141
2142        """
2143        # If it's a duplicate, ignore the second addition
2144        if new_module in self.tracked_neuron_module_vector:
2145            return
2146        self.tracked_neuron_module_vector.append(new_module)
2147        if self.member_vars["doing_pai"]:
2148            PA.set_tracked_params(new_module)
2149
2150    def reset_module_vector(self, net, load_from_restart):
2151        """Clear internal vectors and reset from network.
2152
2153        Parameters
2154        ----------
2155        net : object
2156            The neural network model.
2157        load_from_restart : bool
2158            Whether loading from a restart file.
2159
2160        Returns
2161        -------
2162        None
2163
2164        """
2165        self.neuron_module_vector = []
2166        self.tracked_neuron_module_vector = []
2167        this_list = UPA.get_pai_modules(net, 0)
2168        for module in this_list:
2169            self.add_pai_neuron_module(module, initial_add=load_from_restart)
2170        this_list = UPA.get_tracked_modules(net, 0)
2171        for module in this_list:
2172            self.add_tracked_neuron_module(module, initial_add=load_from_restart)
2173
2174    def reset_vals_for_score_reset(self):
2175        """Reset cycle scores for new cycle.
2176
2177        Parameters
2178        ----------
2179        None
2180
2181        Returns
2182        -------
2183        None
2184            This function does not return a value.
2185        """
2186
2187        if GPA.pc.get_find_best_lr():
2188            self.member_vars["committed_to_initial_rate"] = False
2189            print("Resetting committed to initial rate to False")
2190        # If retaining all dendrties always say that the current dendrites set global best for saving and loading
2191        if GPA.pc.get_retain_all_dendrites():
2192            self.member_vars["current_n_set_global_best"] = True
2193            self.member_vars["global_best_validation_score"] = 0
2194        else:
2195            self.member_vars["current_n_set_global_best"] = False
2196
2197        # Don't reset global best, but do reset current best
2198        self.member_vars["current_best_validation_score"] = 0
2199        self.member_vars["initial_lr_test_epoch_count"] = -1
2200
2201    def set_dendrite_training(self):
2202        """Signal all layers to start dendrite training.
2203
2204        Parameters
2205        ----------
2206        None
2207
2208        Returns
2209        -------
2210        None
2211            This function does not return a value.
2212        """
2213        if GPA.pc.get_verbose():
2214            print("Calling set_dendrite_training")
2215
2216        for layer in self.neuron_module_vector[:]:
2217            worked = layer.set_mode("p")
2218            """
2219            worked is False when a layer was added to the neuron module vector
2220            but then it's never actually been used. This can happen when
2221            you have set a layer to have requires_grad = False or when
2222            you have a module as a member variable but it's not actually
2223            part of the network. Should be moved to be a tracked layer
2224            rather than a neuron layer.
2225            """
2226            if not worked:
2227                self.neuron_module_vector.remove(layer)
2228
2229        for layer in self.tracked_neuron_module_vector[:]:
2230            worked = layer.set_mode("p")
2231
2232        self.create_new_dendrite_module()
2233        self.member_vars["mode"] = "p"
2234        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2235
2236        if GPA.pc.get_learn_dendrites_live():
2237            self.reset_vals_for_score_reset()
2238
2239        self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2240            "current_step_count"
2241        ]
2242
2243        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
2244        GPA.pai_tracker.member_vars["num_cycles"] += 1
2245
2246
2247    def set_neuron_training(self):
2248        """Signal all layers to start neuron training.
2249
2250        Parameters
2251        ----------
2252        None
2253
2254        Returns
2255        -------
2256        None
2257            This function does not return a value.
2258        """
2259        for module in self.neuron_module_vector:
2260            module.set_mode("n")
2261        for module in self.tracked_neuron_module_vector[:]:
2262            module.set_mode("n")
2263
2264        self.member_vars["mode"] = "n"
2265        self.member_vars["num_dendrites_added"] += 1
2266        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2267        self.reset_vals_for_score_reset()
2268
2269        self.member_vars["current_cycle_lr_max_scores"] = []
2270        if GPA.pc.get_learn_dendrites_live():
2271            self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2272                "current_step_count"
2273            ]
2274        GPA.pai_tracker.member_vars["num_cycles"] += 1
2275
2276        if GPA.pc.get_reset_best_score_on_switch():
2277            GPA.pai_tracker.member_vars["current_best_validation_score"] = 0
2278            GPA.pai_tracker.member_vars["running_accuracy"] = 0
2279
2280    def start_epoch(self, internal_call=False):
2281        """Perform steps for when a new training epoch is about to begin.
2282
2283        Parameters
2284        ----------
2285        internal_call : bool, optional
2286            Whether this is an internal call or manual call
2287
2288        Returns
2289        -------
2290        None
2291
2292        Notes
2293        -----
2294        If you ever need to call this manually, set internal_call to False.
2295
2296        """
2297        if self.member_vars["manual_train_switch"] and internal_call:
2298            return
2299
2300        if not internal_call and not self.member_vars["manual_train_switch"]:
2301            self.member_vars["manual_train_switch"] = True
2302            self.saved_time = 0
2303            self.member_vars["num_epochs_run"] = -1
2304            self.member_vars["total_epochs_run"] = -1
2305
2306        end = time.time()
2307        if self.member_vars["manual_train_switch"]:
2308            if self.saved_time != 0:
2309                if self.member_vars["mode"] == "p":
2310                    self.member_vars["p_val_times"].append(end - self.saved_time)
2311                else:
2312                    self.member_vars["n_val_times"].append(end - self.saved_time)
2313
2314        if self.member_vars["mode"] == "p":
2315            for layer in self.neuron_module_vector:
2316                for m in range(0, GPA.pc.get_global_candidates()):
2317                    with torch.no_grad():
2318                        if GPA.pc.get_verbose():
2319                            print(f"Resetting score for {layer.name}")
2320                        # Snapshot best_score before reset so we can compute per-epoch improvement
2321                        layer.dendrite_module.dendrite_values[
2322                            m
2323                        ].epoch_start_best_score.copy_(
2324                            layer.dendrite_module.dendrite_values[
2325                                m
2326                            ].best_score.detach()
2327                        )
2328                        layer.dendrite_module.dendrite_values[
2329                            m
2330                        ].best_score_improved_this_epoch = (
2331                            layer.dendrite_module.dendrite_values[
2332                                m
2333                            ].best_score_improved_this_epoch
2334                            * 0
2335                        )
2336                        layer.dendrite_module.dendrite_values[
2337                            m
2338                        ].nodes_best_improved_this_epoch = (
2339                            layer.dendrite_module.dendrite_values[
2340                                m
2341                            ].nodes_best_improved_this_epoch
2342                            * 0
2343                        )
2344                        layer.dendrite_module.dendrite_values[
2345                            m
2346                        ].nodes_improved_any = (
2347                            layer.dendrite_module.dendrite_values[
2348                                m
2349                            ].nodes_improved_any
2350                            * 0
2351                        )
2352            if GPA.pc.get_perforated_backpropagation():
2353                self.member_vars["best_mean_score_improved_this_epoch"] = 0
2354        self.member_vars["num_epochs_run"] += 1
2355        self.member_vars["total_epochs_run"] = (
2356            self.member_vars["num_epochs_run"] + self.member_vars["overwritten_epochs"]
2357        )
2358        self.saved_time = end
2359
2360    def stop_epoch(self, internal_call=False):
2361        """Perform steps when a training epoch has completed.
2362
2363        Parameters
2364        ----------
2365        internal_call : bool, optional
2366            Whether this is an internal call or manual call
2367
2368        Returns
2369        -------
2370        None
2371
2372        Notes
2373        -----
2374        If you ever need to call this manually, set internal_call to False.
2375
2376        """
2377        end = time.time()
2378        if self.member_vars["manual_train_switch"] and internal_call:
2379            return
2380
2381        if self.member_vars["manual_train_switch"]:
2382            if self.member_vars["mode"] == "p":
2383                self.member_vars["p_train_times"].append(end - self.saved_time)
2384            else:
2385                self.member_vars["n_train_times"].append(end - self.saved_time)
2386        else:
2387            if self.member_vars["mode"] == "p":
2388                self.member_vars["p_epoch_times"].append(end - self.saved_time)
2389            else:
2390                self.member_vars["n_epoch_times"].append(end - self.saved_time)
2391
2392        self.saved_time = end
2393
2394    def initialize(
2395        self,
2396        model,
2397        doing_pai=True,
2398        save_name="PAI",
2399        making_graphs=True,
2400        maximizing_score=True,
2401        num_classes=10000,
2402        values_per_train_epoch=-1,
2403        values_per_val_epoch=-1,
2404        zooming_graph=True,
2405    ):
2406        """Setup the tracker with initial settings.
2407
2408
2409        Parameters
2410        ----------
2411        model : object
2412            The neural network model.
2413        doing_pai : bool, optional
2414            Whether to add dendrites, by default True.
2415        save_name : str, optional
2416            The name under which to save the model.
2417        making_graphs : bool, optional
2418            Whether to make graphs, by default True.
2419        maximizing_score : bool, optional
2420            Whether to maximize the score, by default True.
2421        num_classes : int, optional
2422            The number of classes in the dataset, unused
2423        values_per_train_epoch : int, optional
2424            The number of values to look back for graphing
2425            during training, by default -1 (all values).
2426        values_per_val_epoch : int, optional
2427            The number of values to look back for graphing
2428            during validation, by default -1 (all values).
2429        zooming_graph : bool, optional
2430            Whether to zoom on graphs, by default True.
2431
2432
2433        Returns
2434        -------
2435        nn.Module
2436            Converted model instance configured for the tracker settings.
2437        """
2438        model = UPA.convert_network(model)
2439        self.member_vars["doing_pai"] = doing_pai
2440        self.member_vars["maximizing_score"] = maximizing_score
2441        self.save_name = save_name
2442        self.zooming_graph = zooming_graph
2443        self.making_graphs = making_graphs
2444
2445        if not self.loaded:
2446            self.member_vars["running_accuracy"] = (1.0 / num_classes) * 100
2447
2448        self.values_per_train_epoch = values_per_train_epoch
2449        self.values_per_val_epoch = values_per_val_epoch
2450
2451        if GPA.pc.get_testing_dendrite_capacity():
2452            if not GPA.pc.get_silent():
2453                print("Running a test of Dendrite Capacity.")
2454            GPA.pc.set_switch_mode(GPA.pc.DOING_SWITCH_EVERY_TIME)
2455            self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
2456            GPA.pc.set_retain_all_dendrites(True)
2457            GPA.pc.set_max_dendrite_tries(1000)
2458            GPA.pc.set_max_dendrites(1000)
2459            if GPA.pc.get_perforated_backpropagation():
2460                GPA.pc.set_initial_correlation_batches(1)
2461        else:
2462            if not GPA.pc.get_silent():
2463                print("Running Dendrite Experiment")
2464        return model
2465
2466    def generate_accuracy_plots(self, ax, save_folder, extra_string):
2467        """
2468        Generate plots and csvs for accuracy
2469
2470        Parameters
2471        ----------
2472        ax : object
2473            The matplotlib axis to plot on.
2474        save_folder : str
2475            The folder to save the plots and csvs in.
2476        extra_string : str
2477            An extra string to append to the filenames.
2478
2479        Returns
2480        -------
2481        None
2482
2483        """
2484
2485        # If scores are being saved for epochs that get overwritten, plot them
2486        for list_id in range(len(self.member_vars["overwritten_extras"])):
2487            for extra_id in self.member_vars["overwritten_extras"][list_id]:
2488                ax.plot(
2489                    np.arange(
2490                        len(self.member_vars["overwritten_extras"][list_id][extra_id])
2491                    ),
2492                    self.member_vars["overwritten_extras"][list_id][extra_id],
2493                    "r",
2494                )
2495            ax.plot(
2496                np.arange(len(self.member_vars["overwritten_vals"][list_id])),
2497                self.member_vars["overwritten_vals"][list_id],
2498                "b",
2499            )
2500
2501        # Determine which accuracy vector to use
2502        if GPA.pc.get_drawing_pai():
2503            accuracies = self.member_vars["accuracies"]
2504        else:
2505            accuracies = self.member_vars["n_accuracies"]
2506
2507        # Get pointer to additional scores being saved
2508        extra_scores = self.member_vars["extra_scores"]
2509
2510        # Plot the main accuracy scores
2511        ax.plot(np.arange(len(accuracies)), accuracies, label="Validation Scores")
2512        ax.plot(
2513            np.arange(len(self.member_vars["running_accuracies"])),
2514            self.member_vars["running_accuracies"],
2515            label="Validation Running Scores",
2516        )
2517
2518        # Plot additional scores
2519        for extra_score in extra_scores:
2520            ax.plot(
2521                np.arange(len(extra_scores[extra_score])),
2522                extra_scores[extra_score],
2523                label=extra_score,
2524            )
2525
2526        plt.title(save_folder + "/" + self.save_name + "Scores")
2527        plt.xlabel("Epochs")
2528        plt.ylabel("Score")
2529
2530        # Add point at epoch last improved and best validation score
2531        if GPA.pc.get_drawing_pai():
2532            ax.plot(
2533                self.member_vars["epoch_last_improved"],
2534                self.member_vars["global_best_validation_score"],
2535                "bo",
2536                label="Global best (y)",
2537            )
2538            ax.plot(
2539                self.member_vars["epoch_last_improved"],
2540                accuracies[self.member_vars["epoch_last_improved"]],
2541                "go",
2542                label="Epoch Last Improved",
2543            )
2544        else:
2545            if self.member_vars["mode"] == "n":
2546                missed_time = (
2547                    self.member_vars["num_epochs_run"]
2548                    - self.member_vars["epoch_last_improved"]
2549                )
2550                ax.plot(
2551                    (len(self.member_vars["n_accuracies"]) - 1) - missed_time,
2552                    self.member_vars["n_accuracies"][-(missed_time + 1)],
2553                    "go",
2554                    label="Epoch Last Improved",
2555                )
2556
2557        # Generate csv file for the values graphed
2558        pd1 = pd.DataFrame(
2559            {"Epochs": np.arange(len(accuracies)), "Validation Scores": accuracies}
2560        )
2561        pd2 = pd.DataFrame(
2562            {
2563                "Epochs": np.arange(len(self.member_vars["running_accuracies"])),
2564                "Validation Running Scores": self.member_vars["running_accuracies"],
2565            }
2566        )
2567        pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2568        for extra_score in extra_scores:
2569            pd2 = pd.DataFrame(
2570                {
2571                    "Epochs": np.arange(len(extra_scores[extra_score])),
2572                    extra_score: extra_scores[extra_score],
2573                }
2574            )
2575            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2576        extra_scores_without_graphing = self.member_vars[
2577            "extra_scores_without_graphing"
2578        ]
2579        for extra_score in extra_scores_without_graphing:
2580            pd2 = pd.DataFrame(
2581                {
2582                    "Epochs": np.arange(
2583                        len(extra_scores_without_graphing[extra_score])
2584                    ),
2585                    extra_score: extra_scores_without_graphing[extra_score],
2586                }
2587            )
2588            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2589        pd1.to_csv(
2590            save_folder + "/" + self.save_name + extra_string + "Scores.csv",
2591            index=False,
2592        )
2593        del pd1, pd2
2594
2595        # Set y min and max to zoom in on important part of axis
2596        if (
2597            len(self.member_vars["switch_epochs"]) > 0
2598            and self.member_vars["switch_epochs"][0] > 0
2599            and self.zooming_graph
2600        ):
2601            if GPA.pai_tracker.member_vars["maximizing_score"]:
2602                min_val = np.array(
2603                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2604                ).mean()
2605                for extra_score in extra_scores:
2606                    min_pot = np.array(
2607                        extra_scores[extra_score][
2608                            0 : self.member_vars["switch_epochs"][0]
2609                        ]
2610                    ).mean()
2611                    if min_pot < min_val:
2612                        min_val = min_pot
2613                ax.set_ylim(ymin=min_val)
2614            else:
2615                max_val = np.array(
2616                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2617                ).mean()
2618                for extra_score in extra_scores:
2619                    max_pot = np.array(
2620                        extra_scores[extra_score][
2621                            0 : self.member_vars["switch_epochs"][0]
2622                        ]
2623                    ).mean()
2624                    if max_pot > max_val:
2625                        max_val = max_pot
2626                ax.set_ylim(ymax=max_val)
2627
2628        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2629
2630        # Draw vertical lines for epochs where a dendrite switch occurred
2631        if GPA.pc.get_drawing_pai() and self.member_vars["doing_pai"]:
2632            color = "r"
2633            for switcher in self.member_vars["switch_epochs"]:
2634                plt.axvline(x=switcher, ymin=0, ymax=1, color=color)
2635                if color == "r":
2636                    color = "b"
2637                else:
2638                    color = "r"
2639        else:
2640            for switcher in self.member_vars["n_switch_epochs"]:
2641                plt.axvline(x=switcher, ymin=0, ymax=1, color="b")
2642
2643    def generate_time_plots(self, ax, save_folder, extra_string):
2644        """
2645        Generate plots and csvs for timing
2646
2647        Parameters
2648        ----------
2649        ax : object
2650            The matplotlib axis to plot on.
2651        save_folder : str
2652            The folder to save the plots and csvs in.
2653        extra_string : str
2654            An extra string to append to the filenames.
2655
2656        Returns
2657        -------
2658        None
2659
2660        """
2661        if self.member_vars["manual_train_switch"]:
2662            ax.plot(
2663                np.arange(len(self.member_vars["n_train_times"])),
2664                self.member_vars["n_train_times"],
2665                label="Normal Epoch Train Times",
2666            )
2667            ax.plot(
2668                np.arange(len(self.member_vars["p_train_times"])),
2669                self.member_vars["p_train_times"],
2670                label="PAI Epoch Train Times",
2671            )
2672            ax.plot(
2673                np.arange(len(self.member_vars["n_val_times"])),
2674                self.member_vars["n_val_times"],
2675                label="Normal Epoch Val Times",
2676            )
2677            ax.plot(
2678                np.arange(len(self.member_vars["p_val_times"])),
2679                self.member_vars["p_val_times"],
2680                label="PAI Epoch Val Times",
2681            )
2682
2683            plt.title(
2684                save_folder + "/" + self.save_name + "times (by train() and eval())"
2685            )
2686            plt.xlabel("Iteration")
2687            plt.ylabel("Epoch Time in Seconds ")
2688            ax.set_ylim(ymin=0)
2689            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2690
2691            pd1 = pd.DataFrame(
2692                {
2693                    "Epochs": np.arange(len(self.member_vars["n_train_times"])),
2694                    "Normal Epoch Train Times": self.member_vars["n_train_times"],
2695                }
2696            )
2697            pd2 = pd.DataFrame(
2698                {
2699                    "Epochs": np.arange(len(self.member_vars["p_train_times"])),
2700                    "PAI Epoch Train Times": self.member_vars["p_train_times"],
2701                }
2702            )
2703            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2704
2705            pd2 = pd.DataFrame(
2706                {
2707                    "Epochs": np.arange(len(self.member_vars["n_val_times"])),
2708                    "Normal Epoch Val Times": self.member_vars["n_val_times"],
2709                }
2710            )
2711            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2712
2713            pd2 = pd.DataFrame(
2714                {
2715                    "Epochs": np.arange(len(self.member_vars["p_val_times"])),
2716                    "PAI Epoch Val Times": self.member_vars["p_val_times"],
2717                }
2718            )
2719            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2720
2721            pd1.to_csv(
2722                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2723                index=False,
2724            )
2725            del pd1, pd2
2726        else:
2727            ax.plot(
2728                np.arange(len(self.member_vars["n_epoch_times"])),
2729                self.member_vars["n_epoch_times"],
2730                label="Normal Epoch Times",
2731            )
2732            ax.plot(
2733                np.arange(len(self.member_vars["p_epoch_times"])),
2734                self.member_vars["p_epoch_times"],
2735                label="PAI Epoch Times",
2736            )
2737
2738            plt.title(
2739                save_folder + "/" + self.save_name + "times (by train() and eval())"
2740            )
2741            plt.xlabel("Iteration")
2742            plt.ylabel("Epoch Time in Seconds ")
2743            ax.set_ylim(ymin=0)
2744            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2745
2746            pd1 = pd.DataFrame(
2747                {
2748                    "Epochs": np.arange(len(self.member_vars["n_epoch_times"])),
2749                    "Normal Epoch Times": self.member_vars["n_epoch_times"],
2750                }
2751            )
2752            pd2 = pd.DataFrame(
2753                {
2754                    "Epochs": np.arange(len(self.member_vars["p_epoch_times"])),
2755                    "PAI Epoch Times": self.member_vars["p_epoch_times"],
2756                }
2757            )
2758            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2759
2760            pd1.to_csv(
2761                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2762                index=False,
2763            )
2764            del pd1, pd2
2765
2766        if self.values_per_train_epoch != -1 and self.values_per_val_epoch != -1:
2767            ax2 = ax.twinx()  # Second axes sharing same x-axis
2768            ax2.set_ylabel("Single Datapoint Time in Seconds")
2769
2770            ax2.plot(
2771                np.arange(len(self.member_vars["n_train_times"])),
2772                np.array(self.member_vars["n_train_times"])
2773                / self.values_per_train_epoch,
2774                linestyle="dashed",
2775                label="Normal Train Item Times",
2776            )
2777            ax2.plot(
2778                np.arange(len(self.member_vars["p_train_times"])),
2779                np.array(self.member_vars["p_train_times"])
2780                / self.values_per_train_epoch,
2781                linestyle="dashed",
2782                label="PAI Train Item Times",
2783            )
2784            ax2.plot(
2785                np.arange(len(self.member_vars["n_val_times"])),
2786                np.array(self.member_vars["n_val_times"]) / self.values_per_val_epoch,
2787                linestyle="dashed",
2788                label="Normal Val Item Times",
2789            )
2790            ax2.plot(
2791                np.arange(len(self.member_vars["p_val_times"])),
2792                np.array(self.member_vars["p_val_times"]) / self.values_per_val_epoch,
2793                linestyle="dashed",
2794                label="PAI Val Item Times",
2795            )
2796            ax2.tick_params(axis="y")
2797            ax2.set_ylim(ymin=0)
2798            ax2.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2799
2800    def generate_learning_rate_plots(self, ax, save_folder, extra_string):
2801        """
2802        Generate plots and csvs for learning rate
2803
2804        Parameters
2805        ----------
2806        ax : object
2807            The matplotlib axis to plot on.
2808        save_folder : str
2809            The folder to save the plots and csvs in.
2810        extra_string : str
2811            An extra string to append to the filenames.
2812
2813        Returns
2814        -------
2815        None
2816
2817        """
2818        ax.plot(
2819            np.arange(len(self.member_vars["training_learning_rates"])),
2820            self.member_vars["training_learning_rates"],
2821            label="learning_rate",
2822        )
2823        plt.title(save_folder + "/" + self.save_name + "learning_rate")
2824        plt.xlabel("Epochs")
2825        plt.ylabel("learning_rate")
2826        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2827
2828        pd1 = pd.DataFrame(
2829            {
2830                "Epochs": np.arange(len(self.member_vars["training_learning_rates"])),
2831                "learning_rate": self.member_vars["training_learning_rates"],
2832            }
2833        )
2834        pd1.to_csv(
2835            save_folder + "/" + self.save_name + extra_string + "learning_rate.csv",
2836            index=False,
2837        )
2838        del pd1
2839
2840    def get_current_pb_scores(self):
2841        """
2842        Get the latest best PBScore of each dendrite layer, the same numbers
2843        written to the Best PBScores csv.
2844
2845        Returns
2846        -------
2847        dict[str, Any]
2848            Layer name to score.  Empty outside of dendrite scoring phases,
2849            when no candidate dendrites are being scored.
2850
2851
2852        Parameters
2853        ----------
2854        None
2855
2856        """
2857        if not self.member_vars["doing_pai"]:
2858            return {}
2859        if not GPA.pc.get_perforated_backpropagation():
2860            return {}
2861        # Scores only advance while candidate dendrites are being trained
2862        if (
2863            self.member_vars["mode"] != "p"
2864            and not GPA.pc.get_learn_dendrites_live()
2865        ):
2866            return {}
2867
2868        scores = {}
2869        for layer_id in range(len(self.neuron_module_vector)):
2870            if layer_id >= len(self.member_vars["best_scores"]):
2871                continue
2872            layer_scores = self.member_vars["best_scores"][layer_id]
2873            if len(layer_scores) == 0:
2874                continue
2875            score = layer_scores[-1]
2876            if hasattr(score, "item"):
2877                score = score.item()
2878            score = float(score)
2879            if math.isnan(score) or math.isinf(score):
2880                continue
2881            scores[self.neuron_module_vector[layer_id].name] = score
2882        return scores
2883
2884    def generate_dendrite_learning_plots(self, ax, save_folder, extra_string):
2885        """
2886        Generate dendrite score plots for the tracker.
2887        Also saves csv files associated with the plots.
2888
2889        Parameters
2890        ----------
2891        ax : matplotlib.axes.Axes
2892            Axis used for plotting dendrite-learning curves.
2893        save_folder : str
2894            Directory where plot images and CSV summaries are written.
2895        extra_string : str
2896            Filename suffix used to distinguish this output set.
2897
2898        Returns
2899        -------
2900        None
2901            Saves plots and score CSV files to disk.
2902        """
2903        if self.member_vars["doing_pai"]:
2904            pd1 = None
2905            pd2 = None
2906            num_colors = len(self.neuron_module_vector)
2907
2908            cm = plt.get_cmap("gist_rainbow")
2909            layer_colors = [cm(1.0 * i / max(num_colors, 1)) for i in range(num_colors)]
2910
2911            for layer_id in range(len(self.neuron_module_vector)):
2912                color = layer_colors[layer_id]
2913                ax.plot(
2914                    np.arange(len(self.member_vars["best_scores"][layer_id])),
2915                    self.member_vars["best_scores"][layer_id],
2916                    label=self.neuron_module_vector[layer_id].name,
2917                    color=color,
2918                )
2919
2920                pd2 = pd.DataFrame(
2921                    {
2922                        "Epochs": np.arange(
2923                            len(self.member_vars["best_scores"][layer_id])
2924                        ),
2925                        f"Best ever for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2926                            "best_scores"
2927                        ][
2928                            layer_id
2929                        ],
2930                    }
2931                )
2932
2933                if pd1 is None:
2934                    pd1 = pd2
2935                else:
2936                    pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2937
2938                if len(self.member_vars["current_scores"][layer_id]) != 0:
2939                    ax.plot(
2940                        np.arange(len(self.member_vars["current_scores"][layer_id])),
2941                        self.member_vars["current_scores"][layer_id],
2942                        label=f"Current:{self.neuron_module_vector[layer_id].name}",
2943                        color=color,
2944                        linestyle="--",
2945                    )
2946
2947                pd2 = pd.DataFrame(
2948                    {
2949                        "Epochs": np.arange(
2950                            len(self.member_vars["current_scores"][layer_id])
2951                        ),
2952                        f"Best current for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2953                            "current_scores"
2954                        ][
2955                            layer_id
2956                        ],
2957                    }
2958                )
2959                pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2960
2961            plt.title(save_folder + "/" + self.save_name + " Best PBScores")
2962            plt.xlabel("Epochs")
2963            plt.ylabel("Best PBScore")
2964            ax.legend(
2965                bbox_to_anchor=(1.05, 1),
2966                loc="upper left",
2967                ncol=max(1, math.ceil(len(self.neuron_module_vector) / 30)),
2968            )
2969            for switcher in self.member_vars["p_switch_epochs"]:
2970                plt.axvline(x=switcher, ymin=0, ymax=1, color="r")
2971
2972            if self.member_vars["mode"] == "p":
2973                missed_time = (
2974                    self.member_vars["num_epochs_run"]
2975                    - self.member_vars["epoch_last_improved"]
2976                )
2977                plt.axvline(
2978                    x=(len(self.member_vars["best_scores"][0]) - (missed_time + 1)),
2979                    ymin=0,
2980                    ymax=1,
2981                    color="g",
2982                )
2983
2984            # pd1 here will be none if no PB layers are created
2985            if pd1 is not None:
2986                pd1.to_csv(
2987                    save_folder
2988                    + "/"
2989                    + self.save_name
2990                    + extra_string
2991                    + "Best PBScores.csv",
2992                    index=False,
2993                )
2994            del pd1, pd2
2995
2996    def generate_extra_csv_files(self, save_folder, extra_string):
2997        """
2998        Generate additional csvs
2999
3000        Parameters
3001        ----------
3002        save_folder : str
3003            The folder to save the plots and csvs in.
3004        extra_string : str
3005            An extra string to append to the filenames.
3006
3007        Returns
3008        -------
3009        None
3010
3011        """
3012        pd1 = pd.DataFrame(
3013            {
3014                "Switch Number": np.arange(len(self.member_vars["switch_epochs"])),
3015                "Switch Epoch": self.member_vars["switch_epochs"],
3016            }
3017        )
3018        pd1.to_csv(
3019            save_folder + "/" + self.save_name + extra_string + "switch_epochs.csv",
3020            index=False,
3021        )
3022        del pd1
3023
3024        pd1 = pd.DataFrame(
3025            {
3026                "Switch Number": np.arange(len(self.member_vars["param_counts"])),
3027                "Param Count": self.member_vars["param_counts"],
3028            }
3029        )
3030        pd1.to_csv(
3031            save_folder + "/" + self.save_name + extra_string + "param_counts.csv",
3032            index=False,
3033        )
3034        del pd1
3035
3036        """
3037        Create best_arch_scores.csv file
3038        When working with dendrites there is a tradeoff between additional param count and score improvement.
3039        This file will help track that tradeoff by recording the best scores for all extra_scores
3040        and extra_scores_without_graphing for each architecture version.
3041        The scores recorded here are from the epoch when the best validation score was found
3042        within each switch_epoch boundary.
3043        """
3044        switch_counts = len(self.member_vars["switch_epochs"])
3045        best_valid = []
3046        associated_params = []
3047        
3048        # Initialize dictionaries to store best scores for each extra score type
3049        best_extra_scores = {}
3050        for score_name in self.member_vars["extra_scores"]:
3051            best_extra_scores[score_name] = []
3052        for score_name in self.member_vars["extra_scores_without_graphing"]:
3053            best_extra_scores[score_name] = []
3054
3055        for switch in range(0, switch_counts, 2):
3056            start_index = 0
3057            if switch != 0:
3058                start_index = self.member_vars["switch_epochs"][switch - 1] + 1
3059            end_index = self.member_vars["switch_epochs"][switch] + 1
3060
3061            if GPA.pai_tracker.member_vars["maximizing_score"]:
3062                best_valid_index = start_index + np.argmax(
3063                    self.member_vars["accuracies"][start_index:end_index]
3064                )
3065            else:
3066                best_valid_index = start_index + np.argmin(
3067                    self.member_vars["accuracies"][start_index:end_index]
3068                )
3069
3070            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3071            best_valid.append(best_valid_score)
3072            
3073            # Get corresponding scores from all extra_scores
3074            for score_name in self.member_vars["extra_scores"]:
3075                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3076                    best_extra_scores[score_name].append(
3077                        self.member_vars["extra_scores"][score_name][best_valid_index]
3078                    )
3079                else:
3080                    best_extra_scores[score_name].append(None)
3081            
3082            # Get corresponding scores from all extra_scores_without_graphing
3083            for score_name in self.member_vars["extra_scores_without_graphing"]:
3084                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3085                    best_extra_scores[score_name].append(
3086                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3087                    )
3088                else:
3089                    best_extra_scores[score_name].append(None)
3090            
3091            if self.member_vars["doing_pai"]:
3092                associated_params.append(self.member_vars["param_counts"][switch])
3093            else:
3094                associated_params.append(self.member_vars["param_counts"][-1])
3095
3096        # If in neuron training mode but not the very first epoch
3097        if self.member_vars["mode"] == "n" and (
3098            (len(self.member_vars["switch_epochs"]) == 0)
3099            or (
3100                self.member_vars["switch_epochs"][-1] + 1
3101                != len(self.member_vars["accuracies"])
3102            )
3103        ):
3104            start_index = 0
3105            if len(self.member_vars["switch_epochs"]) != 0:
3106                start_index = self.member_vars["switch_epochs"][-1] + 1
3107
3108            if GPA.pai_tracker.member_vars["maximizing_score"]:
3109                best_valid_index = start_index + np.argmax(
3110                    self.member_vars["accuracies"][start_index:]
3111                )
3112            else:
3113                best_valid_index = start_index + np.argmin(
3114                    self.member_vars["accuracies"][start_index:]
3115                )
3116
3117            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3118            best_valid.append(best_valid_score)
3119            
3120            # Get corresponding scores from all extra_scores
3121            for score_name in self.member_vars["extra_scores"]:
3122                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3123                    best_extra_scores[score_name].append(
3124                        self.member_vars["extra_scores"][score_name][best_valid_index]
3125                    )
3126                else:
3127                    best_extra_scores[score_name].append(None)
3128            
3129            # Get corresponding scores from all extra_scores_without_graphing
3130            for score_name in self.member_vars["extra_scores_without_graphing"]:
3131                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3132                    best_extra_scores[score_name].append(
3133                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3134                    )
3135                else:
3136                    best_extra_scores[score_name].append(None)
3137            
3138            associated_params.append(self.member_vars["param_counts"][-1])
3139
3140        # Build dataframe with all columns
3141        csv_data = {
3142            "Param Counts": associated_params,
3143            "Max Valid Scores": best_valid,
3144        }
3145        
3146        # Add columns for each extra score
3147        for score_name in best_extra_scores:
3148            csv_data[score_name] = best_extra_scores[score_name]
3149        
3150        pd1 = pd.DataFrame(csv_data)
3151        pd1.to_csv(
3152            save_folder + "/" + self.save_name + extra_string + "_best_arch_scores.csv",
3153            index=False,
3154        )
3155        del pd1
3156
3157    def save_graphs(self, extra_string=""):
3158        """
3159        Save graphs and csvs for all the values the tracker records
3160
3161        Parameters
3162        ----------
3163        extra_string : str
3164            An extra string to append to the filenames.
3165
3166        Returns
3167        -------
3168        None
3169
3170        """
3171        # If running DDP only save with rank 0
3172        if "RANK" in os.environ:
3173            if int(os.environ["RANK"]) != 0:
3174                return
3175        if not self.making_graphs:
3176            return
3177
3178        save_folder = "./" + self.save_name + "/"
3179
3180        plt.ioff()
3181        fig = plt.figure(figsize=(28, 14))
3182
3183        # Plot with accuracy scores
3184        ax = plt.subplot(221)
3185        self.generate_accuracy_plots(ax, save_folder, extra_string)
3186
3187        # Plot dendrite learning scores
3188        ax = plt.subplot(222)
3189        self.generate_dendrite_learning_plots(ax, save_folder, extra_string)
3190
3191        if GPA.pc.get_drawing_extra_graphs():
3192            # Plot learning rates for each training epoch
3193            ax = plt.subplot(223)
3194            self.generate_learning_rate_plots(ax, save_folder, extra_string)
3195
3196            # Plot the times for each training epoch
3197            ax = plt.subplot(224)
3198            self.generate_time_plots(ax, save_folder, extra_string)
3199
3200        # Generate extra CSV files
3201        self.generate_extra_csv_files(save_folder, extra_string)
3202
3203        fig.tight_layout()
3204        plt.savefig(save_folder + "/" + self.save_name + extra_string + ".png")
3205        plt.close("all")
3206
3207    def add_loss(self, loss):
3208        """Add loss to tracking vectors.
3209
3210        Parameters
3211        ----------
3212        loss : float or int
3213            The loss value to add.
3214
3215        Returns
3216        -------
3217        None
3218
3219        """
3220        if not isinstance(loss, (float, int)):
3221            loss = loss.item()
3222        self.member_vars["training_loss"].append(loss)
3223
3224    def add_learning_rate(self, learning_rate):
3225        """Add learning rate to tracking vectors.
3226
3227        Parameters
3228        ----------
3229        learning_rate : float or int
3230            The learning rate value to add.
3231
3232        Returns
3233        -------
3234        None
3235
3236        """
3237        if not isinstance(learning_rate, (float, int)):
3238            learning_rate = learning_rate.item()
3239        self.member_vars["training_learning_rates"].append(learning_rate)
3240
3241    def add_extra_score(self, score, extra_score_name):
3242        """Add extra score to tracking vectors.
3243
3244        Parameters
3245        ----------
3246        score : float or int
3247            The score value to add.
3248
3249        extra_score_name : str
3250            The name of the extra score.
3251
3252        Returns
3253        -------
3254        None
3255
3256        """
3257        if not isinstance(score, (float, int)):
3258            try:
3259                score = score.item()
3260            except:
3261                print(
3262                    "Scores added for Perforated Backpropagation should be "
3263                    "float, int, or tensor, yours is a:"
3264                )
3265                print(type(score))
3266                pdb.set_trace()
3267
3268        if GPA.pc.get_verbose():
3269            print(f"Adding extra score {extra_score_name} of {float(score)}")
3270
3271        if extra_score_name not in self.member_vars["extra_scores"]:
3272            self.member_vars["extra_scores"][extra_score_name] = []
3273        self.member_vars["extra_scores"][extra_score_name].append(score)
3274
3275        if self.member_vars["mode"] == "n":
3276            if extra_score_name not in self.member_vars["n_extra_scores"]:
3277                self.member_vars["n_extra_scores"][extra_score_name] = []
3278            self.member_vars["n_extra_scores"][extra_score_name].append(score)
3279
3280    def add_extra_score_without_graphing(self, score, extra_score_name):
3281        """Add extra score without graphing to tracking vectors.
3282
3283        Parameters
3284        ----------
3285        score : float or int
3286            The score value to add.
3287
3288        extra_score_name : str
3289            The name of the extra score.
3290
3291        Returns
3292        -------
3293        None
3294
3295        """
3296        if not isinstance(score, (float, int)):
3297            try:
3298                score = score.item()
3299            except:
3300                print(
3301                    "Scores added for Perforated Backpropagation should be "
3302                    "float, int, or tensor, yours is a:"
3303                )
3304                print(type(score))
3305                print("in add_extra_score_without_graphing")
3306                pdb.set_trace()
3307
3308        if GPA.pc.get_verbose():
3309            print(f"Adding extra score {extra_score_name} of {float(score)}")
3310
3311        if extra_score_name not in self.member_vars["extra_scores_without_graphing"]:
3312            self.member_vars["extra_scores_without_graphing"][extra_score_name] = []
3313        self.member_vars["extra_scores_without_graphing"][extra_score_name].append(
3314            score
3315        )
3316
3317    def add_test_score(self, score, extra_score_name):
3318        """Add test score to tracking vectors.
3319
3320        Parameters
3321        ----------
3322        score : float or int
3323            The score value to add.
3324
3325        extra_score_name : str
3326            The name of the extra score.
3327
3328        Returns
3329        -------
3330        None
3331
3332        Notes
3333        -----
3334        This function is a wrapper around `add_extra_score` that separates
3335        test score for adding to best_arch_scores.csv.
3336
3337        """
3338        self.add_extra_score(score, extra_score_name)
3339
3340        if not isinstance(score, (float, int)):
3341            try:
3342                score = score.item()
3343            except:
3344                print(
3345                    "Scores added for Perforated Backpropagation should be "
3346                    "float, int, or tensor, yours is a:"
3347                )
3348                print(type(score))
3349                print("in add_test_score")
3350                pdb.set_trace()
3351
3352        if GPA.pc.get_verbose():
3353            print(f"Adding test score {extra_score_name} of {float(score)}")
3354        self.member_vars["test_scores"].append(score)
3355
3356    def add_validation_score(self, accuracy, net, force_switch=False):
3357        """Function to add the validation score.
3358
3359        This is complex because it determines neuron and dendrite switching.
3360
3361        Parameters
3362        ----------
3363        accuracy : float or int
3364            The accuracy or loss value to add.
3365        net : object
3366            The neural network model.
3367        force_switch : bool, optional
3368            Whether to force a switch, by default False.
3369
3370        Returns
3371        -------
3372        net : object
3373            The potentially modified neural network model.
3374        training_complete : bool
3375            Whether training is complete.
3376        restructured : bool
3377            Whether the model has been restructured.
3378
3379        Notes
3380        -----
3381        WARNING: Do not call self anywhere in this function. When systems
3382        get loaded the actual tracker you are working with can change.
3383        """
3384
3385        _pai_log("info", f"Adding validation score {accuracy:.8f}")
3386
3387        update_learning_rate()
3388        update_param_count(net)
3389
3390        accuracy = check_input_problems(net, accuracy)
3391
3392        if len(GPA.pai_tracker.member_vars["switch_epochs"]) == 0:
3393            epochs_since_cycle_switch = GPA.pai_tracker.member_vars["num_epochs_run"]
3394        else:
3395            epochs_since_cycle_switch = (
3396                GPA.pai_tracker.member_vars["num_epochs_run"]
3397                - GPA.pai_tracker.member_vars["switch_epochs"][-1]
3398            )
3399
3400        update_running_accuracy(accuracy, epochs_since_cycle_switch)
3401        if GPA.pc.get_perforated_backpropagation():
3402            TPB.update_pb_scores(self)
3403
3404        # Captured before any switch below flips the mode and reloads scores
3405        epoch_pb_scores = self.get_current_pb_scores()
3406
3407        GPA.pai_tracker.stop_epoch(internal_call=True)
3408
3409        # If it is neuron training mode
3410        if (
3411            GPA.pai_tracker.member_vars["mode"] == "n"
3412            or GPA.pc.get_learn_dendrites_live()
3413        ):
3414            check_new_best(net, accuracy, epochs_since_cycle_switch)
3415        elif GPA.pc.get_perforated_backpropagation():
3416            TPB.check_best_pai_score_improvement()
3417
3418        # Save the latest model
3419        if GPA.pc.get_test_saves():
3420            UPA.save_system(net, GPA.pc.get_save_name(), "latest")
3421        if GPA.pc.get_pai_saves():
3422            UPA.pai_save_system(net, GPA.pc.get_save_name(), "latest")
3423
3424        restructuring_status_value = NO_MODEL_UPDATE
3425        # If it is time to switch based on scores and counter or a manual switch
3426        if GPA.pai_tracker.switch_time() or force_switch:
3427            # If testing dendrite capacity switch after enough dendrites added
3428            if (
3429                (GPA.pai_tracker.member_vars["mode"] == "n")
3430                and (GPA.pai_tracker.member_vars["num_dendrites_added"] > 2)
3431                and GPA.pc.get_testing_dendrite_capacity()
3432            ):
3433                GPA.pai_tracker.save_graphs()
3434                _pai_log(
3435                    "info",
3436                    "Successfully added 3 dendrites with GPA.pc.set_testing_dendrite_capacity(True) (default). "
3437                    "You may now set that to False and run a real experiment.",
3438                )
3439                return net, False, True
3440
3441            # If doing neuron training but this dendrite count didn't improve
3442            if (
3443                (GPA.pai_tracker.member_vars["mode"] == "n")
3444                or GPA.pc.get_learn_dendrites_live()
3445            ) and (GPA.pai_tracker.member_vars["current_n_set_global_best"] is False):
3446                new_restructuring_status_value, net = process_no_improvement(net)
3447                # if this was the final try return that training is complete
3448                if new_restructuring_status_value == TRAINING_COMPLETE:
3449                    if _dashboard_emitter is not None:
3450                        _dashboard_emitter.emit_run_end(GPA.pc)
3451                    return net, True, True
3452                else:
3453                    restructuring_status_value = update_restructuring_status(
3454                        restructuring_status_value, new_restructuring_status_value
3455                    )
3456            # Else if did improve, do a normal switch process
3457            else:
3458                if GPA.pc.get_verbose():
3459                    print(
3460                        f"Calling switch_mode with "
3461                        f'{GPA.pai_tracker.member_vars["current_n_set_global_best"]}, '
3462                        f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}, '
3463                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]}, '
3464                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_value"]},'
3465                        f'{GPA.pc.get_max_dendrites()},'
3466                        f'{GPA.pai_tracker.member_vars["num_dendrites_added"]},'
3467                        f'{GPA.pai_tracker.member_vars["num_dendrite_tries"]},'
3468                    )
3469                import pdb; pdb.set_trace
3470                # If the max number of dendrites has been hit or not doing pai and adding dendtites
3471                # then return rather than adding more
3472                if (
3473                    (GPA.pai_tracker.member_vars["mode"] == "n")
3474                    and (
3475                        GPA.pc.get_max_dendrites()
3476                        == GPA.pai_tracker.member_vars["num_dendrites_added"]
3477                    )
3478                ) or (GPA.pai_tracker.member_vars["doing_pai"] is False):
3479                    if GPA.pc.get_verbose():
3480                        print(
3481                            "Max dendrites reached or not doing PAI, finishing training"
3482                        )
3483                    net = process_final_network(net)
3484                    # Increment integrated if we have dendrites (means they're integrated)
3485                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3486                        GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3487                        _pai_log("info", f"Final dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3488                        if _dashboard_emitter is not None:
3489                            _dashboard_emitter.emit_dendrite_added(
3490                                GPA.pc,
3491                                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3492                                num_dendrites_integrated=GPA.pai_tracker.member_vars[
3493                                    "num_dendrites_integrated"
3494                                ],
3495                            )
3496                    if _dashboard_emitter is not None:
3497                        _dashboard_emitter.emit_run_end(GPA.pc)
3498                    return net, True, True
3499
3500                # Otherwise if its neuron training mode reset the counter of failed dendrites
3501                # Check if we should increment integrated count BEFORE change_learning_modes loads old state
3502                should_increment_integrated = False
3503                if GPA.pai_tracker.member_vars["mode"] == "n":
3504                    GPA.pai_tracker.member_vars["num_dendrite_tries"] = 0
3505                    if GPA.pc.get_verbose():
3506                        print(
3507                            "Adding new dendrites without resetting which means "
3508                            "the last ones improved. Resetting num_dendrite_tries"
3509                        )
3510                    # Remember to increment after change_learning_modes (which loads old tracker state)
3511                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3512                        should_increment_integrated = True
3513
3514                GPA.pai_tracker.save_graphs(
3515                    f'_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}'
3516                )
3517
3518                if GPA.pc.get_test_saves():
3519                    UPA.save_system(
3520                        net,
3521                        GPA.pc.get_save_name(),
3522                        f'beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3523                    )
3524                    # Copy current best model from this set of dendrites
3525                    # If running DDP only copy with rank 0
3526                    if "RANK" not in os.environ or int(os.environ["RANK"]) == 0:
3527                        shutil.copyfile(
3528                            f"{GPA.pc.get_save_name()}/best_model.pt",
3529                            f'{GPA.pc.get_save_name()}/best_model_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}.pt',
3530                        )
3531
3532                net = UPA.change_learning_modes(
3533                    net,
3534                    GPA.pc.get_save_name(),
3535                    "best_model",
3536                    GPA.pai_tracker.member_vars["doing_pai"],
3537                )
3538                restructuring_status_value = NETWORK_RESTRUCTURED
3539                
3540                # Now increment after change_learning_modes has loaded the best model
3541                # This ensures the increment persists and doesn't get overwritten
3542                if should_increment_integrated:
3543                    GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3544                    _pai_log("info", f"Dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3545                    if _dashboard_emitter is not None:
3546                        _dashboard_emitter.emit_dendrite_added(
3547                            GPA.pc,
3548                            epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3549                            num_dendrites_integrated=GPA.pai_tracker.member_vars[
3550                                "num_dendrites_integrated"
3551                            ],
3552                        )
3553
3554            # If restructured is true, clear scheduler/optimizer before saving
3555            if restructuring_status_value != NETWORK_RESTRUCTURED:
3556                print(
3557                    "Restructured should always be triggered here, let us know if you encounter this situation"
3558                )
3559                pdb.set_trace()
3560
3561            # Since there is a restructuring optimizer and scheduler must be reinitialized after return
3562            GPA.pai_tracker.clear_optimizer_and_scheduler()
3563
3564            # Save the model from after the switch
3565            UPA.save_system(
3566                net,
3567                GPA.pc.get_save_name(),
3568                f'switch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3569            )
3570
3571        # If not time to switch and you have a scheduler, perform the update step
3572        elif GPA.pai_tracker.member_vars["scheduler"] is not None:
3573            new_restructuring_status_value, net = process_scheduler_update(
3574                net, accuracy, epochs_since_cycle_switch
3575            )
3576            restructuring_status_value = update_restructuring_status(
3577                restructuring_status_value, new_restructuring_status_value
3578            )
3579
3580        GPA.pai_tracker.start_epoch(internal_call=True)
3581        if _dashboard_emitter is not None:
3582            _mv = GPA.pai_tracker.member_vars
3583            _lr = _mv["training_learning_rates"][-1] if _mv["training_learning_rates"] else None
3584            _train_score = _mv["extra_scores"].get("train", [None])[-1]
3585            _n_times = _mv["n_epoch_times"] or [(_mv["n_train_times"][-1] + _mv["n_val_times"][-1]) if (_mv["n_train_times"] and _mv["n_val_times"]) else None]
3586            _p_times = _mv["p_epoch_times"] or [(_mv["p_train_times"][-1] + _mv["p_val_times"][-1]) if (_mv["p_train_times"] and _mv["p_val_times"]) else None]
3587            _dashboard_emitter.emit_epoch(
3588                GPA.pc,
3589                epoch=_mv["num_epochs_run"],
3590                validation_score=accuracy,
3591                learning_rate=_lr,
3592                train_score=_train_score,
3593                normal_time=_n_times[-1],
3594                pai_time=_p_times[-1],
3595                pb_scores=epoch_pb_scores,
3596            )
3597        GPA.pai_tracker.save_graphs()
3598
3599        if restructuring_status_value == NETWORK_RESTRUCTURED:
3600            GPA.pai_tracker.member_vars["epoch_last_improved"] = (
3601                GPA.pai_tracker.member_vars["num_epochs_run"]
3602            )
3603            if GPA.pc.get_verbose():
3604                print(
3605                    f"Setting epoch last improved to "
3606                    f'{GPA.pai_tracker.member_vars["epoch_last_improved"]}'
3607                )
3608
3609            now = datetime.now()
3610            dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
3611
3612            if GPA.pc.get_verbose():
3613                print("Not saving restructure right now")
3614
3615            """
3616            This block of code helped with a save issue with safetensors and huggingface, but it breaks DDP.  
3617            Temporarily removing it to avoid DDP issues, but if you encounter save issues try adding it back in.
3618            for param in net.parameters():
3619                param.data = param.data.contiguous()
3620            """
3621        if GPA.pc.get_verbose():
3622            print(
3623                f"Completed adding score. Restructured is {restructuring_status_value}, "
3624                f"\ncurrent switch list is:"
3625            )
3626            print(GPA.pai_tracker.member_vars["switch_epochs"])
3627
3628        if _dashboard_emitter is not None and restructuring_status_value == NETWORK_RESTRUCTURED:
3629            _param_count = UPA.count_params(net)
3630            _dashboard_emitter.emit_switch(
3631                GPA.pc,
3632                switch_number=GPA.pai_tracker.member_vars["num_dendrites_added"],
3633                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3634                param_count=_param_count,
3635                switch_type=GPA.pai_tracker.member_vars["mode"],
3636            )
3637
3638        # Always False for training complete if nothing triggered that training is over
3639        return net, restructuring_status_value, False
3640
3641    def clear_all_processors(self):
3642        """Clear all processors from modules.
3643
3644        Parameters
3645        ----------
3646        None
3647
3648        Returns
3649        -------
3650        None
3651            This function does not return a value.
3652        """
3653        for module in self.neuron_module_vector:
3654            module.clear_processors()
3655
3656    def create_new_dendrite_module(self):
3657        """Add dendrite module to all neuron modules.
3658
3659        Parameters
3660        ----------
3661        None
3662
3663        Returns
3664        -------
3665        None
3666            This function does not return a value.
3667        """
3668        for module in self.neuron_module_vector:
3669            module.create_new_dendrite_module()
3670
3671    def set_create_dendrite_global(self, fn):
3672        """Call set_create_dendrite(fn) on every tracked PAINeuronModule."""
3673        for module in self.neuron_module_vector:
3674            module.set_create_dendrite(fn)
3675
3676    def set_dendrite_loss_fn_global(self, fn):
3677        """Set the global dendrite loss function used by all dendrite modules."""
3678        from perforatedbp import modules_pbp as MPB
3679        MPB.dendrite_loss_fn = fn
3680
3681    def apply_pb_grads(self):
3682        """Apply perforated backpropagation gradients to all modules.
3683
3684        Parameters
3685        ----------
3686        None
3687
3688        Returns
3689        -------
3690        None
3691            This function does not return a value.
3692        """
3693        if self.member_vars["mode"] == "p":
3694            for module in self.neuron_module_vector:
3695                module.apply_pb_grads()
3696
3697    def apply_pb_zero(self):
3698        """Apply perforated backpropagation zero gradients to all modules.
3699
3700        Parameters
3701        ----------
3702        None
3703
3704        Returns
3705        -------
3706        None
3707            This function does not return a value.
3708        """
3709        if self.member_vars["mode"] == "p":
3710            for module in self.neuron_module_vector:
3711                module.apply_pb_zero()

Manager class that tracks all neuron layers and dendrite layers, controls when new dendrites are added, and communicates signals to modules.

PAINeuronModuleTracker( doing_pai, save_name, making_graphs=True, param_vals_setting=-1, values_per_train_epoch=-1, values_per_val_epoch=-1)
 956    def __init__(
 957        self,
 958        doing_pai,
 959        save_name,
 960        making_graphs=True,
 961        param_vals_setting=-1,
 962        values_per_train_epoch=-1,
 963        values_per_val_epoch=-1,
 964    ):
 965        """Initialize the tracker
 966
 967        Parameters
 968        ----------
 969        doing_pai : bool
 970            Whether or not dendrites should be used.
 971        save_name : str
 972            The base name for saving models and graphs.
 973        making_graphs : bool, optional
 974            Whether or not to generate graphs, by default True.
 975        param_vals_setting : int, optional
 976            Parameter values setting, by default -1.
 977        values_per_train_epoch : int, optional
 978            The number of values to look back for graphing
 979            during training, by default -1 (all values).
 980        values_per_val_epoch : int, optional
 981            The number of values to look back for graphing
 982            during validation, by default -1 (all values).
 983        Returns
 984        -------
 985        None
 986        """
 987
 988        # Dict of member vars and their types for saving
 989        self.member_vars = {}
 990        self.member_var_types = {}
 991
 992        # Whether or not PAI will be running
 993        self.member_vars["doing_pai"] = doing_pai
 994        self.member_var_types["doing_pai"] = "bool"
 995
 996        # How many Dendrites have been added
 997        self.member_vars["num_dendrites_added"] = 0
 998        self.member_var_types["num_dendrites_added"] = "int"
 999
1000        # How many Dendrites have been successfully integrated, does not count currently training dendrites
1001        self.member_vars["num_dendrites_integrated"] = 0
1002        self.member_var_types["num_dendrites_integrated"] = "int"
1003
1004        # How many cycles have been run, *2 or *2+1 of the above
1005        self.member_vars["num_cycles"] = 0
1006        self.member_var_types["num_cycles"] = "int"
1007
1008        # Pointers to all neuron wrapped modules
1009        self.neuron_module_vector = []
1010
1011        # Pointers to all non neuron modules for tracking
1012        self.tracked_neuron_module_vector = []
1013
1014        # Neuron training or dendrite training mode
1015        self.member_vars["mode"] = "n"
1016        self.member_var_types["mode"] = "string"
1017
1018        # Number of epochs run excluding overwritten epochs
1019        self.member_vars["num_epochs_run"] = -1
1020        self.member_var_types["num_epochs_run"] = "int"
1021
1022        # Number including overwritten epochs
1023        self.member_vars["total_epochs_run"] = -1
1024        self.member_var_types["total_epochs_run"] = "int"
1025
1026        # Last epoch that validation/correlation score was improved
1027        self.member_vars["epoch_last_improved"] = 0
1028        self.member_var_types["epoch_last_improved"] = "int"
1029
1030        # Running validation accuracy
1031        self.member_vars["running_accuracy"] = 0
1032        self.member_var_types["running_accuracy"] = "float"
1033
1034        # True if maxing validation, False if minimizing Loss
1035        self.member_vars["maximizing_score"] = True
1036        self.member_var_types["maximizing_score"] = "bool"
1037
1038        # Mode for switching back and forth between learning modes
1039        self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
1040        self.member_var_types["switch_mode"] = "int"
1041
1042        # Epoch of the last switch
1043        self.member_vars["last_switch"] = 0
1044        self.member_var_types["last_switch"] = "int"
1045
1046        # Highest validation score from current cycle
1047        self.member_vars["current_best_validation_score"] = 0
1048        self.member_var_types["current_best_validation_score"] = "float"
1049
1050        # Last epoch where the learning rate was updated
1051        self.member_vars["initial_lr_test_epoch_count"] = -1
1052        self.member_var_types["initial_lr_test_epoch_count"] = "int"
1053
1054        # Highest validation score of full run
1055        self.member_vars["global_best_validation_score"] = 0
1056        self.member_var_types["global_best_validation_score"] = "float"
1057
1058        # List of switch epochs
1059        self.member_vars["switch_epochs"] = []
1060        self.member_var_types["switch_epochs"] = "int array"
1061
1062        # Parameter counts at each network structure
1063        self.member_vars["param_counts"] = []
1064        self.member_var_types["param_counts"] = "int array"
1065
1066        # List of epochs where switch was made to neuron training
1067        self.member_vars["n_switch_epochs"] = []
1068        self.member_var_types["n_switch_epochs"] = "int array"
1069
1070        # List of epochs where switch was made to dendrite training
1071        self.member_vars["p_switch_epochs"] = []
1072        self.member_var_types["p_switch_epochs"] = "int array"
1073
1074        # List of validation accuracies
1075        self.member_vars["accuracies"] = []
1076        self.member_var_types["accuracies"] = "float array"
1077
1078        # List of epochs where score improved for scheduler updates
1079        self.member_vars["last_improved_accuracies"] = []
1080        self.member_var_types["last_improved_accuracies"] = "int array"
1081
1082        # List of test accuracy scores registered
1083        self.member_vars["test_accuracies"] = []
1084        self.member_var_types["test_accuracies"] = "float array"
1085
1086        # List of accuracies registered during neuron training
1087        self.member_vars["n_accuracies"] = []
1088        self.member_var_types["n_accuracies"] = "float array"
1089
1090        # List of accuracies registered during dendrite training
1091        self.member_vars["p_accuracies"] = []
1092        self.member_var_types["p_accuracies"] = "float array"
1093
1094        # Running average accuracies from recent epochs
1095        self.member_vars["running_accuracies"] = []
1096        self.member_var_types["running_accuracies"] = "float array"
1097
1098        # List of additional scores recorded
1099        self.member_vars["extra_scores"] = {}
1100        self.member_var_types["extra_scores"] = "float array dictionary"
1101
1102        # Extra scores not set to be graphed
1103        self.member_vars["extra_scores_without_graphing"] = {}
1104        self.member_var_types["extra_scores_without_graphing"] = (
1105            "float array dictionary"
1106        )
1107
1108        # List of test scores
1109        self.member_vars["test_scores"] = []
1110        self.member_var_types["test_scores"] = "float array"
1111
1112        # Extra scores calculated during neuron training
1113        self.member_vars["n_extra_scores"] = {}
1114        self.member_var_types["n_extra_scores"] = "float array dictionary"
1115
1116        # List of training losses calculated
1117        self.member_vars["training_loss"] = []
1118        self.member_var_types["training_loss"] = "float array"
1119
1120        # List of learning rates at each epoch
1121        self.member_vars["training_learning_rates"] = []
1122        self.member_var_types["training_learning_rates"] = "float array"
1123
1124        # Best dendrite scores
1125        self.member_vars["best_scores"] = []
1126        self.member_var_types["best_scores"] = "float array array"
1127
1128        # Current dendrite scores
1129        self.member_vars["current_scores"] = []
1130        self.member_var_types["current_scores"] = "float array array"
1131
1132        # Times for neuron training epochs
1133        self.member_vars["n_epoch_times"] = []
1134        self.member_var_types["n_epoch_times"] = "float array"
1135
1136        # Timing values
1137        self.member_vars["p_epoch_times"] = []
1138        self.member_var_types["p_epoch_times"] = "float array"
1139        self.member_vars["n_train_times"] = []
1140        self.member_var_types["n_train_times"] = "float array"
1141        self.member_vars["p_train_times"] = []
1142        self.member_var_types["p_train_times"] = "float array"
1143        self.member_vars["n_val_times"] = []
1144        self.member_var_types["n_val_times"] = "float array"
1145        self.member_vars["p_val_times"] = []
1146        self.member_var_types["p_val_times"] = "float array"
1147
1148        # Setting for tracking timing
1149        self.member_vars["manual_train_switch"] = False
1150        self.member_var_types["manual_train_switch"] = "bool"
1151
1152        # Tracking scores overwritten when reloading best model
1153        self.member_vars["overwritten_extras"] = []
1154        self.member_var_types["overwritten_extras"] = "float array dictionary array"
1155        self.member_vars["overwritten_vals"] = []
1156        self.member_var_types["overwritten_vals"] = "float array array"
1157        self.member_vars["overwritten_epochs"] = 0
1158        self.member_var_types["overwritten_epochs"] = "int"
1159
1160        # Setting for determining scores
1161        self.member_vars["param_vals_setting"] = GPA.pc.get_param_vals_setting()
1162        self.member_var_types["param_vals_setting"] = "int"
1163
1164        # Optimizer and scheduler types and instances
1165        self.member_vars["optimizer"] = None
1166        self.member_var_types["optimizer"] = "type"
1167        self.member_vars["scheduler"] = None
1168        self.member_var_types["scheduler"] = "type"
1169        self.member_vars["optimizer_instance"] = None
1170        self.member_var_types["optimizer_instance"] = "empty array"
1171        self.member_vars["scheduler_instance"] = None
1172        self.member_var_types["scheduler_instance"] = "empty array"
1173
1174        # Flag for if the tracker was loaded
1175        self.loaded = False
1176
1177        # flag for 
1178        self.member_vars["step_status"] = STEP_CLEARED
1179        self.member_var_types["step_status"] = "int"
1180
1181
1182        # Settings for tracking learning rates
1183        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
1184        self.member_var_types["current_n_learning_rate_initial_skip_steps"] = "int"
1185        self.member_vars["last_max_learning_rate_steps"] = 0
1186        self.member_var_types["last_max_learning_rate_steps"] = "int"
1187        self.member_vars["last_max_learning_rate_value"] = -1
1188        self.member_var_types["last_max_learning_rate_value"] = "float"
1189        self.member_vars["current_cycle_lr_max_scores"] = []
1190        self.member_var_types["current_cycle_lr_max_scores"] = "float array"
1191        self.member_vars["current_step_count"] = 0
1192        self.member_var_types["current_step_count"] = "int"
1193        self.member_vars["committed_to_initial_rate"] = True
1194        self.member_var_types["committed_to_initial_rate"] = "bool"
1195        self.member_vars["best_mean_score_improved_this_epoch"] = 0
1196        self.member_var_types["best_mean_score_improved_this_epoch"] = "int"
1197
1198        # Flag for if current dendrite achieved highest global score
1199        self.member_vars["current_n_set_global_best"] = True
1200        self.member_var_types["current_n_set_global_best"] = "bool"
1201
1202        # Number of tries adding this dendrite count
1203        self.member_vars["num_dendrite_tries"] = 0
1204        self.member_var_types["num_dendrite_tries"] = "int"
1205
1206        # Count of batches per epoch
1207        self.values_per_train_epoch = values_per_train_epoch
1208        self.values_per_val_epoch = values_per_val_epoch
1209
1210        self.save_name = save_name
1211        self.making_graphs = making_graphs
1212
1213        self.start_time = time.time()
1214        self.saved_time = 0
1215        self.start_epoch(internal_call=True)
1216
1217        if GPA.pc.get_verbose():
1218            print(f'Initializing with switch_mode {self.member_vars["switch_mode"]}')

Initialize the tracker

Parameters
  • doing_pai (bool): Whether or not dendrites should be used.
  • save_name (str): The base name for saving models and graphs.
  • making_graphs (bool, optional): Whether or not to generate graphs, by default True.
  • param_vals_setting (int, optional): Parameter values setting, by default -1.
  • values_per_train_epoch (int, optional): The number of values to look back for graphing during training, by default -1 (all values).
  • values_per_val_epoch (int, optional): The number of values to look back for graphing during validation, by default -1 (all values).
Returns
  • None
member_vars
member_var_types
neuron_module_vector
tracked_neuron_module_vector
loaded
values_per_train_epoch
values_per_val_epoch
save_name
making_graphs
start_time
saved_time
def to_string(self):
1220    def to_string(self):
1221        """Convert tracker values to string for saving with safetensors.
1222
1223        Parameters
1224        ----------
1225        None
1226
1227        Returns
1228        -------
1229        str
1230            Serialized tracker state suitable for storage in a safetensors field.
1231        """
1232
1233        full_string = ""
1234        for var in self.member_vars:
1235            full_string += var + ","
1236            if self.member_vars[var] is None:
1237                full_string += "None"
1238                full_string += "\n"
1239            elif self.member_var_types[var] == "bool":
1240                full_string += str(self.member_vars[var])
1241                full_string += "\n"
1242            elif self.member_var_types[var] in ("int", "float", "string"):
1243                full_string += str(self.member_vars[var])
1244                full_string += "\n"
1245            elif self.member_var_types[var] == "type":
1246                name = (
1247                    self.member_vars[var].__module__
1248                    + "."
1249                    + self.member_vars[var].__name__
1250                )
1251                full_string += str(self.member_vars[var])
1252                full_string += "\n"
1253            elif self.member_var_types[var] == "empty array":
1254                full_string += "[]"
1255                full_string += "\n"
1256            elif self.member_var_types[var] in ("int array", "float array"):
1257                full_string += "\n"
1258                string = ""
1259                for val in self.member_vars[var]:
1260                    string += str(val) + ","
1261                # Remove the last comma
1262                string = string[:-1]
1263                full_string += string
1264                full_string += "\n"
1265            elif self.member_var_types[var] == "float array dictionary array":
1266                full_string += "\n"
1267                for array in self.member_vars[var]:
1268                    for key in array:
1269                        string = key + ","
1270                        for val in array[key]:
1271                            string += str(val) + ","
1272                        # Remove the last comma
1273                        string = string[:-1]
1274                        full_string += string
1275                        full_string += "\n"
1276                    full_string += "endkey"
1277                    full_string += "\n"
1278                full_string += "endarray"
1279                full_string += "\n"
1280            elif self.member_var_types[var] == "float array dictionary":
1281                full_string += "\n"
1282                for key in self.member_vars[var]:
1283                    string = key + ","
1284                    for val in self.member_vars[var][key]:
1285                        string += str(val) + ","
1286                    # Remove the last comma
1287                    string = string[:-1]
1288                    full_string += string
1289                    full_string += "\n"
1290                full_string += "end"
1291                full_string += "\n"
1292            elif self.member_var_types[var] == "float array array":
1293                full_string += "\n"
1294                for array in self.member_vars[var]:
1295                    string = ""
1296                    for val in array:
1297                        string += str(val) + ","
1298                    # Remove the last comma
1299                    string = string[:-1]
1300                    full_string += string
1301                    full_string += "\n"
1302                full_string += "end"
1303                full_string += "\n"
1304            else:
1305                print("Did not find a member variable")
1306                pdb.set_trace()
1307        return full_string

Convert tracker values to string for saving with safetensors.

Parameters
  • None
Returns
  • str: Serialized tracker state suitable for storage in a safetensors field.
def from_string(self, string):
1309    def from_string(self, string):
1310        """Load tracker values from string.
1311
1312        Parameters
1313        ----------
1314        string : str
1315            The string to load from.
1316
1317        Returns
1318        -------
1319        None
1320            This function does not return a value.
1321        """
1322        f = io.StringIO(string)
1323        while True:
1324            line = f.readline()
1325            if not line:
1326                break
1327            vals = line.split(",")
1328            var = vals[0]
1329
1330            if self.member_var_types[var] == "bool":
1331                val = vals[1][:-1]
1332                if val == "True":
1333                    self.member_vars[var] = True
1334                elif val == "False":
1335                    self.member_vars[var] = False
1336                elif val == "1":
1337                    self.member_vars[var] = 1
1338                elif val == "0":
1339                    self.member_vars[var] = 0
1340                else:
1341                    print("Something went wrong with loading")
1342                    pdb.set_trace()
1343            elif self.member_var_types[var] == "int":
1344                val = vals[1]
1345                self.member_vars[var] = int(val)
1346            elif self.member_var_types[var] == "float":
1347                val = vals[1]
1348                self.member_vars[var] = float(val)
1349            elif self.member_var_types[var] == "string":
1350                val = vals[1][:-1]
1351                self.member_vars[var] = val
1352            elif self.member_var_types[var] == "type":
1353                # Ignore loading types, tracker should have them set up
1354                continue
1355            elif self.member_var_types[var] == "empty array":
1356                val = vals[1]
1357                self.member_vars[var] = []
1358            elif self.member_var_types[var] == "int array":
1359                vals = f.readline()[:-1].split(",")
1360                self.member_vars[var] = []
1361                if vals[0] == "":
1362                    continue
1363                for val in vals:
1364                    self.member_vars[var].append(int(val))
1365            elif self.member_var_types[var] == "float array":
1366                vals = f.readline()[:-1].split(",")
1367                self.member_vars[var] = []
1368                if vals[0] == "":
1369                    continue
1370                for val in vals:
1371                    self.member_vars[var].append(float(val))
1372            elif self.member_var_types[var] == "float array dictionary array":
1373                self.member_vars[var] = []
1374                line2 = f.readline()[:-1]
1375                while line2 != "endarray":
1376                    temp = {}
1377                    while line2 != "endkey":
1378                        vals = line2.split(",")
1379                        name = vals[0]
1380                        temp[name] = []
1381                        vals = vals[1:]
1382                        for val in vals:
1383                            temp[name].append(float(val))
1384                        line2 = f.readline()[:-1]
1385                    self.member_vars[var].append(temp)
1386                    line2 = f.readline()[:-1]
1387            elif self.member_var_types[var] == "float array dictionary":
1388                self.member_vars[var] = {}
1389                line2 = f.readline()[:-1]
1390                while line2 != "end":
1391                    vals = line2.split(",")
1392                    name = vals[0]
1393                    self.member_vars[var][name] = []
1394                    vals = vals[1:]
1395                    for val in vals:
1396                        self.member_vars[var][name].append(float(val))
1397                    line2 = f.readline()[:-1]
1398            elif self.member_var_types[var] == "float array array":
1399                self.member_vars[var] = []
1400                line2 = f.readline()[:-1]
1401                while line2 != "end":
1402                    vals = line2.split(",")
1403                    self.member_vars[var].append([])
1404                    if line2:
1405                        for val in vals:
1406                            self.member_vars[var][-1].append(float(val))
1407                    line2 = f.readline()[:-1]
1408            else:
1409                print("Did not find a member variable")
1410
1411                pdb.set_trace()

Load tracker values from string.

Parameters
  • string (str): The string to load from.
Returns
  • None: This function does not return a value.
def from_string_debug(self, string):
1413    def from_string_debug(self, string):
1414        """Debug function to print tracker values from string without loading them.
1415
1416        Parameters
1417        ----------
1418        string : str
1419            The string to debug load from.
1420
1421        Returns
1422        -------
1423        None
1424            This function does not return a value.
1425        """
1426        f = io.StringIO(string)
1427        print("=== DEBUGGING TRACKER VARIABLES ===")
1428
1429        while True:
1430            line = f.readline()
1431            if not line:
1432                break
1433            vals = line.split(",")
1434            var = vals[0]
1435
1436            print(f"\nVariable: {var}")
1437            print(f"Type: {self.member_var_types.get(var, 'UNKNOWN TYPE')}")
1438            print(f"Current value: {self.member_vars.get(var, 'NOT SET')}")
1439
1440            if self.member_var_types.get(var) == "bool":
1441                val = vals[1][:-1]
1442                print(f"Would set to: {val} -> {val == 'True'}")
1443
1444            elif self.member_var_types.get(var) == "int":
1445                val = vals[1]
1446                print(f"Would set to: {int(val)}")
1447
1448            elif self.member_var_types.get(var) == "float":
1449                val = vals[1]
1450                print(f"Would set to: {float(val)}")
1451
1452            elif self.member_var_types.get(var) == "string":
1453                val = vals[1][:-1]
1454                print(f"Would set to: '{val}'")
1455
1456            elif self.member_var_types.get(var) == "type":
1457                print("Would skip (type loading)")
1458
1459            elif self.member_var_types.get(var) == "empty array":
1460                val = vals[1]
1461                print(f"Would set to: [] (empty array)")
1462
1463            elif self.member_var_types.get(var) == "int array":
1464                vals_line = f.readline()[:-1].split(",")
1465                print(f"Would set to int array with {len(vals_line)} elements:")
1466                if vals_line[0] != "":
1467                    print(
1468                        f"  Elements: {vals_line[:5]}{'...' if len(vals_line) > 5 else ''}"
1469                    )
1470                else:
1471                    print("  Empty array")
1472
1473            elif self.member_var_types.get(var) == "float array":
1474                vals_line = f.readline()[:-1].split(",")
1475                print(f"Would set to float array with {len(vals_line)} elements:")
1476                if vals_line[0] != "":
1477                    print(
1478                        f"  Elements: {vals_line[:5]}{'...' if len(vals_line) > 5 else ''}"
1479                    )
1480                else:
1481                    print("  Empty array")
1482
1483            elif self.member_var_types.get(var) == "float array dictionary array":
1484                print("Would process float array dictionary array:")
1485                array_count = 0
1486                line2 = f.readline()[:-1]
1487                while line2 != "endarray":
1488                    key_count = 0
1489                    while line2 != "endkey":
1490                        vals_dict = line2.split(",")
1491                        name = vals_dict[0]
1492                        print(
1493                            f"  Array {array_count}, Key '{name}': {len(vals_dict)-1} elements"
1494                        )
1495                        key_count += 1
1496                        line2 = f.readline()[:-1]
1497                    print(f"  Array {array_count} has {key_count} keys")
1498                    array_count += 1
1499                    line2 = f.readline()[:-1]
1500                print(f"  Total arrays: {array_count}")
1501
1502            elif self.member_var_types.get(var) == "float array dictionary":
1503                print("Would process float array dictionary:")
1504                line2 = f.readline()[:-1]
1505                key_count = 0
1506                while line2 != "end":
1507                    vals_dict = line2.split(",")
1508                    name = vals_dict[0]
1509                    print(f"  Key '{name}': {len(vals_dict)-1} elements")
1510                    key_count += 1
1511                    line2 = f.readline()[:-1]
1512                print(f"  Total keys: {key_count}")
1513
1514            elif self.member_var_types.get(var) == "float array array":
1515                print("Would process float array array:")
1516                line2 = f.readline()[:-1]
1517                array_count = 0
1518                while line2 != "end":
1519                    if line2:
1520                        vals_array = line2.split(",")
1521                        print(f"  Array {array_count}: {len(vals_array)} elements")
1522                    else:
1523                        print(f"  Array {array_count}: empty")
1524                    array_count += 1
1525                    line2 = f.readline()[:-1]
1526                print(f"  Total arrays: {array_count}")
1527
1528            else:
1529                print(f"UNKNOWN TYPE: {self.member_var_types.get(var, 'NOT FOUND')}")
1530
1531        print("\n=== END DEBUG ===")

Debug function to print tracker values from string without loading them.

Parameters
  • string (str): The string to debug load from.
Returns
  • None: This function does not return a value.
def save_tracker_settings(self):
1533    def save_tracker_settings(self):
1534        """Save tracker settings for DistributedDataParallel use.
1535
1536        Saves settings in save_name/array_dims.csv
1537
1538        Parameters
1539        ----------
1540        None
1541        Returns
1542        -------
1543        None
1544
1545        -----
1546        Instructions for use are in API customization.md
1547        """
1548        if not os.path.isdir(self.save_name):
1549            os.makedirs(self.save_name)
1550        f = open(self.save_name + "/array_dims.csv", "w")
1551        for layer in self.neuron_module_vector:
1552            dv = layer.dendrite_module.dendrite_values[0]
1553            shape_str = ",".join(str(s) for s in dv.dendrite_storage_shape.tolist())
1554            f.write(f"{layer.name},{shape_str}\n")
1555        f.close()
1556        if not GPA.pc.get_silent():
1557            print("Tracker settings saved.")
1558            print("You may now delete save_tracker_settings")

Save tracker settings for DistributedDataParallel use.

Saves settings in save_name/array_dims.csv

Parameters
  • None
Returns
  • None
  • -----
  • Instructions for use are in API customization.md
def initialize_tracker_settings(self):
1560    def initialize_tracker_settings(self):
1561        """Initialize tracker settings from saved file.
1562
1563        This function loads tracker settings from a CSV file and applies them
1564        to the layers the tracker is managing.
1565
1566        Parameters
1567        ----------
1568        None
1569
1570        Returns
1571        -------
1572        None
1573
1574        """
1575
1576        channels = {}
1577        if not os.path.exists(self.save_name + "/array_dims.csv"):
1578            print(
1579                "You must call save_tracker_settings before "
1580                "initialize_tracker_settings"
1581            )
1582            print("Follow instructions in customization.md")
1583            pdb.set_trace()
1584        f = open(self.save_name + "/array_dims.csv", "r")
1585        for line in f:
1586            parts = line.strip().split(",")
1587            channels[parts[0]] = [int(s) for s in parts[1:]]
1588        for layer in self.neuron_module_vector:
1589            dv = layer.dendrite_module.dendrite_values[0]
1590            dv.setup_arrays(channels[layer.name])

Initialize tracker settings from saved file.

This function loads tracker settings from a CSV file and applies them to the layers the tracker is managing.

Parameters
  • None
Returns
  • None
def set_optimizer_instance(self, optimizer_instance, additional_optimizers=[]):
1592    def set_optimizer_instance(self, optimizer_instance, additional_optimizers=[]):
1593        """Set optimizer instance directly.
1594
1595        Parameters
1596        ----------
1597        optimizer_instance : object
1598            The optimizer instance to set.
1599
1600        Returns
1601        -------
1602        None
1603
1604        """
1605        # This call must be first before the parameters get filtered.
1606        optimizer_instance.zero_grad()
1607        try:
1608            for param_group in optimizer_instance.param_groups:
1609                if (
1610                    param_group["weight_decay"] > 0
1611                    and GPA.pc.get_weight_decay_accepted() is False
1612                ):
1613                    _pai_log(
1614                        "warning",
1615                        "For PAI training it is recommended to not use weight decay in your optimizer",
1616                    )
1617
1618        except:
1619            pass
1620        self.member_vars["optimizer_instance"] = optimizer_instance
1621        if GPA.pc.get_perforated_backpropagation():
1622            TPB.setup_optimizer_pb(self.member_vars["optimizer_instance"])
1623            for optimizer in additional_optimizers:
1624                TPB.filter_params(optimizer)

Set optimizer instance directly.

Parameters
  • optimizer_instance (object): The optimizer instance to set.
Returns
  • None
def set_optimizer(self, optimizer):
1626    def set_optimizer(self, optimizer):
1627        """Set optimizer type to be initialized later
1628
1629        Parameters
1630        ----------
1631        optimizer : object
1632            The optimizer type to set.
1633
1634        Returns
1635        -------
1636        None
1637
1638        """
1639        self.member_vars["optimizer"] = optimizer

Set optimizer type to be initialized later

Parameters
  • optimizer (object): The optimizer type to set.
Returns
  • None
def set_scheduler(self, scheduler):
1641    def set_scheduler(self, scheduler):
1642        """Set scheduler type to be initialized later
1643
1644        Parameters
1645        ----------
1646        scheduler : object
1647            The scheduler type to set.
1648
1649        Returns
1650        -------
1651        None
1652
1653        """
1654        if scheduler is not torch.optim.lr_scheduler.ReduceLROnPlateau:
1655            if GPA.pc.get_verbose():
1656                print("Not using ReduceLROnPlateau, this is not recommended")
1657        self.member_vars["scheduler"] = scheduler

Set scheduler type to be initialized later

Parameters
  • scheduler (object): The scheduler type to set.
Returns
  • None
def increment_scheduler(self, num_ticks, mode):
1659    def increment_scheduler(self, num_ticks, mode):
1660        """Increment the scheduler a set number of times.
1661
1662        Used for finding best initial learning rate when adding dendrites.
1663
1664        Parameters
1665        ----------
1666        num_ticks : int
1667            The number of scheduler steps to take.
1668        mode : str
1669            The mode for stepping the scheduler. Options are:
1670            - "step_learning_rate": Step based on improved accuracy epochs
1671            - "increment_epoch_count": Step based on total epoch count
1672
1673        Returns
1674        -------
1675        current_steps : int
1676            The number of learning rate changes that occurred.
1677        learning_rate1 : float
1678            The final learning rate after stepping.
1679
1680        """
1681
1682        current_steps = 0
1683        current_ticker = 0
1684
1685        for param_group in GPA.pai_tracker.member_vars[
1686            "optimizer_instance"
1687        ].param_groups:
1688            learning_rate1 = param_group["lr"]
1689
1690        if GPA.pc.get_verbose():
1691            print("Using scheduler:")
1692            print(type(self.member_vars["scheduler_instance"]))
1693
1694        while current_ticker < num_ticks:
1695            if GPA.pc.get_verbose():
1696                print(
1697                    f"Lower start rate initial {learning_rate1} "
1698                    f'stepping {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} times'
1699                )
1700
1701            if (
1702                type(self.member_vars["scheduler_instance"])
1703                is torch.optim.lr_scheduler.ReduceLROnPlateau
1704            ):
1705                if mode == "step_learning_rate":
1706                    # Step with counter as last improved accuracy
1707                    self.member_vars["scheduler_instance"].step(
1708                        metrics=self.member_vars["last_improved_accuracies"][
1709                            GPA.pai_tracker.steps_after_switch() - 1
1710                        ]
1711                    )
1712                elif mode == "increment_epoch_count":
1713                    # Step with improved epoch counts up to current location
1714                    self.member_vars["scheduler_instance"].step(
1715                        metrics=self.member_vars["last_improved_accuracies"][
1716                            -((num_ticks - 1) - current_ticker) - 1
1717                        ]
1718                    )
1719            else:
1720                self.member_vars["scheduler_instance"].step()
1721
1722            for param_group in GPA.pai_tracker.member_vars[
1723                "optimizer_instance"
1724            ].param_groups:
1725                learning_rate2 = param_group["lr"]
1726
1727            if learning_rate2 != learning_rate1:
1728                current_steps += 1
1729                learning_rate1 = learning_rate2
1730                if mode == "step_learning_rate":
1731                    current_ticker += 1
1732                if GPA.pc.get_verbose():
1733                    print(f"1 step {current_steps} to {learning_rate2}")
1734
1735            if mode == "increment_epoch_count":
1736                current_ticker += 1
1737
1738        return current_steps, learning_rate1

Increment the scheduler a set number of times.

Used for finding best initial learning rate when adding dendrites.

Parameters
  • num_ticks (int): The number of scheduler steps to take.
  • mode (str): The mode for stepping the scheduler. Options are:
    • "step_learning_rate": Step based on improved accuracy epochs
    • "increment_epoch_count": Step based on total epoch count
Returns
  • current_steps (int): The number of learning rate changes that occurred.
  • learning_rate1 (float): The final learning rate after stepping.
def setup_optimizer(self, net, opt_args, sched_args=None, parameters=None):
1740    def setup_optimizer(self, net, opt_args, sched_args=None, parameters=None):
1741        """Initialize the optimizer and scheduler when added.
1742
1743        Parameters
1744        ----------
1745        net : object
1746            The neural network model.
1747        opt_args : dict
1748            The arguments for the optimizer.
1749        sched_args : dict, optional
1750            The arguments for the scheduler, by default None.
1751
1752        Returns
1753        -------
1754        optimizer : object
1755            The initialized optimizer instance.
1756        scheduler : object or None
1757            The initialized scheduler instance, or None if no scheduler was set.
1758
1759        """
1760        if "weight_decay" in opt_args and not GPA.pc.get_weight_decay_accepted():
1761            _pai_log(
1762                "warning",
1763                "For PAI training it is recommended to not use weight decay in your optimizer",
1764            )
1765
1766        if ("model" not in opt_args.keys()) and "params" not in opt_args.keys():
1767            print("In setup_optimizer it will be depreciated to not pass in params yourself in the future")
1768            print("please change the settings to include params")
1769            if self.member_vars["mode"] == "n":
1770                if parameters is not None:
1771                    opt_args["params"] = parameters
1772                else:
1773                    opt_args["params"] = filter(lambda p: p.requires_grad, net.parameters())
1774            else:
1775                params = UPA.get_pai_network_params(net)
1776                if parameters is not None:
1777                    # Filter parameters to only those in params, preserving weight_decay
1778                    params_set = set(params)
1779                    filtered_params = []
1780                    for param_group in parameters:
1781                        filtered_group_params = [p for p in param_group["params"] if p in params_set]
1782                        if filtered_group_params:
1783                            filtered_params.append({
1784                                "params": filtered_group_params,
1785                                "weight_decay": param_group["weight_decay"]
1786                            })
1787                    opt_args["params"] = filtered_params
1788                else:
1789                    opt_args["params"] = params
1790        elif "params" in opt_args.keys():
1791            # Check if params is a list of param groups (dicts) or a single param group
1792            params_value = opt_args["params"]
1793            if isinstance(params_value, list) and len(params_value) > 0:
1794                # Check if it's a list of dicts (multiple param groups) or list of tensors (single group)
1795                if isinstance(params_value[0], dict):
1796                    # Multiple param groups format: [{"params": [...], "lr": ...}, ...]
1797                    # Filter each param group for requires_grad
1798                    filtered_param_groups = []
1799                    for param_group in params_value:
1800                        filtered_group_params = [p for p in param_group["params"] if p.requires_grad]
1801                        if filtered_group_params:
1802                            new_group = param_group.copy()
1803                            new_group["params"] = filtered_group_params
1804                            filtered_param_groups.append(new_group)
1805                    opt_args["params"] = filtered_param_groups
1806                else:
1807                    # Single param group format: [tensor1, tensor2, ...] or generator
1808                    # Filter for requires_grad
1809                    opt_args["params"] = [p for p in params_value if p.requires_grad]
1810            elif hasattr(params_value, '__iter__'):
1811                # Handle generators or other iterables
1812                opt_args["params"] = [p for p in params_value if p.requires_grad]
1813
1814        optimizer = self.member_vars["optimizer"](**opt_args)
1815        self.set_optimizer_instance(optimizer)
1816
1817        if self.member_vars["scheduler"] is not None:
1818            # Handle SequentialLR specially
1819            if self.member_vars["scheduler"] is torch.optim.lr_scheduler.SequentialLR:
1820                """
1821                sched_args should be a dict with "schedulers" (list of tuples) and "milestones"
1822                For example:
1823                sequential_schedArgs = {
1824                    "schedulers": [
1825                        (warmup_scheduler_class, warmup_schedArgs),
1826                        (main_scheduler_class, main_schedArgs)
1827                    ],
1828                    "milestones": [switch_epoch]
1829                }
1830                """
1831                schedulers = []
1832                milestones = sched_args.get("milestones", [])
1833                scheduler_configs = sched_args.get("schedulers", [])
1834                
1835                for scheduler_class, scheduler_args in scheduler_configs:
1836                    schedulers.append(scheduler_class(optimizer, **scheduler_args))
1837                
1838                self.member_vars["scheduler_instance"] = torch.optim.lr_scheduler.SequentialLR(
1839                    optimizer, schedulers=schedulers, milestones=milestones
1840                )
1841            else:
1842                self.member_vars["scheduler_instance"] = self.member_vars["scheduler"](
1843                    optimizer, **sched_args
1844                )
1845            current_steps = 0
1846
1847            for param_group in GPA.pai_tracker.member_vars[
1848                "optimizer_instance"
1849            ].param_groups:
1850                learning_rate1 = param_group["lr"]
1851
1852            if GPA.pc.get_verbose():
1853                print(
1854                    f"Resetting scheduler with {GPA.pai_tracker.steps_after_switch()} "
1855                    f"steps and {GPA.pc.get_initial_history_after_switches()} initial ticks to skip"
1856                )
1857
1858            # Find setting of previously used learning rate before adding dendrites
1859            if (
1860                GPA.pai_tracker.member_vars[
1861                    "current_n_learning_rate_initial_skip_steps"
1862                ]
1863                != 0
1864            ):
1865                additional_steps, learning_rate1 = self.increment_scheduler(
1866                    GPA.pai_tracker.member_vars[
1867                        "current_n_learning_rate_initial_skip_steps"
1868                    ],
1869                    "step_learning_rate",
1870                )
1871                current_steps += additional_steps
1872
1873            if self.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
1874                initial = GPA.pc.get_initial_history_after_switches()
1875            else:
1876                initial = 0
1877
1878            if GPA.pai_tracker.steps_after_switch() > initial:
1879                # Minus extra 1 because this gets called after start epoch
1880                additional_steps, learning_rate1 = self.increment_scheduler(
1881                    (GPA.pai_tracker.steps_after_switch() - initial) - 1,
1882                    "increment_epoch_count",
1883                )
1884                current_steps += additional_steps
1885
1886            if GPA.pc.get_verbose():
1887                print(
1888                    f"Scheduler update loop with {current_steps} "
1889                    f"ended with {learning_rate1}"
1890                )
1891                print(
1892                    f"Scheduler ended with {current_steps} steps "
1893                    f"and lr of {learning_rate1}"
1894                )
1895
1896            self.member_vars["current_step_count"] = current_steps
1897            return optimizer, self.member_vars["scheduler_instance"]
1898        else:
1899            return optimizer, None

Initialize the optimizer and scheduler when added.

Parameters
  • net (object): The neural network model.
  • opt_args (dict): The arguments for the optimizer.
  • sched_args (dict, optional): The arguments for the scheduler, by default None.
Returns
  • optimizer (object): The initialized optimizer instance.
  • scheduler (object or None): The initialized scheduler instance, or None if no scheduler was set.
def clear_optimizer_and_scheduler(self):
1901    def clear_optimizer_and_scheduler(self):
1902        """Clear the instances for saving.
1903
1904        Parameters
1905        ----------
1906        None
1907
1908        Returns
1909        -------
1910        None
1911            This function does not return a value.
1912        """
1913        self.member_vars["optimizer_instance"] = None
1914        self.member_vars["scheduler_instance"] = None

Clear the instances for saving.

Parameters
  • None
Returns
  • None: This function does not return a value.
def switch_time(self):
1916    def switch_time(self):
1917        """Determine if it's time to switch between neuron and dendrite training.
1918
1919        Parameters
1920        ----------
1921        None
1922
1923        Returns
1924        -------
1925        bool
1926            True if it's time to switch, False otherwise.
1927
1928        Notes
1929        -----
1930        Based on current settings and history of scores.
1931        """
1932
1933        switch_phrase = "No mode, this should never be the case."
1934        switch_number = GPA.pc.get_n_epochs_to_switch()
1935        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1936            switch_phrase = "DOING_SWITCH_EVERY_TIME"
1937        elif self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY:
1938            switch_phrase = "DOING_HISTORY"
1939        elif self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH:
1940            switch_phrase = "DOING_FIXED_SWITCH"
1941            switch_number = GPA.pc.get_fixed_switch_num()
1942        elif self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1943            switch_phrase = "DOING_NO_SWITCH"
1944        else:
1945            print(
1946                "A switch mode must be set.  Check your settings for GPA.pc.set_switch_mode()."
1947            )
1948            pdb.set_trace()
1949        if not GPA.pc.get_silent():
1950            if(GPA.pc.get_perforated_backpropagation()):
1951                print(
1952                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1953                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1954                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1955                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1956                    f'n: {switch_number}, p: {GPA.pc.get_p_epochs_to_switch()}, '
1957                    f'num_cycles: {self.member_vars["num_cycles"]}'
1958                )
1959            else:
1960                print(
1961                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1962                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1963                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1964                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1965                    f'n: {switch_number}, num_cycles: {self.member_vars["num_cycles"]}'
1966                )
1967            print(
1968                f'  Score tracking: current_n_set_global_best={self.member_vars["current_n_set_global_best"]}, '
1969                f'global_best={self.member_vars["global_best_validation_score"]:.4f}, '
1970                f'current_best={self.member_vars["current_best_validation_score"]:.4f}'
1971            )
1972        if GPA.pc.get_perforated_backpropagation():
1973            # this will fill in epoch last improved
1974            TPB.best_pai_score_improved_this_epoch(self)  ## CLOSED ONLY
1975        if self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1976            if not GPA.pc.get_silent():
1977                print("Returning False - doing no switch mode")
1978            return False
1979
1980        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1981            if not GPA.pc.get_silent():
1982                print("Returning True - switching every time")
1983            return True
1984
1985        # Check if we're in the middle of learning rate optimization
1986        # If so, block ALL switch triggers until committed
1987        if GPA.pc.get_verbose():
1988            print("=== LR Optimization Check ===")
1989            print(f'  mode == "n": {self.member_vars["mode"] == "n"}')
1990            print(f"  get_learn_dendrites_live(): {GPA.pc.get_learn_dendrites_live()}")
1991            print(f'  committed_to_initial_rate: {GPA.pai_tracker.member_vars["committed_to_initial_rate"]}')
1992            print(f"  get_dont_give_up_unless_learning_rate_lowered(): {GPA.pc.get_dont_give_up_unless_learning_rate_lowered()}")
1993            print(f'  current_n_learning_rate_initial_skip_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"]}')
1994            print(f'  last_max_learning_rate_steps: {self.member_vars["last_max_learning_rate_steps"]}')
1995            print(f'  skip_steps < max_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"] < self.member_vars["last_max_learning_rate_steps"]}')
1996            print(f'  scheduler is not None: {self.member_vars["scheduler"] is not None}')
1997            print("=============================")
1998        
1999        if (
2000            ((self.member_vars["mode"] == "n") or GPA.pc.get_learn_dendrites_live())
2001            and (GPA.pai_tracker.member_vars["committed_to_initial_rate"] is False)
2002            and (GPA.pc.get_dont_give_up_unless_learning_rate_lowered())
2003            and (
2004                self.member_vars["current_n_learning_rate_initial_skip_steps"]
2005                <= self.member_vars["last_max_learning_rate_steps"]
2006            )
2007            and self.member_vars["scheduler"] is not None
2008        ):
2009            if not GPA.pc.get_silent():
2010                print(
2011                    f"Returning False - learning rate optimization in progress. "
2012                    f"Not committed yet. Comparing "
2013                    f'initial {self.member_vars["current_n_learning_rate_initial_skip_steps"]} '
2014                    f'to last max {self.member_vars["last_max_learning_rate_steps"]}'
2015                )
2016            return False
2017
2018        if len(self.member_vars["switch_epochs"]) == 0:
2019            this_count = self.member_vars["num_epochs_run"]
2020        else:
2021            this_count = (
2022                self.member_vars["num_epochs_run"]
2023                - self.member_vars["switch_epochs"][-1]
2024            )
2025        cap_switch = False
2026        if GPA.pc.get_perforated_backpropagation():
2027            cap_switch = TPB.check_cap_switch(self, this_count)
2028
2029        if self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY and (
2030            (
2031                (self.member_vars["mode"] == "n")
2032                and (
2033                    self.member_vars["num_epochs_run"]
2034                    - self.member_vars["epoch_last_improved"]
2035                    >= GPA.pc.get_n_epochs_to_switch()
2036                )
2037                and this_count
2038                >= GPA.pc.get_initial_history_after_switches()
2039                + GPA.pc.get_n_epochs_to_switch()
2040            )
2041            or (GPA.pc.get_perforated_backpropagation() and TPB.history_switch(self))
2042            or cap_switch
2043        ):
2044            if not GPA.pc.get_silent():
2045                print("Returning True - History and last improved is hit")
2046            return True
2047
2048        if self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH and (
2049            (
2050                self.member_vars["total_epochs_run"] % GPA.pc.get_fixed_switch_num()
2051                == GPA.pc.get_fixed_switch_num() - 1
2052            )
2053            and self.member_vars["num_epochs_run"]
2054            >= GPA.pc.get_first_fixed_switch_num() - 1
2055        ):
2056            if not GPA.pc.get_silent():
2057                print("Returning True - Fixed switch number is hit")
2058            return True
2059
2060        if not GPA.pc.get_silent():
2061            print("Returning False - no triggers to switch have been hit")
2062        return False

Determine if it's time to switch between neuron and dendrite training.

Parameters
  • None
Returns
  • bool: True if it's time to switch, False otherwise.
Notes

Based on current settings and history of scores.

def steps_after_switch(self):
2064    def steps_after_switch(self):
2065        """Based on settings, return value for steps since a switch.
2066
2067        Different options for param vals setting determine what is returned.
2068
2069        Parameters
2070        ----------
2071        None
2072
2073        Returns
2074        -------
2075        int
2076            The number of epochs since the last switch, or total epochs run,
2077            depending on settings.
2078
2079        """
2080        if self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_TOTAL_EPOCH:
2081            return self.member_vars["num_epochs_run"]
2082        elif (
2083            self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_UPDATE_EPOCH
2084        ):
2085            return self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2086        elif (
2087            self.member_vars["param_vals_setting"]
2088            == GPA.pc.PARAM_VALS_BY_NEURON_EPOCH_START
2089        ):
2090            if self.member_vars["mode"] == "p":
2091                return (
2092                    self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2093                )
2094            else:
2095                return self.member_vars["num_epochs_run"]
2096        else:
2097            print(
2098                f'{self.member_vars["param_vals_setting"]} is not a valid param vals option'
2099            )
2100            pdb.set_trace()

Based on settings, return value for steps since a switch.

Different options for param vals setting determine what is returned.

Parameters
  • None
Returns
  • int: The number of epochs since the last switch, or total epochs run, depending on settings.
def add_pai_neuron_module(self, new_module, initial_add=True):
2102    def add_pai_neuron_module(self, new_module, initial_add=True):
2103        """Add neuron modules to internal vectors.
2104
2105        Parameters
2106        ----------
2107        new_module : object
2108            The new module to add.
2109        initial_add : bool, optional
2110            Whether this is the initial addition rather than loading from file
2111
2112        Returns
2113        -------
2114        None
2115
2116        """
2117
2118        # If it's a duplicate, ignore the second addition
2119        if new_module in self.neuron_module_vector:
2120            return
2121        self.neuron_module_vector.append(new_module)
2122        if self.member_vars["doing_pai"]:
2123            PA.set_wrapped_params(new_module)
2124        if initial_add:
2125            self.member_vars["best_scores"].append([])
2126            self.member_vars["current_scores"].append([])

Add neuron modules to internal vectors.

Parameters
  • new_module (object): The new module to add.
  • initial_add (bool, optional): Whether this is the initial addition rather than loading from file
Returns
  • None
def add_tracked_neuron_module(self, new_module, initial_add=True):
2128    def add_tracked_neuron_module(self, new_module, initial_add=True):
2129        """Add tracked modules to internal vectors
2130
2131        Parameters
2132        ----------
2133        new_module : object
2134            The new module to add.
2135        initial_add : bool, optional
2136            Whether this is the initial addition rather than loading from file
2137
2138        Returns
2139        -------
2140        None
2141
2142        """
2143        # If it's a duplicate, ignore the second addition
2144        if new_module in self.tracked_neuron_module_vector:
2145            return
2146        self.tracked_neuron_module_vector.append(new_module)
2147        if self.member_vars["doing_pai"]:
2148            PA.set_tracked_params(new_module)

Add tracked modules to internal vectors

Parameters
  • new_module (object): The new module to add.
  • initial_add (bool, optional): Whether this is the initial addition rather than loading from file
Returns
  • None
def reset_module_vector(self, net, load_from_restart):
2150    def reset_module_vector(self, net, load_from_restart):
2151        """Clear internal vectors and reset from network.
2152
2153        Parameters
2154        ----------
2155        net : object
2156            The neural network model.
2157        load_from_restart : bool
2158            Whether loading from a restart file.
2159
2160        Returns
2161        -------
2162        None
2163
2164        """
2165        self.neuron_module_vector = []
2166        self.tracked_neuron_module_vector = []
2167        this_list = UPA.get_pai_modules(net, 0)
2168        for module in this_list:
2169            self.add_pai_neuron_module(module, initial_add=load_from_restart)
2170        this_list = UPA.get_tracked_modules(net, 0)
2171        for module in this_list:
2172            self.add_tracked_neuron_module(module, initial_add=load_from_restart)

Clear internal vectors and reset from network.

Parameters
  • net (object): The neural network model.
  • load_from_restart (bool): Whether loading from a restart file.
Returns
  • None
def reset_vals_for_score_reset(self):
2174    def reset_vals_for_score_reset(self):
2175        """Reset cycle scores for new cycle.
2176
2177        Parameters
2178        ----------
2179        None
2180
2181        Returns
2182        -------
2183        None
2184            This function does not return a value.
2185        """
2186
2187        if GPA.pc.get_find_best_lr():
2188            self.member_vars["committed_to_initial_rate"] = False
2189            print("Resetting committed to initial rate to False")
2190        # If retaining all dendrties always say that the current dendrites set global best for saving and loading
2191        if GPA.pc.get_retain_all_dendrites():
2192            self.member_vars["current_n_set_global_best"] = True
2193            self.member_vars["global_best_validation_score"] = 0
2194        else:
2195            self.member_vars["current_n_set_global_best"] = False
2196
2197        # Don't reset global best, but do reset current best
2198        self.member_vars["current_best_validation_score"] = 0
2199        self.member_vars["initial_lr_test_epoch_count"] = -1

Reset cycle scores for new cycle.

Parameters
  • None
Returns
  • None: This function does not return a value.
def set_dendrite_training(self):
2201    def set_dendrite_training(self):
2202        """Signal all layers to start dendrite training.
2203
2204        Parameters
2205        ----------
2206        None
2207
2208        Returns
2209        -------
2210        None
2211            This function does not return a value.
2212        """
2213        if GPA.pc.get_verbose():
2214            print("Calling set_dendrite_training")
2215
2216        for layer in self.neuron_module_vector[:]:
2217            worked = layer.set_mode("p")
2218            """
2219            worked is False when a layer was added to the neuron module vector
2220            but then it's never actually been used. This can happen when
2221            you have set a layer to have requires_grad = False or when
2222            you have a module as a member variable but it's not actually
2223            part of the network. Should be moved to be a tracked layer
2224            rather than a neuron layer.
2225            """
2226            if not worked:
2227                self.neuron_module_vector.remove(layer)
2228
2229        for layer in self.tracked_neuron_module_vector[:]:
2230            worked = layer.set_mode("p")
2231
2232        self.create_new_dendrite_module()
2233        self.member_vars["mode"] = "p"
2234        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2235
2236        if GPA.pc.get_learn_dendrites_live():
2237            self.reset_vals_for_score_reset()
2238
2239        self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2240            "current_step_count"
2241        ]
2242
2243        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
2244        GPA.pai_tracker.member_vars["num_cycles"] += 1

Signal all layers to start dendrite training.

Parameters
  • None
Returns
  • None: This function does not return a value.
def set_neuron_training(self):
2247    def set_neuron_training(self):
2248        """Signal all layers to start neuron training.
2249
2250        Parameters
2251        ----------
2252        None
2253
2254        Returns
2255        -------
2256        None
2257            This function does not return a value.
2258        """
2259        for module in self.neuron_module_vector:
2260            module.set_mode("n")
2261        for module in self.tracked_neuron_module_vector[:]:
2262            module.set_mode("n")
2263
2264        self.member_vars["mode"] = "n"
2265        self.member_vars["num_dendrites_added"] += 1
2266        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2267        self.reset_vals_for_score_reset()
2268
2269        self.member_vars["current_cycle_lr_max_scores"] = []
2270        if GPA.pc.get_learn_dendrites_live():
2271            self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2272                "current_step_count"
2273            ]
2274        GPA.pai_tracker.member_vars["num_cycles"] += 1
2275
2276        if GPA.pc.get_reset_best_score_on_switch():
2277            GPA.pai_tracker.member_vars["current_best_validation_score"] = 0
2278            GPA.pai_tracker.member_vars["running_accuracy"] = 0

Signal all layers to start neuron training.

Parameters
  • None
Returns
  • None: This function does not return a value.
def start_epoch(self, internal_call=False):
2280    def start_epoch(self, internal_call=False):
2281        """Perform steps for when a new training epoch is about to begin.
2282
2283        Parameters
2284        ----------
2285        internal_call : bool, optional
2286            Whether this is an internal call or manual call
2287
2288        Returns
2289        -------
2290        None
2291
2292        Notes
2293        -----
2294        If you ever need to call this manually, set internal_call to False.
2295
2296        """
2297        if self.member_vars["manual_train_switch"] and internal_call:
2298            return
2299
2300        if not internal_call and not self.member_vars["manual_train_switch"]:
2301            self.member_vars["manual_train_switch"] = True
2302            self.saved_time = 0
2303            self.member_vars["num_epochs_run"] = -1
2304            self.member_vars["total_epochs_run"] = -1
2305
2306        end = time.time()
2307        if self.member_vars["manual_train_switch"]:
2308            if self.saved_time != 0:
2309                if self.member_vars["mode"] == "p":
2310                    self.member_vars["p_val_times"].append(end - self.saved_time)
2311                else:
2312                    self.member_vars["n_val_times"].append(end - self.saved_time)
2313
2314        if self.member_vars["mode"] == "p":
2315            for layer in self.neuron_module_vector:
2316                for m in range(0, GPA.pc.get_global_candidates()):
2317                    with torch.no_grad():
2318                        if GPA.pc.get_verbose():
2319                            print(f"Resetting score for {layer.name}")
2320                        # Snapshot best_score before reset so we can compute per-epoch improvement
2321                        layer.dendrite_module.dendrite_values[
2322                            m
2323                        ].epoch_start_best_score.copy_(
2324                            layer.dendrite_module.dendrite_values[
2325                                m
2326                            ].best_score.detach()
2327                        )
2328                        layer.dendrite_module.dendrite_values[
2329                            m
2330                        ].best_score_improved_this_epoch = (
2331                            layer.dendrite_module.dendrite_values[
2332                                m
2333                            ].best_score_improved_this_epoch
2334                            * 0
2335                        )
2336                        layer.dendrite_module.dendrite_values[
2337                            m
2338                        ].nodes_best_improved_this_epoch = (
2339                            layer.dendrite_module.dendrite_values[
2340                                m
2341                            ].nodes_best_improved_this_epoch
2342                            * 0
2343                        )
2344                        layer.dendrite_module.dendrite_values[
2345                            m
2346                        ].nodes_improved_any = (
2347                            layer.dendrite_module.dendrite_values[
2348                                m
2349                            ].nodes_improved_any
2350                            * 0
2351                        )
2352            if GPA.pc.get_perforated_backpropagation():
2353                self.member_vars["best_mean_score_improved_this_epoch"] = 0
2354        self.member_vars["num_epochs_run"] += 1
2355        self.member_vars["total_epochs_run"] = (
2356            self.member_vars["num_epochs_run"] + self.member_vars["overwritten_epochs"]
2357        )
2358        self.saved_time = end

Perform steps for when a new training epoch is about to begin.

Parameters
  • internal_call (bool, optional): Whether this is an internal call or manual call
Returns
  • None
Notes

If you ever need to call this manually, set internal_call to False.

def stop_epoch(self, internal_call=False):
2360    def stop_epoch(self, internal_call=False):
2361        """Perform steps when a training epoch has completed.
2362
2363        Parameters
2364        ----------
2365        internal_call : bool, optional
2366            Whether this is an internal call or manual call
2367
2368        Returns
2369        -------
2370        None
2371
2372        Notes
2373        -----
2374        If you ever need to call this manually, set internal_call to False.
2375
2376        """
2377        end = time.time()
2378        if self.member_vars["manual_train_switch"] and internal_call:
2379            return
2380
2381        if self.member_vars["manual_train_switch"]:
2382            if self.member_vars["mode"] == "p":
2383                self.member_vars["p_train_times"].append(end - self.saved_time)
2384            else:
2385                self.member_vars["n_train_times"].append(end - self.saved_time)
2386        else:
2387            if self.member_vars["mode"] == "p":
2388                self.member_vars["p_epoch_times"].append(end - self.saved_time)
2389            else:
2390                self.member_vars["n_epoch_times"].append(end - self.saved_time)
2391
2392        self.saved_time = end

Perform steps when a training epoch has completed.

Parameters
  • internal_call (bool, optional): Whether this is an internal call or manual call
Returns
  • None
Notes

If you ever need to call this manually, set internal_call to False.

def initialize( self, model, doing_pai=True, save_name='PAI', making_graphs=True, maximizing_score=True, num_classes=10000, values_per_train_epoch=-1, values_per_val_epoch=-1, zooming_graph=True):
2394    def initialize(
2395        self,
2396        model,
2397        doing_pai=True,
2398        save_name="PAI",
2399        making_graphs=True,
2400        maximizing_score=True,
2401        num_classes=10000,
2402        values_per_train_epoch=-1,
2403        values_per_val_epoch=-1,
2404        zooming_graph=True,
2405    ):
2406        """Setup the tracker with initial settings.
2407
2408
2409        Parameters
2410        ----------
2411        model : object
2412            The neural network model.
2413        doing_pai : bool, optional
2414            Whether to add dendrites, by default True.
2415        save_name : str, optional
2416            The name under which to save the model.
2417        making_graphs : bool, optional
2418            Whether to make graphs, by default True.
2419        maximizing_score : bool, optional
2420            Whether to maximize the score, by default True.
2421        num_classes : int, optional
2422            The number of classes in the dataset, unused
2423        values_per_train_epoch : int, optional
2424            The number of values to look back for graphing
2425            during training, by default -1 (all values).
2426        values_per_val_epoch : int, optional
2427            The number of values to look back for graphing
2428            during validation, by default -1 (all values).
2429        zooming_graph : bool, optional
2430            Whether to zoom on graphs, by default True.
2431
2432
2433        Returns
2434        -------
2435        nn.Module
2436            Converted model instance configured for the tracker settings.
2437        """
2438        model = UPA.convert_network(model)
2439        self.member_vars["doing_pai"] = doing_pai
2440        self.member_vars["maximizing_score"] = maximizing_score
2441        self.save_name = save_name
2442        self.zooming_graph = zooming_graph
2443        self.making_graphs = making_graphs
2444
2445        if not self.loaded:
2446            self.member_vars["running_accuracy"] = (1.0 / num_classes) * 100
2447
2448        self.values_per_train_epoch = values_per_train_epoch
2449        self.values_per_val_epoch = values_per_val_epoch
2450
2451        if GPA.pc.get_testing_dendrite_capacity():
2452            if not GPA.pc.get_silent():
2453                print("Running a test of Dendrite Capacity.")
2454            GPA.pc.set_switch_mode(GPA.pc.DOING_SWITCH_EVERY_TIME)
2455            self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
2456            GPA.pc.set_retain_all_dendrites(True)
2457            GPA.pc.set_max_dendrite_tries(1000)
2458            GPA.pc.set_max_dendrites(1000)
2459            if GPA.pc.get_perforated_backpropagation():
2460                GPA.pc.set_initial_correlation_batches(1)
2461        else:
2462            if not GPA.pc.get_silent():
2463                print("Running Dendrite Experiment")
2464        return model

Setup the tracker with initial settings.

Parameters
  • model (object): The neural network model.
  • doing_pai (bool, optional): Whether to add dendrites, by default True.
  • save_name (str, optional): The name under which to save the model.
  • making_graphs (bool, optional): Whether to make graphs, by default True.
  • maximizing_score (bool, optional): Whether to maximize the score, by default True.
  • num_classes (int, optional): The number of classes in the dataset, unused
  • values_per_train_epoch (int, optional): The number of values to look back for graphing during training, by default -1 (all values).
  • values_per_val_epoch (int, optional): The number of values to look back for graphing during validation, by default -1 (all values).
  • zooming_graph (bool, optional): Whether to zoom on graphs, by default True.
Returns
  • nn.Module: Converted model instance configured for the tracker settings.
def generate_accuracy_plots(self, ax, save_folder, extra_string):
2466    def generate_accuracy_plots(self, ax, save_folder, extra_string):
2467        """
2468        Generate plots and csvs for accuracy
2469
2470        Parameters
2471        ----------
2472        ax : object
2473            The matplotlib axis to plot on.
2474        save_folder : str
2475            The folder to save the plots and csvs in.
2476        extra_string : str
2477            An extra string to append to the filenames.
2478
2479        Returns
2480        -------
2481        None
2482
2483        """
2484
2485        # If scores are being saved for epochs that get overwritten, plot them
2486        for list_id in range(len(self.member_vars["overwritten_extras"])):
2487            for extra_id in self.member_vars["overwritten_extras"][list_id]:
2488                ax.plot(
2489                    np.arange(
2490                        len(self.member_vars["overwritten_extras"][list_id][extra_id])
2491                    ),
2492                    self.member_vars["overwritten_extras"][list_id][extra_id],
2493                    "r",
2494                )
2495            ax.plot(
2496                np.arange(len(self.member_vars["overwritten_vals"][list_id])),
2497                self.member_vars["overwritten_vals"][list_id],
2498                "b",
2499            )
2500
2501        # Determine which accuracy vector to use
2502        if GPA.pc.get_drawing_pai():
2503            accuracies = self.member_vars["accuracies"]
2504        else:
2505            accuracies = self.member_vars["n_accuracies"]
2506
2507        # Get pointer to additional scores being saved
2508        extra_scores = self.member_vars["extra_scores"]
2509
2510        # Plot the main accuracy scores
2511        ax.plot(np.arange(len(accuracies)), accuracies, label="Validation Scores")
2512        ax.plot(
2513            np.arange(len(self.member_vars["running_accuracies"])),
2514            self.member_vars["running_accuracies"],
2515            label="Validation Running Scores",
2516        )
2517
2518        # Plot additional scores
2519        for extra_score in extra_scores:
2520            ax.plot(
2521                np.arange(len(extra_scores[extra_score])),
2522                extra_scores[extra_score],
2523                label=extra_score,
2524            )
2525
2526        plt.title(save_folder + "/" + self.save_name + "Scores")
2527        plt.xlabel("Epochs")
2528        plt.ylabel("Score")
2529
2530        # Add point at epoch last improved and best validation score
2531        if GPA.pc.get_drawing_pai():
2532            ax.plot(
2533                self.member_vars["epoch_last_improved"],
2534                self.member_vars["global_best_validation_score"],
2535                "bo",
2536                label="Global best (y)",
2537            )
2538            ax.plot(
2539                self.member_vars["epoch_last_improved"],
2540                accuracies[self.member_vars["epoch_last_improved"]],
2541                "go",
2542                label="Epoch Last Improved",
2543            )
2544        else:
2545            if self.member_vars["mode"] == "n":
2546                missed_time = (
2547                    self.member_vars["num_epochs_run"]
2548                    - self.member_vars["epoch_last_improved"]
2549                )
2550                ax.plot(
2551                    (len(self.member_vars["n_accuracies"]) - 1) - missed_time,
2552                    self.member_vars["n_accuracies"][-(missed_time + 1)],
2553                    "go",
2554                    label="Epoch Last Improved",
2555                )
2556
2557        # Generate csv file for the values graphed
2558        pd1 = pd.DataFrame(
2559            {"Epochs": np.arange(len(accuracies)), "Validation Scores": accuracies}
2560        )
2561        pd2 = pd.DataFrame(
2562            {
2563                "Epochs": np.arange(len(self.member_vars["running_accuracies"])),
2564                "Validation Running Scores": self.member_vars["running_accuracies"],
2565            }
2566        )
2567        pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2568        for extra_score in extra_scores:
2569            pd2 = pd.DataFrame(
2570                {
2571                    "Epochs": np.arange(len(extra_scores[extra_score])),
2572                    extra_score: extra_scores[extra_score],
2573                }
2574            )
2575            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2576        extra_scores_without_graphing = self.member_vars[
2577            "extra_scores_without_graphing"
2578        ]
2579        for extra_score in extra_scores_without_graphing:
2580            pd2 = pd.DataFrame(
2581                {
2582                    "Epochs": np.arange(
2583                        len(extra_scores_without_graphing[extra_score])
2584                    ),
2585                    extra_score: extra_scores_without_graphing[extra_score],
2586                }
2587            )
2588            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2589        pd1.to_csv(
2590            save_folder + "/" + self.save_name + extra_string + "Scores.csv",
2591            index=False,
2592        )
2593        del pd1, pd2
2594
2595        # Set y min and max to zoom in on important part of axis
2596        if (
2597            len(self.member_vars["switch_epochs"]) > 0
2598            and self.member_vars["switch_epochs"][0] > 0
2599            and self.zooming_graph
2600        ):
2601            if GPA.pai_tracker.member_vars["maximizing_score"]:
2602                min_val = np.array(
2603                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2604                ).mean()
2605                for extra_score in extra_scores:
2606                    min_pot = np.array(
2607                        extra_scores[extra_score][
2608                            0 : self.member_vars["switch_epochs"][0]
2609                        ]
2610                    ).mean()
2611                    if min_pot < min_val:
2612                        min_val = min_pot
2613                ax.set_ylim(ymin=min_val)
2614            else:
2615                max_val = np.array(
2616                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2617                ).mean()
2618                for extra_score in extra_scores:
2619                    max_pot = np.array(
2620                        extra_scores[extra_score][
2621                            0 : self.member_vars["switch_epochs"][0]
2622                        ]
2623                    ).mean()
2624                    if max_pot > max_val:
2625                        max_val = max_pot
2626                ax.set_ylim(ymax=max_val)
2627
2628        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2629
2630        # Draw vertical lines for epochs where a dendrite switch occurred
2631        if GPA.pc.get_drawing_pai() and self.member_vars["doing_pai"]:
2632            color = "r"
2633            for switcher in self.member_vars["switch_epochs"]:
2634                plt.axvline(x=switcher, ymin=0, ymax=1, color=color)
2635                if color == "r":
2636                    color = "b"
2637                else:
2638                    color = "r"
2639        else:
2640            for switcher in self.member_vars["n_switch_epochs"]:
2641                plt.axvline(x=switcher, ymin=0, ymax=1, color="b")

Generate plots and csvs for accuracy

Parameters
  • ax (object): The matplotlib axis to plot on.
  • save_folder (str): The folder to save the plots and csvs in.
  • extra_string (str): An extra string to append to the filenames.
Returns
  • None
def generate_time_plots(self, ax, save_folder, extra_string):
2643    def generate_time_plots(self, ax, save_folder, extra_string):
2644        """
2645        Generate plots and csvs for timing
2646
2647        Parameters
2648        ----------
2649        ax : object
2650            The matplotlib axis to plot on.
2651        save_folder : str
2652            The folder to save the plots and csvs in.
2653        extra_string : str
2654            An extra string to append to the filenames.
2655
2656        Returns
2657        -------
2658        None
2659
2660        """
2661        if self.member_vars["manual_train_switch"]:
2662            ax.plot(
2663                np.arange(len(self.member_vars["n_train_times"])),
2664                self.member_vars["n_train_times"],
2665                label="Normal Epoch Train Times",
2666            )
2667            ax.plot(
2668                np.arange(len(self.member_vars["p_train_times"])),
2669                self.member_vars["p_train_times"],
2670                label="PAI Epoch Train Times",
2671            )
2672            ax.plot(
2673                np.arange(len(self.member_vars["n_val_times"])),
2674                self.member_vars["n_val_times"],
2675                label="Normal Epoch Val Times",
2676            )
2677            ax.plot(
2678                np.arange(len(self.member_vars["p_val_times"])),
2679                self.member_vars["p_val_times"],
2680                label="PAI Epoch Val Times",
2681            )
2682
2683            plt.title(
2684                save_folder + "/" + self.save_name + "times (by train() and eval())"
2685            )
2686            plt.xlabel("Iteration")
2687            plt.ylabel("Epoch Time in Seconds ")
2688            ax.set_ylim(ymin=0)
2689            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2690
2691            pd1 = pd.DataFrame(
2692                {
2693                    "Epochs": np.arange(len(self.member_vars["n_train_times"])),
2694                    "Normal Epoch Train Times": self.member_vars["n_train_times"],
2695                }
2696            )
2697            pd2 = pd.DataFrame(
2698                {
2699                    "Epochs": np.arange(len(self.member_vars["p_train_times"])),
2700                    "PAI Epoch Train Times": self.member_vars["p_train_times"],
2701                }
2702            )
2703            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2704
2705            pd2 = pd.DataFrame(
2706                {
2707                    "Epochs": np.arange(len(self.member_vars["n_val_times"])),
2708                    "Normal Epoch Val Times": self.member_vars["n_val_times"],
2709                }
2710            )
2711            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2712
2713            pd2 = pd.DataFrame(
2714                {
2715                    "Epochs": np.arange(len(self.member_vars["p_val_times"])),
2716                    "PAI Epoch Val Times": self.member_vars["p_val_times"],
2717                }
2718            )
2719            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2720
2721            pd1.to_csv(
2722                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2723                index=False,
2724            )
2725            del pd1, pd2
2726        else:
2727            ax.plot(
2728                np.arange(len(self.member_vars["n_epoch_times"])),
2729                self.member_vars["n_epoch_times"],
2730                label="Normal Epoch Times",
2731            )
2732            ax.plot(
2733                np.arange(len(self.member_vars["p_epoch_times"])),
2734                self.member_vars["p_epoch_times"],
2735                label="PAI Epoch Times",
2736            )
2737
2738            plt.title(
2739                save_folder + "/" + self.save_name + "times (by train() and eval())"
2740            )
2741            plt.xlabel("Iteration")
2742            plt.ylabel("Epoch Time in Seconds ")
2743            ax.set_ylim(ymin=0)
2744            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2745
2746            pd1 = pd.DataFrame(
2747                {
2748                    "Epochs": np.arange(len(self.member_vars["n_epoch_times"])),
2749                    "Normal Epoch Times": self.member_vars["n_epoch_times"],
2750                }
2751            )
2752            pd2 = pd.DataFrame(
2753                {
2754                    "Epochs": np.arange(len(self.member_vars["p_epoch_times"])),
2755                    "PAI Epoch Times": self.member_vars["p_epoch_times"],
2756                }
2757            )
2758            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2759
2760            pd1.to_csv(
2761                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2762                index=False,
2763            )
2764            del pd1, pd2
2765
2766        if self.values_per_train_epoch != -1 and self.values_per_val_epoch != -1:
2767            ax2 = ax.twinx()  # Second axes sharing same x-axis
2768            ax2.set_ylabel("Single Datapoint Time in Seconds")
2769
2770            ax2.plot(
2771                np.arange(len(self.member_vars["n_train_times"])),
2772                np.array(self.member_vars["n_train_times"])
2773                / self.values_per_train_epoch,
2774                linestyle="dashed",
2775                label="Normal Train Item Times",
2776            )
2777            ax2.plot(
2778                np.arange(len(self.member_vars["p_train_times"])),
2779                np.array(self.member_vars["p_train_times"])
2780                / self.values_per_train_epoch,
2781                linestyle="dashed",
2782                label="PAI Train Item Times",
2783            )
2784            ax2.plot(
2785                np.arange(len(self.member_vars["n_val_times"])),
2786                np.array(self.member_vars["n_val_times"]) / self.values_per_val_epoch,
2787                linestyle="dashed",
2788                label="Normal Val Item Times",
2789            )
2790            ax2.plot(
2791                np.arange(len(self.member_vars["p_val_times"])),
2792                np.array(self.member_vars["p_val_times"]) / self.values_per_val_epoch,
2793                linestyle="dashed",
2794                label="PAI Val Item Times",
2795            )
2796            ax2.tick_params(axis="y")
2797            ax2.set_ylim(ymin=0)
2798            ax2.legend(bbox_to_anchor=(1.05, 1), loc="upper left")

Generate plots and csvs for timing

Parameters
  • ax (object): The matplotlib axis to plot on.
  • save_folder (str): The folder to save the plots and csvs in.
  • extra_string (str): An extra string to append to the filenames.
Returns
  • None
def generate_learning_rate_plots(self, ax, save_folder, extra_string):
2800    def generate_learning_rate_plots(self, ax, save_folder, extra_string):
2801        """
2802        Generate plots and csvs for learning rate
2803
2804        Parameters
2805        ----------
2806        ax : object
2807            The matplotlib axis to plot on.
2808        save_folder : str
2809            The folder to save the plots and csvs in.
2810        extra_string : str
2811            An extra string to append to the filenames.
2812
2813        Returns
2814        -------
2815        None
2816
2817        """
2818        ax.plot(
2819            np.arange(len(self.member_vars["training_learning_rates"])),
2820            self.member_vars["training_learning_rates"],
2821            label="learning_rate",
2822        )
2823        plt.title(save_folder + "/" + self.save_name + "learning_rate")
2824        plt.xlabel("Epochs")
2825        plt.ylabel("learning_rate")
2826        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2827
2828        pd1 = pd.DataFrame(
2829            {
2830                "Epochs": np.arange(len(self.member_vars["training_learning_rates"])),
2831                "learning_rate": self.member_vars["training_learning_rates"],
2832            }
2833        )
2834        pd1.to_csv(
2835            save_folder + "/" + self.save_name + extra_string + "learning_rate.csv",
2836            index=False,
2837        )
2838        del pd1

Generate plots and csvs for learning rate

Parameters
  • ax (object): The matplotlib axis to plot on.
  • save_folder (str): The folder to save the plots and csvs in.
  • extra_string (str): An extra string to append to the filenames.
Returns
  • None
def get_current_pb_scores(self):
2840    def get_current_pb_scores(self):
2841        """
2842        Get the latest best PBScore of each dendrite layer, the same numbers
2843        written to the Best PBScores csv.
2844
2845        Returns
2846        -------
2847        dict[str, Any]
2848            Layer name to score.  Empty outside of dendrite scoring phases,
2849            when no candidate dendrites are being scored.
2850
2851
2852        Parameters
2853        ----------
2854        None
2855
2856        """
2857        if not self.member_vars["doing_pai"]:
2858            return {}
2859        if not GPA.pc.get_perforated_backpropagation():
2860            return {}
2861        # Scores only advance while candidate dendrites are being trained
2862        if (
2863            self.member_vars["mode"] != "p"
2864            and not GPA.pc.get_learn_dendrites_live()
2865        ):
2866            return {}
2867
2868        scores = {}
2869        for layer_id in range(len(self.neuron_module_vector)):
2870            if layer_id >= len(self.member_vars["best_scores"]):
2871                continue
2872            layer_scores = self.member_vars["best_scores"][layer_id]
2873            if len(layer_scores) == 0:
2874                continue
2875            score = layer_scores[-1]
2876            if hasattr(score, "item"):
2877                score = score.item()
2878            score = float(score)
2879            if math.isnan(score) or math.isinf(score):
2880                continue
2881            scores[self.neuron_module_vector[layer_id].name] = score
2882        return scores

Get the latest best PBScore of each dendrite layer, the same numbers written to the Best PBScores csv.

Returns
  • dict[str, Any]: Layer name to score. Empty outside of dendrite scoring phases, when no candidate dendrites are being scored.
Parameters
  • None
def generate_dendrite_learning_plots(self, ax, save_folder, extra_string):
2884    def generate_dendrite_learning_plots(self, ax, save_folder, extra_string):
2885        """
2886        Generate dendrite score plots for the tracker.
2887        Also saves csv files associated with the plots.
2888
2889        Parameters
2890        ----------
2891        ax : matplotlib.axes.Axes
2892            Axis used for plotting dendrite-learning curves.
2893        save_folder : str
2894            Directory where plot images and CSV summaries are written.
2895        extra_string : str
2896            Filename suffix used to distinguish this output set.
2897
2898        Returns
2899        -------
2900        None
2901            Saves plots and score CSV files to disk.
2902        """
2903        if self.member_vars["doing_pai"]:
2904            pd1 = None
2905            pd2 = None
2906            num_colors = len(self.neuron_module_vector)
2907
2908            cm = plt.get_cmap("gist_rainbow")
2909            layer_colors = [cm(1.0 * i / max(num_colors, 1)) for i in range(num_colors)]
2910
2911            for layer_id in range(len(self.neuron_module_vector)):
2912                color = layer_colors[layer_id]
2913                ax.plot(
2914                    np.arange(len(self.member_vars["best_scores"][layer_id])),
2915                    self.member_vars["best_scores"][layer_id],
2916                    label=self.neuron_module_vector[layer_id].name,
2917                    color=color,
2918                )
2919
2920                pd2 = pd.DataFrame(
2921                    {
2922                        "Epochs": np.arange(
2923                            len(self.member_vars["best_scores"][layer_id])
2924                        ),
2925                        f"Best ever for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2926                            "best_scores"
2927                        ][
2928                            layer_id
2929                        ],
2930                    }
2931                )
2932
2933                if pd1 is None:
2934                    pd1 = pd2
2935                else:
2936                    pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2937
2938                if len(self.member_vars["current_scores"][layer_id]) != 0:
2939                    ax.plot(
2940                        np.arange(len(self.member_vars["current_scores"][layer_id])),
2941                        self.member_vars["current_scores"][layer_id],
2942                        label=f"Current:{self.neuron_module_vector[layer_id].name}",
2943                        color=color,
2944                        linestyle="--",
2945                    )
2946
2947                pd2 = pd.DataFrame(
2948                    {
2949                        "Epochs": np.arange(
2950                            len(self.member_vars["current_scores"][layer_id])
2951                        ),
2952                        f"Best current for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2953                            "current_scores"
2954                        ][
2955                            layer_id
2956                        ],
2957                    }
2958                )
2959                pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2960
2961            plt.title(save_folder + "/" + self.save_name + " Best PBScores")
2962            plt.xlabel("Epochs")
2963            plt.ylabel("Best PBScore")
2964            ax.legend(
2965                bbox_to_anchor=(1.05, 1),
2966                loc="upper left",
2967                ncol=max(1, math.ceil(len(self.neuron_module_vector) / 30)),
2968            )
2969            for switcher in self.member_vars["p_switch_epochs"]:
2970                plt.axvline(x=switcher, ymin=0, ymax=1, color="r")
2971
2972            if self.member_vars["mode"] == "p":
2973                missed_time = (
2974                    self.member_vars["num_epochs_run"]
2975                    - self.member_vars["epoch_last_improved"]
2976                )
2977                plt.axvline(
2978                    x=(len(self.member_vars["best_scores"][0]) - (missed_time + 1)),
2979                    ymin=0,
2980                    ymax=1,
2981                    color="g",
2982                )
2983
2984            # pd1 here will be none if no PB layers are created
2985            if pd1 is not None:
2986                pd1.to_csv(
2987                    save_folder
2988                    + "/"
2989                    + self.save_name
2990                    + extra_string
2991                    + "Best PBScores.csv",
2992                    index=False,
2993                )
2994            del pd1, pd2

Generate dendrite score plots for the tracker. Also saves csv files associated with the plots.

Parameters
  • ax (matplotlib.axes.Axes): Axis used for plotting dendrite-learning curves.
  • save_folder (str): Directory where plot images and CSV summaries are written.
  • extra_string (str): Filename suffix used to distinguish this output set.
Returns
  • None: Saves plots and score CSV files to disk.
def generate_extra_csv_files(self, save_folder, extra_string):
2996    def generate_extra_csv_files(self, save_folder, extra_string):
2997        """
2998        Generate additional csvs
2999
3000        Parameters
3001        ----------
3002        save_folder : str
3003            The folder to save the plots and csvs in.
3004        extra_string : str
3005            An extra string to append to the filenames.
3006
3007        Returns
3008        -------
3009        None
3010
3011        """
3012        pd1 = pd.DataFrame(
3013            {
3014                "Switch Number": np.arange(len(self.member_vars["switch_epochs"])),
3015                "Switch Epoch": self.member_vars["switch_epochs"],
3016            }
3017        )
3018        pd1.to_csv(
3019            save_folder + "/" + self.save_name + extra_string + "switch_epochs.csv",
3020            index=False,
3021        )
3022        del pd1
3023
3024        pd1 = pd.DataFrame(
3025            {
3026                "Switch Number": np.arange(len(self.member_vars["param_counts"])),
3027                "Param Count": self.member_vars["param_counts"],
3028            }
3029        )
3030        pd1.to_csv(
3031            save_folder + "/" + self.save_name + extra_string + "param_counts.csv",
3032            index=False,
3033        )
3034        del pd1
3035
3036        """
3037        Create best_arch_scores.csv file
3038        When working with dendrites there is a tradeoff between additional param count and score improvement.
3039        This file will help track that tradeoff by recording the best scores for all extra_scores
3040        and extra_scores_without_graphing for each architecture version.
3041        The scores recorded here are from the epoch when the best validation score was found
3042        within each switch_epoch boundary.
3043        """
3044        switch_counts = len(self.member_vars["switch_epochs"])
3045        best_valid = []
3046        associated_params = []
3047        
3048        # Initialize dictionaries to store best scores for each extra score type
3049        best_extra_scores = {}
3050        for score_name in self.member_vars["extra_scores"]:
3051            best_extra_scores[score_name] = []
3052        for score_name in self.member_vars["extra_scores_without_graphing"]:
3053            best_extra_scores[score_name] = []
3054
3055        for switch in range(0, switch_counts, 2):
3056            start_index = 0
3057            if switch != 0:
3058                start_index = self.member_vars["switch_epochs"][switch - 1] + 1
3059            end_index = self.member_vars["switch_epochs"][switch] + 1
3060
3061            if GPA.pai_tracker.member_vars["maximizing_score"]:
3062                best_valid_index = start_index + np.argmax(
3063                    self.member_vars["accuracies"][start_index:end_index]
3064                )
3065            else:
3066                best_valid_index = start_index + np.argmin(
3067                    self.member_vars["accuracies"][start_index:end_index]
3068                )
3069
3070            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3071            best_valid.append(best_valid_score)
3072            
3073            # Get corresponding scores from all extra_scores
3074            for score_name in self.member_vars["extra_scores"]:
3075                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3076                    best_extra_scores[score_name].append(
3077                        self.member_vars["extra_scores"][score_name][best_valid_index]
3078                    )
3079                else:
3080                    best_extra_scores[score_name].append(None)
3081            
3082            # Get corresponding scores from all extra_scores_without_graphing
3083            for score_name in self.member_vars["extra_scores_without_graphing"]:
3084                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3085                    best_extra_scores[score_name].append(
3086                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3087                    )
3088                else:
3089                    best_extra_scores[score_name].append(None)
3090            
3091            if self.member_vars["doing_pai"]:
3092                associated_params.append(self.member_vars["param_counts"][switch])
3093            else:
3094                associated_params.append(self.member_vars["param_counts"][-1])
3095
3096        # If in neuron training mode but not the very first epoch
3097        if self.member_vars["mode"] == "n" and (
3098            (len(self.member_vars["switch_epochs"]) == 0)
3099            or (
3100                self.member_vars["switch_epochs"][-1] + 1
3101                != len(self.member_vars["accuracies"])
3102            )
3103        ):
3104            start_index = 0
3105            if len(self.member_vars["switch_epochs"]) != 0:
3106                start_index = self.member_vars["switch_epochs"][-1] + 1
3107
3108            if GPA.pai_tracker.member_vars["maximizing_score"]:
3109                best_valid_index = start_index + np.argmax(
3110                    self.member_vars["accuracies"][start_index:]
3111                )
3112            else:
3113                best_valid_index = start_index + np.argmin(
3114                    self.member_vars["accuracies"][start_index:]
3115                )
3116
3117            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3118            best_valid.append(best_valid_score)
3119            
3120            # Get corresponding scores from all extra_scores
3121            for score_name in self.member_vars["extra_scores"]:
3122                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3123                    best_extra_scores[score_name].append(
3124                        self.member_vars["extra_scores"][score_name][best_valid_index]
3125                    )
3126                else:
3127                    best_extra_scores[score_name].append(None)
3128            
3129            # Get corresponding scores from all extra_scores_without_graphing
3130            for score_name in self.member_vars["extra_scores_without_graphing"]:
3131                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3132                    best_extra_scores[score_name].append(
3133                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3134                    )
3135                else:
3136                    best_extra_scores[score_name].append(None)
3137            
3138            associated_params.append(self.member_vars["param_counts"][-1])
3139
3140        # Build dataframe with all columns
3141        csv_data = {
3142            "Param Counts": associated_params,
3143            "Max Valid Scores": best_valid,
3144        }
3145        
3146        # Add columns for each extra score
3147        for score_name in best_extra_scores:
3148            csv_data[score_name] = best_extra_scores[score_name]
3149        
3150        pd1 = pd.DataFrame(csv_data)
3151        pd1.to_csv(
3152            save_folder + "/" + self.save_name + extra_string + "_best_arch_scores.csv",
3153            index=False,
3154        )
3155        del pd1

Generate additional csvs

Parameters
  • save_folder (str): The folder to save the plots and csvs in.
  • extra_string (str): An extra string to append to the filenames.
Returns
  • None
def save_graphs(self, extra_string=''):
3157    def save_graphs(self, extra_string=""):
3158        """
3159        Save graphs and csvs for all the values the tracker records
3160
3161        Parameters
3162        ----------
3163        extra_string : str
3164            An extra string to append to the filenames.
3165
3166        Returns
3167        -------
3168        None
3169
3170        """
3171        # If running DDP only save with rank 0
3172        if "RANK" in os.environ:
3173            if int(os.environ["RANK"]) != 0:
3174                return
3175        if not self.making_graphs:
3176            return
3177
3178        save_folder = "./" + self.save_name + "/"
3179
3180        plt.ioff()
3181        fig = plt.figure(figsize=(28, 14))
3182
3183        # Plot with accuracy scores
3184        ax = plt.subplot(221)
3185        self.generate_accuracy_plots(ax, save_folder, extra_string)
3186
3187        # Plot dendrite learning scores
3188        ax = plt.subplot(222)
3189        self.generate_dendrite_learning_plots(ax, save_folder, extra_string)
3190
3191        if GPA.pc.get_drawing_extra_graphs():
3192            # Plot learning rates for each training epoch
3193            ax = plt.subplot(223)
3194            self.generate_learning_rate_plots(ax, save_folder, extra_string)
3195
3196            # Plot the times for each training epoch
3197            ax = plt.subplot(224)
3198            self.generate_time_plots(ax, save_folder, extra_string)
3199
3200        # Generate extra CSV files
3201        self.generate_extra_csv_files(save_folder, extra_string)
3202
3203        fig.tight_layout()
3204        plt.savefig(save_folder + "/" + self.save_name + extra_string + ".png")
3205        plt.close("all")

Save graphs and csvs for all the values the tracker records

Parameters
  • extra_string (str): An extra string to append to the filenames.
Returns
  • None
def add_loss(self, loss):
3207    def add_loss(self, loss):
3208        """Add loss to tracking vectors.
3209
3210        Parameters
3211        ----------
3212        loss : float or int
3213            The loss value to add.
3214
3215        Returns
3216        -------
3217        None
3218
3219        """
3220        if not isinstance(loss, (float, int)):
3221            loss = loss.item()
3222        self.member_vars["training_loss"].append(loss)

Add loss to tracking vectors.

Parameters
  • loss (float or int): The loss value to add.
Returns
  • None
def add_learning_rate(self, learning_rate):
3224    def add_learning_rate(self, learning_rate):
3225        """Add learning rate to tracking vectors.
3226
3227        Parameters
3228        ----------
3229        learning_rate : float or int
3230            The learning rate value to add.
3231
3232        Returns
3233        -------
3234        None
3235
3236        """
3237        if not isinstance(learning_rate, (float, int)):
3238            learning_rate = learning_rate.item()
3239        self.member_vars["training_learning_rates"].append(learning_rate)

Add learning rate to tracking vectors.

Parameters
  • learning_rate (float or int): The learning rate value to add.
Returns
  • None
def add_extra_score(self, score, extra_score_name):
3241    def add_extra_score(self, score, extra_score_name):
3242        """Add extra score to tracking vectors.
3243
3244        Parameters
3245        ----------
3246        score : float or int
3247            The score value to add.
3248
3249        extra_score_name : str
3250            The name of the extra score.
3251
3252        Returns
3253        -------
3254        None
3255
3256        """
3257        if not isinstance(score, (float, int)):
3258            try:
3259                score = score.item()
3260            except:
3261                print(
3262                    "Scores added for Perforated Backpropagation should be "
3263                    "float, int, or tensor, yours is a:"
3264                )
3265                print(type(score))
3266                pdb.set_trace()
3267
3268        if GPA.pc.get_verbose():
3269            print(f"Adding extra score {extra_score_name} of {float(score)}")
3270
3271        if extra_score_name not in self.member_vars["extra_scores"]:
3272            self.member_vars["extra_scores"][extra_score_name] = []
3273        self.member_vars["extra_scores"][extra_score_name].append(score)
3274
3275        if self.member_vars["mode"] == "n":
3276            if extra_score_name not in self.member_vars["n_extra_scores"]:
3277                self.member_vars["n_extra_scores"][extra_score_name] = []
3278            self.member_vars["n_extra_scores"][extra_score_name].append(score)

Add extra score to tracking vectors.

Parameters
  • score (float or int): The score value to add.
  • extra_score_name (str): The name of the extra score.
Returns
  • None
def add_extra_score_without_graphing(self, score, extra_score_name):
3280    def add_extra_score_without_graphing(self, score, extra_score_name):
3281        """Add extra score without graphing to tracking vectors.
3282
3283        Parameters
3284        ----------
3285        score : float or int
3286            The score value to add.
3287
3288        extra_score_name : str
3289            The name of the extra score.
3290
3291        Returns
3292        -------
3293        None
3294
3295        """
3296        if not isinstance(score, (float, int)):
3297            try:
3298                score = score.item()
3299            except:
3300                print(
3301                    "Scores added for Perforated Backpropagation should be "
3302                    "float, int, or tensor, yours is a:"
3303                )
3304                print(type(score))
3305                print("in add_extra_score_without_graphing")
3306                pdb.set_trace()
3307
3308        if GPA.pc.get_verbose():
3309            print(f"Adding extra score {extra_score_name} of {float(score)}")
3310
3311        if extra_score_name not in self.member_vars["extra_scores_without_graphing"]:
3312            self.member_vars["extra_scores_without_graphing"][extra_score_name] = []
3313        self.member_vars["extra_scores_without_graphing"][extra_score_name].append(
3314            score
3315        )

Add extra score without graphing to tracking vectors.

Parameters
  • score (float or int): The score value to add.
  • extra_score_name (str): The name of the extra score.
Returns
  • None
def add_test_score(self, score, extra_score_name):
3317    def add_test_score(self, score, extra_score_name):
3318        """Add test score to tracking vectors.
3319
3320        Parameters
3321        ----------
3322        score : float or int
3323            The score value to add.
3324
3325        extra_score_name : str
3326            The name of the extra score.
3327
3328        Returns
3329        -------
3330        None
3331
3332        Notes
3333        -----
3334        This function is a wrapper around `add_extra_score` that separates
3335        test score for adding to best_arch_scores.csv.
3336
3337        """
3338        self.add_extra_score(score, extra_score_name)
3339
3340        if not isinstance(score, (float, int)):
3341            try:
3342                score = score.item()
3343            except:
3344                print(
3345                    "Scores added for Perforated Backpropagation should be "
3346                    "float, int, or tensor, yours is a:"
3347                )
3348                print(type(score))
3349                print("in add_test_score")
3350                pdb.set_trace()
3351
3352        if GPA.pc.get_verbose():
3353            print(f"Adding test score {extra_score_name} of {float(score)}")
3354        self.member_vars["test_scores"].append(score)

Add test score to tracking vectors.

Parameters
  • score (float or int): The score value to add.
  • extra_score_name (str): The name of the extra score.
Returns
  • None
Notes

This function is a wrapper around add_extra_score that separates test score for adding to best_arch_scores.csv.

def add_validation_score(self, accuracy, net, force_switch=False):
3356    def add_validation_score(self, accuracy, net, force_switch=False):
3357        """Function to add the validation score.
3358
3359        This is complex because it determines neuron and dendrite switching.
3360
3361        Parameters
3362        ----------
3363        accuracy : float or int
3364            The accuracy or loss value to add.
3365        net : object
3366            The neural network model.
3367        force_switch : bool, optional
3368            Whether to force a switch, by default False.
3369
3370        Returns
3371        -------
3372        net : object
3373            The potentially modified neural network model.
3374        training_complete : bool
3375            Whether training is complete.
3376        restructured : bool
3377            Whether the model has been restructured.
3378
3379        Notes
3380        -----
3381        WARNING: Do not call self anywhere in this function. When systems
3382        get loaded the actual tracker you are working with can change.
3383        """
3384
3385        _pai_log("info", f"Adding validation score {accuracy:.8f}")
3386
3387        update_learning_rate()
3388        update_param_count(net)
3389
3390        accuracy = check_input_problems(net, accuracy)
3391
3392        if len(GPA.pai_tracker.member_vars["switch_epochs"]) == 0:
3393            epochs_since_cycle_switch = GPA.pai_tracker.member_vars["num_epochs_run"]
3394        else:
3395            epochs_since_cycle_switch = (
3396                GPA.pai_tracker.member_vars["num_epochs_run"]
3397                - GPA.pai_tracker.member_vars["switch_epochs"][-1]
3398            )
3399
3400        update_running_accuracy(accuracy, epochs_since_cycle_switch)
3401        if GPA.pc.get_perforated_backpropagation():
3402            TPB.update_pb_scores(self)
3403
3404        # Captured before any switch below flips the mode and reloads scores
3405        epoch_pb_scores = self.get_current_pb_scores()
3406
3407        GPA.pai_tracker.stop_epoch(internal_call=True)
3408
3409        # If it is neuron training mode
3410        if (
3411            GPA.pai_tracker.member_vars["mode"] == "n"
3412            or GPA.pc.get_learn_dendrites_live()
3413        ):
3414            check_new_best(net, accuracy, epochs_since_cycle_switch)
3415        elif GPA.pc.get_perforated_backpropagation():
3416            TPB.check_best_pai_score_improvement()
3417
3418        # Save the latest model
3419        if GPA.pc.get_test_saves():
3420            UPA.save_system(net, GPA.pc.get_save_name(), "latest")
3421        if GPA.pc.get_pai_saves():
3422            UPA.pai_save_system(net, GPA.pc.get_save_name(), "latest")
3423
3424        restructuring_status_value = NO_MODEL_UPDATE
3425        # If it is time to switch based on scores and counter or a manual switch
3426        if GPA.pai_tracker.switch_time() or force_switch:
3427            # If testing dendrite capacity switch after enough dendrites added
3428            if (
3429                (GPA.pai_tracker.member_vars["mode"] == "n")
3430                and (GPA.pai_tracker.member_vars["num_dendrites_added"] > 2)
3431                and GPA.pc.get_testing_dendrite_capacity()
3432            ):
3433                GPA.pai_tracker.save_graphs()
3434                _pai_log(
3435                    "info",
3436                    "Successfully added 3 dendrites with GPA.pc.set_testing_dendrite_capacity(True) (default). "
3437                    "You may now set that to False and run a real experiment.",
3438                )
3439                return net, False, True
3440
3441            # If doing neuron training but this dendrite count didn't improve
3442            if (
3443                (GPA.pai_tracker.member_vars["mode"] == "n")
3444                or GPA.pc.get_learn_dendrites_live()
3445            ) and (GPA.pai_tracker.member_vars["current_n_set_global_best"] is False):
3446                new_restructuring_status_value, net = process_no_improvement(net)
3447                # if this was the final try return that training is complete
3448                if new_restructuring_status_value == TRAINING_COMPLETE:
3449                    if _dashboard_emitter is not None:
3450                        _dashboard_emitter.emit_run_end(GPA.pc)
3451                    return net, True, True
3452                else:
3453                    restructuring_status_value = update_restructuring_status(
3454                        restructuring_status_value, new_restructuring_status_value
3455                    )
3456            # Else if did improve, do a normal switch process
3457            else:
3458                if GPA.pc.get_verbose():
3459                    print(
3460                        f"Calling switch_mode with "
3461                        f'{GPA.pai_tracker.member_vars["current_n_set_global_best"]}, '
3462                        f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}, '
3463                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]}, '
3464                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_value"]},'
3465                        f'{GPA.pc.get_max_dendrites()},'
3466                        f'{GPA.pai_tracker.member_vars["num_dendrites_added"]},'
3467                        f'{GPA.pai_tracker.member_vars["num_dendrite_tries"]},'
3468                    )
3469                import pdb; pdb.set_trace
3470                # If the max number of dendrites has been hit or not doing pai and adding dendtites
3471                # then return rather than adding more
3472                if (
3473                    (GPA.pai_tracker.member_vars["mode"] == "n")
3474                    and (
3475                        GPA.pc.get_max_dendrites()
3476                        == GPA.pai_tracker.member_vars["num_dendrites_added"]
3477                    )
3478                ) or (GPA.pai_tracker.member_vars["doing_pai"] is False):
3479                    if GPA.pc.get_verbose():
3480                        print(
3481                            "Max dendrites reached or not doing PAI, finishing training"
3482                        )
3483                    net = process_final_network(net)
3484                    # Increment integrated if we have dendrites (means they're integrated)
3485                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3486                        GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3487                        _pai_log("info", f"Final dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3488                        if _dashboard_emitter is not None:
3489                            _dashboard_emitter.emit_dendrite_added(
3490                                GPA.pc,
3491                                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3492                                num_dendrites_integrated=GPA.pai_tracker.member_vars[
3493                                    "num_dendrites_integrated"
3494                                ],
3495                            )
3496                    if _dashboard_emitter is not None:
3497                        _dashboard_emitter.emit_run_end(GPA.pc)
3498                    return net, True, True
3499
3500                # Otherwise if its neuron training mode reset the counter of failed dendrites
3501                # Check if we should increment integrated count BEFORE change_learning_modes loads old state
3502                should_increment_integrated = False
3503                if GPA.pai_tracker.member_vars["mode"] == "n":
3504                    GPA.pai_tracker.member_vars["num_dendrite_tries"] = 0
3505                    if GPA.pc.get_verbose():
3506                        print(
3507                            "Adding new dendrites without resetting which means "
3508                            "the last ones improved. Resetting num_dendrite_tries"
3509                        )
3510                    # Remember to increment after change_learning_modes (which loads old tracker state)
3511                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3512                        should_increment_integrated = True
3513
3514                GPA.pai_tracker.save_graphs(
3515                    f'_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}'
3516                )
3517
3518                if GPA.pc.get_test_saves():
3519                    UPA.save_system(
3520                        net,
3521                        GPA.pc.get_save_name(),
3522                        f'beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3523                    )
3524                    # Copy current best model from this set of dendrites
3525                    # If running DDP only copy with rank 0
3526                    if "RANK" not in os.environ or int(os.environ["RANK"]) == 0:
3527                        shutil.copyfile(
3528                            f"{GPA.pc.get_save_name()}/best_model.pt",
3529                            f'{GPA.pc.get_save_name()}/best_model_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}.pt',
3530                        )
3531
3532                net = UPA.change_learning_modes(
3533                    net,
3534                    GPA.pc.get_save_name(),
3535                    "best_model",
3536                    GPA.pai_tracker.member_vars["doing_pai"],
3537                )
3538                restructuring_status_value = NETWORK_RESTRUCTURED
3539                
3540                # Now increment after change_learning_modes has loaded the best model
3541                # This ensures the increment persists and doesn't get overwritten
3542                if should_increment_integrated:
3543                    GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3544                    _pai_log("info", f"Dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3545                    if _dashboard_emitter is not None:
3546                        _dashboard_emitter.emit_dendrite_added(
3547                            GPA.pc,
3548                            epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3549                            num_dendrites_integrated=GPA.pai_tracker.member_vars[
3550                                "num_dendrites_integrated"
3551                            ],
3552                        )
3553
3554            # If restructured is true, clear scheduler/optimizer before saving
3555            if restructuring_status_value != NETWORK_RESTRUCTURED:
3556                print(
3557                    "Restructured should always be triggered here, let us know if you encounter this situation"
3558                )
3559                pdb.set_trace()
3560
3561            # Since there is a restructuring optimizer and scheduler must be reinitialized after return
3562            GPA.pai_tracker.clear_optimizer_and_scheduler()
3563
3564            # Save the model from after the switch
3565            UPA.save_system(
3566                net,
3567                GPA.pc.get_save_name(),
3568                f'switch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3569            )
3570
3571        # If not time to switch and you have a scheduler, perform the update step
3572        elif GPA.pai_tracker.member_vars["scheduler"] is not None:
3573            new_restructuring_status_value, net = process_scheduler_update(
3574                net, accuracy, epochs_since_cycle_switch
3575            )
3576            restructuring_status_value = update_restructuring_status(
3577                restructuring_status_value, new_restructuring_status_value
3578            )
3579
3580        GPA.pai_tracker.start_epoch(internal_call=True)
3581        if _dashboard_emitter is not None:
3582            _mv = GPA.pai_tracker.member_vars
3583            _lr = _mv["training_learning_rates"][-1] if _mv["training_learning_rates"] else None
3584            _train_score = _mv["extra_scores"].get("train", [None])[-1]
3585            _n_times = _mv["n_epoch_times"] or [(_mv["n_train_times"][-1] + _mv["n_val_times"][-1]) if (_mv["n_train_times"] and _mv["n_val_times"]) else None]
3586            _p_times = _mv["p_epoch_times"] or [(_mv["p_train_times"][-1] + _mv["p_val_times"][-1]) if (_mv["p_train_times"] and _mv["p_val_times"]) else None]
3587            _dashboard_emitter.emit_epoch(
3588                GPA.pc,
3589                epoch=_mv["num_epochs_run"],
3590                validation_score=accuracy,
3591                learning_rate=_lr,
3592                train_score=_train_score,
3593                normal_time=_n_times[-1],
3594                pai_time=_p_times[-1],
3595                pb_scores=epoch_pb_scores,
3596            )
3597        GPA.pai_tracker.save_graphs()
3598
3599        if restructuring_status_value == NETWORK_RESTRUCTURED:
3600            GPA.pai_tracker.member_vars["epoch_last_improved"] = (
3601                GPA.pai_tracker.member_vars["num_epochs_run"]
3602            )
3603            if GPA.pc.get_verbose():
3604                print(
3605                    f"Setting epoch last improved to "
3606                    f'{GPA.pai_tracker.member_vars["epoch_last_improved"]}'
3607                )
3608
3609            now = datetime.now()
3610            dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
3611
3612            if GPA.pc.get_verbose():
3613                print("Not saving restructure right now")
3614
3615            """
3616            This block of code helped with a save issue with safetensors and huggingface, but it breaks DDP.  
3617            Temporarily removing it to avoid DDP issues, but if you encounter save issues try adding it back in.
3618            for param in net.parameters():
3619                param.data = param.data.contiguous()
3620            """
3621        if GPA.pc.get_verbose():
3622            print(
3623                f"Completed adding score. Restructured is {restructuring_status_value}, "
3624                f"\ncurrent switch list is:"
3625            )
3626            print(GPA.pai_tracker.member_vars["switch_epochs"])
3627
3628        if _dashboard_emitter is not None and restructuring_status_value == NETWORK_RESTRUCTURED:
3629            _param_count = UPA.count_params(net)
3630            _dashboard_emitter.emit_switch(
3631                GPA.pc,
3632                switch_number=GPA.pai_tracker.member_vars["num_dendrites_added"],
3633                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3634                param_count=_param_count,
3635                switch_type=GPA.pai_tracker.member_vars["mode"],
3636            )
3637
3638        # Always False for training complete if nothing triggered that training is over
3639        return net, restructuring_status_value, False

Function to add the validation score.

This is complex because it determines neuron and dendrite switching.

Parameters
  • accuracy (float or int): The accuracy or loss value to add.
  • net (object): The neural network model.
  • force_switch (bool, optional): Whether to force a switch, by default False.
Returns
  • net (object): The potentially modified neural network model.
  • training_complete (bool): Whether training is complete.
  • restructured (bool): Whether the model has been restructured.
Notes

WARNING: Do not call self anywhere in this function. When systems get loaded the actual tracker you are working with can change.

def clear_all_processors(self):
3641    def clear_all_processors(self):
3642        """Clear all processors from modules.
3643
3644        Parameters
3645        ----------
3646        None
3647
3648        Returns
3649        -------
3650        None
3651            This function does not return a value.
3652        """
3653        for module in self.neuron_module_vector:
3654            module.clear_processors()

Clear all processors from modules.

Parameters
  • None
Returns
  • None: This function does not return a value.
def create_new_dendrite_module(self):
3656    def create_new_dendrite_module(self):
3657        """Add dendrite module to all neuron modules.
3658
3659        Parameters
3660        ----------
3661        None
3662
3663        Returns
3664        -------
3665        None
3666            This function does not return a value.
3667        """
3668        for module in self.neuron_module_vector:
3669            module.create_new_dendrite_module()

Add dendrite module to all neuron modules.

Parameters
  • None
Returns
  • None: This function does not return a value.
def set_create_dendrite_global(self, fn):
3671    def set_create_dendrite_global(self, fn):
3672        """Call set_create_dendrite(fn) on every tracked PAINeuronModule."""
3673        for module in self.neuron_module_vector:
3674            module.set_create_dendrite(fn)

Call set_create_dendrite(fn) on every tracked PAINeuronModule.

def set_dendrite_loss_fn_global(self, fn):
3676    def set_dendrite_loss_fn_global(self, fn):
3677        """Set the global dendrite loss function used by all dendrite modules."""
3678        from perforatedbp import modules_pbp as MPB
3679        MPB.dendrite_loss_fn = fn

Set the global dendrite loss function used by all dendrite modules.

def apply_pb_grads(self):
3681    def apply_pb_grads(self):
3682        """Apply perforated backpropagation gradients to all modules.
3683
3684        Parameters
3685        ----------
3686        None
3687
3688        Returns
3689        -------
3690        None
3691            This function does not return a value.
3692        """
3693        if self.member_vars["mode"] == "p":
3694            for module in self.neuron_module_vector:
3695                module.apply_pb_grads()

Apply perforated backpropagation gradients to all modules.

Parameters
  • None
Returns
  • None: This function does not return a value.
def apply_pb_zero(self):
3697    def apply_pb_zero(self):
3698        """Apply perforated backpropagation zero gradients to all modules.
3699
3700        Parameters
3701        ----------
3702        None
3703
3704        Returns
3705        -------
3706        None
3707            This function does not return a value.
3708        """
3709        if self.member_vars["mode"] == "p":
3710            for module in self.neuron_module_vector:
3711                module.apply_pb_zero()

Apply perforated backpropagation zero gradients to all modules.

Parameters
  • None
Returns
  • None: This function does not return a value.