如何在PyTorch中生成用于优化任务的参数化SU(n)酉矩阵并确保其酉性
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

