Skip to content

Commit c87132d

Browse files
committed
change imaml/neumann to higher implementation
1 parent 693dc3d commit c87132d

2 files changed

Lines changed: 64 additions & 71 deletions

File tree

gbml/imaml.py

Lines changed: 32 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -13,33 +13,28 @@ class iMAML(GBML):
1313
def __init__(self, args):
1414
super().__init__(args)
1515
self._init_net()
16-
self.network_aux = type(self.network)(args).cuda()
17-
self.network_aux.load_state_dict(self.network.state_dict())
1816
self._init_opt()
19-
self.inner_optimizer = torch.optim.SGD(self.network_aux.parameters(), lr=self.args.inner_lr)
2017
self.lamb = 2.0
2118
self.n_cg = 3
2219
return None
2320

2421
@torch.enable_grad()
25-
def inner_loop(self, train_input, train_target):
26-
self.network_aux.zero_grad()
27-
train_logit = self.network_aux(train_input)
22+
def inner_loop(self, fmodel, diffopt, train_input, train_target):
23+
24+
train_logit = fmodel(train_input)
2825
inner_loss = F.cross_entropy(train_logit, train_target)
29-
inner_loss += (self.lamb/2.) * ((torch.nn.utils.parameters_to_vector(self.network_aux.parameters())-torch.nn.utils.parameters_to_vector(self.network.parameters()).detach())**2).sum()
30-
inner_loss.backward()
31-
self.inner_optimizer.step()
26+
inner_loss += (self.lamb/2.) * ((torch.nn.utils.parameters_to_vector(self.network.parameters())-torch.nn.utils.parameters_to_vector(self.network.parameters()).detach())**2).sum()
27+
diffopt.step(inner_loss)
28+
3229
return None
3330

34-
@torch.enable_grad()
35-
def cg(self, in_grad, outer_grad):
36-
in_grad = torch.nn.utils.parameters_to_vector(in_grad)
37-
outer_grad = torch.nn.utils.parameters_to_vector(outer_grad)
31+
@torch.no_grad()
32+
def cg(self, in_grad, outer_grad, params):
3833
x = outer_grad.clone().detach()
39-
r = outer_grad.clone().detach() - self.hv_prod(in_grad, x)
34+
r = outer_grad.clone().detach() - self.hv_prod(in_grad, x, params)
4035
p = r.clone().detach()
4136
for i in range(self.n_cg):
42-
Ap = self.hv_prod(in_grad, p)
37+
Ap = self.hv_prod(in_grad, p, params)
4338
alpha = (r @ r)/(p @ Ap)
4439
x = x + alpha * p
4540
r_new = r - alpha * Ap
@@ -59,9 +54,9 @@ def vec_to_grad(self, vec):
5954
pointer += num_param
6055
return res
6156

62-
def hv_prod(self, in_grad, x):
63-
scalar = in_grad @ x.detach()
64-
hv = torch.autograd.grad(scalar, self.network_aux.parameters(), retain_graph=True)
57+
@torch.enable_grad()
58+
def hv_prod(self, in_grad, x, params):
59+
hv = torch.autograd.grad(in_grad, params, retain_graph=True, grad_outputs=x)
6560
hv = torch.nn.utils.parameters_to_vector(hv).detach()
6661
# precondition with identity matrix
6762
return hv/self.lamb + x
@@ -77,27 +72,28 @@ def outer_loop(self, batch, is_train):
7772

7873
for (train_input, train_target, test_input, test_target) in zip(train_inputs, train_targets, test_inputs, test_targets):
7974

80-
self.network_aux.load_state_dict(self.network.state_dict())
75+
with higher.innerloop_ctx(self.network, self.inner_optimizer, track_higher_grads=False) as (fmodel, diffopt):
8176

82-
for step in range(self.args.n_inner):
83-
self.inner_loop(train_input, train_target)
84-
85-
train_logit = self.network_aux(train_input)
86-
in_loss = F.cross_entropy(train_logit, train_target)
77+
for step in range(self.args.n_inner):
78+
self.inner_loop(fmodel, diffopt, train_input, train_target)
79+
80+
train_logit = fmodel(train_input)
81+
in_loss = F.cross_entropy(train_logit, train_target)
8782

88-
test_logit = self.network_aux(test_input)
89-
outer_loss = F.cross_entropy(test_logit, test_target)
90-
loss_log += outer_loss.item()/self.batch_size
83+
test_logit = fmodel(test_input)
84+
outer_loss = F.cross_entropy(test_logit, test_target)
85+
loss_log += outer_loss.item()/self.batch_size
9186

92-
with torch.no_grad():
93-
acc_log += get_accuracy(test_logit, test_target).item()/self.batch_size
94-
95-
if is_train:
96-
in_grad = torch.autograd.grad(in_loss, self.network_aux.parameters(), create_graph=True)
97-
outer_grad = torch.autograd.grad(outer_loss, self.network_aux.parameters())
98-
implicit_grad = self.cg(in_grad, outer_grad)
99-
grad_list.append(implicit_grad)
100-
loss_list.append(outer_loss.item())
87+
with torch.no_grad():
88+
acc_log += get_accuracy(test_logit, test_target).item()/self.batch_size
89+
90+
if is_train:
91+
params = list(fmodel.parameters(time=-1))
92+
in_grad = torch.nn.utils.parameters_to_vector(torch.autograd.grad(in_loss, params, create_graph=True))
93+
outer_grad = torch.nn.utils.parameters_to_vector(torch.autograd.grad(outer_loss, params))
94+
implicit_grad = self.cg(in_grad, outer_grad, params)
95+
grad_list.append(implicit_grad)
96+
loss_list.append(outer_loss.item())
10197

