Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 53 additions & 0 deletions tests/test_adahessian_complex.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
"""Complex Adahessian estimates curvature on both components."""
import torch
import torch_optimizer as optim


def test_real_path_is_unchanged():
torch.manual_seed(1)
param = torch.nn.Parameter(torch.tensor([0.5, -0.2, 0.3]))
opt = optim.Adahessian([param], lr=0.05)
for _ in range(2):
opt.zero_grad(set_to_none=True)
loss = (param ** 2).sum()
loss.backward(create_graph=True)
opt.step()
expected = torch.tensor(
[0.43060994148254395, -0.17224396765232086, 0.2583659589290619]
)
assert torch.allclose(param.detach(), expected)


def test_complex_matches_the_real_algorithm():
start = torch.tensor([0.5 + 0.1j, -0.2 + 0.3j])
torch.manual_seed(2)
param = torch.nn.Parameter(start.clone())
opt = optim.Adahessian([param], lr=0.05)
opt.zero_grad(set_to_none=True)
loss = (param.real ** 2 + param.imag ** 2).sum()
loss.backward(create_graph=True)
opt.step()

torch.manual_seed(2)
real = torch.nn.Parameter(torch.view_as_real(start).clone())
real_opt = optim.Adahessian([real], lr=0.05)
real_opt.zero_grad(set_to_none=True)
real_loss = (real ** 2).sum()
real_loss.backward(create_graph=True)
real_opt.step()
assert torch.equal(torch.view_as_real(param.detach()), real.detach())


def test_complex_matrix_steps():
torch.manual_seed(3)
param = torch.nn.Parameter(
torch.tensor([[0.2 + 0.1j, -0.3 + 0.0j], [0.4 - 0.2j, 0.1 + 0.5j]])
)
before = param.detach().clone()
opt = optim.Adahessian([param], lr=0.05)
opt.zero_grad(set_to_none=True)
loss = (param.real ** 2 + param.imag ** 2).sum()
loss.backward(create_graph=True)
opt.step()
assert param.dtype == torch.complex64
assert not torch.equal(param.detach(), before)
73 changes: 56 additions & 17 deletions torch_optimizer/adahessian.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,14 +97,24 @@ def get_trace(self, params: Params, grads: Grads) -> List[torch.Tensor]:
)
raise RuntimeError(msg.format(i))

v = [
2
* torch.randint_like(
p, high=2, memory_format=torch.preserve_format
)
- 1
for p in params
]
v = []
for p in params:
if torch.is_complex(p):
# randint_like rejects complex. Signs on both components,
# packed back so the vector matches the complex gradient.
real = torch.view_as_real(p)
signs = torch.randint_like(real, high=2)
v.append(
torch.view_as_complex((2 * signs - 1).to(dtype=real.dtype))
)
else:
v.append(
2
* torch.randint_like(
p, high=2, memory_format=torch.preserve_format
)
- 1
)

# this is for distributed setting with single node and multi-gpus,
# for multi nodes setting, we have not support it yet.
Expand All @@ -114,6 +124,23 @@ def get_trace(self, params: Params, grads: Grads) -> List[torch.Tensor]:

hutchinson_trace = []
for hv in hvs:
if torch.is_complex(hv):
# The last axis is the two components, not a spatial axis.
# abs removes the ±1 sign. A conv kernel still averages its
# spatial dims, matching the real 4D case.
hv_r = torch.view_as_real(hv).abs()
rank = hv_r.ndim - 1
if rank <= 2:
tmp_output = hv_r
elif rank == 4:
tmp_output = hv_r.mean(dim=[2, 3], keepdim=True)
else:
raise RuntimeError(
"Adahessian does not support complex tensors of rank "
f"{rank}"
)
hutchinson_trace.append(tmp_output)
continue
param_size = hv.size()
if len(param_size) <= 2: # for 0/1/2D tensor
# Hessian diagonal block size is 1 here.
Expand Down Expand Up @@ -165,10 +192,11 @@ def step(self, closure: OptLossClosure = None) -> OptFloat:
# State initialization
if len(state) == 0:
state["step"] = 0
# Exponential moving average of gradient values
state["exp_avg"] = torch.zeros_like(p.data)
# Exponential moving average of Hessian diagonal square values
state["exp_hessian_diag_sq"] = torch.zeros_like(p.data)
# Complex moments follow the real and imaginary parts.
# Real parameters keep a state tensor shaped like p.
moment = torch.view_as_real(p.data) if torch.is_complex(p) else p.data
state["exp_avg"] = torch.zeros_like(moment)
state["exp_hessian_diag_sq"] = torch.zeros_like(moment)

exp_avg, exp_hessian_diag_sq = (
state["exp_avg"],
Expand All @@ -180,7 +208,11 @@ def step(self, closure: OptLossClosure = None) -> OptFloat:
state["step"] += 1

# Decay the first and second moment running average coefficient
exp_avg.mul_(beta1).add_(grad.detach_(), alpha=1 - beta1)
grad.detach_()
grad_for_avg = (
torch.view_as_real(grad) if torch.is_complex(p) else grad
)
exp_avg.mul_(beta1).add_(grad_for_avg, alpha=1 - beta1)
exp_hessian_diag_sq.mul_(beta2).addcmul_(
hut_trace, hut_trace, value=1 - beta2
)
Expand All @@ -196,9 +228,16 @@ def step(self, closure: OptLossClosure = None) -> OptFloat:
).add_(group["eps"])

# make update
p.data = p.data - group["lr"] * (
exp_avg / bias_correction1 / denom
+ group["weight_decay"] * p.data
)
step_dir = exp_avg / bias_correction1 / denom
if torch.is_complex(p):
real_p = torch.view_as_real(p.data)
real_p.sub_(
group["lr"]
* (step_dir + group["weight_decay"] * real_p)
)
else:
p.data = p.data - group["lr"] * (
step_dir + group["weight_decay"] * p.data
)

return loss