summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py35
1 files changed, 29 insertions, 6 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py
index 3c0de00..162a8be 100644
--- a/experiments/conv_run.py
+++ b/experiments/conv_run.py
@@ -12,9 +12,11 @@ import time
import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
-from sdil.conv import (CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig,
+from sdil.conv import (CIFARHierarchicalFAResNet, CIFARLocalResNet,
+ CIFARSDILResNet, ConvSDILConfig,
conv_alignment_report, conv_apical_calibration_step,
- conv_local_step, evaluate_conv)
+ conv_hierarchical_alignment_report,
+ conv_hierarchical_step, conv_local_step, evaluate_conv)
from sdil.data import DATA_DIR, get_cifar_image_splits
@@ -101,6 +103,16 @@ def build(args):
bn_momentum=args.bn_momentum, bn_eps=args.bn_eps)
if args.mode == "bp":
return CIFARLocalResNet(**common), None
+ if args.mode == "hfa":
+ net = CIFARHierarchicalFAResNet(
+ **common, feedback_seed=args.apical_seed,
+ feedback_scale=args.a_scale)
+ config = ConvSDILConfig(
+ eta=args.lr, eta_output=args.output_lr, eta_A=0.0,
+ momentum=args.momentum, weight_decay=args.weight_decay,
+ learn_A=False, learn_P=False)
+ config.validate()
+ return net, config
net = CIFARSDILResNet(
**common, a_scale=args.a_scale, apical_seed=args.apical_seed,
vectorizer_mode=args.vectorizer_mode)
@@ -213,7 +225,7 @@ def run(args):
"schema_version": 1,
"protocol_family": "oral_a_cifar_local_resnet_development",
"calibration_metric_space": (
- None if config is None else {
+ None if config is None or args.mode == "hfa" else {
"unit_targets": "full_hidden_field",
"channel_subspace": "channel_basis_moments",
"vectorizer_subspace": "vectorizer_parameter_gradients",
@@ -238,6 +250,8 @@ def run(args):
"vectorizer_mode": getattr(net, "vectorizer_mode", None),
"fixed_traffic_coefficients": getattr(
net, "n_fixed_traffic_coefficients", 0),
+ "fixed_feedback_parameters": getattr(
+ net, "n_fixed_feedback_parameters", 0),
},
"epochs": [],
}
@@ -328,6 +342,10 @@ def run(args):
x, y, lr, momentum=args.momentum,
weight_decay=args.weight_decay)
did_perturb = False
+ elif args.mode == "hfa":
+ result = conv_hierarchical_step(net, x, y, config)
+ loss = result["loss"]
+ did_perturb = False
else:
result = conv_local_step(
net, x, y, config, step, generator=perturb_generator)
@@ -404,8 +422,11 @@ def run(args):
probe = min(args.alignment_probe, train.x.shape[0])
sync(args.device)
diagnostic_start = time.time()
- diagnostics = conv_alignment_report(
- net, train.x[:probe], train.y[:probe], config)
+ diagnostics = (conv_hierarchical_alignment_report(
+ net, train.x[:probe], train.y[:probe])
+ if args.mode == "hfa" else
+ conv_alignment_report(
+ net, train.x[:probe], train.y[:probe], config))
sync(args.device)
diagnostics["wall_s"] = time.time() - diagnostic_start
values = diagnostics["teaching_negative_gradient_cosine"]
@@ -449,7 +470,9 @@ def run(args):
def parse_args():
parser = argparse.ArgumentParser()
- parser.add_argument("--mode", choices=("bp", "dfa", "sdil", "nodepert"), required=True)
+ parser.add_argument(
+ "--mode", choices=("bp", "dfa", "hfa", "sdil", "nodepert"),
+ required=True)
parser.add_argument("--out", required=True)
parser.add_argument("--device", default="cpu")
parser.add_argument("--data_dir", default=DATA_DIR)