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
14 changes: 10 additions & 4 deletions src/src_method/_sweep.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
from itertools import count
from math import prod
from time import perf_counter_ns
from typing import TYPE_CHECKING, NamedTuple
from typing import TYPE_CHECKING, Any, NamedTuple

from opt_einsum import contract_expression, get_symbol

Expand Down Expand Up @@ -101,7 +101,8 @@ def sweep(
*,
cutoff: float = 0.0,
dtype: DTypeLike,
) -> list[NDArray]:
to_host: bool = True,
) -> list[Any]:
"""Contract and compress a stack in ket form with one SRC sweep.

Args:
Expand All @@ -114,10 +115,12 @@ def sweep(
xp: Array module (``numpy`` or ``cupy``).
cutoff: Relative singular-value cutoff for adaptive bond truncation.
dtype: The data type of the sketches.
to_host: If ``True`` (default), the result is copied to the host as numpy
arrays. If ``False``, it stays on the device of ``xp``.

Returns:
The site arrays of the compressed train in right-canonical form, as numpy
arrays.
arrays if ``to_host`` is ``True``, otherwise as arrays of ``xp``..
"""
depth = len(layers)
n_sites = len(layers[0])
Expand Down Expand Up @@ -157,4 +160,7 @@ def sweep(
eta = [first.reshape(1, *first.shape[depth:]), *reversed(eta_reversed)]
logger.debug("Right-to-left sweep: %.3f s", (perf_counter_ns() - tms) * 1e-9)

return [to_numpy(site) for site in unpad(eta, kind)]
sites_out = unpad(eta, kind)
if to_host:
return [to_numpy(site) for site in sites_out]
return list(sites_out)
8 changes: 6 additions & 2 deletions src/src_method/apply.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from __future__ import annotations

from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any

from ._tensor_train import infer_kind
from .stack import src
Expand All @@ -28,7 +28,8 @@ def apply(
dtype: DTypeLike | None = None,
seed: int | None = None,
device: str = "cpu",
) -> list[NDArray]:
to_host: bool = True,
) -> list[Any]:
"""Applies the Successive Randomized Compression (SRC) algorithm.

Equivalent to ``src(left_tensor, right_tensor, ...)`` restricted to an MPO on
Expand All @@ -54,6 +55,8 @@ def apply(
seed: An optional seed for the random number generator.
device: ``"cpu"`` (default, numpy) or ``"gpu"`` (cupy). Requires
the optional ``cupy`` dependency for GPU execution.
to_host: If ``True`` (default), results are returned as numpy arrays on
the host. Set to ``False`` to keep tensors on the compute device.

Returns:
The site arrays of the compressed tensor network (MPS or MPO).
Expand Down Expand Up @@ -85,4 +88,5 @@ def apply(
dtype=dtype,
seed=seed,
device=device,
to_host=to_host,
)
15 changes: 12 additions & 3 deletions src/src_method/compress.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from __future__ import annotations

from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any

from .stack import src

Expand All @@ -26,7 +26,8 @@ def compress(
dtype: DTypeLike | None = None,
seed: int | None = None,
device: str = "cpu",
) -> list[NDArray]:
to_host: bool = True,
) -> list[Any]:
"""Applies the Successive Randomized Compression (SRC) algorithm.

