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            f.write(
1552                f"{layer.name},{layer.dendrite_module.dendrite_values[0].out_channels}\n"
1553            )
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            channels[line.split(",")[0]] = int(line.split(",")[1])
1586        for layer in self.neuron_module_vector:
1587            layer.dendrite_module.dendrite_values[0].setup_arrays(channels[layer.name])
1588
1589    def set_optimizer_instance(self, optimizer_instance, additional_optimizers=[]):
1590        """Set optimizer instance directly.
1591
1592        Parameters
1593        ----------
1594        optimizer_instance : object
1595            The optimizer instance to set.
1596
1597        Returns
1598        -------
1599        None
1600
1601        """
1602
1603        try:
1604            for param_group in optimizer_instance.param_groups:
1605                if (
1606                    param_group["weight_decay"] > 0
1607                    and GPA.pc.get_weight_decay_accepted() is False
1608                ):
1609                    _pai_log(
1610                        "warning",
1611                        "For PAI training it is recommended to not use weight decay in your optimizer",
1612                    )
1613
1614        except:
1615            pass
1616        self.member_vars["optimizer_instance"] = optimizer_instance
1617        if GPA.pc.get_perforated_backpropagation():
1618            TPB.setup_optimizer_pb(self.member_vars["optimizer_instance"])
1619            for optimizer in additional_optimizers:
1620                TPB.filter_params(optimizer)
1621        optimizer_instance.zero_grad()
1622
1623    def set_optimizer(self, optimizer):
1624        """Set optimizer type to be initialized later
1625
1626        Parameters
1627        ----------
1628        optimizer : object
1629            The optimizer type to set.
1630
1631        Returns
1632        -------
1633        None
1634
1635        """
1636        self.member_vars["optimizer"] = optimizer
1637
1638    def set_scheduler(self, scheduler):
1639        """Set scheduler type to be initialized later
1640
1641        Parameters
1642        ----------
1643        scheduler : object
1644            The scheduler type to set.
1645
1646        Returns
1647        -------
1648        None
1649
1650        """
1651        if scheduler is not torch.optim.lr_scheduler.ReduceLROnPlateau:
1652            if GPA.pc.get_verbose():
1653                print("Not using ReduceLROnPlateau, this is not recommended")
1654        self.member_vars["scheduler"] = scheduler
1655
1656    def increment_scheduler(self, num_ticks, mode):
1657        """Increment the scheduler a set number of times.
1658
1659        Used for finding best initial learning rate when adding dendrites.
1660
1661        Parameters
1662        ----------
1663        num_ticks : int
1664            The number of scheduler steps to take.
1665        mode : str
1666            The mode for stepping the scheduler. Options are:
1667            - "step_learning_rate": Step based on improved accuracy epochs
1668            - "increment_epoch_count": Step based on total epoch count
1669
1670        Returns
1671        -------
1672        current_steps : int
1673            The number of learning rate changes that occurred.
1674        learning_rate1 : float
1675            The final learning rate after stepping.
1676
1677        """
1678
1679        current_steps = 0
1680        current_ticker = 0
1681
1682        for param_group in GPA.pai_tracker.member_vars[
1683            "optimizer_instance"
1684        ].param_groups:
1685            learning_rate1 = param_group["lr"]
1686
1687        if GPA.pc.get_verbose():
1688            print("Using scheduler:")
1689            print(type(self.member_vars["scheduler_instance"]))
1690
1691        while current_ticker < num_ticks:
1692            if GPA.pc.get_verbose():
1693                print(
1694                    f"Lower start rate initial {learning_rate1} "
1695                    f'stepping {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} times'
1696                )
1697
1698            if (
1699                type(self.member_vars["scheduler_instance"])
1700                is torch.optim.lr_scheduler.ReduceLROnPlateau
1701            ):
1702                if mode == "step_learning_rate":
1703                    # Step with counter as last improved accuracy
1704                    self.member_vars["scheduler_instance"].step(
1705                        metrics=self.member_vars["last_improved_accuracies"][
1706                            GPA.pai_tracker.steps_after_switch() - 1
1707                        ]
1708                    )
1709                elif mode == "increment_epoch_count":
1710                    # Step with improved epoch counts up to current location
1711                    self.member_vars["scheduler_instance"].step(
1712                        metrics=self.member_vars["last_improved_accuracies"][
1713                            -((num_ticks - 1) - current_ticker) - 1
1714                        ]
1715                    )
1716            else:
1717                self.member_vars["scheduler_instance"].step()
1718
1719            for param_group in GPA.pai_tracker.member_vars[
1720                "optimizer_instance"
1721            ].param_groups:
1722                learning_rate2 = param_group["lr"]
1723
1724            if learning_rate2 != learning_rate1:
1725                current_steps += 1
1726                learning_rate1 = learning_rate2
1727                if mode == "step_learning_rate":
1728                    current_ticker += 1
1729                if GPA.pc.get_verbose():
1730                    print(f"1 step {current_steps} to {learning_rate2}")
1731
1732            if mode == "increment_epoch_count":
1733                current_ticker += 1
1734
1735        return current_steps, learning_rate1
1736
1737    def setup_optimizer(self, net, opt_args, sched_args=None, parameters=None):
1738        """Initialize the optimizer and scheduler when added.
1739
1740        Parameters
1741        ----------
1742        net : object
1743            The neural network model.
1744        opt_args : dict
1745            The arguments for the optimizer.
1746        sched_args : dict, optional
1747            The arguments for the scheduler, by default None.
1748
1749        Returns
1750        -------
1751        optimizer : object
1752            The initialized optimizer instance.
1753        scheduler : object, optional
1754            The initialized scheduler instance, if a scheduler was set.
1755
1756        """
1757        if "weight_decay" in opt_args and not GPA.pc.get_weight_decay_accepted():
1758            _pai_log(
1759                "warning",
1760                "For PAI training it is recommended to not use weight decay in your optimizer",
1761            )
1762
1763        if ("model" not in opt_args.keys()) and "params" not in opt_args.keys():
1764            print("In setup_optimizer it will be depreciated to not pass in params yourself in the future")
1765            print("please change the settings to include params")
1766            if self.member_vars["mode"] == "n":
1767                if parameters is not None:
1768                    opt_args["params"] = parameters
1769                else:
1770                    opt_args["params"] = filter(lambda p: p.requires_grad, net.parameters())
1771            else:
1772                params = UPA.get_pai_network_params(net)
1773                if parameters is not None:
1774                    # Filter parameters to only those in params, preserving weight_decay
1775                    params_set = set(params)
1776                    filtered_params = []
1777                    for param_group in parameters:
1778                        filtered_group_params = [p for p in param_group["params"] if p in params_set]
1779                        if filtered_group_params:
1780                            filtered_params.append({
1781                                "params": filtered_group_params,
1782                                "weight_decay": param_group["weight_decay"]
1783                            })
1784                    opt_args["params"] = filtered_params
1785                else:
1786                    opt_args["params"] = params
1787        elif "params" in opt_args.keys():
1788            # Check if params is a list of param groups (dicts) or a single param group
1789            params_value = opt_args["params"]
1790            if isinstance(params_value, list) and len(params_value) > 0:
1791                # Check if it's a list of dicts (multiple param groups) or list of tensors (single group)
1792                if isinstance(params_value[0], dict):
1793                    # Multiple param groups format: [{"params": [...], "lr": ...}, ...]
1794                    # Filter each param group for requires_grad
1795                    filtered_param_groups = []
1796                    for param_group in params_value:
1797                        filtered_group_params = [p for p in param_group["params"] if p.requires_grad]
1798                        if filtered_group_params:
1799                            new_group = param_group.copy()
1800                            new_group["params"] = filtered_group_params
1801                            filtered_param_groups.append(new_group)
1802                    opt_args["params"] = filtered_param_groups
1803                else:
1804                    # Single param group format: [tensor1, tensor2, ...] or generator
1805                    # Filter for requires_grad
1806                    opt_args["params"] = [p for p in params_value if p.requires_grad]
1807            elif hasattr(params_value, '__iter__'):
1808                # Handle generators or other iterables
1809                opt_args["params"] = [p for p in params_value if p.requires_grad]
1810
1811        optimizer = self.member_vars["optimizer"](**opt_args)
1812        self.set_optimizer_instance(optimizer)
1813
1814        if self.member_vars["scheduler"] is not None:
1815            # Handle SequentialLR specially
1816            if self.member_vars["scheduler"] is torch.optim.lr_scheduler.SequentialLR:
1817                """
1818                sched_args should be a dict with "schedulers" (list of tuples) and "milestones"
1819                For example:
1820                sequential_schedArgs = {
1821                    "schedulers": [
1822                        (warmup_scheduler_class, warmup_schedArgs),
1823                        (main_scheduler_class, main_schedArgs)
1824                    ],
1825                    "milestones": [switch_epoch]
1826                }
1827                """
1828                schedulers = []
1829                milestones = sched_args.get("milestones", [])
1830                scheduler_configs = sched_args.get("schedulers", [])
1831                
1832                for scheduler_class, scheduler_args in scheduler_configs:
1833                    schedulers.append(scheduler_class(optimizer, **scheduler_args))
1834                
1835                self.member_vars["scheduler_instance"] = torch.optim.lr_scheduler.SequentialLR(
1836                    optimizer, schedulers=schedulers, milestones=milestones
1837                )
1838            else:
1839                self.member_vars["scheduler_instance"] = self.member_vars["scheduler"](
1840                    optimizer, **sched_args
1841                )
1842            current_steps = 0
1843
1844            for param_group in GPA.pai_tracker.member_vars[
1845                "optimizer_instance"
1846            ].param_groups:
1847                learning_rate1 = param_group["lr"]
1848
1849            if GPA.pc.get_verbose():
1850                print(
1851                    f"Resetting scheduler with {GPA.pai_tracker.steps_after_switch()} "
1852                    f"steps and {GPA.pc.get_initial_history_after_switches()} initial ticks to skip"
1853                )
1854
1855            # Find setting of previously used learning rate before adding dendrites
1856            if (
1857                GPA.pai_tracker.member_vars[
1858                    "current_n_learning_rate_initial_skip_steps"
1859                ]
1860                != 0
1861            ):
1862                additional_steps, learning_rate1 = self.increment_scheduler(
1863                    GPA.pai_tracker.member_vars[
1864                        "current_n_learning_rate_initial_skip_steps"
1865                    ],
1866                    "step_learning_rate",
1867                )
1868                current_steps += additional_steps
1869
1870            if self.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
1871                initial = GPA.pc.get_initial_history_after_switches()
1872            else:
1873                initial = 0
1874
1875            if GPA.pai_tracker.steps_after_switch() > initial:
1876                # Minus extra 1 because this gets called after start epoch
1877                additional_steps, learning_rate1 = self.increment_scheduler(
1878                    (GPA.pai_tracker.steps_after_switch() - initial) - 1,
1879                    "increment_epoch_count",
1880                )
1881                current_steps += additional_steps
1882
1883            if GPA.pc.get_verbose():
1884                print(
1885                    f"Scheduler update loop with {current_steps} "
1886                    f"ended with {learning_rate1}"
1887                )
1888                print(
1889                    f"Scheduler ended with {current_steps} steps "
1890                    f"and lr of {learning_rate1}"
1891                )
1892
1893            self.member_vars["current_step_count"] = current_steps
1894            return optimizer, self.member_vars["scheduler_instance"]
1895        else:
1896            return optimizer
1897
1898    def clear_optimizer_and_scheduler(self):
1899        """Clear the instances for saving.
1900
1901        Parameters
1902        ----------
1903        None
1904
1905        Returns
1906        -------
1907        None
1908            This function does not return a value.
1909        """
1910        self.member_vars["optimizer_instance"] = None
1911        self.member_vars["scheduler_instance"] = None
1912
1913    def switch_time(self):
1914        """Determine if it's time to switch between neuron and dendrite training.
1915
1916        Parameters
1917        ----------
1918        None
1919
1920        Returns
1921        -------
1922        bool
1923            True if it's time to switch, False otherwise.
1924
1925        Notes
1926        -----
1927        Based on current settings and history of scores.
1928        """
1929
1930        switch_phrase = "No mode, this should never be the case."
1931        switch_number = GPA.pc.get_n_epochs_to_switch()
1932        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1933            switch_phrase = "DOING_SWITCH_EVERY_TIME"
1934        elif self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY:
1935            switch_phrase = "DOING_HISTORY"
1936        elif self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH:
1937            switch_phrase = "DOING_FIXED_SWITCH"
1938            switch_number = GPA.pc.get_fixed_switch_num()
1939        elif self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1940            switch_phrase = "DOING_NO_SWITCH"
1941        else:
1942            print(
1943                "A switch mode must be set.  Check your settings for GPA.pc.set_switch_mode()."
1944            )
1945            pdb.set_trace()
1946        if not GPA.pc.get_silent():
1947            if(GPA.pc.get_perforated_backpropagation()):
1948                print(
1949                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1950                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1951                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1952                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1953                    f'n: {switch_number}, p: {GPA.pc.get_p_epochs_to_switch()}, '
1954                    f'num_cycles: {self.member_vars["num_cycles"]}'
1955                )
1956            else:
1957                print(
1958                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1959                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1960                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1961                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1962                    f'n: {switch_number}, num_cycles: {self.member_vars["num_cycles"]}'
1963                )
1964            print(
1965                f'  Score tracking: current_n_set_global_best={self.member_vars["current_n_set_global_best"]}, '
1966                f'global_best={self.member_vars["global_best_validation_score"]:.4f}, '
1967                f'current_best={self.member_vars["current_best_validation_score"]:.4f}'
1968            )
1969        if GPA.pc.get_perforated_backpropagation():
1970            # this will fill in epoch last improved
1971            TPB.best_pai_score_improved_this_epoch(self)  ## CLOSED ONLY
1972        if self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1973            if not GPA.pc.get_silent():
1974                print("Returning False - doing no switch mode")
1975            return False
1976
1977        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1978            if not GPA.pc.get_silent():
1979                print("Returning True - switching every time")
1980            return True
1981
1982        # Check if we're in the middle of learning rate optimization
1983        # If so, block ALL switch triggers until committed
1984        if GPA.pc.get_verbose():
1985            print("=== LR Optimization Check ===")
1986            print(f'  mode == "n": {self.member_vars["mode"] == "n"}')
1987            print(f"  get_learn_dendrites_live(): {GPA.pc.get_learn_dendrites_live()}")
1988            print(f'  committed_to_initial_rate: {GPA.pai_tracker.member_vars["committed_to_initial_rate"]}')
1989            print(f"  get_dont_give_up_unless_learning_rate_lowered(): {GPA.pc.get_dont_give_up_unless_learning_rate_lowered()}")
1990            print(f'  current_n_learning_rate_initial_skip_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"]}')
1991            print(f'  last_max_learning_rate_steps: {self.member_vars["last_max_learning_rate_steps"]}')
1992            print(f'  skip_steps < max_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"] < self.member_vars["last_max_learning_rate_steps"]}')
1993            print(f'  scheduler is not None: {self.member_vars["scheduler"] is not None}')
1994            print("=============================")
1995        
1996        if (
1997            ((self.member_vars["mode"] == "n") or GPA.pc.get_learn_dendrites_live())
1998            and (GPA.pai_tracker.member_vars["committed_to_initial_rate"] is False)
1999            and (GPA.pc.get_dont_give_up_unless_learning_rate_lowered())
2000            and (
2001                self.member_vars["current_n_learning_rate_initial_skip_steps"]
2002                <= self.member_vars["last_max_learning_rate_steps"]
2003            )
2004            and self.member_vars["scheduler"] is not None
2005        ):
2006            if not GPA.pc.get_silent():
2007                print(
2008                    f"Returning False - learning rate optimization in progress. "
2009                    f"Not committed yet. Comparing "
2010                    f'initial {self.member_vars["current_n_learning_rate_initial_skip_steps"]} '
2011                    f'to last max {self.member_vars["last_max_learning_rate_steps"]}'
2012                )
2013            return False
2014
2015        if len(self.member_vars["switch_epochs"]) == 0:
2016            this_count = self.member_vars["num_epochs_run"]
2017        else:
2018            this_count = (
2019                self.member_vars["num_epochs_run"]
2020                - self.member_vars["switch_epochs"][-1]
2021            )
2022        cap_switch = False
2023        if GPA.pc.get_perforated_backpropagation():
2024            cap_switch = TPB.check_cap_switch(self, this_count)
2025
2026        if self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY and (
2027            (
2028                (self.member_vars["mode"] == "n")
2029                and (
2030                    self.member_vars["num_epochs_run"]
2031                    - self.member_vars["epoch_last_improved"]
2032                    >= GPA.pc.get_n_epochs_to_switch()
2033                )
2034                and this_count
2035                >= GPA.pc.get_initial_history_after_switches()
2036                + GPA.pc.get_n_epochs_to_switch()
2037            )
2038            or (GPA.pc.get_perforated_backpropagation() and TPB.history_switch(self))
2039            or cap_switch
2040        ):
2041            if not GPA.pc.get_silent():
2042                print("Returning True - History and last improved is hit")
2043            return True
2044
2045        if self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH and (
2046            (
2047                self.member_vars["total_epochs_run"] % GPA.pc.get_fixed_switch_num()
2048                == GPA.pc.get_fixed_switch_num() - 1
2049            )
2050            and self.member_vars["num_epochs_run"]
2051            >= GPA.pc.get_first_fixed_switch_num() - 1
2052        ):
2053            if not GPA.pc.get_silent():
2054                print("Returning True - Fixed switch number is hit")
2055            return True
2056
2057        if not GPA.pc.get_silent():
2058            print("Returning False - no triggers to switch have been hit")
2059        return False
2060
2061    def steps_after_switch(self):
2062        """Based on settings, return value for steps since a switch.
2063
2064        Different options for param vals setting determine what is returned.
2065
2066        Parameters
2067        ----------
2068        None
2069
2070        Returns
2071        -------
2072        int
2073            The number of epochs since the last switch, or total epochs run,
2074            depending on settings.
2075
2076        """
2077        if self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_TOTAL_EPOCH:
2078            return self.member_vars["num_epochs_run"]
2079        elif (
2080            self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_UPDATE_EPOCH
2081        ):
2082            return self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2083        elif (
2084            self.member_vars["param_vals_setting"]
2085            == GPA.pc.PARAM_VALS_BY_NEURON_EPOCH_START
2086        ):
2087            if self.member_vars["mode"] == "p":
2088                return (
2089                    self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2090                )
2091            else:
2092                return self.member_vars["num_epochs_run"]
2093        else:
2094            print(
2095                f'{self.member_vars["param_vals_setting"]} is not a valid param vals option'
2096            )
2097            pdb.set_trace()
2098
2099    def add_pai_neuron_module(self, new_module, initial_add=True):
2100        """Add neuron modules to internal vectors.
2101
2102        Parameters
2103        ----------
2104        new_module : object
2105            The new module to add.
2106        initial_add : bool, optional
2107            Whether this is the initial addition rather than loading from file
2108
2109        Returns
2110        -------
2111        None
2112
2113        """
2114
2115        # If it's a duplicate, ignore the second addition
2116        if new_module in self.neuron_module_vector:
2117            return
2118        self.neuron_module_vector.append(new_module)
2119        if self.member_vars["doing_pai"]:
2120            PA.set_wrapped_params(new_module)
2121        if initial_add:
2122            self.member_vars["best_scores"].append([])
2123            self.member_vars["current_scores"].append([])
2124
2125    def add_tracked_neuron_module(self, new_module, initial_add=True):
2126        """Add tracked modules to internal vectors
2127
2128        Parameters
2129        ----------
2130        new_module : object
2131            The new module to add.
2132        initial_add : bool, optional
2133            Whether this is the initial addition rather than loading from file
2134
2135        Returns
2136        -------
2137        None
2138
2139        """
2140        # If it's a duplicate, ignore the second addition
2141        if new_module in self.tracked_neuron_module_vector:
2142            return
2143        self.tracked_neuron_module_vector.append(new_module)
2144        if self.member_vars["doing_pai"]:
2145            PA.set_tracked_params(new_module)
2146
2147    def reset_module_vector(self, net, load_from_restart):
2148        """Clear internal vectors and reset from network.
2149
2150        Parameters
2151        ----------
2152        net : object
2153            The neural network model.
2154        load_from_restart : bool
2155            Whether loading from a restart file.
2156
2157        Returns
2158        -------
2159        None
2160
2161        """
2162        self.neuron_module_vector = []
2163        self.tracked_neuron_module_vector = []
2164        this_list = UPA.get_pai_modules(net, 0)
2165        for module in this_list:
2166            self.add_pai_neuron_module(module, initial_add=load_from_restart)
2167        this_list = UPA.get_tracked_modules(net, 0)
2168        for module in this_list:
2169            self.add_tracked_neuron_module(module, initial_add=load_from_restart)
2170
2171    def reset_vals_for_score_reset(self):
2172        """Reset cycle scores for new cycle.
2173
2174        Parameters
2175        ----------
2176        None
2177
2178        Returns
2179        -------
2180        None
2181            This function does not return a value.
2182        """
2183
2184        if GPA.pc.get_find_best_lr():
2185            self.member_vars["committed_to_initial_rate"] = False
2186            print("Resetting committed to initial rate to False")
2187        # If retaining all dendrties always say that the current dendrites set global best for saving and loading
2188        if GPA.pc.get_retain_all_dendrites():
2189            self.member_vars["current_n_set_global_best"] = True
2190            self.member_vars["global_best_validation_score"] = 0
2191        else:
2192            self.member_vars["current_n_set_global_best"] = False
2193
2194        # Don't reset global best, but do reset current best
2195        self.member_vars["current_best_validation_score"] = 0
2196        self.member_vars["initial_lr_test_epoch_count"] = -1
2197
2198    def set_dendrite_training(self):
2199        """Signal all layers to start dendrite training.
2200
2201        Parameters
2202        ----------
2203        None
2204
2205        Returns
2206        -------
2207        None
2208            This function does not return a value.
2209        """
2210        if GPA.pc.get_verbose():
2211            print("Calling set_dendrite_training")
2212
2213        for layer in self.neuron_module_vector[:]:
2214            worked = layer.set_mode("p")
2215            """
2216            worked is False when a layer was added to the neuron module vector
2217            but then it's never actually been used. This can happen when
2218            you have set a layer to have requires_grad = False or when
2219            you have a module as a member variable but it's not actually
2220            part of the network. Should be moved to be a tracked layer
2221            rather than a neuron layer.
2222            """
2223            if not worked:
2224                self.neuron_module_vector.remove(layer)
2225
2226        for layer in self.tracked_neuron_module_vector[:]:
2227            worked = layer.set_mode("p")
2228
2229        self.create_new_dendrite_module()
2230        self.member_vars["mode"] = "p"
2231        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2232
2233        if GPA.pc.get_learn_dendrites_live():
2234            self.reset_vals_for_score_reset()
2235
2236        self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2237            "current_step_count"
2238        ]
2239
2240        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
2241        GPA.pai_tracker.member_vars["num_cycles"] += 1
2242
2243
2244    def set_neuron_training(self):
2245        """Signal all layers to start neuron training.
2246
2247        Parameters
2248        ----------
2249        None
2250
2251        Returns
2252        -------
2253        None
2254            This function does not return a value.
2255        """
2256        for module in self.neuron_module_vector:
2257            module.set_mode("n")
2258        for module in self.tracked_neuron_module_vector[:]:
2259            module.set_mode("n")
2260
2261        self.member_vars["mode"] = "n"
2262        self.member_vars["num_dendrites_added"] += 1
2263        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2264        self.reset_vals_for_score_reset()
2265
2266        self.member_vars["current_cycle_lr_max_scores"] = []
2267        if GPA.pc.get_learn_dendrites_live():
2268            self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2269                "current_step_count"
2270            ]
2271        GPA.pai_tracker.member_vars["num_cycles"] += 1
2272
2273        if GPA.pc.get_reset_best_score_on_switch():
2274            GPA.pai_tracker.member_vars["current_best_validation_score"] = 0
2275            GPA.pai_tracker.member_vars["running_accuracy"] = 0
2276
2277    def start_epoch(self, internal_call=False):
2278        """Perform steps for when a new training epoch is about to begin.
2279
2280        Parameters
2281        ----------
2282        internal_call : bool, optional
2283            Whether this is an internal call or manual call
2284
2285        Returns
2286        -------
2287        None
2288
2289        Notes
2290        -----
2291        If you ever need to call this manually, set internal_call to False.
2292
2293        """
2294        if self.member_vars["manual_train_switch"] and internal_call:
2295            return
2296
2297        if not internal_call and not self.member_vars["manual_train_switch"]:
2298            self.member_vars["manual_train_switch"] = True
2299            self.saved_time = 0
2300            self.member_vars["num_epochs_run"] = -1
2301            self.member_vars["total_epochs_run"] = -1
2302
2303        end = time.time()
2304        if self.member_vars["manual_train_switch"]:
2305            if self.saved_time != 0:
2306                if self.member_vars["mode"] == "p":
2307                    self.member_vars["p_val_times"].append(end - self.saved_time)
2308                else:
2309                    self.member_vars["n_val_times"].append(end - self.saved_time)
2310
2311        if self.member_vars["mode"] == "p":
2312            for layer in self.neuron_module_vector:
2313                for m in range(0, GPA.pc.get_global_candidates()):
2314                    with torch.no_grad():
2315                        if GPA.pc.get_verbose():
2316                            print(f"Resetting score for {layer.name}")
2317                        # Snapshot best_score before reset so we can compute per-epoch improvement
2318                        layer.dendrite_module.dendrite_values[
2319                            m
2320                        ].epoch_start_best_score.copy_(
2321                            layer.dendrite_module.dendrite_values[
2322                                m
2323                            ].best_score.detach()
2324                        )
2325                        layer.dendrite_module.dendrite_values[
2326                            m
2327                        ].best_score_improved_this_epoch = (
2328                            layer.dendrite_module.dendrite_values[
2329                                m
2330                            ].best_score_improved_this_epoch
2331                            * 0
2332                        )
2333                        layer.dendrite_module.dendrite_values[
2334                            m
2335                        ].nodes_best_improved_this_epoch = (
2336                            layer.dendrite_module.dendrite_values[
2337                                m
2338                            ].nodes_best_improved_this_epoch
2339                            * 0
2340                        )
2341                        layer.dendrite_module.dendrite_values[
2342                            m
2343                        ].nodes_improved_any = (
2344                            layer.dendrite_module.dendrite_values[
2345                                m
2346                            ].nodes_improved_any
2347                            * 0
2348                        )
2349            if GPA.pc.get_perforated_backpropagation():
2350                self.member_vars["best_mean_score_improved_this_epoch"] = 0
2351        self.member_vars["num_epochs_run"] += 1
2352        self.member_vars["total_epochs_run"] = (
2353            self.member_vars["num_epochs_run"] + self.member_vars["overwritten_epochs"]
2354        )
2355        self.saved_time = end
2356
2357    def stop_epoch(self, internal_call=False):
2358        """Perform steps when a training epoch has completed.
2359
2360        Parameters
2361        ----------
2362        internal_call : bool, optional
2363            Whether this is an internal call or manual call
2364
2365        Returns
2366        -------
2367        None
2368
2369        Notes
2370        -----
2371        If you ever need to call this manually, set internal_call to False.
2372
2373        """
2374        end = time.time()
2375        if self.member_vars["manual_train_switch"] and internal_call:
2376            return
2377
2378        if self.member_vars["manual_train_switch"]:
2379            if self.member_vars["mode"] == "p":
2380                self.member_vars["p_train_times"].append(end - self.saved_time)
2381            else:
2382                self.member_vars["n_train_times"].append(end - self.saved_time)
2383        else:
2384            if self.member_vars["mode"] == "p":
2385                self.member_vars["p_epoch_times"].append(end - self.saved_time)
2386            else:
2387                self.member_vars["n_epoch_times"].append(end - self.saved_time)
2388
2389        self.saved_time = end
2390
2391    def initialize(
2392        self,
2393        model,
2394        doing_pai=True,
2395        save_name="PAI",
2396        making_graphs=True,
2397        maximizing_score=True,
2398        num_classes=10000,
2399        values_per_train_epoch=-1,
2400        values_per_val_epoch=-1,
2401        zooming_graph=True,
2402    ):
2403        """Setup the tracker with initial settings.
2404
2405
2406        Parameters
2407        ----------
2408        model : object
2409            The neural network model.
2410        doing_pai : bool, optional
2411            Whether to add dendrites, by default True.
2412        save_name : str, optional
2413            The name under which to save the model.
2414        making_graphs : bool, optional
2415            Whether to make graphs, by default True.
2416        maximizing_score : bool, optional
2417            Whether to maximize the score, by default True.
2418        num_classes : int, optional
2419            The number of classes in the dataset, unused
2420        values_per_train_epoch : int, optional
2421            The number of values to look back for graphing
2422            during training, by default -1 (all values).
2423        values_per_val_epoch : int, optional
2424            The number of values to look back for graphing
2425            during validation, by default -1 (all values).
2426        zooming_graph : bool, optional
2427            Whether to zoom on graphs, by default True.
2428
2429
2430        Returns
2431        -------
2432        nn.Module
2433            Converted model instance configured for the tracker settings.
2434        """
2435        model = UPA.convert_network(model)
2436        self.member_vars["doing_pai"] = doing_pai
2437        self.member_vars["maximizing_score"] = maximizing_score
2438        self.save_name = save_name
2439        self.zooming_graph = zooming_graph
2440        self.making_graphs = making_graphs
2441
2442        if not self.loaded:
2443            self.member_vars["running_accuracy"] = (1.0 / num_classes) * 100
2444
2445        self.values_per_train_epoch = values_per_train_epoch
2446        self.values_per_val_epoch = values_per_val_epoch
2447
2448        if GPA.pc.get_testing_dendrite_capacity():
2449            if not GPA.pc.get_silent():
2450                print("Running a test of Dendrite Capacity.")
2451            GPA.pc.set_switch_mode(GPA.pc.DOING_SWITCH_EVERY_TIME)
2452            self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
2453            GPA.pc.set_retain_all_dendrites(True)
2454            GPA.pc.set_max_dendrite_tries(1000)
2455            GPA.pc.set_max_dendrites(1000)
2456            if GPA.pc.get_perforated_backpropagation():
2457                GPA.pc.set_initial_correlation_batches(1)
2458        else:
2459            if not GPA.pc.get_silent():
2460                print("Running Dendrite Experiment")
2461        return model
2462
2463    def generate_accuracy_plots(self, ax, save_folder, extra_string):
2464        """
2465        Generate plots and csvs for accuracy
2466
2467        Parameters
2468        ----------
2469        ax : object
2470            The matplotlib axis to plot on.
2471        save_folder : str
2472            The folder to save the plots and csvs in.
2473        extra_string : str
2474            An extra string to append to the filenames.
2475
2476        Returns
2477        -------
2478        None
2479
2480        """
2481
2482        # If scores are being saved for epochs that get overwritten, plot them
2483        for list_id in range(len(self.member_vars["overwritten_extras"])):
2484            for extra_id in self.member_vars["overwritten_extras"][list_id]:
2485                ax.plot(
2486                    np.arange(
2487                        len(self.member_vars["overwritten_extras"][list_id][extra_id])
2488                    ),
2489                    self.member_vars["overwritten_extras"][list_id][extra_id],
2490                    "r",
2491                )
2492            ax.plot(
2493                np.arange(len(self.member_vars["overwritten_vals"][list_id])),
2494                self.member_vars["overwritten_vals"][list_id],
2495                "b",
2496            )
2497
2498        # Determine which accuracy vector to use
2499        if GPA.pc.get_drawing_pai():
2500            accuracies = self.member_vars["accuracies"]
2501        else:
2502            accuracies = self.member_vars["n_accuracies"]
2503
2504        # Get pointer to additional scores being saved
2505        extra_scores = self.member_vars["extra_scores"]
2506
2507        # Plot the main accuracy scores
2508        ax.plot(np.arange(len(accuracies)), accuracies, label="Validation Scores")
2509        ax.plot(
2510            np.arange(len(self.member_vars["running_accuracies"])),
2511            self.member_vars["running_accuracies"],
2512            label="Validation Running Scores",
2513        )
2514
2515        # Plot additional scores
2516        for extra_score in extra_scores:
2517            ax.plot(
2518                np.arange(len(extra_scores[extra_score])),
2519                extra_scores[extra_score],
2520                label=extra_score,
2521            )
2522
2523        plt.title(save_folder + "/" + self.save_name + "Scores")
2524        plt.xlabel("Epochs")
2525        plt.ylabel("Score")
2526
2527        # Add point at epoch last improved and best validation score
2528        if GPA.pc.get_drawing_pai():
2529            ax.plot(
2530                self.member_vars["epoch_last_improved"],
2531                self.member_vars["global_best_validation_score"],
2532                "bo",
2533                label="Global best (y)",
2534            )
2535            ax.plot(
2536                self.member_vars["epoch_last_improved"],
2537                accuracies[self.member_vars["epoch_last_improved"]],
2538                "go",
2539                label="Epoch Last Improved",
2540            )
2541        else:
2542            if self.member_vars["mode"] == "n":
2543                missed_time = (
2544                    self.member_vars["num_epochs_run"]
2545                    - self.member_vars["epoch_last_improved"]
2546                )
2547                ax.plot(
2548                    (len(self.member_vars["n_accuracies"]) - 1) - missed_time,
2549                    self.member_vars["n_accuracies"][-(missed_time + 1)],
2550                    "go",
2551                    label="Epoch Last Improved",
2552                )
2553
2554        # Generate csv file for the values graphed
2555        pd1 = pd.DataFrame(
2556            {"Epochs": np.arange(len(accuracies)), "Validation Scores": accuracies}
2557        )
2558        pd2 = pd.DataFrame(
2559            {
2560                "Epochs": np.arange(len(self.member_vars["running_accuracies"])),
2561                "Validation Running Scores": self.member_vars["running_accuracies"],
2562            }
2563        )
2564        pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2565        for extra_score in extra_scores:
2566            pd2 = pd.DataFrame(
2567                {
2568                    "Epochs": np.arange(len(extra_scores[extra_score])),
2569                    extra_score: extra_scores[extra_score],
2570                }
2571            )
2572            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2573        extra_scores_without_graphing = self.member_vars[
2574            "extra_scores_without_graphing"
2575        ]
2576        for extra_score in extra_scores_without_graphing:
2577            pd2 = pd.DataFrame(
2578                {
2579                    "Epochs": np.arange(
2580                        len(extra_scores_without_graphing[extra_score])
2581                    ),
2582                    extra_score: extra_scores_without_graphing[extra_score],
2583                }
2584            )
2585            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2586        pd1.to_csv(
2587            save_folder + "/" + self.save_name + extra_string + "Scores.csv",
2588            index=False,
2589        )
2590        del pd1, pd2
2591
2592        # Set y min and max to zoom in on important part of axis
2593        if (
2594            len(self.member_vars["switch_epochs"]) > 0
2595            and self.member_vars["switch_epochs"][0] > 0
2596            and self.zooming_graph
2597        ):
2598            if GPA.pai_tracker.member_vars["maximizing_score"]:
2599                min_val = np.array(
2600                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2601                ).mean()
2602                for extra_score in extra_scores:
2603                    min_pot = np.array(
2604                        extra_scores[extra_score][
2605                            0 : self.member_vars["switch_epochs"][0]
2606                        ]
2607                    ).mean()
2608                    if min_pot < min_val:
2609                        min_val = min_pot
2610                ax.set_ylim(ymin=min_val)
2611            else:
2612                max_val = np.array(
2613                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2614                ).mean()
2615                for extra_score in extra_scores:
2616                    max_pot = np.array(
2617                        extra_scores[extra_score][
2618                            0 : self.member_vars["switch_epochs"][0]
2619                        ]
2620                    ).mean()
2621                    if max_pot > max_val:
2622                        max_val = max_pot
2623                ax.set_ylim(ymax=max_val)
2624
2625        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2626
2627        # Draw vertical lines for epochs where a dendrite switch occurred
2628        if GPA.pc.get_drawing_pai() and self.member_vars["doing_pai"]:
2629            color = "r"
2630            for switcher in self.member_vars["switch_epochs"]:
2631                plt.axvline(x=switcher, ymin=0, ymax=1, color=color)
2632                if color == "r":
2633                    color = "b"
2634                else:
2635                    color = "r"
2636        else:
2637            for switcher in self.member_vars["n_switch_epochs"]:
2638                plt.axvline(x=switcher, ymin=0, ymax=1, color="b")
2639
2640    def generate_time_plots(self, ax, save_folder, extra_string):
2641        """
2642        Generate plots and csvs for timing
2643
2644        Parameters
2645        ----------
2646        ax : object
2647            The matplotlib axis to plot on.
2648        save_folder : str
2649            The folder to save the plots and csvs in.
2650        extra_string : str
2651            An extra string to append to the filenames.
2652
2653        Returns
2654        -------
2655        None
2656
2657        """
2658        if self.member_vars["manual_train_switch"]:
2659            ax.plot(
2660                np.arange(len(self.member_vars["n_train_times"])),
2661                self.member_vars["n_train_times"],
2662                label="Normal Epoch Train Times",
2663            )
2664            ax.plot(
2665                np.arange(len(self.member_vars["p_train_times"])),
2666                self.member_vars["p_train_times"],
2667                label="PAI Epoch Train Times",
2668            )
2669            ax.plot(
2670                np.arange(len(self.member_vars["n_val_times"])),
2671                self.member_vars["n_val_times"],
2672                label="Normal Epoch Val Times",
2673            )
2674            ax.plot(
2675                np.arange(len(self.member_vars["p_val_times"])),
2676                self.member_vars["p_val_times"],
2677                label="PAI Epoch Val Times",
2678            )
2679
2680            plt.title(
2681                save_folder + "/" + self.save_name + "times (by train() and eval())"
2682            )
2683            plt.xlabel("Iteration")
2684            plt.ylabel("Epoch Time in Seconds ")
2685            ax.set_ylim(ymin=0)
2686            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2687
2688            pd1 = pd.DataFrame(
2689                {
2690                    "Epochs": np.arange(len(self.member_vars["n_train_times"])),
2691                    "Normal Epoch Train Times": self.member_vars["n_train_times"],
2692                }
2693            )
2694            pd2 = pd.DataFrame(
2695                {
2696                    "Epochs": np.arange(len(self.member_vars["p_train_times"])),
2697                    "PAI Epoch Train Times": self.member_vars["p_train_times"],
2698                }
2699            )
2700            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2701
2702            pd2 = pd.DataFrame(
2703                {
2704                    "Epochs": np.arange(len(self.member_vars["n_val_times"])),
2705                    "Normal Epoch Val Times": self.member_vars["n_val_times"],
2706                }
2707            )
2708            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2709
2710            pd2 = pd.DataFrame(
2711                {
2712                    "Epochs": np.arange(len(self.member_vars["p_val_times"])),
2713                    "PAI Epoch Val Times": self.member_vars["p_val_times"],
2714                }
2715            )
2716            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2717
2718            pd1.to_csv(
2719                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2720                index=False,
2721            )
2722            del pd1, pd2
2723        else:
2724            ax.plot(
2725                np.arange(len(self.member_vars["n_epoch_times"])),
2726                self.member_vars["n_epoch_times"],
2727                label="Normal Epoch Times",
2728            )
2729            ax.plot(
2730                np.arange(len(self.member_vars["p_epoch_times"])),
2731                self.member_vars["p_epoch_times"],
2732                label="PAI Epoch Times",
2733            )
2734
2735            plt.title(
2736                save_folder + "/" + self.save_name + "times (by train() and eval())"
2737            )
2738            plt.xlabel("Iteration")
2739            plt.ylabel("Epoch Time in Seconds ")
2740            ax.set_ylim(ymin=0)
2741            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2742
2743            pd1 = pd.DataFrame(
2744                {
2745                    "Epochs": np.arange(len(self.member_vars["n_epoch_times"])),
2746                    "Normal Epoch Times": self.member_vars["n_epoch_times"],
2747                }
2748            )
2749            pd2 = pd.DataFrame(
2750                {
2751                    "Epochs": np.arange(len(self.member_vars["p_epoch_times"])),
2752                    "PAI Epoch Times": self.member_vars["p_epoch_times"],
2753                }
2754            )
2755            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2756
2757            pd1.to_csv(
2758                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2759                index=False,
2760            )
2761            del pd1, pd2
2762
2763        if self.values_per_train_epoch != -1 and self.values_per_val_epoch != -1:
2764            ax2 = ax.twinx()  # Second axes sharing same x-axis
2765            ax2.set_ylabel("Single Datapoint Time in Seconds")
2766
2767            ax2.plot(
2768                np.arange(len(self.member_vars["n_train_times"])),
2769                np.array(self.member_vars["n_train_times"])
2770                / self.values_per_train_epoch,
2771                linestyle="dashed",
2772                label="Normal Train Item Times",
2773            )
2774            ax2.plot(
2775                np.arange(len(self.member_vars["p_train_times"])),
2776                np.array(self.member_vars["p_train_times"])
2777                / self.values_per_train_epoch,
2778                linestyle="dashed",
2779                label="PAI Train Item Times",
2780            )
2781            ax2.plot(
2782                np.arange(len(self.member_vars["n_val_times"])),
2783                np.array(self.member_vars["n_val_times"]) / self.values_per_val_epoch,
2784                linestyle="dashed",
2785                label="Normal Val Item Times",
2786            )
2787            ax2.plot(
2788                np.arange(len(self.member_vars["p_val_times"])),
2789                np.array(self.member_vars["p_val_times"]) / self.values_per_val_epoch,
2790                linestyle="dashed",
2791                label="PAI Val Item Times",
2792            )
2793            ax2.tick_params(axis="y")
2794            ax2.set_ylim(ymin=0)
2795            ax2.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2796
2797    def generate_learning_rate_plots(self, ax, save_folder, extra_string):
2798        """
2799        Generate plots and csvs for learning rate
2800
2801        Parameters
2802        ----------
2803        ax : object
2804            The matplotlib axis to plot on.
2805        save_folder : str
2806            The folder to save the plots and csvs in.
2807        extra_string : str
2808            An extra string to append to the filenames.
2809
2810        Returns
2811        -------
2812        None
2813
2814        """
2815        ax.plot(
2816            np.arange(len(self.member_vars["training_learning_rates"])),
2817            self.member_vars["training_learning_rates"],
2818            label="learning_rate",
2819        )
2820        plt.title(save_folder + "/" + self.save_name + "learning_rate")
2821        plt.xlabel("Epochs")
2822        plt.ylabel("learning_rate")
2823        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2824
2825        pd1 = pd.DataFrame(
2826            {
2827                "Epochs": np.arange(len(self.member_vars["training_learning_rates"])),
2828                "learning_rate": self.member_vars["training_learning_rates"],
2829            }
2830        )
2831        pd1.to_csv(
2832            save_folder + "/" + self.save_name + extra_string + "learning_rate.csv",
2833            index=False,
2834        )
2835        del pd1
2836
2837    def get_current_pb_scores(self):
2838        """
2839        Get the latest best PBScore of each dendrite layer, the same numbers
2840        written to the Best PBScores csv.
2841
2842        Returns
2843        -------
2844        dict[str, Any]
2845            Layer name to score.  Empty outside of dendrite scoring phases,
2846            when no candidate dendrites are being scored.
2847
2848
2849        Parameters
2850        ----------
2851        None
2852
2853        """
2854        if not self.member_vars["doing_pai"]:
2855            return {}
2856        if not GPA.pc.get_perforated_backpropagation():
2857            return {}
2858        # Scores only advance while candidate dendrites are being trained
2859        if (
2860            self.member_vars["mode"] != "p"
2861            and not GPA.pc.get_learn_dendrites_live()
2862        ):
2863            return {}
2864
2865        scores = {}
2866        for layer_id in range(len(self.neuron_module_vector)):
2867            if layer_id >= len(self.member_vars["best_scores"]):
2868                continue
2869            layer_scores = self.member_vars["best_scores"][layer_id]
2870            if len(layer_scores) == 0:
2871                continue
2872            score = layer_scores[-1]
2873            if hasattr(score, "item"):
2874                score = score.item()
2875            score = float(score)
2876            if math.isnan(score) or math.isinf(score):
2877                continue
2878            scores[self.neuron_module_vector[layer_id].name] = score
2879        return scores
2880
2881    def generate_dendrite_learning_plots(self, ax, save_folder, extra_string):
2882        """
2883        Generate dendrite score plots for the tracker.
2884        Also saves csv files associated with the plots.
2885
2886        Parameters
2887        ----------
2888        ax : matplotlib.axes.Axes
2889            Axis used for plotting dendrite-learning curves.
2890        save_folder : str
2891            Directory where plot images and CSV summaries are written.
2892        extra_string : str
2893            Filename suffix used to distinguish this output set.
2894
2895        Returns
2896        -------
2897        None
2898            Saves plots and score CSV files to disk.
2899        """
2900        if self.member_vars["doing_pai"]:
2901            pd1 = None
2902            pd2 = None
2903            num_colors = len(self.neuron_module_vector)
2904
2905            if (
2906                len(self.neuron_module_vector) > 0
2907                and len(self.member_vars["current_scores"][0]) != 0
2908            ):
2909                num_colors *= 2
2910
2911            cm = plt.get_cmap("gist_rainbow")
2912            ax.set_prop_cycle(
2913                "color", [cm(1.0 * i / num_colors) for i in range(num_colors)]
2914            )
2915
2916            for layer_id in range(len(self.neuron_module_vector)):
2917                ax.plot(
2918                    np.arange(len(self.member_vars["best_scores"][layer_id])),
2919                    self.member_vars["best_scores"][layer_id],
2920                    label=self.neuron_module_vector[layer_id].name,
2921                )
2922
2923                pd2 = pd.DataFrame(
2924                    {
2925                        "Epochs": np.arange(
2926                            len(self.member_vars["best_scores"][layer_id])
2927                        ),
2928                        f"Best ever for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2929                            "best_scores"
2930                        ][
2931                            layer_id
2932                        ],
2933                    }
2934                )
2935
2936                if pd1 is None:
2937                    pd1 = pd2
2938                else:
2939                    pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2940
2941                if len(self.member_vars["current_scores"][layer_id]) != 0:
2942                    ax.plot(
2943                        np.arange(len(self.member_vars["current_scores"][layer_id])),
2944                        self.member_vars["current_scores"][layer_id],
2945                        label=f"Current:{self.neuron_module_vector[layer_id].name}",
2946                    )
2947
2948                pd2 = pd.DataFrame(
2949                    {
2950                        "Epochs": np.arange(
2951                            len(self.member_vars["current_scores"][layer_id])
2952                        ),
2953                        f"Best current for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2954                            "current_scores"
2955                        ][
2956                            layer_id
2957                        ],
2958                    }
2959                )
2960                pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2961
2962            plt.title(save_folder + "/" + self.save_name + " Best PBScores")
2963            plt.xlabel("Epochs")
2964            plt.ylabel("Best PBScore")
2965            ax.legend(
2966                bbox_to_anchor=(1.05, 1),
2967                loc="upper left",
2968                ncol=max(1, math.ceil(len(self.neuron_module_vector) / 30)),
2969            )
2970            for switcher in self.member_vars["p_switch_epochs"]:
2971                plt.axvline(x=switcher, ymin=0, ymax=1, color="r")
2972
2973            if self.member_vars["mode"] == "p":
2974                missed_time = (
2975                    self.member_vars["num_epochs_run"]
2976                    - self.member_vars["epoch_last_improved"]
2977                )
2978                plt.axvline(
2979                    x=(len(self.member_vars["best_scores"][0]) - (missed_time + 1)),
2980                    ymin=0,
2981                    ymax=1,
2982                    color="g",
2983                )
2984
2985            # pd1 here will be none if no PB layers are created
2986            if pd1 is not None:
2987                pd1.to_csv(
2988                    save_folder
2989                    + "/"
2990                    + self.save_name
2991                    + extra_string
2992                    + "Best PBScores.csv",
2993                    index=False,
2994                )
2995            del pd1, pd2
2996
2997    def generate_extra_csv_files(self, save_folder, extra_string):
2998        """
2999        Generate additional csvs
3000
3001        Parameters
3002        ----------
3003        save_folder : str
3004            The folder to save the plots and csvs in.
3005        extra_string : str
3006            An extra string to append to the filenames.
3007
3008        Returns
3009        -------
3010        None
3011
3012        """
3013        pd1 = pd.DataFrame(
3014            {
3015                "Switch Number": np.arange(len(self.member_vars["switch_epochs"])),
3016                "Switch Epoch": self.member_vars["switch_epochs"],
3017            }
3018        )
3019        pd1.to_csv(
3020            save_folder + "/" + self.save_name + extra_string + "switch_epochs.csv",
3021            index=False,
3022        )
3023        del pd1
3024
3025        pd1 = pd.DataFrame(
3026            {
3027                "Switch Number": np.arange(len(self.member_vars["param_counts"])),
3028                "Param Count": self.member_vars["param_counts"],
3029            }
3030        )
3031        pd1.to_csv(
3032            save_folder + "/" + self.save_name + extra_string + "param_counts.csv",
3033            index=False,
3034        )
3035        del pd1
3036
3037        """
3038        Create best_arch_scores.csv file
3039        When working with dendrites there is a tradeoff between additional param count and score improvement.
3040        This file will help track that tradeoff by recording the best scores for all extra_scores
3041        and extra_scores_without_graphing for each architecture version.
3042        The scores recorded here are from the epoch when the best validation score was found
3043        within each switch_epoch boundary.
3044        """
3045        switch_counts = len(self.member_vars["switch_epochs"])
3046        best_valid = []
3047        associated_params = []
3048        
3049        # Initialize dictionaries to store best scores for each extra score type
3050        best_extra_scores = {}
3051        for score_name in self.member_vars["extra_scores"]:
3052            best_extra_scores[score_name] = []
3053        for score_name in self.member_vars["extra_scores_without_graphing"]:
3054            best_extra_scores[score_name] = []
3055
3056        for switch in range(0, switch_counts, 2):
3057            start_index = 0
3058            if switch != 0:
3059                start_index = self.member_vars["switch_epochs"][switch - 1] + 1
3060            end_index = self.member_vars["switch_epochs"][switch] + 1
3061
3062            if GPA.pai_tracker.member_vars["maximizing_score"]:
3063                best_valid_index = start_index + np.argmax(
3064                    self.member_vars["accuracies"][start_index:end_index]
3065                )
3066            else:
3067                best_valid_index = start_index + np.argmin(
3068                    self.member_vars["accuracies"][start_index:end_index]
3069                )
3070
3071            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3072            best_valid.append(best_valid_score)
3073            
3074            # Get corresponding scores from all extra_scores
3075            for score_name in self.member_vars["extra_scores"]:
3076                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3077                    best_extra_scores[score_name].append(
3078                        self.member_vars["extra_scores"][score_name][best_valid_index]
3079                    )
3080                else:
3081                    best_extra_scores[score_name].append(None)
3082            
3083            # Get corresponding scores from all extra_scores_without_graphing
3084            for score_name in self.member_vars["extra_scores_without_graphing"]:
3085                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3086                    best_extra_scores[score_name].append(
3087                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3088                    )
3089                else:
3090                    best_extra_scores[score_name].append(None)
3091            
3092            if self.member_vars["doing_pai"]:
3093                associated_params.append(self.member_vars["param_counts"][switch])
3094            else:
3095                associated_params.append(self.member_vars["param_counts"][-1])
3096
3097        # If in neuron training mode but not the very first epoch
3098        if self.member_vars["mode"] == "n" and (
3099            (len(self.member_vars["switch_epochs"]) == 0)
3100            or (
3101                self.member_vars["switch_epochs"][-1] + 1
3102                != len(self.member_vars["accuracies"])
3103            )
3104        ):
3105            start_index = 0
3106            if len(self.member_vars["switch_epochs"]) != 0:
3107                start_index = self.member_vars["switch_epochs"][-1] + 1
3108
3109            if GPA.pai_tracker.member_vars["maximizing_score"]:
3110                best_valid_index = start_index + np.argmax(
3111                    self.member_vars["accuracies"][start_index:]
3112                )
3113            else:
3114                best_valid_index = start_index + np.argmin(
3115                    self.member_vars["accuracies"][start_index:]
3116                )
3117
3118            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3119            best_valid.append(best_valid_score)
3120            
3121            # Get corresponding scores from all extra_scores
3122            for score_name in self.member_vars["extra_scores"]:
3123                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3124                    best_extra_scores[score_name].append(
3125                        self.member_vars["extra_scores"][score_name][best_valid_index]
3126                    )
3127                else:
3128                    best_extra_scores[score_name].append(None)
3129            
3130            # Get corresponding scores from all extra_scores_without_graphing
3131            for score_name in self.member_vars["extra_scores_without_graphing"]:
3132                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3133                    best_extra_scores[score_name].append(
3134                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3135                    )
3136                else:
3137                    best_extra_scores[score_name].append(None)
3138            
3139            associated_params.append(self.member_vars["param_counts"][-1])
3140
3141        # Build dataframe with all columns
3142        csv_data = {
3143            "Param Counts": associated_params,
3144            "Max Valid Scores": best_valid,
3145        }
3146        
3147        # Add columns for each extra score
3148        for score_name in best_extra_scores:
3149            csv_data[score_name] = best_extra_scores[score_name]
3150        
3151        pd1 = pd.DataFrame(csv_data)
3152        pd1.to_csv(
3153            save_folder + "/" + self.save_name + extra_string + "_best_arch_scores.csv",
3154            index=False,
3155        )
3156        del pd1
3157
3158    def save_graphs(self, extra_string=""):
3159        """
3160        Save graphs and csvs for all the values the tracker records
3161
3162        Parameters
3163        ----------
3164        extra_string : str
3165            An extra string to append to the filenames.
3166
3167        Returns
3168        -------
3169        None
3170
3171        """
3172        # If running DDP only save with rank 0
3173        if "RANK" in os.environ:
3174            if int(os.environ["RANK"]) != 0:
3175                return
3176        if not self.making_graphs:
3177            return
3178
3179        save_folder = "./" + self.save_name + "/"
3180
3181        plt.ioff()
3182        fig = plt.figure(figsize=(28, 14))
3183
3184        # Plot with accuracy scores
3185        ax = plt.subplot(221)
3186        self.generate_accuracy_plots(ax, save_folder, extra_string)
3187
3188        # Plot dendrite learning scores
3189        ax = plt.subplot(222)
3190        self.generate_dendrite_learning_plots(ax, save_folder, extra_string)
3191
3192        if GPA.pc.get_drawing_extra_graphs():
3193            # Plot learning rates for each training epoch
3194            ax = plt.subplot(223)
3195            self.generate_learning_rate_plots(ax, save_folder, extra_string)
3196
3197            # Plot the times for each training epoch
3198            ax = plt.subplot(224)
3199            self.generate_time_plots(ax, save_folder, extra_string)
3200
3201        # Generate extra CSV files
3202        self.generate_extra_csv_files(save_folder, extra_string)
3203
3204        fig.tight_layout()
3205        plt.savefig(save_folder + "/" + self.save_name + extra_string + ".png")
3206        plt.close("all")
3207
3208    def add_loss(self, loss):
3209        """Add loss to tracking vectors.
3210
3211        Parameters
3212        ----------
3213        loss : float or int
3214            The loss value to add.
3215
3216        Returns
3217        -------
3218        None
3219
3220        """
3221        if not isinstance(loss, (float, int)):
3222            loss = loss.item()
3223        self.member_vars["training_loss"].append(loss)
3224
3225    def add_learning_rate(self, learning_rate):
3226        """Add learning rate to tracking vectors.
3227
3228        Parameters
3229        ----------
3230        learning_rate : float or int
3231            The learning rate value to add.
3232
3233        Returns
3234        -------
3235        None
3236
3237        """
3238        if not isinstance(learning_rate, (float, int)):
3239            learning_rate = learning_rate.item()
3240        self.member_vars["training_learning_rates"].append(learning_rate)
3241
3242    def add_extra_score(self, score, extra_score_name):
3243        """Add extra score to tracking vectors.
3244
3245        Parameters
3246        ----------
3247        score : float or int
3248            The score value to add.
3249
3250        extra_score_name : str
3251            The name of the extra score.
3252
3253        Returns
3254        -------
3255        None
3256
3257        """
3258        if not isinstance(score, (float, int)):
3259            try:
3260                score = score.item()
3261            except:
3262                print(
3263                    "Scores added for Perforated Backpropagation should be "
3264                    "float, int, or tensor, yours is a:"
3265                )
3266                print(type(score))
3267                pdb.set_trace()
3268
3269        if GPA.pc.get_verbose():
3270            print(f"Adding extra score {extra_score_name} of {float(score)}")
3271
3272        if extra_score_name not in self.member_vars["extra_scores"]:
3273            self.member_vars["extra_scores"][extra_score_name] = []
3274        self.member_vars["extra_scores"][extra_score_name].append(score)
3275
3276        if self.member_vars["mode"] == "n":
3277            if extra_score_name not in self.member_vars["n_extra_scores"]:
3278                self.member_vars["n_extra_scores"][extra_score_name] = []
3279            self.member_vars["n_extra_scores"][extra_score_name].append(score)
3280
3281    def add_extra_score_without_graphing(self, score, extra_score_name):
3282        """Add extra score without graphing to tracking vectors.
3283
3284        Parameters
3285        ----------
3286        score : float or int
3287            The score value to add.
3288
3289        extra_score_name : str
3290            The name of the extra score.
3291
3292        Returns
3293        -------
3294        None
3295
3296        """
3297        if not isinstance(score, (float, int)):
3298            try:
3299                score = score.item()
3300            except:
3301                print(
3302                    "Scores added for Perforated Backpropagation should be "
3303                    "float, int, or tensor, yours is a:"
3304                )
3305                print(type(score))
3306                print("in add_extra_score_without_graphing")
3307                pdb.set_trace()
3308
3309        if GPA.pc.get_verbose():
3310            print(f"Adding extra score {extra_score_name} of {float(score)}")
3311
3312        if extra_score_name not in self.member_vars["extra_scores_without_graphing"]:
3313            self.member_vars["extra_scores_without_graphing"][extra_score_name] = []
3314        self.member_vars["extra_scores_without_graphing"][extra_score_name].append(
3315            score
3316        )
3317
3318    def add_test_score(self, score, extra_score_name):
3319        """Add test score to tracking vectors.
3320
3321        Parameters
3322        ----------
3323        score : float or int
3324            The score value to add.
3325
3326        extra_score_name : str
3327            The name of the extra score.
3328
3329        Returns
3330        -------
3331        None
3332
3333        Notes
3334        -----
3335        This function is a wrapper around `add_extra_score` that separates
3336        test score for adding to best_arch_scores.csv.
3337
3338        """
3339        self.add_extra_score(score, extra_score_name)
3340
3341        if not isinstance(score, (float, int)):
3342            try:
3343                score = score.item()
3344            except:
3345                print(
3346                    "Scores added for Perforated Backpropagation should be "
3347                    "float, int, or tensor, yours is a:"
3348                )
3349                print(type(score))
3350                print("in add_test_score")
3351                pdb.set_trace()
3352
3353        if GPA.pc.get_verbose():
3354            print(f"Adding test score {extra_score_name} of {float(score)}")
3355        self.member_vars["test_scores"].append(score)
3356
3357    def add_validation_score(self, accuracy, net, force_switch=False):
3358        """Function to add the validation score.
3359
3360        This is complex because it determines neuron and dendrite switching.
3361
3362        Parameters
3363        ----------
3364        accuracy : float or int
3365            The accuracy or loss value to add.
3366        net : object
3367            The neural network model.
3368        force_switch : bool, optional
3369            Whether to force a switch, by default False.
3370
3371        Returns
3372        -------
3373        net : object
3374            The potentially modified neural network model.
3375        training_complete : bool
3376            Whether training is complete.
3377        restructured : bool
3378            Whether the model has been restructured.
3379
3380        Notes
3381        -----
3382        WARNING: Do not call self anywhere in this function. When systems
3383        get loaded the actual tracker you are working with can change.
3384        """
3385
3386        _pai_log("info", f"Adding validation score {accuracy:.8f}")
3387
3388        update_learning_rate()
3389        update_param_count(net)
3390
3391        accuracy = check_input_problems(net, accuracy)
3392
3393        if len(GPA.pai_tracker.member_vars["switch_epochs"]) == 0:
3394            epochs_since_cycle_switch = GPA.pai_tracker.member_vars["num_epochs_run"]
3395        else:
3396            epochs_since_cycle_switch = (
3397                GPA.pai_tracker.member_vars["num_epochs_run"]
3398                - GPA.pai_tracker.member_vars["switch_epochs"][-1]
3399            )
3400
3401        update_running_accuracy(accuracy, epochs_since_cycle_switch)
3402        if GPA.pc.get_perforated_backpropagation():
3403            TPB.update_pb_scores(self)
3404
3405        # Captured before any switch below flips the mode and reloads scores
3406        epoch_pb_scores = self.get_current_pb_scores()
3407
3408        GPA.pai_tracker.stop_epoch(internal_call=True)
3409
3410        # If it is neuron training mode
3411        if (
3412            GPA.pai_tracker.member_vars["mode"] == "n"
3413            or GPA.pc.get_learn_dendrites_live()
3414        ):
3415            check_new_best(net, accuracy, epochs_since_cycle_switch)
3416        elif GPA.pc.get_perforated_backpropagation():
3417            TPB.check_best_pai_score_improvement()
3418
3419        # Save the latest model
3420        if GPA.pc.get_test_saves():
3421            UPA.save_system(net, GPA.pc.get_save_name(), "latest")
3422        if GPA.pc.get_pai_saves():
3423            UPA.pai_save_system(net, GPA.pc.get_save_name(), "latest")
3424
3425        restructuring_status_value = NO_MODEL_UPDATE
3426        # If it is time to switch based on scores and counter or a manual switch
3427        if GPA.pai_tracker.switch_time() or force_switch:
3428            # If testing dendrite capacity switch after enough dendrites added
3429            if (
3430                (GPA.pai_tracker.member_vars["mode"] == "n")
3431                and (GPA.pai_tracker.member_vars["num_dendrites_added"] > 2)
3432                and GPA.pc.get_testing_dendrite_capacity()
3433            ):
3434                GPA.pai_tracker.save_graphs()
3435                _pai_log(
3436                    "info",
3437                    "Successfully added 3 dendrites with GPA.pc.set_testing_dendrite_capacity(True) (default). "
3438                    "You may now set that to False and run a real experiment.",
3439                )
3440                return net, False, True
3441
3442            # If doing neuron training but this dendrite count didn't improve
3443            if (
3444                (GPA.pai_tracker.member_vars["mode"] == "n")
3445                or GPA.pc.get_learn_dendrites_live()
3446            ) and (GPA.pai_tracker.member_vars["current_n_set_global_best"] is False):
3447                new_restructuring_status_value, net = process_no_improvement(net)
3448                # if this was the final try return that training is complete
3449                if new_restructuring_status_value == TRAINING_COMPLETE:
3450                    if _dashboard_emitter is not None:
3451                        _dashboard_emitter.emit_run_end(GPA.pc)
3452                    return net, True, True
3453                else:
3454                    restructuring_status_value = update_restructuring_status(
3455                        restructuring_status_value, new_restructuring_status_value
3456                    )
3457            # Else if did improve, do a normal switch process
3458            else:
3459                if GPA.pc.get_verbose():
3460                    print(
3461                        f"Calling switch_mode with "
3462                        f'{GPA.pai_tracker.member_vars["current_n_set_global_best"]}, '
3463                        f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}, '
3464                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]}, '
3465                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_value"]},'
3466                        f'{GPA.pc.get_max_dendrites()},'
3467                        f'{GPA.pai_tracker.member_vars["num_dendrites_added"]},'
3468                        f'{GPA.pai_tracker.member_vars["num_dendrite_tries"]},'
3469                    )
3470                import pdb; pdb.set_trace
3471                # If the max number of dendrites has been hit or not doing pai and adding dendtites
3472                # then return rather than adding more
3473                if (
3474                    (GPA.pai_tracker.member_vars["mode"] == "n")
3475                    and (
3476                        GPA.pc.get_max_dendrites()
3477                        == GPA.pai_tracker.member_vars["num_dendrites_added"]
3478                    )
3479                ) or (GPA.pai_tracker.member_vars["doing_pai"] is False):
3480                    if GPA.pc.get_verbose():
3481                        print(
3482                            "Max dendrites reached or not doing PAI, finishing training"
3483                        )
3484                    net = process_final_network(net)
3485                    # Increment integrated if we have dendrites (means they're integrated)
3486                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3487                        GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3488                        _pai_log("info", f"Final dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3489                        if _dashboard_emitter is not None:
3490                            _dashboard_emitter.emit_dendrite_added(
3491                                GPA.pc,
3492                                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3493                                num_dendrites_integrated=GPA.pai_tracker.member_vars[
3494                                    "num_dendrites_integrated"
3495                                ],
3496                            )
3497                    if _dashboard_emitter is not None:
3498                        _dashboard_emitter.emit_run_end(GPA.pc)
3499                    return net, True, True
3500
3501                # Otherwise if its neuron training mode reset the counter of failed dendrites
3502                # Check if we should increment integrated count BEFORE change_learning_modes loads old state
3503                should_increment_integrated = False
3504                if GPA.pai_tracker.member_vars["mode"] == "n":
3505                    GPA.pai_tracker.member_vars["num_dendrite_tries"] = 0
3506                    if GPA.pc.get_verbose():
3507                        print(
3508                            "Adding new dendrites without resetting which means "
3509                            "the last ones improved. Resetting num_dendrite_tries"
3510                        )
3511                    # Remember to increment after change_learning_modes (which loads old tracker state)
3512                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3513                        should_increment_integrated = True
3514
3515                GPA.pai_tracker.save_graphs(
3516                    f'_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}'
3517                )
3518
3519                if GPA.pc.get_test_saves():
3520                    UPA.save_system(
3521                        net,
3522                        GPA.pc.get_save_name(),
3523                        f'beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3524                    )
3525                    # Copy current best model from this set of dendrites
3526                    # If running DDP only copy with rank 0
3527                    if "RANK" not in os.environ or int(os.environ["RANK"]) == 0:
3528                        shutil.copyfile(
3529                            f"{GPA.pc.get_save_name()}/best_model.pt",
3530                            f'{GPA.pc.get_save_name()}/best_model_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}.pt',
3531                        )
3532
3533                net = UPA.change_learning_modes(
3534                    net,
3535                    GPA.pc.get_save_name(),
3536                    "best_model",
3537                    GPA.pai_tracker.member_vars["doing_pai"],
3538                )
3539                restructuring_status_value = NETWORK_RESTRUCTURED
3540                
3541                # Now increment after change_learning_modes has loaded the best model
3542                # This ensures the increment persists and doesn't get overwritten
3543                if should_increment_integrated:
3544                    GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3545                    _pai_log("info", f"Dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3546                    if _dashboard_emitter is not None:
3547                        _dashboard_emitter.emit_dendrite_added(
3548                            GPA.pc,
3549                            epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3550                            num_dendrites_integrated=GPA.pai_tracker.member_vars[
3551                                "num_dendrites_integrated"
3552                            ],
3553                        )
3554
3555            # If restructured is true, clear scheduler/optimizer before saving
3556            if restructuring_status_value != NETWORK_RESTRUCTURED:
3557                print(
3558                    "Restructured should always be triggered here, let us know if you encounter this situation"
3559                )
3560                pdb.set_trace()
3561
3562            # Since there is a restructuring optimizer and scheduler must be reinitialized after return
3563            GPA.pai_tracker.clear_optimizer_and_scheduler()
3564
3565            # Save the model from after the switch
3566            UPA.save_system(
3567                net,
3568                GPA.pc.get_save_name(),
3569                f'switch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3570            )
3571
3572        # If not time to switch and you have a scheduler, perform the update step
3573        elif GPA.pai_tracker.member_vars["scheduler"] is not None:
3574            new_restructuring_status_value, net = process_scheduler_update(
3575                net, accuracy, epochs_since_cycle_switch
3576            )
3577            restructuring_status_value = update_restructuring_status(
3578                restructuring_status_value, new_restructuring_status_value
3579            )
3580
3581        GPA.pai_tracker.start_epoch(internal_call=True)
3582        if _dashboard_emitter is not None:
3583            _mv = GPA.pai_tracker.member_vars
3584            _lr = _mv["training_learning_rates"][-1] if _mv["training_learning_rates"] else None
3585            _train_score = _mv["extra_scores"].get("train", [None])[-1]
3586            _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]
3587            _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]
3588            _dashboard_emitter.emit_epoch(
3589                GPA.pc,
3590                epoch=_mv["num_epochs_run"],
3591                validation_score=accuracy,
3592                learning_rate=_lr,
3593                train_score=_train_score,
3594                normal_time=_n_times[-1],
3595                pai_time=_p_times[-1],
3596                pb_scores=epoch_pb_scores,
3597            )
3598        GPA.pai_tracker.save_graphs()
3599
3600        if restructuring_status_value == NETWORK_RESTRUCTURED:
3601            GPA.pai_tracker.member_vars["epoch_last_improved"] = (
3602                GPA.pai_tracker.member_vars["num_epochs_run"]
3603            )
3604            if GPA.pc.get_verbose():
3605                print(
3606                    f"Setting epoch last improved to "
3607                    f'{GPA.pai_tracker.member_vars["epoch_last_improved"]}'
3608                )
3609
3610            now = datetime.now()
3611            dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
3612
3613            if GPA.pc.get_verbose():
3614                print("Not saving restructure right now")
3615
3616            """
3617            This block of code helped with a save issue with safetensors and huggingface, but it breaks DDP.  
3618            Temporarily removing it to avoid DDP issues, but if you encounter save issues try adding it back in.
3619            for param in net.parameters():
3620                param.data = param.data.contiguous()
3621            """
3622        if GPA.pc.get_verbose():
3623            print(
3624                f"Completed adding score. Restructured is {restructuring_status_value}, "
3625                f"\ncurrent switch list is:"
3626            )
3627            print(GPA.pai_tracker.member_vars["switch_epochs"])
3628
3629        if _dashboard_emitter is not None and restructuring_status_value == NETWORK_RESTRUCTURED:
3630            _param_count = UPA.count_params(net)
3631            _dashboard_emitter.emit_switch(
3632                GPA.pc,
3633                switch_number=GPA.pai_tracker.member_vars["num_dendrites_added"],
3634                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3635                param_count=_param_count,
3636                switch_type=GPA.pai_tracker.member_vars["mode"],
3637            )
3638
3639        # Always False for training complete if nothing triggered that training is over
3640        return net, restructuring_status_value, False
3641
3642    def clear_all_processors(self):
3643        """Clear all processors from modules.
3644
3645        Parameters
3646        ----------
3647        None
3648
3649        Returns
3650        -------
3651        None
3652            This function does not return a value.
3653        """
3654        for module in self.neuron_module_vector:
3655            module.clear_processors()
3656
3657    def create_new_dendrite_module(self):
3658        """Add dendrite module to all neuron modules.
3659
3660        Parameters
3661        ----------
3662        None
3663
3664        Returns
3665        -------
3666        None
3667            This function does not return a value.
3668        """
3669        for module in self.neuron_module_vector:
3670            module.create_new_dendrite_module()
3671
3672    def apply_pb_grads(self):
3673        """Apply perforated backpropagation gradients to all modules.
3674
3675        Parameters
3676        ----------
3677        None
3678
3679        Returns
3680        -------
3681        None
3682            This function does not return a value.
3683        """
3684        if self.member_vars["mode"] == "p":
3685            for module in self.neuron_module_vector:
3686                module.apply_pb_grads()
3687
3688    def apply_pb_zero(self):
3689        """Apply perforated backpropagation zero gradients to all modules.
3690
3691        Parameters
3692        ----------
3693        None
3694
3695        Returns
3696        -------
3697        None
3698            This function does not return a value.
3699        """
3700        if self.member_vars["mode"] == "p":
3701            for module in self.neuron_module_vector:
3702                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            f.write(
1553                f"{layer.name},{layer.dendrite_module.dendrite_values[0].out_channels}\n"
1554            )
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            channels[line.split(",")[0]] = int(line.split(",")[1])
1587        for layer in self.neuron_module_vector:
1588            layer.dendrite_module.dendrite_values[0].setup_arrays(channels[layer.name])
1589
1590    def set_optimizer_instance(self, optimizer_instance, additional_optimizers=[]):
1591        """Set optimizer instance directly.
1592
1593        Parameters
1594        ----------
1595        optimizer_instance : object
1596            The optimizer instance to set.
1597
1598        Returns
1599        -------
1600        None
1601
1602        """
1603
1604        try:
1605            for param_group in optimizer_instance.param_groups:
1606                if (
1607                    param_group["weight_decay"] > 0
1608                    and GPA.pc.get_weight_decay_accepted() is False
1609                ):
1610                    _pai_log(
1611                        "warning",
1612                        "For PAI training it is recommended to not use weight decay in your optimizer",
1613                    )
1614
1615        except:
1616            pass
1617        self.member_vars["optimizer_instance"] = optimizer_instance
1618        if GPA.pc.get_perforated_backpropagation():
1619            TPB.setup_optimizer_pb(self.member_vars["optimizer_instance"])
1620            for optimizer in additional_optimizers:
1621                TPB.filter_params(optimizer)
1622        optimizer_instance.zero_grad()
1623
1624    def set_optimizer(self, optimizer):
1625        """Set optimizer type to be initialized later
1626
1627        Parameters
1628        ----------
1629        optimizer : object
1630            The optimizer type to set.
1631
1632        Returns
1633        -------
1634        None
1635
1636        """
1637        self.member_vars["optimizer"] = optimizer
1638
1639    def set_scheduler(self, scheduler):
1640        """Set scheduler type to be initialized later
1641
1642        Parameters
1643        ----------
1644        scheduler : object
1645            The scheduler type to set.
1646
1647        Returns
1648        -------
1649        None
1650
1651        """
1652        if scheduler is not torch.optim.lr_scheduler.ReduceLROnPlateau:
1653            if GPA.pc.get_verbose():
1654                print("Not using ReduceLROnPlateau, this is not recommended")
1655        self.member_vars["scheduler"] = scheduler
1656
1657    def increment_scheduler(self, num_ticks, mode):
1658        """Increment the scheduler a set number of times.
1659
1660        Used for finding best initial learning rate when adding dendrites.
1661
1662        Parameters
1663        ----------
1664        num_ticks : int
1665            The number of scheduler steps to take.
1666        mode : str
1667            The mode for stepping the scheduler. Options are:
1668            - "step_learning_rate": Step based on improved accuracy epochs
1669            - "increment_epoch_count": Step based on total epoch count
1670
1671        Returns
1672        -------
1673        current_steps : int
1674            The number of learning rate changes that occurred.
1675        learning_rate1 : float
1676            The final learning rate after stepping.
1677
1678        """
1679
1680        current_steps = 0
1681        current_ticker = 0
1682
1683        for param_group in GPA.pai_tracker.member_vars[
1684            "optimizer_instance"
1685        ].param_groups:
1686            learning_rate1 = param_group["lr"]
1687
1688        if GPA.pc.get_verbose():
1689            print("Using scheduler:")
1690            print(type(self.member_vars["scheduler_instance"]))
1691
1692        while current_ticker < num_ticks:
1693            if GPA.pc.get_verbose():
1694                print(
1695                    f"Lower start rate initial {learning_rate1} "
1696                    f'stepping {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} times'
1697                )
1698
1699            if (
1700                type(self.member_vars["scheduler_instance"])
1701                is torch.optim.lr_scheduler.ReduceLROnPlateau
1702            ):
1703                if mode == "step_learning_rate":
1704                    # Step with counter as last improved accuracy
1705                    self.member_vars["scheduler_instance"].step(
1706                        metrics=self.member_vars["last_improved_accuracies"][
1707                            GPA.pai_tracker.steps_after_switch() - 1
1708                        ]
1709                    )
1710                elif mode == "increment_epoch_count":
1711                    # Step with improved epoch counts up to current location
1712                    self.member_vars["scheduler_instance"].step(
1713                        metrics=self.member_vars["last_improved_accuracies"][
1714                            -((num_ticks - 1) - current_ticker) - 1
1715                        ]
1716                    )
1717            else:
1718                self.member_vars["scheduler_instance"].step()
1719
1720            for param_group in GPA.pai_tracker.member_vars[
1721                "optimizer_instance"
1722            ].param_groups:
1723                learning_rate2 = param_group["lr"]
1724
1725            if learning_rate2 != learning_rate1:
1726                current_steps += 1
1727                learning_rate1 = learning_rate2
1728                if mode == "step_learning_rate":
1729                    current_ticker += 1
1730                if GPA.pc.get_verbose():
1731                    print(f"1 step {current_steps} to {learning_rate2}")
1732
1733            if mode == "increment_epoch_count":
1734                current_ticker += 1
1735
1736        return current_steps, learning_rate1
1737
1738    def setup_optimizer(self, net, opt_args, sched_args=None, parameters=None):
1739        """Initialize the optimizer and scheduler when added.
1740
1741        Parameters
1742        ----------
1743        net : object
1744            The neural network model.
1745        opt_args : dict
1746            The arguments for the optimizer.
1747        sched_args : dict, optional
1748            The arguments for the scheduler, by default None.
1749
1750        Returns
1751        -------
1752        optimizer : object
1753            The initialized optimizer instance.
1754        scheduler : object, optional
1755            The initialized scheduler instance, if a scheduler was set.
1756
1757        """
1758        if "weight_decay" in opt_args and not GPA.pc.get_weight_decay_accepted():
1759            _pai_log(
1760                "warning",
1761                "For PAI training it is recommended to not use weight decay in your optimizer",
1762            )
1763
1764        if ("model" not in opt_args.keys()) and "params" not in opt_args.keys():
1765            print("In setup_optimizer it will be depreciated to not pass in params yourself in the future")
1766            print("please change the settings to include params")
1767            if self.member_vars["mode"] == "n":
1768                if parameters is not None:
1769                    opt_args["params"] = parameters
1770                else:
1771                    opt_args["params"] = filter(lambda p: p.requires_grad, net.parameters())
1772            else:
1773                params = UPA.get_pai_network_params(net)
1774                if parameters is not None:
1775                    # Filter parameters to only those in params, preserving weight_decay
1776                    params_set = set(params)
1777                    filtered_params = []
1778                    for param_group in parameters:
1779                        filtered_group_params = [p for p in param_group["params"] if p in params_set]
1780                        if filtered_group_params:
1781                            filtered_params.append({
1782                                "params": filtered_group_params,
1783                                "weight_decay": param_group["weight_decay"]
1784                            })
1785                    opt_args["params"] = filtered_params
1786                else:
1787                    opt_args["params"] = params
1788        elif "params" in opt_args.keys():
1789            # Check if params is a list of param groups (dicts) or a single param group
1790            params_value = opt_args["params"]
1791            if isinstance(params_value, list) and len(params_value) > 0:
1792                # Check if it's a list of dicts (multiple param groups) or list of tensors (single group)
1793                if isinstance(params_value[0], dict):
1794                    # Multiple param groups format: [{"params": [...], "lr": ...}, ...]
1795                    # Filter each param group for requires_grad
1796                    filtered_param_groups = []
1797                    for param_group in params_value:
1798                        filtered_group_params = [p for p in param_group["params"] if p.requires_grad]
1799                        if filtered_group_params:
1800                            new_group = param_group.copy()
1801                            new_group["params"] = filtered_group_params
1802                            filtered_param_groups.append(new_group)
1803                    opt_args["params"] = filtered_param_groups
1804                else:
1805                    # Single param group format: [tensor1, tensor2, ...] or generator
1806                    # Filter for requires_grad
1807                    opt_args["params"] = [p for p in params_value if p.requires_grad]
1808            elif hasattr(params_value, '__iter__'):
1809                # Handle generators or other iterables
1810                opt_args["params"] = [p for p in params_value if p.requires_grad]
1811
1812        optimizer = self.member_vars["optimizer"](**opt_args)
1813        self.set_optimizer_instance(optimizer)
1814
1815        if self.member_vars["scheduler"] is not None:
1816            # Handle SequentialLR specially
1817            if self.member_vars["scheduler"] is torch.optim.lr_scheduler.SequentialLR:
1818                """
1819                sched_args should be a dict with "schedulers" (list of tuples) and "milestones"
1820                For example:
1821                sequential_schedArgs = {
1822                    "schedulers": [
1823                        (warmup_scheduler_class, warmup_schedArgs),
1824                        (main_scheduler_class, main_schedArgs)
1825                    ],
1826                    "milestones": [switch_epoch]
1827                }
1828                """
1829                schedulers = []
1830                milestones = sched_args.get("milestones", [])
1831                scheduler_configs = sched_args.get("schedulers", [])
1832                
1833                for scheduler_class, scheduler_args in scheduler_configs:
1834                    schedulers.append(scheduler_class(optimizer, **scheduler_args))
1835                
1836                self.member_vars["scheduler_instance"] = torch.optim.lr_scheduler.SequentialLR(
1837                    optimizer, schedulers=schedulers, milestones=milestones
1838                )
1839            else:
1840                self.member_vars["scheduler_instance"] = self.member_vars["scheduler"](
1841                    optimizer, **sched_args
1842                )
1843            current_steps = 0
1844
1845            for param_group in GPA.pai_tracker.member_vars[
1846                "optimizer_instance"
1847            ].param_groups:
1848                learning_rate1 = param_group["lr"]
1849
1850            if GPA.pc.get_verbose():
1851                print(
1852                    f"Resetting scheduler with {GPA.pai_tracker.steps_after_switch()} "
1853                    f"steps and {GPA.pc.get_initial_history_after_switches()} initial ticks to skip"
1854                )
1855
1856            # Find setting of previously used learning rate before adding dendrites
1857            if (
1858                GPA.pai_tracker.member_vars[
1859                    "current_n_learning_rate_initial_skip_steps"
1860                ]
1861                != 0
1862            ):
1863                additional_steps, learning_rate1 = self.increment_scheduler(
1864                    GPA.pai_tracker.member_vars[
1865                        "current_n_learning_rate_initial_skip_steps"
1866                    ],
1867                    "step_learning_rate",
1868                )
1869                current_steps += additional_steps
1870
1871            if self.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
1872                initial = GPA.pc.get_initial_history_after_switches()
1873            else:
1874                initial = 0
1875
1876            if GPA.pai_tracker.steps_after_switch() > initial:
1877                # Minus extra 1 because this gets called after start epoch
1878                additional_steps, learning_rate1 = self.increment_scheduler(
1879                    (GPA.pai_tracker.steps_after_switch() - initial) - 1,
1880                    "increment_epoch_count",
1881                )
1882                current_steps += additional_steps
1883
1884            if GPA.pc.get_verbose():
1885                print(
1886                    f"Scheduler update loop with {current_steps} "
1887                    f"ended with {learning_rate1}"
1888                )
1889                print(
1890                    f"Scheduler ended with {current_steps} steps "
1891                    f"and lr of {learning_rate1}"
1892                )
1893
1894            self.member_vars["current_step_count"] = current_steps
1895            return optimizer, self.member_vars["scheduler_instance"]
1896        else:
1897            return optimizer
1898
1899    def clear_optimizer_and_scheduler(self):
1900        """Clear the instances for saving.
1901
1902        Parameters
1903        ----------
1904        None
1905
1906        Returns
1907        -------
1908        None
1909            This function does not return a value.
1910        """
1911        self.member_vars["optimizer_instance"] = None
1912        self.member_vars["scheduler_instance"] = None
1913
1914    def switch_time(self):
1915        """Determine if it's time to switch between neuron and dendrite training.
1916
1917        Parameters
1918        ----------
1919        None
1920
1921        Returns
1922        -------
1923        bool
1924            True if it's time to switch, False otherwise.
1925
1926        Notes
1927        -----
1928        Based on current settings and history of scores.
1929        """
1930
1931        switch_phrase = "No mode, this should never be the case."
1932        switch_number = GPA.pc.get_n_epochs_to_switch()
1933        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1934            switch_phrase = "DOING_SWITCH_EVERY_TIME"
1935        elif self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY:
1936            switch_phrase = "DOING_HISTORY"
1937        elif self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH:
1938            switch_phrase = "DOING_FIXED_SWITCH"
1939            switch_number = GPA.pc.get_fixed_switch_num()
1940        elif self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1941            switch_phrase = "DOING_NO_SWITCH"
1942        else:
1943            print(
1944                "A switch mode must be set.  Check your settings for GPA.pc.set_switch_mode()."
1945            )
1946            pdb.set_trace()
1947        if not GPA.pc.get_silent():
1948            if(GPA.pc.get_perforated_backpropagation()):
1949                print(
1950                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1951                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1952                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1953                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1954                    f'n: {switch_number}, p: {GPA.pc.get_p_epochs_to_switch()}, '
1955                    f'num_cycles: {self.member_vars["num_cycles"]}'
1956                )
1957            else:
1958                print(
1959                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1960                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1961                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1962                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1963                    f'n: {switch_number}, num_cycles: {self.member_vars["num_cycles"]}'
1964                )
1965            print(
1966                f'  Score tracking: current_n_set_global_best={self.member_vars["current_n_set_global_best"]}, '
1967                f'global_best={self.member_vars["global_best_validation_score"]:.4f}, '
1968                f'current_best={self.member_vars["current_best_validation_score"]:.4f}'
1969            )
1970        if GPA.pc.get_perforated_backpropagation():
1971            # this will fill in epoch last improved
1972            TPB.best_pai_score_improved_this_epoch(self)  ## CLOSED ONLY
1973        if self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1974            if not GPA.pc.get_silent():
1975                print("Returning False - doing no switch mode")
1976            return False
1977
1978        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1979            if not GPA.pc.get_silent():
1980                print("Returning True - switching every time")
1981            return True
1982
1983        # Check if we're in the middle of learning rate optimization
1984        # If so, block ALL switch triggers until committed
1985        if GPA.pc.get_verbose():
1986            print("=== LR Optimization Check ===")
1987            print(f'  mode == "n": {self.member_vars["mode"] == "n"}')
1988            print(f"  get_learn_dendrites_live(): {GPA.pc.get_learn_dendrites_live()}")
1989            print(f'  committed_to_initial_rate: {GPA.pai_tracker.member_vars["committed_to_initial_rate"]}')
1990            print(f"  get_dont_give_up_unless_learning_rate_lowered(): {GPA.pc.get_dont_give_up_unless_learning_rate_lowered()}")
1991            print(f'  current_n_learning_rate_initial_skip_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"]}')
1992            print(f'  last_max_learning_rate_steps: {self.member_vars["last_max_learning_rate_steps"]}')
1993            print(f'  skip_steps < max_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"] < self.member_vars["last_max_learning_rate_steps"]}')
1994            print(f'  scheduler is not None: {self.member_vars["scheduler"] is not None}')
1995            print("=============================")
1996        
1997        if (
1998            ((self.member_vars["mode"] == "n") or GPA.pc.get_learn_dendrites_live())
1999            and (GPA.pai_tracker.member_vars["committed_to_initial_rate"] is False)
2000            and (GPA.pc.get_dont_give_up_unless_learning_rate_lowered())
2001            and (
2002                self.member_vars["current_n_learning_rate_initial_skip_steps"]
2003                <= self.member_vars["last_max_learning_rate_steps"]
2004            )
2005            and self.member_vars["scheduler"] is not None
2006        ):
2007            if not GPA.pc.get_silent():
2008                print(
2009                    f"Returning False - learning rate optimization in progress. "
2010                    f"Not committed yet. Comparing "
2011                    f'initial {self.member_vars["current_n_learning_rate_initial_skip_steps"]} '
2012                    f'to last max {self.member_vars["last_max_learning_rate_steps"]}'
2013                )
2014            return False
2015
2016        if len(self.member_vars["switch_epochs"]) == 0:
2017            this_count = self.member_vars["num_epochs_run"]
2018        else:
2019            this_count = (
2020                self.member_vars["num_epochs_run"]
2021                - self.member_vars["switch_epochs"][-1]
2022            )
2023        cap_switch = False
2024        if GPA.pc.get_perforated_backpropagation():
2025            cap_switch = TPB.check_cap_switch(self, this_count)
2026
2027        if self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY and (
2028            (
2029                (self.member_vars["mode"] == "n")
2030                and (
2031                    self.member_vars["num_epochs_run"]
2032                    - self.member_vars["epoch_last_improved"]
2033                    >= GPA.pc.get_n_epochs_to_switch()
2034                )
2035                and this_count
2036                >= GPA.pc.get_initial_history_after_switches()
2037                + GPA.pc.get_n_epochs_to_switch()
2038            )
2039            or (GPA.pc.get_perforated_backpropagation() and TPB.history_switch(self))
2040            or cap_switch
2041        ):
2042            if not GPA.pc.get_silent():
2043                print("Returning True - History and last improved is hit")
2044            return True
2045
2046        if self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH and (
2047            (
2048                self.member_vars["total_epochs_run"] % GPA.pc.get_fixed_switch_num()
2049                == GPA.pc.get_fixed_switch_num() - 1
2050            )
2051            and self.member_vars["num_epochs_run"]
2052            >= GPA.pc.get_first_fixed_switch_num() - 1
2053        ):
2054            if not GPA.pc.get_silent():
2055                print("Returning True - Fixed switch number is hit")
2056            return True
2057
2058        if not GPA.pc.get_silent():
2059            print("Returning False - no triggers to switch have been hit")
2060        return False
2061
2062    def steps_after_switch(self):
2063        """Based on settings, return value for steps since a switch.
2064
2065        Different options for param vals setting determine what is returned.
2066
2067        Parameters
2068        ----------
2069        None
2070
2071        Returns
2072        -------
2073        int
2074            The number of epochs since the last switch, or total epochs run,
2075            depending on settings.
2076
2077        """
2078        if self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_TOTAL_EPOCH:
2079            return self.member_vars["num_epochs_run"]
2080        elif (
2081            self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_UPDATE_EPOCH
2082        ):
2083            return self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2084        elif (
2085            self.member_vars["param_vals_setting"]
2086            == GPA.pc.PARAM_VALS_BY_NEURON_EPOCH_START
2087        ):
2088            if self.member_vars["mode"] == "p":
2089                return (
2090                    self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2091                )
2092            else:
2093                return self.member_vars["num_epochs_run"]
2094        else:
2095            print(
2096                f'{self.member_vars["param_vals_setting"]} is not a valid param vals option'
2097            )
2098            pdb.set_trace()
2099
2100    def add_pai_neuron_module(self, new_module, initial_add=True):
2101        """Add neuron modules to internal vectors.
2102
2103        Parameters
2104        ----------
2105        new_module : object
2106            The new module to add.
2107        initial_add : bool, optional
2108            Whether this is the initial addition rather than loading from file
2109
2110        Returns
2111        -------
2112        None
2113
2114        """
2115
2116        # If it's a duplicate, ignore the second addition
2117        if new_module in self.neuron_module_vector:
2118            return
2119        self.neuron_module_vector.append(new_module)
2120        if self.member_vars["doing_pai"]:
2121            PA.set_wrapped_params(new_module)
2122        if initial_add:
2123            self.member_vars["best_scores"].append([])
2124            self.member_vars["current_scores"].append([])
2125
2126    def add_tracked_neuron_module(self, new_module, initial_add=True):
2127        """Add tracked modules to internal vectors
2128
2129        Parameters
2130        ----------
2131        new_module : object
2132            The new module to add.
2133        initial_add : bool, optional
2134            Whether this is the initial addition rather than loading from file
2135
2136        Returns
2137        -------
2138        None
2139
2140        """
2141        # If it's a duplicate, ignore the second addition
2142        if new_module in self.tracked_neuron_module_vector:
2143            return
2144        self.tracked_neuron_module_vector.append(new_module)
2145        if self.member_vars["doing_pai"]:
2146            PA.set_tracked_params(new_module)
2147
2148    def reset_module_vector(self, net, load_from_restart):
2149        """Clear internal vectors and reset from network.
2150
2151        Parameters
2152        ----------
2153        net : object
2154            The neural network model.
2155        load_from_restart : bool
2156            Whether loading from a restart file.
2157
2158        Returns
2159        -------
2160        None
2161
2162        """
2163        self.neuron_module_vector = []
2164        self.tracked_neuron_module_vector = []
2165        this_list = UPA.get_pai_modules(net, 0)
2166        for module in this_list:
2167            self.add_pai_neuron_module(module, initial_add=load_from_restart)
2168        this_list = UPA.get_tracked_modules(net, 0)
2169        for module in this_list:
2170            self.add_tracked_neuron_module(module, initial_add=load_from_restart)
2171
2172    def reset_vals_for_score_reset(self):
2173        """Reset cycle scores for new cycle.
2174
2175        Parameters
2176        ----------
2177        None
2178
2179        Returns
2180        -------
2181        None
2182            This function does not return a value.
2183        """
2184
2185        if GPA.pc.get_find_best_lr():
2186            self.member_vars["committed_to_initial_rate"] = False
2187            print("Resetting committed to initial rate to False")
2188        # If retaining all dendrties always say that the current dendrites set global best for saving and loading
2189        if GPA.pc.get_retain_all_dendrites():
2190            self.member_vars["current_n_set_global_best"] = True
2191            self.member_vars["global_best_validation_score"] = 0
2192        else:
2193            self.member_vars["current_n_set_global_best"] = False
2194
2195        # Don't reset global best, but do reset current best
2196        self.member_vars["current_best_validation_score"] = 0
2197        self.member_vars["initial_lr_test_epoch_count"] = -1
2198
2199    def set_dendrite_training(self):
2200        """Signal all layers to start dendrite training.
2201
2202        Parameters
2203        ----------
2204        None
2205
2206        Returns
2207        -------
2208        None
2209            This function does not return a value.
2210        """
2211        if GPA.pc.get_verbose():
2212            print("Calling set_dendrite_training")
2213
2214        for layer in self.neuron_module_vector[:]:
2215            worked = layer.set_mode("p")
2216            """
2217            worked is False when a layer was added to the neuron module vector
2218            but then it's never actually been used. This can happen when
2219            you have set a layer to have requires_grad = False or when
2220            you have a module as a member variable but it's not actually
2221            part of the network. Should be moved to be a tracked layer
2222            rather than a neuron layer.
2223            """
2224            if not worked:
2225                self.neuron_module_vector.remove(layer)
2226
2227        for layer in self.tracked_neuron_module_vector[:]:
2228            worked = layer.set_mode("p")
2229
2230        self.create_new_dendrite_module()
2231        self.member_vars["mode"] = "p"
2232        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2233
2234        if GPA.pc.get_learn_dendrites_live():
2235            self.reset_vals_for_score_reset()
2236
2237        self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2238            "current_step_count"
2239        ]
2240
2241        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
2242        GPA.pai_tracker.member_vars["num_cycles"] += 1
2243
2244
2245    def set_neuron_training(self):
2246        """Signal all layers to start neuron training.
2247
2248        Parameters
2249        ----------
2250        None
2251
2252        Returns
2253        -------
2254        None
2255            This function does not return a value.
2256        """
2257        for module in self.neuron_module_vector:
2258            module.set_mode("n")
2259        for module in self.tracked_neuron_module_vector[:]:
2260            module.set_mode("n")
2261
2262        self.member_vars["mode"] = "n"
2263        self.member_vars["num_dendrites_added"] += 1
2264        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2265        self.reset_vals_for_score_reset()
2266
2267        self.member_vars["current_cycle_lr_max_scores"] = []
2268        if GPA.pc.get_learn_dendrites_live():
2269            self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2270                "current_step_count"
2271            ]
2272        GPA.pai_tracker.member_vars["num_cycles"] += 1
2273
2274        if GPA.pc.get_reset_best_score_on_switch():
2275            GPA.pai_tracker.member_vars["current_best_validation_score"] = 0
2276            GPA.pai_tracker.member_vars["running_accuracy"] = 0
2277
2278    def start_epoch(self, internal_call=False):
2279        """Perform steps for when a new training epoch is about to begin.
2280
2281        Parameters
2282        ----------
2283        internal_call : bool, optional
2284            Whether this is an internal call or manual call
2285
2286        Returns
2287        -------
2288        None
2289
2290        Notes
2291        -----
2292        If you ever need to call this manually, set internal_call to False.
2293
2294        """
2295        if self.member_vars["manual_train_switch"] and internal_call:
2296            return
2297
2298        if not internal_call and not self.member_vars["manual_train_switch"]:
2299            self.member_vars["manual_train_switch"] = True
2300            self.saved_time = 0
2301            self.member_vars["num_epochs_run"] = -1
2302            self.member_vars["total_epochs_run"] = -1
2303
2304        end = time.time()
2305        if self.member_vars["manual_train_switch"]:
2306            if self.saved_time != 0:
2307                if self.member_vars["mode"] == "p":
2308                    self.member_vars["p_val_times"].append(end - self.saved_time)
2309                else:
2310                    self.member_vars["n_val_times"].append(end - self.saved_time)
2311
2312        if self.member_vars["mode"] == "p":
2313            for layer in self.neuron_module_vector:
2314                for m in range(0, GPA.pc.get_global_candidates()):
2315                    with torch.no_grad():
2316                        if GPA.pc.get_verbose():
2317                            print(f"Resetting score for {layer.name}")
2318                        # Snapshot best_score before reset so we can compute per-epoch improvement
2319                        layer.dendrite_module.dendrite_values[
2320                            m
2321                        ].epoch_start_best_score.copy_(
2322                            layer.dendrite_module.dendrite_values[
2323                                m
2324                            ].best_score.detach()
2325                        )
2326                        layer.dendrite_module.dendrite_values[
2327                            m
2328                        ].best_score_improved_this_epoch = (
2329                            layer.dendrite_module.dendrite_values[
2330                                m
2331                            ].best_score_improved_this_epoch
2332                            * 0
2333                        )
2334                        layer.dendrite_module.dendrite_values[
2335                            m
2336                        ].nodes_best_improved_this_epoch = (
2337                            layer.dendrite_module.dendrite_values[
2338                                m
2339                            ].nodes_best_improved_this_epoch
2340                            * 0
2341                        )
2342                        layer.dendrite_module.dendrite_values[
2343                            m
2344                        ].nodes_improved_any = (
2345                            layer.dendrite_module.dendrite_values[
2346                                m
2347                            ].nodes_improved_any
2348                            * 0
2349                        )
2350            if GPA.pc.get_perforated_backpropagation():
2351                self.member_vars["best_mean_score_improved_this_epoch"] = 0
2352        self.member_vars["num_epochs_run"] += 1
2353        self.member_vars["total_epochs_run"] = (
2354            self.member_vars["num_epochs_run"] + self.member_vars["overwritten_epochs"]
2355        )
2356        self.saved_time = end
2357
2358    def stop_epoch(self, internal_call=False):
2359        """Perform steps when a training epoch has completed.
2360
2361        Parameters
2362        ----------
2363        internal_call : bool, optional
2364            Whether this is an internal call or manual call
2365
2366        Returns
2367        -------
2368        None
2369
2370        Notes
2371        -----
2372        If you ever need to call this manually, set internal_call to False.
2373
2374        """
2375        end = time.time()
2376        if self.member_vars["manual_train_switch"] and internal_call:
2377            return
2378
2379        if self.member_vars["manual_train_switch"]:
2380            if self.member_vars["mode"] == "p":
2381                self.member_vars["p_train_times"].append(end - self.saved_time)
2382            else:
2383                self.member_vars["n_train_times"].append(end - self.saved_time)
2384        else:
2385            if self.member_vars["mode"] == "p":
2386                self.member_vars["p_epoch_times"].append(end - self.saved_time)
2387            else:
2388                self.member_vars["n_epoch_times"].append(end - self.saved_time)
2389
2390        self.saved_time = end
2391
2392    def initialize(
2393        self,
2394        model,
2395        doing_pai=True,
2396        save_name="PAI",
2397        making_graphs=True,
2398        maximizing_score=True,
2399        num_classes=10000,
2400        values_per_train_epoch=-1,
2401        values_per_val_epoch=-1,
2402        zooming_graph=True,
2403    ):
2404        """Setup the tracker with initial settings.
2405
2406
2407        Parameters
2408        ----------
2409        model : object
2410            The neural network model.
2411        doing_pai : bool, optional
2412            Whether to add dendrites, by default True.
2413        save_name : str, optional
2414            The name under which to save the model.
2415        making_graphs : bool, optional
2416            Whether to make graphs, by default True.
2417        maximizing_score : bool, optional
2418            Whether to maximize the score, by default True.
2419        num_classes : int, optional
2420            The number of classes in the dataset, unused
2421        values_per_train_epoch : int, optional
2422            The number of values to look back for graphing
2423            during training, by default -1 (all values).
2424        values_per_val_epoch : int, optional
2425            The number of values to look back for graphing
2426            during validation, by default -1 (all values).
2427        zooming_graph : bool, optional
2428            Whether to zoom on graphs, by default True.
2429
2430
2431        Returns
2432        -------
2433        nn.Module
2434            Converted model instance configured for the tracker settings.
2435        """
2436        model = UPA.convert_network(model)
2437        self.member_vars["doing_pai"] = doing_pai
2438        self.member_vars["maximizing_score"] = maximizing_score
2439        self.save_name = save_name
2440        self.zooming_graph = zooming_graph
2441        self.making_graphs = making_graphs
2442
2443        if not self.loaded:
2444            self.member_vars["running_accuracy"] = (1.0 / num_classes) * 100
2445
2446        self.values_per_train_epoch = values_per_train_epoch
2447        self.values_per_val_epoch = values_per_val_epoch
2448
2449        if GPA.pc.get_testing_dendrite_capacity():
2450            if not GPA.pc.get_silent():
2451                print("Running a test of Dendrite Capacity.")
2452            GPA.pc.set_switch_mode(GPA.pc.DOING_SWITCH_EVERY_TIME)
2453            self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
2454            GPA.pc.set_retain_all_dendrites(True)
2455            GPA.pc.set_max_dendrite_tries(1000)
2456            GPA.pc.set_max_dendrites(1000)
2457            if GPA.pc.get_perforated_backpropagation():
2458                GPA.pc.set_initial_correlation_batches(1)
2459        else:
2460            if not GPA.pc.get_silent():
2461                print("Running Dendrite Experiment")
2462        return model
2463
2464    def generate_accuracy_plots(self, ax, save_folder, extra_string):
2465        """
2466        Generate plots and csvs for accuracy
2467
2468        Parameters
2469        ----------
2470        ax : object
2471            The matplotlib axis to plot on.
2472        save_folder : str
2473            The folder to save the plots and csvs in.
2474        extra_string : str
2475            An extra string to append to the filenames.
2476
2477        Returns
2478        -------
2479        None
2480
2481        """
2482
2483        # If scores are being saved for epochs that get overwritten, plot them
2484        for list_id in range(len(self.member_vars["overwritten_extras"])):
2485            for extra_id in self.member_vars["overwritten_extras"][list_id]:
2486                ax.plot(
2487                    np.arange(
2488                        len(self.member_vars["overwritten_extras"][list_id][extra_id])
2489                    ),
2490                    self.member_vars["overwritten_extras"][list_id][extra_id],
2491                    "r",
2492                )
2493            ax.plot(
2494                np.arange(len(self.member_vars["overwritten_vals"][list_id])),
2495                self.member_vars["overwritten_vals"][list_id],
2496                "b",
2497            )
2498
2499        # Determine which accuracy vector to use
2500        if GPA.pc.get_drawing_pai():
2501            accuracies = self.member_vars["accuracies"]
2502        else:
2503            accuracies = self.member_vars["n_accuracies"]
2504
2505        # Get pointer to additional scores being saved
2506        extra_scores = self.member_vars["extra_scores"]
2507
2508        # Plot the main accuracy scores
2509        ax.plot(np.arange(len(accuracies)), accuracies, label="Validation Scores")
2510        ax.plot(
2511            np.arange(len(self.member_vars["running_accuracies"])),
2512            self.member_vars["running_accuracies"],
2513            label="Validation Running Scores",
2514        )
2515
2516        # Plot additional scores
2517        for extra_score in extra_scores:
2518            ax.plot(
2519                np.arange(len(extra_scores[extra_score])),
2520                extra_scores[extra_score],
2521                label=extra_score,
2522            )
2523
2524        plt.title(save_folder + "/" + self.save_name + "Scores")
2525        plt.xlabel("Epochs")
2526        plt.ylabel("Score")
2527
2528        # Add point at epoch last improved and best validation score
2529        if GPA.pc.get_drawing_pai():
2530            ax.plot(
2531                self.member_vars["epoch_last_improved"],
2532                self.member_vars["global_best_validation_score"],
2533                "bo",
2534                label="Global best (y)",
2535            )
2536            ax.plot(
2537                self.member_vars["epoch_last_improved"],
2538                accuracies[self.member_vars["epoch_last_improved"]],
2539                "go",
2540                label="Epoch Last Improved",
2541            )
2542        else:
2543            if self.member_vars["mode"] == "n":
2544                missed_time = (
2545                    self.member_vars["num_epochs_run"]
2546                    - self.member_vars["epoch_last_improved"]
2547                )
2548                ax.plot(
2549                    (len(self.member_vars["n_accuracies"]) - 1) - missed_time,
2550                    self.member_vars["n_accuracies"][-(missed_time + 1)],
2551                    "go",
2552                    label="Epoch Last Improved",
2553                )
2554
2555        # Generate csv file for the values graphed
2556        pd1 = pd.DataFrame(
2557            {"Epochs": np.arange(len(accuracies)), "Validation Scores": accuracies}
2558        )
2559        pd2 = pd.DataFrame(
2560            {
2561                "Epochs": np.arange(len(self.member_vars["running_accuracies"])),
2562                "Validation Running Scores": self.member_vars["running_accuracies"],
2563            }
2564        )
2565        pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2566        for extra_score in extra_scores:
2567            pd2 = pd.DataFrame(
2568                {
2569                    "Epochs": np.arange(len(extra_scores[extra_score])),
2570                    extra_score: extra_scores[extra_score],
2571                }
2572            )
2573            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2574        extra_scores_without_graphing = self.member_vars[
2575            "extra_scores_without_graphing"
2576        ]
2577        for extra_score in extra_scores_without_graphing:
2578            pd2 = pd.DataFrame(
2579                {
2580                    "Epochs": np.arange(
2581                        len(extra_scores_without_graphing[extra_score])
2582                    ),
2583                    extra_score: extra_scores_without_graphing[extra_score],
2584                }
2585            )
2586            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2587        pd1.to_csv(
2588            save_folder + "/" + self.save_name + extra_string + "Scores.csv",
2589            index=False,
2590        )
2591        del pd1, pd2
2592
2593        # Set y min and max to zoom in on important part of axis
2594        if (
2595            len(self.member_vars["switch_epochs"]) > 0
2596            and self.member_vars["switch_epochs"][0] > 0
2597            and self.zooming_graph
2598        ):
2599            if GPA.pai_tracker.member_vars["maximizing_score"]:
2600                min_val = np.array(
2601                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2602                ).mean()
2603                for extra_score in extra_scores:
2604                    min_pot = np.array(
2605                        extra_scores[extra_score][
2606                            0 : self.member_vars["switch_epochs"][0]
2607                        ]
2608                    ).mean()
2609                    if min_pot < min_val:
2610                        min_val = min_pot
2611                ax.set_ylim(ymin=min_val)
2612            else:
2613                max_val = np.array(
2614                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2615                ).mean()
2616                for extra_score in extra_scores:
2617                    max_pot = np.array(
2618                        extra_scores[extra_score][
2619                            0 : self.member_vars["switch_epochs"][0]
2620                        ]
2621                    ).mean()
2622                    if max_pot > max_val:
2623                        max_val = max_pot
2624                ax.set_ylim(ymax=max_val)
2625
2626        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2627
2628        # Draw vertical lines for epochs where a dendrite switch occurred
2629        if GPA.pc.get_drawing_pai() and self.member_vars["doing_pai"]:
2630            color = "r"
2631            for switcher in self.member_vars["switch_epochs"]:
2632                plt.axvline(x=switcher, ymin=0, ymax=1, color=color)
2633                if color == "r":
2634                    color = "b"
2635                else:
2636                    color = "r"
2637        else:
2638            for switcher in self.member_vars["n_switch_epochs"]:
2639                plt.axvline(x=switcher, ymin=0, ymax=1, color="b")
2640
2641    def generate_time_plots(self, ax, save_folder, extra_string):
2642        """
2643        Generate plots and csvs for timing
2644
2645        Parameters
2646        ----------
2647        ax : object
2648            The matplotlib axis to plot on.
2649        save_folder : str
2650            The folder to save the plots and csvs in.
2651        extra_string : str
2652            An extra string to append to the filenames.
2653
2654        Returns
2655        -------
2656        None
2657
2658        """
2659        if self.member_vars["manual_train_switch"]:
2660            ax.plot(
2661                np.arange(len(self.member_vars["n_train_times"])),
2662                self.member_vars["n_train_times"],
2663                label="Normal Epoch Train Times",
2664            )
2665            ax.plot(
2666                np.arange(len(self.member_vars["p_train_times"])),
2667                self.member_vars["p_train_times"],
2668                label="PAI Epoch Train Times",
2669            )
2670            ax.plot(
2671                np.arange(len(self.member_vars["n_val_times"])),
2672                self.member_vars["n_val_times"],
2673                label="Normal Epoch Val Times",
2674            )
2675            ax.plot(
2676                np.arange(len(self.member_vars["p_val_times"])),
2677                self.member_vars["p_val_times"],
2678                label="PAI Epoch Val Times",
2679            )
2680
2681            plt.title(
2682                save_folder + "/" + self.save_name + "times (by train() and eval())"
2683            )
2684            plt.xlabel("Iteration")
2685            plt.ylabel("Epoch Time in Seconds ")
2686            ax.set_ylim(ymin=0)
2687            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2688
2689            pd1 = pd.DataFrame(
2690                {
2691                    "Epochs": np.arange(len(self.member_vars["n_train_times"])),
2692                    "Normal Epoch Train Times": self.member_vars["n_train_times"],
2693                }
2694            )
2695            pd2 = pd.DataFrame(
2696                {
2697                    "Epochs": np.arange(len(self.member_vars["p_train_times"])),
2698                    "PAI Epoch Train Times": self.member_vars["p_train_times"],
2699                }
2700            )
2701            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2702
2703            pd2 = pd.DataFrame(
2704                {
2705                    "Epochs": np.arange(len(self.member_vars["n_val_times"])),
2706                    "Normal Epoch Val Times": self.member_vars["n_val_times"],
2707                }
2708            )
2709            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2710
2711            pd2 = pd.DataFrame(
2712                {
2713                    "Epochs": np.arange(len(self.member_vars["p_val_times"])),
2714                    "PAI Epoch Val Times": self.member_vars["p_val_times"],
2715                }
2716            )
2717            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2718
2719            pd1.to_csv(
2720                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2721                index=False,
2722            )
2723            del pd1, pd2
2724        else:
2725            ax.plot(
2726                np.arange(len(self.member_vars["n_epoch_times"])),
2727                self.member_vars["n_epoch_times"],
2728                label="Normal Epoch Times",
2729            )
2730            ax.plot(
2731                np.arange(len(self.member_vars["p_epoch_times"])),
2732                self.member_vars["p_epoch_times"],
2733                label="PAI Epoch Times",
2734            )
2735
2736            plt.title(
2737                save_folder + "/" + self.save_name + "times (by train() and eval())"
2738            )
2739            plt.xlabel("Iteration")
2740            plt.ylabel("Epoch Time in Seconds ")
2741            ax.set_ylim(ymin=0)
2742            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2743
2744            pd1 = pd.DataFrame(
2745                {
2746                    "Epochs": np.arange(len(self.member_vars["n_epoch_times"])),
2747                    "Normal Epoch Times": self.member_vars["n_epoch_times"],
2748                }
2749            )
2750            pd2 = pd.DataFrame(
2751                {
2752                    "Epochs": np.arange(len(self.member_vars["p_epoch_times"])),
2753                    "PAI Epoch Times": self.member_vars["p_epoch_times"],
2754                }
2755            )
2756            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2757
2758            pd1.to_csv(
2759                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2760                index=False,
2761            )
2762            del pd1, pd2
2763
2764        if self.values_per_train_epoch != -1 and self.values_per_val_epoch != -1:
2765            ax2 = ax.twinx()  # Second axes sharing same x-axis
2766            ax2.set_ylabel("Single Datapoint Time in Seconds")
2767
2768            ax2.plot(
2769                np.arange(len(self.member_vars["n_train_times"])),
2770                np.array(self.member_vars["n_train_times"])
2771                / self.values_per_train_epoch,
2772                linestyle="dashed",
2773                label="Normal Train Item Times",
2774            )
2775            ax2.plot(
2776                np.arange(len(self.member_vars["p_train_times"])),
2777                np.array(self.member_vars["p_train_times"])
2778                / self.values_per_train_epoch,
2779                linestyle="dashed",
2780                label="PAI Train Item Times",
2781            )
2782            ax2.plot(
2783                np.arange(len(self.member_vars["n_val_times"])),
2784                np.array(self.member_vars["n_val_times"]) / self.values_per_val_epoch,
2785                linestyle="dashed",
2786                label="Normal Val Item Times",
2787            )
2788            ax2.plot(
2789                np.arange(len(self.member_vars["p_val_times"])),
2790                np.array(self.member_vars["p_val_times"]) / self.values_per_val_epoch,
2791                linestyle="dashed",
2792                label="PAI Val Item Times",
2793            )
2794            ax2.tick_params(axis="y")
2795            ax2.set_ylim(ymin=0)
2796            ax2.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2797
2798    def generate_learning_rate_plots(self, ax, save_folder, extra_string):
2799        """
2800        Generate plots and csvs for learning rate
2801
2802        Parameters
2803        ----------
2804        ax : object
2805            The matplotlib axis to plot on.
2806        save_folder : str
2807            The folder to save the plots and csvs in.
2808        extra_string : str
2809            An extra string to append to the filenames.
2810
2811        Returns
2812        -------
2813        None
2814
2815        """
2816        ax.plot(
2817            np.arange(len(self.member_vars["training_learning_rates"])),
2818            self.member_vars["training_learning_rates"],
2819            label="learning_rate",
2820        )
2821        plt.title(save_folder + "/" + self.save_name + "learning_rate")
2822        plt.xlabel("Epochs")
2823        plt.ylabel("learning_rate")
2824        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2825
2826        pd1 = pd.DataFrame(
2827            {
2828                "Epochs": np.arange(len(self.member_vars["training_learning_rates"])),
2829                "learning_rate": self.member_vars["training_learning_rates"],
2830            }
2831        )
2832        pd1.to_csv(
2833            save_folder + "/" + self.save_name + extra_string + "learning_rate.csv",
2834            index=False,
2835        )
2836        del pd1
2837
2838    def get_current_pb_scores(self):
2839        """
2840        Get the latest best PBScore of each dendrite layer, the same numbers
2841        written to the Best PBScores csv.
2842
2843        Returns
2844        -------
2845        dict[str, Any]
2846            Layer name to score.  Empty outside of dendrite scoring phases,
2847            when no candidate dendrites are being scored.
2848
2849
2850        Parameters
2851        ----------
2852        None
2853
2854        """
2855        if not self.member_vars["doing_pai"]:
2856            return {}
2857        if not GPA.pc.get_perforated_backpropagation():
2858            return {}
2859        # Scores only advance while candidate dendrites are being trained
2860        if (
2861            self.member_vars["mode"] != "p"
2862            and not GPA.pc.get_learn_dendrites_live()
2863        ):
2864            return {}
2865
2866        scores = {}
2867        for layer_id in range(len(self.neuron_module_vector)):
2868            if layer_id >= len(self.member_vars["best_scores"]):
2869                continue
2870            layer_scores = self.member_vars["best_scores"][layer_id]
2871            if len(layer_scores) == 0:
2872                continue
2873            score = layer_scores[-1]
2874            if hasattr(score, "item"):
2875                score = score.item()
2876            score = float(score)
2877            if math.isnan(score) or math.isinf(score):
2878                continue
2879            scores[self.neuron_module_vector[layer_id].name] = score
2880        return scores
2881
2882    def generate_dendrite_learning_plots(self, ax, save_folder, extra_string):
2883        """
2884        Generate dendrite score plots for the tracker.
2885        Also saves csv files associated with the plots.
2886
2887        Parameters
2888        ----------
2889        ax : matplotlib.axes.Axes
2890            Axis used for plotting dendrite-learning curves.
2891        save_folder : str
2892            Directory where plot images and CSV summaries are written.
2893        extra_string : str
2894            Filename suffix used to distinguish this output set.
2895
2896        Returns
2897        -------
2898        None
2899            Saves plots and score CSV files to disk.
2900        """
2901        if self.member_vars["doing_pai"]:
2902            pd1 = None
2903            pd2 = None
2904            num_colors = len(self.neuron_module_vector)
2905
2906            if (
2907                len(self.neuron_module_vector) > 0
2908                and len(self.member_vars["current_scores"][0]) != 0
2909            ):
2910                num_colors *= 2
2911
2912            cm = plt.get_cmap("gist_rainbow")
2913            ax.set_prop_cycle(
2914                "color", [cm(1.0 * i / num_colors) for i in range(num_colors)]
2915            )
2916
2917            for layer_id in range(len(self.neuron_module_vector)):
2918                ax.plot(
2919                    np.arange(len(self.member_vars["best_scores"][layer_id])),
2920                    self.member_vars["best_scores"][layer_id],
2921                    label=self.neuron_module_vector[layer_id].name,
2922                )
2923
2924                pd2 = pd.DataFrame(
2925                    {
2926                        "Epochs": np.arange(
2927                            len(self.member_vars["best_scores"][layer_id])
2928                        ),
2929                        f"Best ever for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2930                            "best_scores"
2931                        ][
2932                            layer_id
2933                        ],
2934                    }
2935                )
2936
2937                if pd1 is None:
2938                    pd1 = pd2
2939                else:
2940                    pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2941
2942                if len(self.member_vars["current_scores"][layer_id]) != 0:
2943                    ax.plot(
2944                        np.arange(len(self.member_vars["current_scores"][layer_id])),
2945                        self.member_vars["current_scores"][layer_id],
2946                        label=f"Current:{self.neuron_module_vector[layer_id].name}",
2947                    )
2948
2949                pd2 = pd.DataFrame(
2950                    {
2951                        "Epochs": np.arange(
2952                            len(self.member_vars["current_scores"][layer_id])
2953                        ),
2954                        f"Best current for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2955                            "current_scores"
2956                        ][
2957                            layer_id
2958                        ],
2959                    }
2960                )
2961                pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2962
2963            plt.title(save_folder + "/" + self.save_name + " Best PBScores")
2964            plt.xlabel("Epochs")
2965            plt.ylabel("Best PBScore")
2966            ax.legend(
2967                bbox_to_anchor=(1.05, 1),
2968                loc="upper left",
2969                ncol=max(1, math.ceil(len(self.neuron_module_vector) / 30)),
2970            )
2971            for switcher in self.member_vars["p_switch_epochs"]:
2972                plt.axvline(x=switcher, ymin=0, ymax=1, color="r")
2973
2974            if self.member_vars["mode"] == "p":
2975                missed_time = (
2976                    self.member_vars["num_epochs_run"]
2977                    - self.member_vars["epoch_last_improved"]
2978                )
2979                plt.axvline(
2980                    x=(len(self.member_vars["best_scores"][0]) - (missed_time + 1)),
2981                    ymin=0,
2982                    ymax=1,
2983                    color="g",
2984                )
2985
2986            # pd1 here will be none if no PB layers are created
2987            if pd1 is not None:
2988                pd1.to_csv(
2989                    save_folder
2990                    + "/"
2991                    + self.save_name
2992                    + extra_string
2993                    + "Best PBScores.csv",
2994                    index=False,
2995                )
2996            del pd1, pd2
2997
2998    def generate_extra_csv_files(self, save_folder, extra_string):
2999        """
3000        Generate additional csvs
3001
3002        Parameters
3003        ----------
3004        save_folder : str
3005            The folder to save the plots and csvs in.
3006        extra_string : str
3007            An extra string to append to the filenames.
3008
3009        Returns
3010        -------
3011        None
3012
3013        """
3014        pd1 = pd.DataFrame(
3015            {
3016                "Switch Number": np.arange(len(self.member_vars["switch_epochs"])),
3017                "Switch Epoch": self.member_vars["switch_epochs"],
3018            }
3019        )
3020        pd1.to_csv(
3021            save_folder + "/" + self.save_name + extra_string + "switch_epochs.csv",
3022            index=False,
3023        )
3024        del pd1
3025
3026        pd1 = pd.DataFrame(
3027            {
3028                "Switch Number": np.arange(len(self.member_vars["param_counts"])),
3029                "Param Count": self.member_vars["param_counts"],
3030            }
3031        )
3032        pd1.to_csv(
3033            save_folder + "/" + self.save_name + extra_string + "param_counts.csv",
3034            index=False,
3035        )
3036        del pd1
3037
3038        """
3039        Create best_arch_scores.csv file
3040        When working with dendrites there is a tradeoff between additional param count and score improvement.
3041        This file will help track that tradeoff by recording the best scores for all extra_scores
3042        and extra_scores_without_graphing for each architecture version.
3043        The scores recorded here are from the epoch when the best validation score was found
3044        within each switch_epoch boundary.
3045        """
3046        switch_counts = len(self.member_vars["switch_epochs"])
3047        best_valid = []
3048        associated_params = []
3049        
3050        # Initialize dictionaries to store best scores for each extra score type
3051        best_extra_scores = {}
3052        for score_name in self.member_vars["extra_scores"]:
3053            best_extra_scores[score_name] = []
3054        for score_name in self.member_vars["extra_scores_without_graphing"]:
3055            best_extra_scores[score_name] = []
3056
3057        for switch in range(0, switch_counts, 2):
3058            start_index = 0
3059            if switch != 0:
3060                start_index = self.member_vars["switch_epochs"][switch - 1] + 1
3061            end_index = self.member_vars["switch_epochs"][switch] + 1
3062
3063            if GPA.pai_tracker.member_vars["maximizing_score"]:
3064                best_valid_index = start_index + np.argmax(
3065                    self.member_vars["accuracies"][start_index:end_index]
3066                )
3067            else:
3068                best_valid_index = start_index + np.argmin(
3069                    self.member_vars["accuracies"][start_index:end_index]
3070                )
3071
3072            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3073            best_valid.append(best_valid_score)
3074            
3075            # Get corresponding scores from all extra_scores
3076            for score_name in self.member_vars["extra_scores"]:
3077                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3078                    best_extra_scores[score_name].append(
3079                        self.member_vars["extra_scores"][score_name][best_valid_index]
3080                    )
3081                else:
3082                    best_extra_scores[score_name].append(None)
3083            
3084            # Get corresponding scores from all extra_scores_without_graphing
3085            for score_name in self.member_vars["extra_scores_without_graphing"]:
3086                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3087                    best_extra_scores[score_name].append(
3088                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3089                    )
3090                else:
3091                    best_extra_scores[score_name].append(None)
3092            
3093            if self.member_vars["doing_pai"]:
3094                associated_params.append(self.member_vars["param_counts"][switch])
3095            else:
3096                associated_params.append(self.member_vars["param_counts"][-1])
3097
3098        # If in neuron training mode but not the very first epoch
3099        if self.member_vars["mode"] == "n" and (
3100            (len(self.member_vars["switch_epochs"]) == 0)
3101            or (
3102                self.member_vars["switch_epochs"][-1] + 1
3103                != len(self.member_vars["accuracies"])
3104            )
3105        ):
3106            start_index = 0
3107            if len(self.member_vars["switch_epochs"]) != 0:
3108                start_index = self.member_vars["switch_epochs"][-1] + 1
3109
3110            if GPA.pai_tracker.member_vars["maximizing_score"]:
3111                best_valid_index = start_index + np.argmax(
3112                    self.member_vars["accuracies"][start_index:]
3113                )
3114            else:
3115                best_valid_index = start_index + np.argmin(
3116                    self.member_vars["accuracies"][start_index:]
3117                )
3118
3119            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3120            best_valid.append(best_valid_score)
3121            
3122            # Get corresponding scores from all extra_scores
3123            for score_name in self.member_vars["extra_scores"]:
3124                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3125                    best_extra_scores[score_name].append(
3126                        self.member_vars["extra_scores"][score_name][best_valid_index]
3127                    )
3128                else:
3129                    best_extra_scores[score_name].append(None)
3130            
3131            # Get corresponding scores from all extra_scores_without_graphing
3132            for score_name in self.member_vars["extra_scores_without_graphing"]:
3133                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3134                    best_extra_scores[score_name].append(
3135                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3136                    )
3137                else:
3138                    best_extra_scores[score_name].append(None)
3139            
3140            associated_params.append(self.member_vars["param_counts"][-1])
3141
3142        # Build dataframe with all columns
3143        csv_data = {
3144            "Param Counts": associated_params,
3145            "Max Valid Scores": best_valid,
3146        }
3147        
3148        # Add columns for each extra score
3149        for score_name in best_extra_scores:
3150            csv_data[score_name] = best_extra_scores[score_name]
3151        
3152        pd1 = pd.DataFrame(csv_data)
3153        pd1.to_csv(
3154            save_folder + "/" + self.save_name + extra_string + "_best_arch_scores.csv",
3155            index=False,
3156        )
3157        del pd1
3158
3159    def save_graphs(self, extra_string=""):
3160        """
3161        Save graphs and csvs for all the values the tracker records
3162
3163        Parameters
3164        ----------
3165        extra_string : str
3166            An extra string to append to the filenames.
3167
3168        Returns
3169        -------
3170        None
3171
3172        """
3173        # If running DDP only save with rank 0
3174        if "RANK" in os.environ:
3175            if int(os.environ["RANK"]) != 0:
3176                return
3177        if not self.making_graphs:
3178            return
3179
3180        save_folder = "./" + self.save_name + "/"
3181
3182        plt.ioff()
3183        fig = plt.figure(figsize=(28, 14))
3184
3185        # Plot with accuracy scores
3186        ax = plt.subplot(221)
3187        self.generate_accuracy_plots(ax, save_folder, extra_string)
3188
3189        # Plot dendrite learning scores
3190        ax = plt.subplot(222)
3191        self.generate_dendrite_learning_plots(ax, save_folder, extra_string)
3192
3193        if GPA.pc.get_drawing_extra_graphs():
3194            # Plot learning rates for each training epoch
3195            ax = plt.subplot(223)
3196            self.generate_learning_rate_plots(ax, save_folder, extra_string)
3197
3198            # Plot the times for each training epoch
3199            ax = plt.subplot(224)
3200            self.generate_time_plots(ax, save_folder, extra_string)
3201
3202        # Generate extra CSV files
3203        self.generate_extra_csv_files(save_folder, extra_string)
3204
3205        fig.tight_layout()
3206        plt.savefig(save_folder + "/" + self.save_name + extra_string + ".png")
3207        plt.close("all")
3208
3209    def add_loss(self, loss):
3210        """Add loss to tracking vectors.
3211
3212        Parameters
3213        ----------
3214        loss : float or int
3215            The loss value to add.
3216
3217        Returns
3218        -------
3219        None
3220
3221        """
3222        if not isinstance(loss, (float, int)):
3223            loss = loss.item()
3224        self.member_vars["training_loss"].append(loss)
3225
3226    def add_learning_rate(self, learning_rate):
3227        """Add learning rate to tracking vectors.
3228
3229        Parameters
3230        ----------
3231        learning_rate : float or int
3232            The learning rate value to add.
3233
3234        Returns
3235        -------
3236        None
3237
3238        """
3239        if not isinstance(learning_rate, (float, int)):
3240            learning_rate = learning_rate.item()
3241        self.member_vars["training_learning_rates"].append(learning_rate)
3242
3243    def add_extra_score(self, score, extra_score_name):
3244        """Add extra score to tracking vectors.
3245
3246        Parameters
3247        ----------
3248        score : float or int
3249            The score value to add.
3250
3251        extra_score_name : str
3252            The name of the extra score.
3253
3254        Returns
3255        -------
3256        None
3257
3258        """
3259        if not isinstance(score, (float, int)):
3260            try:
3261                score = score.item()
3262            except:
3263                print(
3264                    "Scores added for Perforated Backpropagation should be "
3265                    "float, int, or tensor, yours is a:"
3266                )
3267                print(type(score))
3268                pdb.set_trace()
3269
3270        if GPA.pc.get_verbose():
3271            print(f"Adding extra score {extra_score_name} of {float(score)}")
3272
3273        if extra_score_name not in self.member_vars["extra_scores"]:
3274            self.member_vars["extra_scores"][extra_score_name] = []
3275        self.member_vars["extra_scores"][extra_score_name].append(score)
3276
3277        if self.member_vars["mode"] == "n":
3278            if extra_score_name not in self.member_vars["n_extra_scores"]:
3279                self.member_vars["n_extra_scores"][extra_score_name] = []
3280            self.member_vars["n_extra_scores"][extra_score_name].append(score)
3281
3282    def add_extra_score_without_graphing(self, score, extra_score_name):
3283        """Add extra score without graphing to tracking vectors.
3284
3285        Parameters
3286        ----------
3287        score : float or int
3288            The score value to add.
3289
3290        extra_score_name : str
3291            The name of the extra score.
3292
3293        Returns
3294        -------
3295        None
3296
3297        """
3298        if not isinstance(score, (float, int)):
3299            try:
3300                score = score.item()
3301            except:
3302                print(
3303                    "Scores added for Perforated Backpropagation should be "
3304                    "float, int, or tensor, yours is a:"
3305                )
3306                print(type(score))
3307                print("in add_extra_score_without_graphing")
3308                pdb.set_trace()
3309
3310        if GPA.pc.get_verbose():
3311            print(f"Adding extra score {extra_score_name} of {float(score)}")
3312
3313        if extra_score_name not in self.member_vars["extra_scores_without_graphing"]:
3314            self.member_vars["extra_scores_without_graphing"][extra_score_name] = []
3315        self.member_vars["extra_scores_without_graphing"][extra_score_name].append(
3316            score
3317        )
3318
3319    def add_test_score(self, score, extra_score_name):
3320        """Add test score to tracking vectors.
3321
3322        Parameters
3323        ----------
3324        score : float or int
3325            The score value to add.
3326
3327        extra_score_name : str
3328            The name of the extra score.
3329
3330        Returns
3331        -------
3332        None
3333
3334        Notes
3335        -----
3336        This function is a wrapper around `add_extra_score` that separates
3337        test score for adding to best_arch_scores.csv.
3338
3339        """
3340        self.add_extra_score(score, extra_score_name)
3341
3342        if not isinstance(score, (float, int)):
3343            try:
3344                score = score.item()
3345            except:
3346                print(
3347                    "Scores added for Perforated Backpropagation should be "
3348                    "float, int, or tensor, yours is a:"
3349                )
3350                print(type(score))
3351                print("in add_test_score")
3352                pdb.set_trace()
3353
3354        if GPA.pc.get_verbose():
3355            print(f"Adding test score {extra_score_name} of {float(score)}")
3356        self.member_vars["test_scores"].append(score)
3357
3358    def add_validation_score(self, accuracy, net, force_switch=False):
3359        """Function to add the validation score.
3360
3361        This is complex because it determines neuron and dendrite switching.
3362
3363        Parameters
3364        ----------
3365        accuracy : float or int
3366            The accuracy or loss value to add.
3367        net : object
3368            The neural network model.
3369        force_switch : bool, optional
3370            Whether to force a switch, by default False.
3371
3372        Returns
3373        -------
3374        net : object
3375            The potentially modified neural network model.
3376        training_complete : bool
3377            Whether training is complete.
3378        restructured : bool
3379            Whether the model has been restructured.
3380
3381        Notes
3382        -----
3383        WARNING: Do not call self anywhere in this function. When systems
3384        get loaded the actual tracker you are working with can change.
3385        """
3386
3387        _pai_log("info", f"Adding validation score {accuracy:.8f}")
3388
3389        update_learning_rate()
3390        update_param_count(net)
3391
3392        accuracy = check_input_problems(net, accuracy)
3393
3394        if len(GPA.pai_tracker.member_vars["switch_epochs"]) == 0:
3395            epochs_since_cycle_switch = GPA.pai_tracker.member_vars["num_epochs_run"]
3396        else:
3397            epochs_since_cycle_switch = (
3398                GPA.pai_tracker.member_vars["num_epochs_run"]
3399                - GPA.pai_tracker.member_vars["switch_epochs"][-1]
3400            )
3401
3402        update_running_accuracy(accuracy, epochs_since_cycle_switch)
3403        if GPA.pc.get_perforated_backpropagation():
3404            TPB.update_pb_scores(self)
3405
3406        # Captured before any switch below flips the mode and reloads scores
3407        epoch_pb_scores = self.get_current_pb_scores()
3408
3409        GPA.pai_tracker.stop_epoch(internal_call=True)
3410
3411        # If it is neuron training mode
3412        if (
3413            GPA.pai_tracker.member_vars["mode"] == "n"
3414            or GPA.pc.get_learn_dendrites_live()
3415        ):
3416            check_new_best(net, accuracy, epochs_since_cycle_switch)
3417        elif GPA.pc.get_perforated_backpropagation():
3418            TPB.check_best_pai_score_improvement()
3419
3420        # Save the latest model
3421        if GPA.pc.get_test_saves():
3422            UPA.save_system(net, GPA.pc.get_save_name(), "latest")
3423        if GPA.pc.get_pai_saves():
3424            UPA.pai_save_system(net, GPA.pc.get_save_name(), "latest")
3425
3426        restructuring_status_value = NO_MODEL_UPDATE
3427        # If it is time to switch based on scores and counter or a manual switch
3428        if GPA.pai_tracker.switch_time() or force_switch:
3429            # If testing dendrite capacity switch after enough dendrites added
3430            if (
3431                (GPA.pai_tracker.member_vars["mode"] == "n")
3432                and (GPA.pai_tracker.member_vars["num_dendrites_added"] > 2)
3433                and GPA.pc.get_testing_dendrite_capacity()
3434            ):
3435                GPA.pai_tracker.save_graphs()
3436                _pai_log(
3437                    "info",
3438                    "Successfully added 3 dendrites with GPA.pc.set_testing_dendrite_capacity(True) (default). "
3439                    "You may now set that to False and run a real experiment.",
3440                )
3441                return net, False, True
3442
3443            # If doing neuron training but this dendrite count didn't improve
3444            if (
3445                (GPA.pai_tracker.member_vars["mode"] == "n")
3446                or GPA.pc.get_learn_dendrites_live()
3447            ) and (GPA.pai_tracker.member_vars["current_n_set_global_best"] is False):
3448                new_restructuring_status_value, net = process_no_improvement(net)
3449                # if this was the final try return that training is complete
3450                if new_restructuring_status_value == TRAINING_COMPLETE:
3451                    if _dashboard_emitter is not None:
3452                        _dashboard_emitter.emit_run_end(GPA.pc)
3453                    return net, True, True
3454                else:
3455                    restructuring_status_value = update_restructuring_status(
3456                        restructuring_status_value, new_restructuring_status_value
3457                    )
3458            # Else if did improve, do a normal switch process
3459            else:
3460                if GPA.pc.get_verbose():
3461                    print(
3462                        f"Calling switch_mode with "
3463                        f'{GPA.pai_tracker.member_vars["current_n_set_global_best"]}, '
3464                        f'{GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]}, '
3465                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_steps"]}, '
3466                        f'{GPA.pai_tracker.member_vars["last_max_learning_rate_value"]},'
3467                        f'{GPA.pc.get_max_dendrites()},'
3468                        f'{GPA.pai_tracker.member_vars["num_dendrites_added"]},'
3469                        f'{GPA.pai_tracker.member_vars["num_dendrite_tries"]},'
3470                    )
3471                import pdb; pdb.set_trace
3472                # If the max number of dendrites has been hit or not doing pai and adding dendtites
3473                # then return rather than adding more
3474                if (
3475                    (GPA.pai_tracker.member_vars["mode"] == "n")
3476                    and (
3477                        GPA.pc.get_max_dendrites()
3478                        == GPA.pai_tracker.member_vars["num_dendrites_added"]
3479                    )
3480                ) or (GPA.pai_tracker.member_vars["doing_pai"] is False):
3481                    if GPA.pc.get_verbose():
3482                        print(
3483                            "Max dendrites reached or not doing PAI, finishing training"
3484                        )
3485                    net = process_final_network(net)
3486                    # Increment integrated if we have dendrites (means they're integrated)
3487                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3488                        GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3489                        _pai_log("info", f"Final dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3490                        if _dashboard_emitter is not None:
3491                            _dashboard_emitter.emit_dendrite_added(
3492                                GPA.pc,
3493                                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3494                                num_dendrites_integrated=GPA.pai_tracker.member_vars[
3495                                    "num_dendrites_integrated"
3496                                ],
3497                            )
3498                    if _dashboard_emitter is not None:
3499                        _dashboard_emitter.emit_run_end(GPA.pc)
3500                    return net, True, True
3501
3502                # Otherwise if its neuron training mode reset the counter of failed dendrites
3503                # Check if we should increment integrated count BEFORE change_learning_modes loads old state
3504                should_increment_integrated = False
3505                if GPA.pai_tracker.member_vars["mode"] == "n":
3506                    GPA.pai_tracker.member_vars["num_dendrite_tries"] = 0
3507                    if GPA.pc.get_verbose():
3508                        print(
3509                            "Adding new dendrites without resetting which means "
3510                            "the last ones improved. Resetting num_dendrite_tries"
3511                        )
3512                    # Remember to increment after change_learning_modes (which loads old tracker state)
3513                    if GPA.pai_tracker.member_vars["num_dendrites_added"] > 0:
3514                        should_increment_integrated = True
3515
3516                GPA.pai_tracker.save_graphs(
3517                    f'_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}'
3518                )
3519
3520                if GPA.pc.get_test_saves():
3521                    UPA.save_system(
3522                        net,
3523                        GPA.pc.get_save_name(),
3524                        f'beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3525                    )
3526                    # Copy current best model from this set of dendrites
3527                    # If running DDP only copy with rank 0
3528                    if "RANK" not in os.environ or int(os.environ["RANK"]) == 0:
3529                        shutil.copyfile(
3530                            f"{GPA.pc.get_save_name()}/best_model.pt",
3531                            f'{GPA.pc.get_save_name()}/best_model_beforeSwitch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}.pt',
3532                        )
3533
3534                net = UPA.change_learning_modes(
3535                    net,
3536                    GPA.pc.get_save_name(),
3537                    "best_model",
3538                    GPA.pai_tracker.member_vars["doing_pai"],
3539                )
3540                restructuring_status_value = NETWORK_RESTRUCTURED
3541                
3542                # Now increment after change_learning_modes has loaded the best model
3543                # This ensures the increment persists and doesn't get overwritten
3544                if should_increment_integrated:
3545                    GPA.pai_tracker.member_vars["num_dendrites_integrated"] += 1
3546                    _pai_log("info", f"Dendrites successfully integrated! Total integrated: {GPA.pai_tracker.member_vars['num_dendrites_integrated']}")
3547                    if _dashboard_emitter is not None:
3548                        _dashboard_emitter.emit_dendrite_added(
3549                            GPA.pc,
3550                            epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3551                            num_dendrites_integrated=GPA.pai_tracker.member_vars[
3552                                "num_dendrites_integrated"
3553                            ],
3554                        )
3555
3556            # If restructured is true, clear scheduler/optimizer before saving
3557            if restructuring_status_value != NETWORK_RESTRUCTURED:
3558                print(
3559                    "Restructured should always be triggered here, let us know if you encounter this situation"
3560                )
3561                pdb.set_trace()
3562
3563            # Since there is a restructuring optimizer and scheduler must be reinitialized after return
3564            GPA.pai_tracker.clear_optimizer_and_scheduler()
3565
3566            # Save the model from after the switch
3567            UPA.save_system(
3568                net,
3569                GPA.pc.get_save_name(),
3570                f'switch_{len(GPA.pai_tracker.member_vars["switch_epochs"])}',
3571            )
3572
3573        # If not time to switch and you have a scheduler, perform the update step
3574        elif GPA.pai_tracker.member_vars["scheduler"] is not None:
3575            new_restructuring_status_value, net = process_scheduler_update(
3576                net, accuracy, epochs_since_cycle_switch
3577            )
3578            restructuring_status_value = update_restructuring_status(
3579                restructuring_status_value, new_restructuring_status_value
3580            )
3581
3582        GPA.pai_tracker.start_epoch(internal_call=True)
3583        if _dashboard_emitter is not None:
3584            _mv = GPA.pai_tracker.member_vars
3585            _lr = _mv["training_learning_rates"][-1] if _mv["training_learning_rates"] else None
3586            _train_score = _mv["extra_scores"].get("train", [None])[-1]
3587            _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]
3588            _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]
3589            _dashboard_emitter.emit_epoch(
3590                GPA.pc,
3591                epoch=_mv["num_epochs_run"],
3592                validation_score=accuracy,
3593                learning_rate=_lr,
3594                train_score=_train_score,
3595                normal_time=_n_times[-1],
3596                pai_time=_p_times[-1],
3597                pb_scores=epoch_pb_scores,
3598            )
3599        GPA.pai_tracker.save_graphs()
3600
3601        if restructuring_status_value == NETWORK_RESTRUCTURED:
3602            GPA.pai_tracker.member_vars["epoch_last_improved"] = (
3603                GPA.pai_tracker.member_vars["num_epochs_run"]
3604            )
3605            if GPA.pc.get_verbose():
3606                print(
3607                    f"Setting epoch last improved to "
3608                    f'{GPA.pai_tracker.member_vars["epoch_last_improved"]}'
3609                )
3610
3611            now = datetime.now()
3612            dt_string = now.strftime("_%d.%m.%Y.%H.%M.%S")
3613
3614            if GPA.pc.get_verbose():
3615                print("Not saving restructure right now")
3616
3617            """
3618            This block of code helped with a save issue with safetensors and huggingface, but it breaks DDP.  
3619            Temporarily removing it to avoid DDP issues, but if you encounter save issues try adding it back in.
3620            for param in net.parameters():
3621                param.data = param.data.contiguous()
3622            """
3623        if GPA.pc.get_verbose():
3624            print(
3625                f"Completed adding score. Restructured is {restructuring_status_value}, "
3626                f"\ncurrent switch list is:"
3627            )
3628            print(GPA.pai_tracker.member_vars["switch_epochs"])
3629
3630        if _dashboard_emitter is not None and restructuring_status_value == NETWORK_RESTRUCTURED:
3631            _param_count = UPA.count_params(net)
3632            _dashboard_emitter.emit_switch(
3633                GPA.pc,
3634                switch_number=GPA.pai_tracker.member_vars["num_dendrites_added"],
3635                epoch=GPA.pai_tracker.member_vars["num_epochs_run"],
3636                param_count=_param_count,
3637                switch_type=GPA.pai_tracker.member_vars["mode"],
3638            )
3639
3640        # Always False for training complete if nothing triggered that training is over
3641        return net, restructuring_status_value, False
3642
3643    def clear_all_processors(self):
3644        """Clear all processors from modules.
3645
3646        Parameters
3647        ----------
3648        None
3649
3650        Returns
3651        -------
3652        None
3653            This function does not return a value.
3654        """
3655        for module in self.neuron_module_vector:
3656            module.clear_processors()
3657
3658    def create_new_dendrite_module(self):
3659        """Add dendrite module to all neuron modules.
3660
3661        Parameters
3662        ----------
3663        None
3664
3665        Returns
3666        -------
3667        None
3668            This function does not return a value.
3669        """
3670        for module in self.neuron_module_vector:
3671            module.create_new_dendrite_module()
3672
3673    def apply_pb_grads(self):
3674        """Apply perforated backpropagation gradients to all modules.
3675
3676        Parameters
3677        ----------
3678        None
3679
3680        Returns
3681        -------
3682        None
3683            This function does not return a value.
3684        """
3685        if self.member_vars["mode"] == "p":
3686            for module in self.neuron_module_vector:
3687                module.apply_pb_grads()
3688
3689    def apply_pb_zero(self):
3690        """Apply perforated backpropagation zero gradients to all modules.
3691
3692        Parameters
3693        ----------
3694        None
3695
3696        Returns
3697        -------
3698        None
3699            This function does not return a value.
3700        """
3701        if self.member_vars["mode"] == "p":
3702            for module in self.neuron_module_vector:
3703                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            f.write(
1553                f"{layer.name},{layer.dendrite_module.dendrite_values[0].out_channels}\n"
1554            )
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            channels[line.split(",")[0]] = int(line.split(",")[1])
1587        for layer in self.neuron_module_vector:
1588            layer.dendrite_module.dendrite_values[0].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=[]):
1590    def set_optimizer_instance(self, optimizer_instance, additional_optimizers=[]):
1591        """Set optimizer instance directly.
1592
1593        Parameters
1594        ----------
1595        optimizer_instance : object
1596            The optimizer instance to set.
1597
1598        Returns
1599        -------
1600        None
1601
1602        """
1603
1604        try:
1605            for param_group in optimizer_instance.param_groups:
1606                if (
1607                    param_group["weight_decay"] > 0
1608                    and GPA.pc.get_weight_decay_accepted() is False
1609                ):
1610                    _pai_log(
1611                        "warning",
1612                        "For PAI training it is recommended to not use weight decay in your optimizer",
1613                    )
1614
1615        except:
1616            pass
1617        self.member_vars["optimizer_instance"] = optimizer_instance
1618        if GPA.pc.get_perforated_backpropagation():
1619            TPB.setup_optimizer_pb(self.member_vars["optimizer_instance"])
1620            for optimizer in additional_optimizers:
1621                TPB.filter_params(optimizer)
1622        optimizer_instance.zero_grad()