10298
if is_train:
10399
self.outer_optimizer.zero_grad()

gbml/neumann.py

Lines changed: 32 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -13,29 +13,25 @@ class Neumann(GBML):
1313
def __init__(self, args):
1414
super().__init__(args)
1515
self._init_net()
16-
self.network_aux = type(self.network)(args).cuda()
17-
self.network_aux.load_state_dict(self.network.state_dict())
1816
self._init_opt()
19-
self.inner_optimizer = torch.optim.SGD(self.network_aux.parameters(), lr=self.args.inner_lr)
20-
self.n_series = 5
17+
self.n_series = 3
2118
return None
2219

2320
@torch.enable_grad()
24-
def inner_loop(self, train_input, train_target):
25-
self.network_aux.zero_grad()
26-
train_logit = self.network_aux(train_input)
21+
def inner_loop(self, fmodel, diffopt, train_input, train_target):
22+
23+
train_logit = fmodel(train_input)
2724
inner_loss = F.cross_entropy(train_logit, train_target)
28-
inner_loss.backward()
29-
self.inner_optimizer.step()
25+
diffopt.step(inner_loss)
26+
3027
return None
3128

32-
@torch.enable_grad()
33-
def neumann_approx(self, in_grad, outer_grad):
34-
in_grad = torch.nn.utils.parameters_to_vector(in_grad)
35-
outer_grad = torch.nn.utils.parameters_to_vector(outer_grad)
29+
@torch.no_grad()
30+
def neumann_approx(self, in_grad, outer_grad, params):
31+
3632
x = outer_grad.clone().detach()
3733
for i in range(self.n_series):
38-
outer_grad = self.hv_prod(in_grad, outer_grad)
34+
outer_grad = self.hv_prod(in_grad, outer_grad, params)
3935
x = x + outer_grad
4036
return self.vec_to_grad(x)
4137

@@ -48,9 +44,9 @@ def vec_to_grad(self, vec):
4844
pointer += num_param
4945
return res
5046

51-
def hv_prod(self, in_grad, x):
52-
scalar = in_grad @ x.detach()
53-
hv = torch.autograd.grad(scalar, self.network_aux.parameters(), retain_graph=True)
47+
@torch.enable_grad()
48+
def hv_prod(self, in_grad, x, params):
49+
hv = torch.autograd.grad(in_grad, params, retain_graph=True, grad_outputs=x)
5450
hv = torch.nn.utils.parameters_to_vector(hv)
5551
hv = (-1.*self.args.inner_lr) * hv # scale for regularization
5652
return hv.detach()
@@ -66,27 +62,28 @@ def outer_loop(self, batch, is_train):
6662

6763
for (train_input, train_target, test_input, test_target) in zip(train_inputs, train_targets, test_inputs, test_targets):
6864

69-
self.network_aux.load_state_dict(self.network.state_dict())
65+
with higher.innerloop_ctx(self.network, self.inner_optimizer, track_higher_grads=False) as (fmodel, diffopt):
7066

71-
for step in range(self.args.n_inner):
72-
self.inner_loop(train_input, train_target)
73-
74-
train_logit = self.network_aux(train_input)
75-
in_loss = F.cross_entropy(train_logit, train_target)
67+
for step in range(self.args.n_inner):
68+
self.inner_loop(fmodel, diffopt, train_input, train_target)
69+
70+
train_logit = fmodel(train_input)
71+
in_loss = F.cross_entropy(train_logit, train_target)
7672

77-
test_logit = self.network_aux(test_input)
78-
outer_loss = F.cross_entropy(test_logit, test_target)
79-
loss_log += outer_loss.item()/self.batch_size
73+
test_logit = fmodel(test_input)
74+
outer_loss = F.cross_entropy(test_logit, test_target)
75+
loss_log += outer_loss.item()/self.batch_size
8076

81-
with torch.no_grad():
82-
acc_log += get_accuracy(test_logit, test_target).item()/self.batch_size
83-
84-
if is_train:
85-
in_grad = torch.autograd.grad(in_loss, self.network_aux.parameters(), create_graph=True)
86-
outer_grad = torch.autograd.grad(outer_loss, self.network_aux.parameters())
87-
implicit_grad = self.neumann_approx(in_grad, outer_grad)
88-
grad_list.append(implicit_grad)
89-
loss_list.append(outer_loss.item())
77+
with torch.no_grad():
78+
acc_log += get_accuracy(test_logit, test_target).item()/self.batch_size
79+
80+
if is_train:
81+
params = list(fmodel.parameters(time=-1))
82+
in_grad = torch.nn.utils.parameters_to_vector(torch.autograd.grad(in_loss, params, create_graph=True))
83+
outer_grad = torch.nn.utils.parameters_to_vector(torch.autograd.grad(outer_loss, params))
84+
implicit_grad = self.neumann_approx(in_grad, outer_grad, params)
85+
grad_list.append(implicit_grad)
86+
loss_list.append(outer_loss.item())
9087

9188
if is_train:
9289
self.outer_optimizer.zero_grad()

0 commit comments

Comments
 (0)