You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch中生成用于优化任务的参数化SU(n)酉矩阵并确保其酉性

Generating & Optimizing SU(n) Matrices in PyTorch

Great question! Working with unitary matrices—especially SU(n) matrices (special unitary, determinant = 1, parameterized by n²-1 variables)—in PyTorch for optimization tasks can feel tricky at first, but there are two solid approaches that work reliably. Let’s break them down:

1. Lie Algebra Exponential Mapping (Best for Exact Parameterization)

SU(n) matrices are the exponential of traceless anti-Hermitian matrices (the Lie algebra of SU(n)). This is the most direct way to get exactly n²-1 free parameters, since the Lie algebra has dimension n²-1. Here’s how to implement it:

How it works:

  • First, create a traceless anti-Hermitian matrix using n²-1 real parameters. An anti-Hermitian matrix satisfies ( A^\dagger = -A ) (conjugate transpose equals negative itself), and being traceless means ( \text{tr}(A) = 0 ).
  • Take the matrix exponential of this matrix: ( U = \exp(A) ). By definition, this will be a unitary matrix, and since ( \det(\exp(A)) = \exp(\text{tr}(A)) = 1 ), it’s guaranteed to be in SU(n).

PyTorch Code Example:

import torch

def create_su_n_params_and_matrix(n: int):
    # Generate exactly n²-1 trainable parameters
    num_params = n * n - 1
    params = torch.randn(num_params, requires_grad=True, dtype=torch.float32)
    
    # Build traceless anti-Hermitian matrix
    A = torch.zeros(n, n, dtype=torch.complex64)
    param_idx = 0
    
    # Fill off-diagonal elements (anti-Hermitian constraint)
    for i in range(n):
        for j in range(n):
            if i == j:
                continue
            elif i < j:
                # Use two real params for complex value (real + imag*i)
                real_part = params[param_idx]
                imag_part = params[param_idx + 1]
                A[i, j] = real_part + 1j * imag_part
                # Anti-Hermitian: conjugate transpose = -A
                A[j, i] = -real_part + 1j * imag_part
                param_idx += 2
    
    # Adjust diagonal to ensure trace is 0
    diag_sum = torch.trace(A).real
    for i in range(n):
        A[i, i] -= diag_sum / n
    
    # Exponentiate to get SU(n) matrix
    U = torch.matrix_exp(A)
    return params, U

# Test with n=3
params, U = create_su_n_params_and_matrix(3)

# Verify properties
print("Is unitary?", torch.allclose(U @ U.conj().T, torch.eye(3, dtype=torch.complex64)))
print("Determinant is 1?", torch.allclose(torch.det(U), torch.tensor(1.0 + 0.0j)))

When optimizing, you’ll update the params tensor directly—every time you re-compute ( U = \exp(A) ), it will stay in SU(n) automatically. No extra projection steps needed!

2. QR Decomposition Projection (Best for Post-Optimization Correction)

If you’re optimizing a matrix that might drift away from unitarity (e.g., starting with a random matrix and applying gradient updates), you can project it back to SU(n) using QR decomposition. This is simpler to implement but doesn’t enforce the n²-1 parameter count upfront—use it when you need a quick way to "fix" a matrix during training.

How it works:

  • Take any complex matrix, run QR decomposition to get a unitary matrix ( Q ) (from U(n)).
  • Adjust ( Q ) by a phase factor to make its determinant 1, converting it to SU(n).

PyTorch Code Example:

def project_to_su_n(mat: torch.Tensor) -> torch.Tensor:
    # Project to U(n) via QR decomposition
    Q, _ = torch.linalg.qr(mat)
    # Compute phase factor to make determinant = 1
    det_Q = torch.det(Q)
    phase = det_Q ** (-1 / mat.shape[0])
    # Apply phase to get SU(n) matrix
    Q_su = Q * phase
    return Q_su

# Example usage during optimization
n = 3
# Start with a random complex matrix (not unitary)
mat = torch.randn(n, n, dtype=torch.complex64, requires_grad=True)

# After some gradient updates...
optimizer = torch.optim.SGD([mat], lr=0.01)
optimizer.zero_grad()
# (your loss calculation here)
# loss.backward()
# optimizer.step()

# Project back to SU(n)
U_su = project_to_su_n(mat)

# Verify properties
print("Is unitary after projection?", torch.allclose(U_su @ U_su.conj().T, torch.eye(n, dtype=torch.complex64)))
print("Determinant is 1 after projection?", torch.allclose(torch.det(U_su), torch.tensor(1.0 + 0.0j)))

Key Notes

  • Use Lie Algebra for Exact Parameterization: If you need to strictly work with n²-1 parameters (no extra degrees of freedom), the exponential mapping is your go-to. It’s clean and ensures SU(n) compliance at every step.
  • Use QR Projection for Flexibility: If your optimization setup starts with a full matrix and you just need to enforce unitarity/SU(n) after updates, QR projection is quick and easy to integrate.

内容的提问来源于stack exchange,提问作者youryanmi

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.27 16:33:13