From f79dbcded5a5638332d316ec8489eb23f534bcbd Mon Sep 17 00:00:00 2001 From: shanyu Date: Sat, 26 Sep 2026 00:21:39 +0800 Subject: [PATCH] Estimate the Adahessian trace on both complex components randint_like rejects a complex parameter, so Adahessian raised RuntimeError before the Hutchinson vector existed (issue 458). Signs are drawn on the real and imaginary parts and packed back into a complex vector. The trace and the moments stay on those two components, and a real parameter still follows the old update. --- tests/test_adahessian_complex.py | 53 +++++++++++++++++++++++ torch_optimizer/adahessian.py | 73 ++++++++++++++++++++++++-------- 2 files changed, 109 insertions(+), 17 deletions(-) create mode 100644 tests/test_adahessian_complex.py diff --git a/tests/test_adahessian_complex.py b/tests/test_adahessian_complex.py new file mode 100644 index 00000000..734ca36f --- /dev/null +++ b/tests/test_adahessian_complex.py @@ -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) diff --git a/torch_optimizer/adahessian.py b/torch_optimizer/adahessian.py index 6c836481..1c76cf64 100644 --- a/torch_optimizer/adahessian.py +++ b/torch_optimizer/adahessian.py @@ -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. @@ -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. @@ -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"], @@ -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 ) @@ -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