Set optimizer instance directly.

Parameters
  • optimizer_instance (object): The optimizer instance to set.
Returns
  • None
def set_optimizer(self, optimizer):
1624    def set_optimizer(self, optimizer):
1625        """Set optimizer type to be initialized later
1626
1627        Parameters
1628        ----------
1629        optimizer : object
1630            The optimizer type to set.
1631
1632        Returns
1633        -------
1634        None
1635
1636        """
1637        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):
1639    def set_scheduler(self, scheduler):
1640        """Set scheduler type to be initialized later
1641
1642        Parameters
1643        ----------
1644        scheduler : object
1645            The scheduler type to set.
1646
1647        Returns
1648        -------
1649        None
1650
1651        """
1652        if scheduler is not torch.optim.lr_scheduler.ReduceLROnPlateau:
1653            if GPA.pc.get_verbose():
1654                print("Not using ReduceLROnPlateau, this is not recommended")
1655        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):
1657    def increment_scheduler(self, num_ticks, mode):
1658        """Increment the scheduler a set number of times.
1659
1660        Used for finding best initial learning rate when adding dendrites.
1661
1662        Parameters
1663        ----------
1664        num_ticks : int
1665            The number of scheduler steps to take.
1666        mode : str
1667            The mode for stepping the scheduler. Options are:
1668            - "step_learning_rate": Step based on improved accuracy epochs
1669            - "increment_epoch_count": Step based on total epoch count
1670
1671        Returns
1672        -------
1673        current_steps : int
1674            The number of learning rate changes that occurred.
1675        learning_rate1 : float
1676            The final learning rate after stepping.
1677
1678        """
1679
1680        current_steps = 0
1681        current_ticker = 0
1682
1683        for param_group in GPA.pai_tracker.member_vars[
1684            "optimizer_instance"
1685        ].param_groups:
1686            learning_rate1 = param_group["lr"]
1687
1688        if GPA.pc.get_verbose():
1689            print("Using scheduler:")
1690            print(type(self.member_vars["scheduler_instance"]))
1691
1692        while current_ticker < num_ticks:
1693            if GPA.pc.get_verbose():
1694                print(
1695                    f"Lower start rate initial {learning_rate1} "
1696                    f'stepping {GPA.pai_tracker.member_vars["current_n_learning_rate_initial_skip_steps"]} times'
1697                )
1698
1699            if (
1700                type(self.member_vars["scheduler_instance"])
1701                is torch.optim.lr_scheduler.ReduceLROnPlateau
1702            ):
1703                if mode == "step_learning_rate":
1704                    # Step with counter as last improved accuracy
1705                    self.member_vars["scheduler_instance"].step(
1706                        metrics=self.member_vars["last_improved_accuracies"][
1707                            GPA.pai_tracker.steps_after_switch() - 1
1708                        ]
1709                    )
1710                elif mode == "increment_epoch_count":
1711                    # Step with improved epoch counts up to current location
1712                    self.member_vars["scheduler_instance"].step(
1713                        metrics=self.member_vars["last_improved_accuracies"][
1714                            -((num_ticks - 1) - current_ticker) - 1
1715                        ]
1716                    )
1717            else:
1718                self.member_vars["scheduler_instance"].step()
1719
1720            for param_group in GPA.pai_tracker.member_vars[
1721                "optimizer_instance"
1722            ].param_groups:
1723                learning_rate2 = param_group["lr"]
1724
1725            if learning_rate2 != learning_rate1:
1726                current_steps += 1
1727                learning_rate1 = learning_rate2
1728                if mode == "step_learning_rate":
1729                    current_ticker += 1
1730                if GPA.pc.get_verbose():
1731                    print(f"1 step {current_steps} to {learning_rate2}")
1732
1733            if mode == "increment_epoch_count":
1734                current_ticker += 1
1735
1736        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):
1738    def setup_optimizer(self, net, opt_args, sched_args=None, parameters=None):
1739        """Initialize the optimizer and scheduler when added.
1740
1741        Parameters
1742        ----------
1743        net : object
1744            The neural network model.
1745        opt_args : dict
1746            The arguments for the optimizer.
1747        sched_args : dict, optional
1748            The arguments for the scheduler, by default None.
1749
1750        Returns
1751        -------
1752        optimizer : object
1753            The initialized optimizer instance.
1754        scheduler : object, optional
1755            The initialized scheduler instance, if a scheduler was set.
1756
1757        """
1758        if "weight_decay" in opt_args and not GPA.pc.get_weight_decay_accepted():
1759            _pai_log(
1760                "warning",
1761                "For PAI training it is recommended to not use weight decay in your optimizer",
1762            )
1763
1764        if ("model" not in opt_args.keys()) and "params" not in opt_args.keys():
1765            print("In setup_optimizer it will be depreciated to not pass in params yourself in the future")
1766            print("please change the settings to include params")
1767            if self.member_vars["mode"] == "n":
1768                if parameters is not None:
1769                    opt_args["params"] = parameters
1770                else:
1771                    opt_args["params"] = filter(lambda p: p.requires_grad, net.parameters())
1772            else:
1773                params = UPA.get_pai_network_params(net)
1774                if parameters is not None:
1775                    # Filter parameters to only those in params, preserving weight_decay
1776                    params_set = set(params)
1777                    filtered_params = []
1778                    for param_group in parameters:
1779                        filtered_group_params = [p for p in param_group["params"] if p in params_set]
1780                        if filtered_group_params:
1781                            filtered_params.append({
1782                                "params": filtered_group_params,
1783                                "weight_decay": param_group["weight_decay"]
1784                            })
1785                    opt_args["params"] = filtered_params
1786                else:
1787                    opt_args["params"] = params
1788        elif "params" in opt_args.keys():
1789            # Check if params is a list of param groups (dicts) or a single param group
1790            params_value = opt_args["params"]
1791            if isinstance(params_value, list) and len(params_value) > 0:
1792                # Check if it's a list of dicts (multiple param groups) or list of tensors (single group)
1793                if isinstance(params_value[0], dict):
1794                    # Multiple param groups format: [{"params": [...], "lr": ...}, ...]
1795                    # Filter each param group for requires_grad
1796                    filtered_param_groups = []
1797                    for param_group in params_value:
1798                        filtered_group_params = [p for p in param_group["params"] if p.requires_grad]
1799                        if filtered_group_params:
1800                            new_group = param_group.copy()
1801                            new_group["params"] = filtered_group_params
1802                            filtered_param_groups.append(new_group)
1803                    opt_args["params"] = filtered_param_groups
1804                else:
1805                    # Single param group format: [tensor1, tensor2, ...] or generator
1806                    # Filter for requires_grad
1807                    opt_args["params"] = [p for p in params_value if p.requires_grad]
1808            elif hasattr(params_value, '__iter__'):
1809                # Handle generators or other iterables
1810                opt_args["params"] = [p for p in params_value if p.requires_grad]
1811
1812        optimizer = self.member_vars["optimizer"](**opt_args)
1813        self.set_optimizer_instance(optimizer)
1814
1815        if self.member_vars["scheduler"] is not None:
1816            # Handle SequentialLR specially
1817            if self.member_vars["scheduler"] is torch.optim.lr_scheduler.SequentialLR:
1818                """
1819                sched_args should be a dict with "schedulers" (list of tuples) and "milestones"
1820                For example:
1821                sequential_schedArgs = {
1822                    "schedulers": [
1823                        (warmup_scheduler_class, warmup_schedArgs),
1824                        (main_scheduler_class, main_schedArgs)
1825                    ],
1826                    "milestones": [switch_epoch]
1827                }
1828                """
1829                schedulers = []
1830                milestones = sched_args.get("milestones", [])
1831                scheduler_configs = sched_args.get("schedulers", [])
1832                
1833                for scheduler_class, scheduler_args in scheduler_configs:
1834                    schedulers.append(scheduler_class(optimizer, **scheduler_args))
1835                
1836                self.member_vars["scheduler_instance"] = torch.optim.lr_scheduler.SequentialLR(
1837                    optimizer, schedulers=schedulers, milestones=milestones
1838                )
1839            else:
1840                self.member_vars["scheduler_instance"] = self.member_vars["scheduler"](
1841                    optimizer, **sched_args
1842                )
1843            current_steps = 0
1844
1845            for param_group in GPA.pai_tracker.member_vars[
1846                "optimizer_instance"
1847            ].param_groups:
1848                learning_rate1 = param_group["lr"]
1849
1850            if GPA.pc.get_verbose():
1851                print(
1852                    f"Resetting scheduler with {GPA.pai_tracker.steps_after_switch()} "
1853                    f"steps and {GPA.pc.get_initial_history_after_switches()} initial ticks to skip"
1854                )
1855
1856            # Find setting of previously used learning rate before adding dendrites
1857            if (
1858                GPA.pai_tracker.member_vars[
1859                    "current_n_learning_rate_initial_skip_steps"
1860                ]
1861                != 0
1862            ):
1863                additional_steps, learning_rate1 = self.increment_scheduler(
1864                    GPA.pai_tracker.member_vars[
1865                        "current_n_learning_rate_initial_skip_steps"
1866                    ],
1867                    "step_learning_rate",
1868                )
1869                current_steps += additional_steps
1870
1871            if self.member_vars["mode"] == "n" or GPA.pc.get_learn_dendrites_live():
1872                initial = GPA.pc.get_initial_history_after_switches()
1873            else:
1874                initial = 0
1875
1876            if GPA.pai_tracker.steps_after_switch() > initial:
1877                # Minus extra 1 because this gets called after start epoch
1878                additional_steps, learning_rate1 = self.increment_scheduler(
1879                    (GPA.pai_tracker.steps_after_switch() - initial) - 1,
1880                    "increment_epoch_count",
1881                )
1882                current_steps += additional_steps
1883
1884            if GPA.pc.get_verbose():
1885                print(
1886                    f"Scheduler update loop with {current_steps} "
1887                    f"ended with {learning_rate1}"
1888                )
1889                print(
1890                    f"Scheduler ended with {current_steps} steps "
1891                    f"and lr of {learning_rate1}"
1892                )
1893
1894            self.member_vars["current_step_count"] = current_steps
1895            return optimizer, self.member_vars["scheduler_instance"]
1896        else:
1897            return optimizer

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, optional): The initialized scheduler instance, if a scheduler was set.
def clear_optimizer_and_scheduler(self):
1899    def clear_optimizer_and_scheduler(self):
1900        """Clear the instances for saving.
1901
1902        Parameters
1903        ----------
1904        None
1905
1906        Returns
1907        -------
1908        None
1909            This function does not return a value.
1910        """
1911        self.member_vars["optimizer_instance"] = None
1912        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):
1914    def switch_time(self):
1915        """Determine if it's time to switch between neuron and dendrite training.
1916
1917        Parameters
1918        ----------
1919        None
1920
1921        Returns
1922        -------
1923        bool
1924            True if it's time to switch, False otherwise.
1925
1926        Notes
1927        -----
1928        Based on current settings and history of scores.
1929        """
1930
1931        switch_phrase = "No mode, this should never be the case."
1932        switch_number = GPA.pc.get_n_epochs_to_switch()
1933        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1934            switch_phrase = "DOING_SWITCH_EVERY_TIME"
1935        elif self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY:
1936            switch_phrase = "DOING_HISTORY"
1937        elif self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH:
1938            switch_phrase = "DOING_FIXED_SWITCH"
1939            switch_number = GPA.pc.get_fixed_switch_num()
1940        elif self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1941            switch_phrase = "DOING_NO_SWITCH"
1942        else:
1943            print(
1944                "A switch mode must be set.  Check your settings for GPA.pc.set_switch_mode()."
1945            )
1946            pdb.set_trace()
1947        if not GPA.pc.get_silent():
1948            if(GPA.pc.get_perforated_backpropagation()):
1949                print(
1950                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1951                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1952                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1953                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1954                    f'n: {switch_number}, p: {GPA.pc.get_p_epochs_to_switch()}, '
1955                    f'num_cycles: {self.member_vars["num_cycles"]}'
1956                )
1957            else:
1958                print(
1959                    f'Checking PAI switch with mode {self.member_vars["mode"]}, '
1960                    f'switch mode {switch_phrase}, epoch {self.member_vars["num_epochs_run"]}, '
1961                    f'last improved epoch {self.member_vars["epoch_last_improved"]}, '
1962                    f'total epochs {self.member_vars["total_epochs_run"]}, '
1963                    f'n: {switch_number}, num_cycles: {self.member_vars["num_cycles"]}'
1964                )
1965            print(
1966                f'  Score tracking: current_n_set_global_best={self.member_vars["current_n_set_global_best"]}, '
1967                f'global_best={self.member_vars["global_best_validation_score"]:.4f}, '
1968                f'current_best={self.member_vars["current_best_validation_score"]:.4f}'
1969            )
1970        if GPA.pc.get_perforated_backpropagation():
1971            # this will fill in epoch last improved
1972            TPB.best_pai_score_improved_this_epoch(self)  ## CLOSED ONLY
1973        if self.member_vars["switch_mode"] == GPA.pc.DOING_NO_SWITCH:
1974            if not GPA.pc.get_silent():
1975                print("Returning False - doing no switch mode")
1976            return False
1977
1978        if self.member_vars["switch_mode"] == GPA.pc.DOING_SWITCH_EVERY_TIME:
1979            if not GPA.pc.get_silent():
1980                print("Returning True - switching every time")
1981            return True
1982
1983        # Check if we're in the middle of learning rate optimization
1984        # If so, block ALL switch triggers until committed
1985        if GPA.pc.get_verbose():
1986            print("=== LR Optimization Check ===")
1987            print(f'  mode == "n": {self.member_vars["mode"] == "n"}')
1988            print(f"  get_learn_dendrites_live(): {GPA.pc.get_learn_dendrites_live()}")
1989            print(f'  committed_to_initial_rate: {GPA.pai_tracker.member_vars["committed_to_initial_rate"]}')
1990            print(f"  get_dont_give_up_unless_learning_rate_lowered(): {GPA.pc.get_dont_give_up_unless_learning_rate_lowered()}")
1991            print(f'  current_n_learning_rate_initial_skip_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"]}')
1992            print(f'  last_max_learning_rate_steps: {self.member_vars["last_max_learning_rate_steps"]}')
1993            print(f'  skip_steps < max_steps: {self.member_vars["current_n_learning_rate_initial_skip_steps"] < self.member_vars["last_max_learning_rate_steps"]}')
1994            print(f'  scheduler is not None: {self.member_vars["scheduler"] is not None}')
1995            print("=============================")
1996        
1997        if (
1998            ((self.member_vars["mode"] == "n") or GPA.pc.get_learn_dendrites_live())
1999            and (GPA.pai_tracker.member_vars["committed_to_initial_rate"] is False)
2000            and (GPA.pc.get_dont_give_up_unless_learning_rate_lowered())
2001            and (
2002                self.member_vars["current_n_learning_rate_initial_skip_steps"]
2003                <= self.member_vars["last_max_learning_rate_steps"]
2004            )
2005            and self.member_vars["scheduler"] is not None
2006        ):
2007            if not GPA.pc.get_silent():
2008                print(
2009                    f"Returning False - learning rate optimization in progress. "
2010                    f"Not committed yet. Comparing "
2011                    f'initial {self.member_vars["current_n_learning_rate_initial_skip_steps"]} '
2012                    f'to last max {self.member_vars["last_max_learning_rate_steps"]}'
2013                )
2014            return False
2015
2016        if len(self.member_vars["switch_epochs"]) == 0:
2017            this_count = self.member_vars["num_epochs_run"]
2018        else:
2019            this_count = (
2020                self.member_vars["num_epochs_run"]
2021                - self.member_vars["switch_epochs"][-1]
2022            )
2023        cap_switch = False
2024        if GPA.pc.get_perforated_backpropagation():
2025            cap_switch = TPB.check_cap_switch(self, this_count)
2026
2027        if self.member_vars["switch_mode"] == GPA.pc.DOING_HISTORY and (
2028            (
2029                (self.member_vars["mode"] == "n")
2030                and (
2031                    self.member_vars["num_epochs_run"]
2032                    - self.member_vars["epoch_last_improved"]
2033                    >= GPA.pc.get_n_epochs_to_switch()
2034                )
2035                and this_count
2036                >= GPA.pc.get_initial_history_after_switches()
2037                + GPA.pc.get_n_epochs_to_switch()
2038            )
2039            or (GPA.pc.get_perforated_backpropagation() and TPB.history_switch(self))
2040            or cap_switch
2041        ):
2042            if not GPA.pc.get_silent():
2043                print("Returning True - History and last improved is hit")
2044            return True
2045
2046        if self.member_vars["switch_mode"] == GPA.pc.DOING_FIXED_SWITCH and (
2047            (
2048                self.member_vars["total_epochs_run"] % GPA.pc.get_fixed_switch_num()
2049                == GPA.pc.get_fixed_switch_num() - 1
2050            )
2051            and self.member_vars["num_epochs_run"]
2052            >= GPA.pc.get_first_fixed_switch_num() - 1
2053        ):
2054            if not GPA.pc.get_silent():
2055                print("Returning True - Fixed switch number is hit")
2056            return True
2057
2058        if not GPA.pc.get_silent():
2059            print("Returning False - no triggers to switch have been hit")
2060        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):
2062    def steps_after_switch(self):
2063        """Based on settings, return value for steps since a switch.
2064
2065        Different options for param vals setting determine what is returned.
2066
2067        Parameters
2068        ----------
2069        None
2070
2071        Returns
2072        -------
2073        int
2074            The number of epochs since the last switch, or total epochs run,
2075            depending on settings.
2076
2077        """
2078        if self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_TOTAL_EPOCH:
2079            return self.member_vars["num_epochs_run"]
2080        elif (
2081            self.member_vars["param_vals_setting"] == GPA.pc.PARAM_VALS_BY_UPDATE_EPOCH
2082        ):
2083            return self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2084        elif (
2085            self.member_vars["param_vals_setting"]
2086            == GPA.pc.PARAM_VALS_BY_NEURON_EPOCH_START
2087        ):
2088            if self.member_vars["mode"] == "p":
2089                return (
2090                    self.member_vars["num_epochs_run"] - self.member_vars["last_switch"]
2091                )
2092            else:
2093                return self.member_vars["num_epochs_run"]
2094        else:
2095            print(
2096                f'{self.member_vars["param_vals_setting"]} is not a valid param vals option'
2097            )
2098            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):
2100    def add_pai_neuron_module(self, new_module, initial_add=True):
2101        """Add neuron modules to internal vectors.
2102
2103        Parameters
2104        ----------
2105        new_module : object
2106            The new module to add.
2107        initial_add : bool, optional
2108            Whether this is the initial addition rather than loading from file
2109
2110        Returns
2111        -------
2112        None
2113
2114        """
2115
2116        # If it's a duplicate, ignore the second addition
2117        if new_module in self.neuron_module_vector:
2118            return
2119        self.neuron_module_vector.append(new_module)
2120        if self.member_vars["doing_pai"]:
2121            PA.set_wrapped_params(new_module)
2122        if initial_add:
2123            self.member_vars["best_scores"].append([])
2124            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):
2126    def add_tracked_neuron_module(self, new_module, initial_add=True):
2127        """Add tracked modules to internal vectors
2128
2129        Parameters
2130        ----------
2131        new_module : object
2132            The new module to add.
2133        initial_add : bool, optional
2134            Whether this is the initial addition rather than loading from file
2135
2136        Returns
2137        -------
2138        None
2139
2140        """
2141        # If it's a duplicate, ignore the second addition
2142        if new_module in self.tracked_neuron_module_vector:
2143            return
2144        self.tracked_neuron_module_vector.append(new_module)
2145        if self.member_vars["doing_pai"]:
2146            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):
2148    def reset_module_vector(self, net, load_from_restart):
2149        """Clear internal vectors and reset from network.
2150
2151        Parameters
2152        ----------
2153        net : object
2154            The neural network model.
2155        load_from_restart : bool
2156            Whether loading from a restart file.
2157
2158        Returns
2159        -------
2160        None
2161
2162        """
2163        self.neuron_module_vector = []
2164        self.tracked_neuron_module_vector = []
2165        this_list = UPA.get_pai_modules(net, 0)
2166        for module in this_list:
2167            self.add_pai_neuron_module(module, initial_add=load_from_restart)
2168        this_list = UPA.get_tracked_modules(net, 0)
2169        for module in this_list:
2170            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):
2172    def reset_vals_for_score_reset(self):
2173        """Reset cycle scores for new cycle.
2174
2175        Parameters
2176        ----------
2177        None
2178
2179        Returns
2180        -------
2181        None
2182            This function does not return a value.
2183        """
2184
2185        if GPA.pc.get_find_best_lr():
2186            self.member_vars["committed_to_initial_rate"] = False
2187            print("Resetting committed to initial rate to False")
2188        # If retaining all dendrties always say that the current dendrites set global best for saving and loading
2189        if GPA.pc.get_retain_all_dendrites():
2190            self.member_vars["current_n_set_global_best"] = True
2191            self.member_vars["global_best_validation_score"] = 0
2192        else:
2193            self.member_vars["current_n_set_global_best"] = False
2194
2195        # Don't reset global best, but do reset current best
2196        self.member_vars["current_best_validation_score"] = 0
2197        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):
2199    def set_dendrite_training(self):
2200        """Signal all layers to start dendrite training.
2201
2202        Parameters
2203        ----------
2204        None
2205
2206        Returns
2207        -------
2208        None
2209            This function does not return a value.
2210        """
2211        if GPA.pc.get_verbose():
2212            print("Calling set_dendrite_training")
2213
2214        for layer in self.neuron_module_vector[:]:
2215            worked = layer.set_mode("p")
2216            """
2217            worked is False when a layer was added to the neuron module vector
2218            but then it's never actually been used. This can happen when
2219            you have set a layer to have requires_grad = False or when
2220            you have a module as a member variable but it's not actually
2221            part of the network. Should be moved to be a tracked layer
2222            rather than a neuron layer.
2223            """
2224            if not worked:
2225                self.neuron_module_vector.remove(layer)
2226
2227        for layer in self.tracked_neuron_module_vector[:]:
2228            worked = layer.set_mode("p")
2229
2230        self.create_new_dendrite_module()
2231        self.member_vars["mode"] = "p"
2232        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2233
2234        if GPA.pc.get_learn_dendrites_live():
2235            self.reset_vals_for_score_reset()
2236
2237        self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2238            "current_step_count"
2239        ]
2240
2241        GPA.pai_tracker.member_vars["current_cycle_lr_max_scores"] = []
2242        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):
2245    def set_neuron_training(self):
2246        """Signal all layers to start neuron training.
2247
2248        Parameters
2249        ----------
2250        None
2251
2252        Returns
2253        -------
2254        None
2255            This function does not return a value.
2256        """
2257        for module in self.neuron_module_vector:
2258            module.set_mode("n")
2259        for module in self.tracked_neuron_module_vector[:]:
2260            module.set_mode("n")
2261
2262        self.member_vars["mode"] = "n"
2263        self.member_vars["num_dendrites_added"] += 1
2264        self.member_vars["current_n_learning_rate_initial_skip_steps"] = 0
2265        self.reset_vals_for_score_reset()
2266
2267        self.member_vars["current_cycle_lr_max_scores"] = []
2268        if GPA.pc.get_learn_dendrites_live():
2269            self.member_vars["last_max_learning_rate_steps"] = self.member_vars[
2270                "current_step_count"
2271            ]
2272        GPA.pai_tracker.member_vars["num_cycles"] += 1
2273
2274        if GPA.pc.get_reset_best_score_on_switch():
2275            GPA.pai_tracker.member_vars["current_best_validation_score"] = 0
2276            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):
2278    def start_epoch(self, internal_call=False):
2279        """Perform steps for when a new training epoch is about to begin.
2280
2281        Parameters
2282        ----------
2283        internal_call : bool, optional
2284            Whether this is an internal call or manual call
2285
2286        Returns
2287        -------
2288        None
2289
2290        Notes
2291        -----
2292        If you ever need to call this manually, set internal_call to False.
2293
2294        """
2295        if self.member_vars["manual_train_switch"] and internal_call:
2296            return
2297
2298        if not internal_call and not self.member_vars["manual_train_switch"]:
2299            self.member_vars["manual_train_switch"] = True
2300            self.saved_time = 0
2301            self.member_vars["num_epochs_run"] = -1
2302            self.member_vars["total_epochs_run"] = -1
2303
2304        end = time.time()
2305        if self.member_vars["manual_train_switch"]:
2306            if self.saved_time != 0:
2307                if self.member_vars["mode"] == "p":
2308                    self.member_vars["p_val_times"].append(end - self.saved_time)
2309                else:
2310                    self.member_vars["n_val_times"].append(end - self.saved_time)
2311
2312        if self.member_vars["mode"] == "p":
2313            for layer in self.neuron_module_vector:
2314                for m in range(0, GPA.pc.get_global_candidates()):
2315                    with torch.no_grad():
2316                        if GPA.pc.get_verbose():
2317                            print(f"Resetting score for {layer.name}")
2318                        # Snapshot best_score before reset so we can compute per-epoch improvement
2319                        layer.dendrite_module.dendrite_values[
2320                            m
2321                        ].epoch_start_best_score.copy_(
2322                            layer.dendrite_module.dendrite_values[
2323                                m
2324                            ].best_score.detach()
2325                        )
2326                        layer.dendrite_module.dendrite_values[
2327                            m
2328                        ].best_score_improved_this_epoch = (
2329                            layer.dendrite_module.dendrite_values[
2330                                m
2331                            ].best_score_improved_this_epoch
2332                            * 0
2333                        )
2334                        layer.dendrite_module.dendrite_values[
2335                            m
2336                        ].nodes_best_improved_this_epoch = (
2337                            layer.dendrite_module.dendrite_values[
2338                                m
2339                            ].nodes_best_improved_this_epoch
2340                            * 0
2341                        )
2342                        layer.dendrite_module.dendrite_values[
2343                            m
2344                        ].nodes_improved_any = (
2345                            layer.dendrite_module.dendrite_values[
2346                                m
2347                            ].nodes_improved_any
2348                            * 0
2349                        )
2350            if GPA.pc.get_perforated_backpropagation():
2351                self.member_vars["best_mean_score_improved_this_epoch"] = 0
2352        self.member_vars["num_epochs_run"] += 1
2353        self.member_vars["total_epochs_run"] = (
2354            self.member_vars["num_epochs_run"] + self.member_vars["overwritten_epochs"]
2355        )
2356        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):
2358    def stop_epoch(self, internal_call=False):
2359        """Perform steps when a training epoch has completed.
2360
2361        Parameters
2362        ----------
2363        internal_call : bool, optional
2364            Whether this is an internal call or manual call
2365
2366        Returns
2367        -------
2368        None
2369
2370        Notes
2371        -----
2372        If you ever need to call this manually, set internal_call to False.
2373
2374        """
2375        end = time.time()
2376        if self.member_vars["manual_train_switch"] and internal_call:
2377            return
2378
2379        if self.member_vars["manual_train_switch"]:
2380            if self.member_vars["mode"] == "p":
2381                self.member_vars["p_train_times"].append(end - self.saved_time)
2382            else:
2383                self.member_vars["n_train_times"].append(end - self.saved_time)
2384        else:
2385            if self.member_vars["mode"] == "p":
2386                self.member_vars["p_epoch_times"].append(end - self.saved_time)
2387            else:
2388                self.member_vars["n_epoch_times"].append(end - self.saved_time)
2389
2390        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):
2392    def initialize(
2393        self,
2394        model,
2395        doing_pai=True,
2396        save_name="PAI",
2397        making_graphs=True,
2398        maximizing_score=True,
2399        num_classes=10000,
2400        values_per_train_epoch=-1,
2401        values_per_val_epoch=-1,
2402        zooming_graph=True,
2403    ):
2404        """Setup the tracker with initial settings.
2405
2406
2407        Parameters
2408        ----------
2409        model : object
2410            The neural network model.
2411        doing_pai : bool, optional
2412            Whether to add dendrites, by default True.
2413        save_name : str, optional
2414            The name under which to save the model.
2415        making_graphs : bool, optional
2416            Whether to make graphs, by default True.
2417        maximizing_score : bool, optional
2418            Whether to maximize the score, by default True.
2419        num_classes : int, optional
2420            The number of classes in the dataset, unused
2421        values_per_train_epoch : int, optional
2422            The number of values to look back for graphing
2423            during training, by default -1 (all values).
2424        values_per_val_epoch : int, optional
2425            The number of values to look back for graphing
2426            during validation, by default -1 (all values).
2427        zooming_graph : bool, optional
2428            Whether to zoom on graphs, by default True.
2429
2430
2431        Returns
2432        -------
2433        nn.Module
2434            Converted model instance configured for the tracker settings.
2435        """
2436        model = UPA.convert_network(model)
2437        self.member_vars["doing_pai"] = doing_pai
2438        self.member_vars["maximizing_score"] = maximizing_score
2439        self.save_name = save_name
2440        self.zooming_graph = zooming_graph
2441        self.making_graphs = making_graphs
2442
2443        if not self.loaded:
2444            self.member_vars["running_accuracy"] = (1.0 / num_classes) * 100
2445
2446        self.values_per_train_epoch = values_per_train_epoch
2447        self.values_per_val_epoch = values_per_val_epoch
2448
2449        if GPA.pc.get_testing_dendrite_capacity():
2450            if not GPA.pc.get_silent():
2451                print("Running a test of Dendrite Capacity.")
2452            GPA.pc.set_switch_mode(GPA.pc.DOING_SWITCH_EVERY_TIME)
2453            self.member_vars["switch_mode"] = GPA.pc.get_switch_mode()
2454            GPA.pc.set_retain_all_dendrites(True)
2455            GPA.pc.set_max_dendrite_tries(1000)
2456            GPA.pc.set_max_dendrites(1000)
2457            if GPA.pc.get_perforated_backpropagation():
2458                GPA.pc.set_initial_correlation_batches(1)
2459        else:
2460            if not GPA.pc.get_silent():
2461                print("Running Dendrite Experiment")
2462        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):
2464    def generate_accuracy_plots(self, ax, save_folder, extra_string):
2465        """
2466        Generate plots and csvs for accuracy
2467
2468        Parameters
2469        ----------
2470        ax : object
2471            The matplotlib axis to plot on.
2472        save_folder : str
2473            The folder to save the plots and csvs in.
2474        extra_string : str
2475            An extra string to append to the filenames.
2476
2477        Returns
2478        -------
2479        None
2480
2481        """
2482
2483        # If scores are being saved for epochs that get overwritten, plot them
2484        for list_id in range(len(self.member_vars["overwritten_extras"])):
2485            for extra_id in self.member_vars["overwritten_extras"][list_id]:
2486                ax.plot(
2487                    np.arange(
2488                        len(self.member_vars["overwritten_extras"][list_id][extra_id])
2489                    ),
2490                    self.member_vars["overwritten_extras"][list_id][extra_id],
2491                    "r",
2492                )
2493            ax.plot(
2494                np.arange(len(self.member_vars["overwritten_vals"][list_id])),
2495                self.member_vars["overwritten_vals"][list_id],
2496                "b",
2497            )
2498
2499        # Determine which accuracy vector to use
2500        if GPA.pc.get_drawing_pai():
2501            accuracies = self.member_vars["accuracies"]
2502        else:
2503            accuracies = self.member_vars["n_accuracies"]
2504
2505        # Get pointer to additional scores being saved
2506        extra_scores = self.member_vars["extra_scores"]
2507
2508        # Plot the main accuracy scores
2509        ax.plot(np.arange(len(accuracies)), accuracies, label="Validation Scores")
2510        ax.plot(
2511            np.arange(len(self.member_vars["running_accuracies"])),
2512            self.member_vars["running_accuracies"],
2513            label="Validation Running Scores",
2514        )
2515
2516        # Plot additional scores
2517        for extra_score in extra_scores:
2518            ax.plot(
2519                np.arange(len(extra_scores[extra_score])),
2520                extra_scores[extra_score],
2521                label=extra_score,
2522            )
2523
2524        plt.title(save_folder + "/" + self.save_name + "Scores")
2525        plt.xlabel("Epochs")
2526        plt.ylabel("Score")
2527
2528        # Add point at epoch last improved and best validation score
2529        if GPA.pc.get_drawing_pai():
2530            ax.plot(
2531                self.member_vars["epoch_last_improved"],
2532                self.member_vars["global_best_validation_score"],
2533                "bo",
2534                label="Global best (y)",
2535            )
2536            ax.plot(
2537                self.member_vars["epoch_last_improved"],
2538                accuracies[self.member_vars["epoch_last_improved"]],
2539                "go",
2540                label="Epoch Last Improved",
2541            )
2542        else:
2543            if self.member_vars["mode"] == "n":
2544                missed_time = (
2545                    self.member_vars["num_epochs_run"]
2546                    - self.member_vars["epoch_last_improved"]
2547                )
2548                ax.plot(
2549                    (len(self.member_vars["n_accuracies"]) - 1) - missed_time,
2550                    self.member_vars["n_accuracies"][-(missed_time + 1)],
2551                    "go",
2552                    label="Epoch Last Improved",
2553                )
2554
2555        # Generate csv file for the values graphed
2556        pd1 = pd.DataFrame(
2557            {"Epochs": np.arange(len(accuracies)), "Validation Scores": accuracies}
2558        )
2559        pd2 = pd.DataFrame(
2560            {
2561                "Epochs": np.arange(len(self.member_vars["running_accuracies"])),
2562                "Validation Running Scores": self.member_vars["running_accuracies"],
2563            }
2564        )
2565        pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2566        for extra_score in extra_scores:
2567            pd2 = pd.DataFrame(
2568                {
2569                    "Epochs": np.arange(len(extra_scores[extra_score])),
2570                    extra_score: extra_scores[extra_score],
2571                }
2572            )
2573            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2574        extra_scores_without_graphing = self.member_vars[
2575            "extra_scores_without_graphing"
2576        ]
2577        for extra_score in extra_scores_without_graphing:
2578            pd2 = pd.DataFrame(
2579                {
2580                    "Epochs": np.arange(
2581                        len(extra_scores_without_graphing[extra_score])
2582                    ),
2583                    extra_score: extra_scores_without_graphing[extra_score],
2584                }
2585            )
2586            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2587        pd1.to_csv(
2588            save_folder + "/" + self.save_name + extra_string + "Scores.csv",
2589            index=False,
2590        )
2591        del pd1, pd2
2592
2593        # Set y min and max to zoom in on important part of axis
2594        if (
2595            len(self.member_vars["switch_epochs"]) > 0
2596            and self.member_vars["switch_epochs"][0] > 0
2597            and self.zooming_graph
2598        ):
2599            if GPA.pai_tracker.member_vars["maximizing_score"]:
2600                min_val = np.array(
2601                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2602                ).mean()
2603                for extra_score in extra_scores:
2604                    min_pot = np.array(
2605                        extra_scores[extra_score][
2606                            0 : self.member_vars["switch_epochs"][0]
2607                        ]
2608                    ).mean()
2609                    if min_pot < min_val:
2610                        min_val = min_pot
2611                ax.set_ylim(ymin=min_val)
2612            else:
2613                max_val = np.array(
2614                    accuracies[0 : self.member_vars["switch_epochs"][0]]
2615                ).mean()
2616                for extra_score in extra_scores:
2617                    max_pot = np.array(
2618                        extra_scores[extra_score][
2619                            0 : self.member_vars["switch_epochs"][0]
2620                        ]
2621                    ).mean()
2622                    if max_pot > max_val:
2623                        max_val = max_pot
2624                ax.set_ylim(ymax=max_val)
2625
2626        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2627
2628        # Draw vertical lines for epochs where a dendrite switch occurred
2629        if GPA.pc.get_drawing_pai() and self.member_vars["doing_pai"]:
2630            color = "r"
2631            for switcher in self.member_vars["switch_epochs"]:
2632                plt.axvline(x=switcher, ymin=0, ymax=1, color=color)
2633                if color == "r":
2634                    color = "b"
2635                else:
2636                    color = "r"
2637        else:
2638            for switcher in self.member_vars["n_switch_epochs"]:
2639                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):
2641    def generate_time_plots(self, ax, save_folder, extra_string):
2642        """
2643        Generate plots and csvs for timing
2644
2645        Parameters
2646        ----------
2647        ax : object
2648            The matplotlib axis to plot on.
2649        save_folder : str
2650            The folder to save the plots and csvs in.
2651        extra_string : str
2652            An extra string to append to the filenames.
2653
2654        Returns
2655        -------
2656        None
2657
2658        """
2659        if self.member_vars["manual_train_switch"]:
2660            ax.plot(
2661                np.arange(len(self.member_vars["n_train_times"])),
2662                self.member_vars["n_train_times"],
2663                label="Normal Epoch Train Times",
2664            )
2665            ax.plot(
2666                np.arange(len(self.member_vars["p_train_times"])),
2667                self.member_vars["p_train_times"],
2668                label="PAI Epoch Train Times",
2669            )
2670            ax.plot(
2671                np.arange(len(self.member_vars["n_val_times"])),
2672                self.member_vars["n_val_times"],
2673                label="Normal Epoch Val Times",
2674            )
2675            ax.plot(
2676                np.arange(len(self.member_vars["p_val_times"])),
2677                self.member_vars["p_val_times"],
2678                label="PAI Epoch Val Times",
2679            )
2680
2681            plt.title(
2682                save_folder + "/" + self.save_name + "times (by train() and eval())"
2683            )
2684            plt.xlabel("Iteration")
2685            plt.ylabel("Epoch Time in Seconds ")
2686            ax.set_ylim(ymin=0)
2687            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2688
2689            pd1 = pd.DataFrame(
2690                {
2691                    "Epochs": np.arange(len(self.member_vars["n_train_times"])),
2692                    "Normal Epoch Train Times": self.member_vars["n_train_times"],
2693                }
2694            )
2695            pd2 = pd.DataFrame(
2696                {
2697                    "Epochs": np.arange(len(self.member_vars["p_train_times"])),
2698                    "PAI Epoch Train Times": self.member_vars["p_train_times"],
2699                }
2700            )
2701            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2702
2703            pd2 = pd.DataFrame(
2704                {
2705                    "Epochs": np.arange(len(self.member_vars["n_val_times"])),
2706                    "Normal Epoch Val Times": self.member_vars["n_val_times"],
2707                }
2708            )
2709            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2710
2711            pd2 = pd.DataFrame(
2712                {
2713                    "Epochs": np.arange(len(self.member_vars["p_val_times"])),
2714                    "PAI Epoch Val Times": self.member_vars["p_val_times"],
2715                }
2716            )
2717            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2718
2719            pd1.to_csv(
2720                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2721                index=False,
2722            )
2723            del pd1, pd2
2724        else:
2725            ax.plot(
2726                np.arange(len(self.member_vars["n_epoch_times"])),
2727                self.member_vars["n_epoch_times"],
2728                label="Normal Epoch Times",
2729            )
2730            ax.plot(
2731                np.arange(len(self.member_vars["p_epoch_times"])),
2732                self.member_vars["p_epoch_times"],
2733                label="PAI Epoch Times",
2734            )
2735
2736            plt.title(
2737                save_folder + "/" + self.save_name + "times (by train() and eval())"
2738            )
2739            plt.xlabel("Iteration")
2740            plt.ylabel("Epoch Time in Seconds ")
2741            ax.set_ylim(ymin=0)
2742            ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2743
2744            pd1 = pd.DataFrame(
2745                {
2746                    "Epochs": np.arange(len(self.member_vars["n_epoch_times"])),
2747                    "Normal Epoch Times": self.member_vars["n_epoch_times"],
2748                }
2749            )
2750            pd2 = pd.DataFrame(
2751                {
2752                    "Epochs": np.arange(len(self.member_vars["p_epoch_times"])),
2753                    "PAI Epoch Times": self.member_vars["p_epoch_times"],
2754                }
2755            )
2756            pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2757
2758            pd1.to_csv(
2759                save_folder + "/" + self.save_name + extra_string + "Times.csv",
2760                index=False,
2761            )
2762            del pd1, pd2
2763
2764        if self.values_per_train_epoch != -1 and self.values_per_val_epoch != -1:
2765            ax2 = ax.twinx()  # Second axes sharing same x-axis
2766            ax2.set_ylabel("Single Datapoint Time in Seconds")
2767
2768            ax2.plot(
2769                np.arange(len(self.member_vars["n_train_times"])),
2770                np.array(self.member_vars["n_train_times"])
2771                / self.values_per_train_epoch,
2772                linestyle="dashed",
2773                label="Normal Train Item Times",
2774            )
2775            ax2.plot(
2776                np.arange(len(self.member_vars["p_train_times"])),
2777                np.array(self.member_vars["p_train_times"])
2778                / self.values_per_train_epoch,
2779                linestyle="dashed",
2780                label="PAI Train Item Times",
2781            )
2782            ax2.plot(
2783                np.arange(len(self.member_vars["n_val_times"])),
2784                np.array(self.member_vars["n_val_times"]) / self.values_per_val_epoch,
2785                linestyle="dashed",
2786                label="Normal Val Item Times",
2787            )
2788            ax2.plot(
2789                np.arange(len(self.member_vars["p_val_times"])),
2790                np.array(self.member_vars["p_val_times"]) / self.values_per_val_epoch,
2791                linestyle="dashed",
2792                label="PAI Val Item Times",
2793            )
2794            ax2.tick_params(axis="y")
2795            ax2.set_ylim(ymin=0)
2796            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):
2798    def generate_learning_rate_plots(self, ax, save_folder, extra_string):
2799        """
2800        Generate plots and csvs for learning rate
2801
2802        Parameters
2803        ----------
2804        ax : object
2805            The matplotlib axis to plot on.
2806        save_folder : str
2807            The folder to save the plots and csvs in.
2808        extra_string : str
2809            An extra string to append to the filenames.
2810
2811        Returns
2812        -------
2813        None
2814
2815        """
2816        ax.plot(
2817            np.arange(len(self.member_vars["training_learning_rates"])),
2818            self.member_vars["training_learning_rates"],
2819            label="learning_rate",
2820        )
2821        plt.title(save_folder + "/" + self.save_name + "learning_rate")
2822        plt.xlabel("Epochs")
2823        plt.ylabel("learning_rate")
2824        ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left")
2825
2826        pd1 = pd.DataFrame(
2827            {
2828                "Epochs": np.arange(len(self.member_vars["training_learning_rates"])),
2829                "learning_rate": self.member_vars["training_learning_rates"],
2830            }
2831        )
2832        pd1.to_csv(
2833            save_folder + "/" + self.save_name + extra_string + "learning_rate.csv",
2834            index=False,
2835        )
2836        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):
2838    def get_current_pb_scores(self):
2839        """
2840        Get the latest best PBScore of each dendrite layer, the same numbers
2841        written to the Best PBScores csv.
2842
2843        Returns
2844        -------
2845        dict[str, Any]
2846            Layer name to score.  Empty outside of dendrite scoring phases,
2847            when no candidate dendrites are being scored.
2848
2849
2850        Parameters
2851        ----------
2852        None
2853
2854        """
2855        if not self.member_vars["doing_pai"]:
2856            return {}
2857        if not GPA.pc.get_perforated_backpropagation():
2858            return {}
2859        # Scores only advance while candidate dendrites are being trained
2860        if (
2861            self.member_vars["mode"] != "p"
2862            and not GPA.pc.get_learn_dendrites_live()
2863        ):
2864            return {}
2865
2866        scores = {}
2867        for layer_id in range(len(self.neuron_module_vector)):
2868            if layer_id >= len(self.member_vars["best_scores"]):
2869                continue
2870            layer_scores = self.member_vars["best_scores"][layer_id]
2871            if len(layer_scores) == 0:
2872                continue
2873            score = layer_scores[-1]
2874            if hasattr(score, "item"):
2875                score = score.item()
2876            score = float(score)
2877            if math.isnan(score) or math.isinf(score):
2878                continue
2879            scores[self.neuron_module_vector[layer_id].name] = score
2880        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):
2882    def generate_dendrite_learning_plots(self, ax, save_folder, extra_string):
2883        """
2884        Generate dendrite score plots for the tracker.
2885        Also saves csv files associated with the plots.
2886
2887        Parameters
2888        ----------
2889        ax : matplotlib.axes.Axes
2890            Axis used for plotting dendrite-learning curves.
2891        save_folder : str
2892            Directory where plot images and CSV summaries are written.
2893        extra_string : str
2894            Filename suffix used to distinguish this output set.
2895
2896        Returns
2897        -------
2898        None
2899            Saves plots and score CSV files to disk.
2900        """
2901        if self.member_vars["doing_pai"]:
2902            pd1 = None
2903            pd2 = None
2904            num_colors = len(self.neuron_module_vector)
2905
2906            if (
2907                len(self.neuron_module_vector) > 0
2908                and len(self.member_vars["current_scores"][0]) != 0
2909            ):
2910                num_colors *= 2
2911
2912            cm = plt.get_cmap("gist_rainbow")
2913            ax.set_prop_cycle(
2914                "color", [cm(1.0 * i / num_colors) for i in range(num_colors)]
2915            )
2916
2917            for layer_id in range(len(self.neuron_module_vector)):
2918                ax.plot(
2919                    np.arange(len(self.member_vars["best_scores"][layer_id])),
2920                    self.member_vars["best_scores"][layer_id],
2921                    label=self.neuron_module_vector[layer_id].name,
2922                )
2923
2924                pd2 = pd.DataFrame(
2925                    {
2926                        "Epochs": np.arange(
2927                            len(self.member_vars["best_scores"][layer_id])
2928                        ),
2929                        f"Best ever for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2930                            "best_scores"
2931                        ][
2932                            layer_id
2933                        ],
2934                    }
2935                )
2936
2937                if pd1 is None:
2938                    pd1 = pd2
2939                else:
2940                    pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2941
2942                if len(self.member_vars["current_scores"][layer_id]) != 0:
2943                    ax.plot(
2944                        np.arange(len(self.member_vars["current_scores"][layer_id])),
2945                        self.member_vars["current_scores"][layer_id],
2946                        label=f"Current:{self.neuron_module_vector[layer_id].name}",
2947                    )
2948
2949                pd2 = pd.DataFrame(
2950                    {
2951                        "Epochs": np.arange(
2952                            len(self.member_vars["current_scores"][layer_id])
2953                        ),
2954                        f"Best current for all nodes Layer {self.neuron_module_vector[layer_id].name}": self.member_vars[
2955                            "current_scores"
2956                        ][
2957                            layer_id
2958                        ],
2959                    }
2960                )
2961                pd1 = pd.concat([pd1, pd.DataFrame(pd2)], ignore_index=True)
2962
2963            plt.title(save_folder + "/" + self.save_name + " Best PBScores")
2964            plt.xlabel("Epochs")
2965            plt.ylabel("Best PBScore")
2966            ax.legend(
2967                bbox_to_anchor=(1.05, 1),
2968                loc="upper left",
2969                ncol=max(1, math.ceil(len(self.neuron_module_vector) / 30)),
2970            )
2971            for switcher in self.member_vars["p_switch_epochs"]:
2972                plt.axvline(x=switcher, ymin=0, ymax=1, color="r")
2973
2974            if self.member_vars["mode"] == "p":
2975                missed_time = (
2976                    self.member_vars["num_epochs_run"]
2977                    - self.member_vars["epoch_last_improved"]
2978                )
2979                plt.axvline(
2980                    x=(len(self.member_vars["best_scores"][0]) - (missed_time + 1)),
2981                    ymin=0,
2982                    ymax=1,
2983                    color="g",
2984                )
2985
2986            # pd1 here will be none if no PB layers are created
2987            if pd1 is not None:
2988                pd1.to_csv(
2989                    save_folder
2990                    + "/"
2991                    + self.save_name
2992                    + extra_string
2993                    + "Best PBScores.csv",
2994                    index=False,
2995                )
2996            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):
2998    def generate_extra_csv_files(self, save_folder, extra_string):
2999        """
3000        Generate additional csvs
3001
3002        Parameters
3003        ----------
3004        save_folder : str
3005            The folder to save the plots and csvs in.
3006        extra_string : str
3007            An extra string to append to the filenames.
3008
3009        Returns
3010        -------
3011        None
3012
3013        """
3014        pd1 = pd.DataFrame(
3015            {
3016                "Switch Number": np.arange(len(self.member_vars["switch_epochs"])),
3017                "Switch Epoch": self.member_vars["switch_epochs"],
3018            }
3019        )
3020        pd1.to_csv(
3021            save_folder + "/" + self.save_name + extra_string + "switch_epochs.csv",
3022            index=False,
3023        )
3024        del pd1
3025
3026        pd1 = pd.DataFrame(
3027            {
3028                "Switch Number": np.arange(len(self.member_vars["param_counts"])),
3029                "Param Count": self.member_vars["param_counts"],
3030            }
3031        )
3032        pd1.to_csv(
3033            save_folder + "/" + self.save_name + extra_string + "param_counts.csv",
3034            index=False,
3035        )
3036        del pd1
3037
3038        """
3039        Create best_arch_scores.csv file
3040        When working with dendrites there is a tradeoff between additional param count and score improvement.
3041        This file will help track that tradeoff by recording the best scores for all extra_scores
3042        and extra_scores_without_graphing for each architecture version.
3043        The scores recorded here are from the epoch when the best validation score was found
3044        within each switch_epoch boundary.
3045        """
3046        switch_counts = len(self.member_vars["switch_epochs"])
3047        best_valid = []
3048        associated_params = []
3049        
3050        # Initialize dictionaries to store best scores for each extra score type
3051        best_extra_scores = {}
3052        for score_name in self.member_vars["extra_scores"]:
3053            best_extra_scores[score_name] = []
3054        for score_name in self.member_vars["extra_scores_without_graphing"]:
3055            best_extra_scores[score_name] = []
3056
3057        for switch in range(0, switch_counts, 2):
3058            start_index = 0
3059            if switch != 0:
3060                start_index = self.member_vars["switch_epochs"][switch - 1] + 1
3061            end_index = self.member_vars["switch_epochs"][switch] + 1
3062
3063            if GPA.pai_tracker.member_vars["maximizing_score"]:
3064                best_valid_index = start_index + np.argmax(
3065                    self.member_vars["accuracies"][start_index:end_index]
3066                )
3067            else:
3068                best_valid_index = start_index + np.argmin(
3069                    self.member_vars["accuracies"][start_index:end_index]
3070                )
3071
3072            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3073            best_valid.append(best_valid_score)
3074            
3075            # Get corresponding scores from all extra_scores
3076            for score_name in self.member_vars["extra_scores"]:
3077                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3078                    best_extra_scores[score_name].append(
3079                        self.member_vars["extra_scores"][score_name][best_valid_index]
3080                    )
3081                else:
3082                    best_extra_scores[score_name].append(None)
3083            
3084            # Get corresponding scores from all extra_scores_without_graphing
3085            for score_name in self.member_vars["extra_scores_without_graphing"]:
3086                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3087                    best_extra_scores[score_name].append(
3088                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3089                    )
3090                else:
3091                    best_extra_scores[score_name].append(None)
3092            
3093            if self.member_vars["doing_pai"]:
3094                associated_params.append(self.member_vars["param_counts"][switch])
3095            else:
3096                associated_params.append(self.member_vars["param_counts"][-1])
3097
3098        # If in neuron training mode but not the very first epoch
3099        if self.member_vars["mode"] == "n" and (
3100            (len(self.member_vars["switch_epochs"]) == 0)
3101            or (
3102                self.member_vars["switch_epochs"][-1] + 1
3103                != len(self.member_vars["accuracies"])
3104            )
3105        ):
3106            start_index = 0
3107            if len(self.member_vars["switch_epochs"]) != 0:
3108                start_index = self.member_vars["switch_epochs"][-1] + 1
3109
3110            if GPA.pai_tracker.member_vars["maximizing_score"]:
3111                best_valid_index = start_index + np.argmax(
3112                    self.member_vars["accuracies"][start_index:]
3113                )
3114            else:
3115                best_valid_index = start_index + np.argmin(
3116                    self.member_vars["accuracies"][start_index:]
3117                )
3118
3119            best_valid_score = self.member_vars["accuracies"][best_valid_index]
3120            best_valid.append(best_valid_score)
3121            
3122            # Get corresponding scores from all extra_scores
3123            for score_name in self.member_vars["extra_scores"]:
3124                if best_valid_index < len(self.member_vars["extra_scores"][score_name]):
3125                    best_extra_scores[score_name].append(
3126                        self.member_vars["extra_scores"][score_name][best_valid_index]
3127                    )
3128                else:
3129                    best_extra_scores[score_name].append(None)
3130            
3131            # Get corresponding scores from all extra_scores_without_graphing
3132            for score_name in self.member_vars["extra_scores_without_graphing"]:
3133                if best_valid_index < len(self.member_vars["extra_scores_without_graphing"][score_name]):
3134                    best_extra_scores[score_name].append(
3135                        self.member_vars["extra_scores_without_graphing"][score_name][best_valid_index]
3136                    )
3137                else:
3138                    best_extra_scores[score_name].append(None)
3139            
3140            associated_params.append(self.member_vars["param_counts"][-1])
3141
3142        # Build dataframe with all columns
3143        csv_data = {
3144            "Param Counts": associated_params,
3145            "Max Valid Scores": best_valid,
3146        }
3147        
3148        # Add columns for each extra score
3149        for score_name in best_extra_scores:
3150            csv_data[score_name] = best_extra_scores[score_name]
3151        
3152        pd1 = pd.DataFrame(csv_data)
3153        pd1.to_csv(
3154            save_folder + "/" + self.save_name + extra_string + "_best_arch_scores.csv",
3155            index=False,
3156        )
3157        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=''):
3159    def save_graphs(self, extra_string=""):
3160        """
3161        Save graphs and csvs for all the values the tracker records
3162
3163        Parameters
3164        ----------
3165        extra_string : str
3166            An extra string to append to the filenames.
3167
3168        Returns
3169        -------
3170        None
3171
3172        """
3173        # If running DDP only save with rank 0
3174        if "RANK" in os.environ:
3175            if int(os.environ["RANK"]) != 0:
3176                return
3177        if not self.making_graphs:
3178            return
3179
3180        save_folder = "./" + self.save_name + "/"
3181
3182        plt.ioff()
3183        fig = plt.figure(figsize=(28, 14))
3184
3185        # Plot with accuracy scores
3186        ax = plt.subplot(221)
3187        self.generate_accuracy_plots(ax, save_folder, extra_string)
3188
3189        # Plot dendrite learning scores
3190        ax = plt.subplot(222)
3191        self.generate_dendrite_learning_plots(ax, save_folder, extra_string)
3192
3193        if GPA.pc.get_drawing_extra_graphs():
3194            # Plot learning rates for each training epoch
3195            ax = plt.subplot(223)
3196            self.generate_learning_rate_plots(ax, save_folder, extra_string)
3197
3198            # Plot the times for each training epoch
3199            ax = plt.subplot(224)
3200            self.generate_time_plots(ax, save_folder, extra_string)
3201
3202        # Generate extra CSV files
3203        self.generate_extra_csv_files(save_folder, extra_string)
3204
3205        fig.tight_layout()
3206        plt.savefig(save_folder + "/" + self.save_name + extra_string + ".png")
3207        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):
3209    def add_loss(self, loss):
3210        """Add loss to tracking vectors.
3211
3212        Parameters
3213        ----------
3214        loss : float or int
3215            The loss value to add.
3216
3217        Returns
3218        -------
3219        None
3220
3221        """
3222        if not isinstance(loss, (float, int)):
3223            loss = loss.item()
3224        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):
3226    def add_learning_rate(self, learning_rate):
3227        """Add learning rate to tracking vectors.
3228
3229        Parameters
3230        ----------
3231        learning_rate : float or int
3232            The learning rate value to add.
3233
3234        Returns
3235        -------
3236        None
3237
3238        """
3239        if not isinstance(learning_rate, (float, int)):
3240            learning_rate = learning_rate.item()
3241        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):
3243    def add_extra_score(self, score, extra_score_name):
3244        """Add extra score to tracking vectors.
3245
3246        Parameters
3247        ----------
3248        score : float or int
3249            The score value to add.
3250
3251        extra_score_name : str
3252            The name of the extra score.
3253
3254        Returns
3255        -------
3256        None
3257
3258        """
3259        if not isinstance(score, (float, int)):
3260            try:
3261                score = score.item()
3262            except:
3263                print(
3264                    "Scores added for Perforated Backpropagation should be "
3265                    "float, int, or tensor, yours is a:"
3266                )
3267                print(type(score))
3268                pdb.set_trace()
3269
3270        if GPA.pc.get_verbose():
3271            print(f"Adding extra score {extra_score_name} of {float(score)}")
3272
3273        if extra_score_name not in self.member_vars["extra_scores"]:
3274            self.member_vars["extra_scores"][extra_score_name] = []
3275        self.member_vars["extra_scores"][extra_score_name].append(score)
3276
3277        if self.member_vars["mode"] == "n":
3278            if extra_score_name not in self.member_vars["n_extra_scores"]:
3279                self.member_vars["n_extra_scores"][extra_score_name] = []
3280            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):
3282    def add_extra_score_without_graphing(self, score, extra_score_name):
3283        """Add extra score without graphing to tracking vectors.
3284
3285        Parameters
3286        ----------
3287        score : float or int
3288            The score value to add.
3289
3290        extra_score_name : str
3291            The name of the extra score.
3292
3293        Returns
3294        -------
3295        None
3296
3297        """
3298        if not isinstance(score, (float, int)):
3299            try:
3300                score = score.item()
3301            except:
3302                print(
3303                    "Scores added for Perforated Backpropagation should be "
3304                    "float, int, or tensor, yours is a:"
3305                )
3306                print(type(score))
3307                print("in add_extra_score_without_graphing")
3308                pdb.set_trace()
3309
3310        if GPA.pc.get_verbose():
3311            print(f"Adding extra score {extra_score_name} of {float(score)}")
3312
3313        if extra_score_name not in self.member_vars["extra_scores_without_graphing"]:
3314            self.member_vars["extra_scores_without_graphing"][extra_score_name] = []
3315        self.member_vars["extra_scores_without_graphing"][extra_score_name].append(
3316            score
3317        )

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

Add dendrite module to all neuron modules.

Parameters
  • None
Returns
  • None: This function does not return a value.
def apply_pb_grads(self):
3673    def apply_pb_grads(self):
3674        """Apply perforated backpropagation gradients to all modules.
3675
3676        Parameters
3677        ----------
3678        None
3679
3680        Returns
3681        -------
3682        None
3683            This function does not return a value.
3684        """
3685        if self.member_vars["mode"] == "p":
3686            for module in self.neuron_module_vector:
3687                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):
3689    def apply_pb_zero(self):
3690        """Apply perforated backpropagation zero gradients to all modules.
3691
3692        Parameters
3693        ----------
3694        None
3695
3696        Returns
3697        -------
3698        None
3699            This function does not return a value.
3700        """
3701        if self.member_vars["mode"] == "p":
3702            for module in self.neuron_module_vector:
3703                module.apply_pb_zero()

Apply perforated backpropagation zero gradients to all modules.

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