Equivalent to ``src(tensor, ...)``; see `src_method.stack.src` for the
Expand All @@ -50,6 +51,8 @@ def compress(
seed: An optional seed for the random number generator.
device: ``"cpu"`` (default, numpy) or ``"gpu"`` (cupy). Requires
the optional ``cupy`` dependency for GPU execution.
to_host: If ``True`` (default), results are returned as numpy arrays on
the host. Set to ``False`` to keep tensors on the compute device.

Returns:
The site arrays of the compressed tensor network (MPS or MPO).
Expand All @@ -64,5 +67,11 @@ def compress(
ImportError: If ``device="gpu"`` but cupy is not installed.
"""
return src(
tensor, chi_out=chi_out, cutoff=cutoff, dtype=dtype, seed=seed, device=device
tensor,
chi_out=chi_out,
cutoff=cutoff,
dtype=dtype,
seed=seed,
device=device,
to_host=to_host,
)
13 changes: 9 additions & 4 deletions src/src_method/stack.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from __future__ import annotations

import logging
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any

from ._sweep import sweep
from ._tensor_train import (
Expand Down Expand Up @@ -40,7 +40,8 @@ def src(
dtype: DTypeLike | None = None,
seed: int | None = None,
device: str = "cpu",
) -> list[NDArray]:
to_host: bool = True,
) -> list[Any]:
"""Contract a stack of tensor trains and compress the result with SRC.

The stack ``src(T_1, T_2, ..., T_m)`` is the product ``T_1 T_2 ... T_m`` in
Expand Down Expand Up @@ -80,10 +81,12 @@ def src(
seed: An optional seed for the random number generator.
device: ``"cpu"`` (default, numpy) or ``"gpu"`` (cupy). Requires
the optional ``cupy`` dependency for GPU execution.
to_host: If ``True`` (default), results are returned as numpy arrays on
the host. Set to ``False`` to keep tensors on the compute device.

Returns:
The site arrays of the compressed train (MPS or MPO), in right-canonical
form, as numpy arrays (host-side, whatever the ``device``).
form, as numpy arrays if ``to_host`` is ``True``, otherwise as arrays of ``xp``.

Raises:
TypeError: If ``chi_out`` is not an integer, if a train has an unrecognised
Expand All @@ -105,7 +108,8 @@ def src(
if n_sites < MIN_SRC_SITES:
check_exact_supported(n_sites)
logger.warning(LOG_WARN_SMALL)
return exact_stack(layers, chi_out, kind)
result = exact_stack(layers, chi_out, kind)
return result if to_host else [xp.asarray(t) for t in result]
Comment on lines +111 to +112

logger.debug(
"Starting SRC: n_sites=%d, depth=%d, output=%s, device=%s",
Expand All @@ -122,6 +126,7 @@ def src(
xp,
cutoff=cutoff,
dtype=sketch_dtype(dtype, *layers),
to_host=to_host,
)
logger.debug("SRC complete")
return result
7 changes: 2 additions & 5 deletions src/src_method/utils/linalg.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,16 +40,13 @@ def truncated_qr(
return xp.linalg.qr(matrix)[0]

Q, R = xp.linalg.qr(matrix.T if transpose else matrix)
R_np = R.get() if hasattr(R, "get") else np.asarray(R)
U, S, _ = np.linalg.svd(R_np.T if transpose else R_np, full_matrices=False)
U, S, _ = xp.linalg.svd(R.T if transpose else R, full_matrices=False)
rank = max(1, int((cutoff * S[0] <= S).sum()))

if transpose:
Q_trunc = U[:, :rank]
else:
U_trunc = xp.asarray(U[:, :rank]) if xp is not np else U[:, :rank]
U_trunc = U[:, :rank]
Q_trunc = Q @ U_trunc

if xp is not np:
return xp.asarray(Q_trunc)
return Q_trunc
112 changes: 112 additions & 0 deletions tests/test_gpu_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,3 +181,115 @@ def test_src_stack_gpu_matches_reference(device: str) -> None:
ref = H1.apply(H2.apply(psi, compress=False), compress=False)

np.testing.assert_allclose(ref.distance(out), 0.0, atol=1e-6)


def test_apply_mpo_mps_to_host_false_returns_cupy() -> None:
"""With to_host=False, apply() must return cupy arrays on GPU."""
n_sites, phys_dim, chi_out = 5, 2, 8
H = qtn.MPO_rand(n_sites, bond_dim=4, phys_dim=phys_dim, dtype=np.complex128)
psi = qtn.MPS_rand_state(
n_sites, bond_dim=chi_out, phys_dim=phys_dim, dtype=np.complex128
)

result = apply(
H.arrays,
psi.arrays,
chi_out=chi_out,
dtype=np.complex128,
seed=0,
device="gpu",
to_host=False,
)
assert all(isinstance(t, cupy.ndarray) for t in result)


def test_compress_mpo_to_host_false_returns_cupy() -> None:
"""With to_host=False, compress() must return cupy arrays on GPU."""
n_sites, phys_dim, chi_out = 5, 2, 8
H = qtn.MPO_rand(n_sites, bond_dim=4, phys_dim=phys_dim, dtype=np.complex128)

result = compress(
H.arrays,
chi_out=chi_out,
dtype=np.complex128,
seed=0,
device="gpu",
to_host=False,
)
assert all(isinstance(t, cupy.ndarray) for t in result)


def test_chained_apply_no_host_roundtrip() -> None:
"""Chained GPU calls with to_host=False must work and match CPU results."""
n_sites, phys_dim, chi_out = 5, 2, 8
H1 = qtn.MPO_rand(n_sites, bond_dim=4, phys_dim=phys_dim, dtype=np.complex128)
H2 = qtn.MPO_rand(n_sites, bond_dim=4, phys_dim=phys_dim, dtype=np.complex128)
psi = qtn.MPS_rand_state(
n_sites, bond_dim=chi_out, phys_dim=phys_dim, dtype=np.complex128
)

# 1. Chained GPU apply
intermediate_gpu = apply(
H1.arrays,
psi.arrays,
chi_out=chi_out,
dtype=np.complex128,
seed=0,
device="gpu",
to_host=False,
)
assert isinstance(intermediate_gpu[0], cupy.ndarray)

final_gpu = apply(
H2.arrays,
intermediate_gpu,
chi_out=chi_out,
dtype=np.complex128,
seed=0,
device="gpu",
to_host=False,
)
assert isinstance(final_gpu[0], cupy.ndarray)

# 2. Chained CPU apply for numerical comparison
intermediate_cpu = apply(
H1.arrays,
psi.arrays,
chi_out=chi_out,
dtype=np.complex128,
seed=0,
device="cpu",
to_host=True,
)
final_cpu = apply(
H2.arrays,
intermediate_cpu,
chi_out=chi_out,
dtype=np.complex128,
seed=0,
device="cpu",
to_host=True,
)

# Bring GPU result to host and compare
final_gpu_host = [t.get() for t in final_gpu]

qtn_cpu = as_mps(final_cpu)
qtn_gpu = as_mps(final_gpu_host)
np.testing.assert_allclose(qtn_cpu.distance(qtn_gpu), 0.0, atol=1e-6)


def test_to_host_true_backward_compatible() -> None:
"""to_host=True (implicit default) must return numpy arrays even when device='gpu'."""
n_sites, phys_dim, chi_out = 5, 2, 8
H = qtn.MPO_rand(n_sites, bond_dim=4, phys_dim=phys_dim, dtype=np.complex128)

result = compress(
H.arrays,
chi_out=chi_out,
dtype=np.complex128,
seed=0,
device="gpu",
# to_host=True is the implicit default here
)
assert all(isinstance(t, np.ndarray) for t in result)
Loading