Skip to content
Draft
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
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
__pycache__
.ipynb*
*Manifest.toml
LocalPreferences.toml
.vscode
experimental
refs
Expand Down
3 changes: 3 additions & 0 deletions Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ Enzyme = "7da242da-08ed-463a-9acd-ee780be4f1d9"
EnzymeTestUtils = "12d8515a-0907-448a-8884-5fe00fdf1c5a"
FiniteDifferences = "26cc04aa-876d-5657-8c51-4c34ba976000"
GPUArrays = "0c68f7d7-f131-5f86-a1c3-88cf8149b2d7"
JLD2 = "033835bb-8acc-5ee8-8aae-3f567f8a3819"
Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6"

[extensions]
Expand All @@ -43,6 +44,7 @@ TensorKitEnzymeExt = "Enzyme"
TensorKitEnzymeTestUtilsExt = "EnzymeTestUtils"
TensorKitFiniteDifferencesExt = "FiniteDifferences"
TensorKitGPUArraysExt = "GPUArrays"
TensorKitJLD2Ext = "JLD2"
TensorKitMooncakeExt = "Mooncake"

[compat]
Expand All @@ -55,6 +57,7 @@ Enzyme = "0.13.195"
EnzymeTestUtils = "0.2.8"
FiniteDifferences = "0.12"
GPUArrays = "11.4.1"
JLD2 = "0.6"
LRUCache = "1.6"
LinearAlgebra = "1"
MatrixAlgebraKit = "0.6.9"
Expand Down
2 changes: 2 additions & 0 deletions docs/src/Changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@ When releasing a new version, move the "Unreleased" changes to a new version sec

### Added

- Versioned, fusion-tree-based `save_tensor` and `load_tensor` support for `TensorMap`, `DiagonalTensorMap`, and `BraidingTensor` objects through the optional JLD2-based `TensorKitJLD2Ext` extension.

### Changed
- For sector types with `GenericUnit` such that colorings are not unique, `GradedSpace`, `ProductSpace` and `HomSpace` now check for this compatibility. In particular, this prevents the construction of `TensorMap`s with incompatible colorings, which previously either errored or produced empty tensors inconsistently. ([#515](https://github.com/QuantumKitHub/TensorKit.jl/pull/515))

Expand Down
6 changes: 6 additions & 0 deletions docs/src/lib/tensors.md
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,12 @@ AdjointTensorMap
BraidingTensor
```

Tensor maps can be stored and restored with:
```@docs
save_tensor
load_tensor
```

Of those, `TensorMap` provides the generic instantiation of our tensor concept. It supports various constructors, which are discussed in the next subsection.

Furthermore, some aliases are provided for convenience:
Expand Down
26 changes: 19 additions & 7 deletions docs/src/man/tensors.md
Original file line number Diff line number Diff line change
Expand Up @@ -427,12 +427,24 @@ f1, f2 = first(fusiontrees(t))
t[f1,f2]
```

## [Reading and writing tensors: `Dict` conversion](@id ss_tensor_readwrite)
## [Reading and writing tensors](@id ss_tensor_readwrite)

There are no custom or dedicated methods for reading, writing or storing `TensorMap`s, however, there is the possibility to convert a `t::AbstractTensorMap` into a `Dict`, simply as `convert(Dict, t)`.
The backward conversion `convert(TensorMap, dict)` will return a tensor that is equal to `t`, i.e. `t == convert(TensorMap, convert(Dict, t))`.
TensorKit provides [`save_tensor`](@ref) and [`load_tensor`](@ref) for storing one tensor map in a versioned JLD2 file.
Install JLD2 and load it with `using JLD2` to activate these functions through the `TensorKitJLD2Ext` extension.

This conversion relies on that the string representation of objects such as `VectorSpace`, `FusionTree` or `Sector` should be such that it represents valid code to recreate the object.
Hence, we store information about the domain and codomain of the tensor, and the sector associated with each data block, as a `String` obtained with `repr`.
This provides the flexibility to still change the internal structure of such objects, without this breaking the ability to load older data files.
The resulting dictionary can then be stored using any of the provided Julia packages such as [JLD.jl](https://github.com/JuliaIO/JLD.jl), [JLD2.jl](https://github.com/JuliaIO/JLD2.jl), [BSON.jl](https://github.com/JuliaIO/BSON.jl), [JSON.jl](https://github.com/JuliaIO/JSON.jl), ...
```julia
using JLD2

filename = "tensor.jld2"
save_tensor(filename, t)
t′ = load_tensor(filename)
```

`TensorMap`, `DiagonalTensorMap`, and `BraidingTensor` retain their semantic types, while numerical storage is copied to a CPU `Vector` when saving and loading.
Dense numerical segments are labeled by the semantic fields of their codomain and domain fusion trees, so loading does not depend on fusion-tree or block iteration order.
The compact data of `DiagonalTensorMap` and the structural description of `BraidingTensor` are stored without materializing dense blocks.
A lazy `AdjointTensorMap` must be materialized explicitly before saving, for example with `save_tensor(filename, convert(TensorMap, t'))`.
TensorKit does not add a filename extension and replaces an existing file at the requested path.

The older `convert(Dict, t)` and `convert(TensorMap, dict)` workflow remains available for compatibility.
That representation stores spaces and block sectors as strings and does not preserve specialized tensor-map types, so it is no longer recommended for new files.
18 changes: 18 additions & 0 deletions ext/TensorKitJLD2Ext/TensorKitJLD2Ext.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
module TensorKitJLD2Ext

using TensorKit
using TensorKit: AdjointTensorMap
import TensorKit: save_tensor, load_tensor
import JLD2

# TensorMap IO
#=============#

const TENSORMAP_FILE_FORMAT = "TensorKit.AbstractTensorMap"
const TENSORMAP_FILE_VERSION = UInt16(1)

include("fusiontrees.jl")
include("records.jl")
include("io.jl")

end
118 changes: 118 additions & 0 deletions ext/TensorKitJLD2Ext/fusiontrees.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
const _FUSIONTREE_TABLE_FIELDS = (:uncoupled, :coupled, :isdual, :innerlines, :vertices)

"""Return the position of a fusion tree, adding it to the table when necessary."""
function _intern_fusiontree!(trees::AbstractVector, tree::FusionTree)
index = findfirst(==(tree), trees)
if isnothing(index)
push!(trees, tree)
return length(trees)
end
return index
end

"""Encode fusion trees as columnar arrays of their semantic fields."""
function _encode_fusiontrees(trees::AbstractVector, ::Type{I}, numlegs::Int) where {I <: Sector}
numtrees = length(trees)
numinner = max(0, numlegs - 2)
numvertices = max(0, numlegs - 1)
uncoupled = Matrix{I}(undef, numlegs, numtrees)
coupled = Vector{I}(undef, numtrees)
isdual = falses(numlegs, numtrees)
innerlines = Matrix{I}(undef, numinner, numtrees)
vertices = Matrix{Int}(undef, numvertices, numtrees)
for (column, tree) in enumerate(trees)
length(tree.uncoupled) == numlegs ||
error("inconsistent fusion-tree leg count while saving")
length(tree.innerlines) == numinner ||
error("inconsistent fusion-tree inner-line count while saving")
length(tree.vertices) == numvertices ||
error("inconsistent fusion-tree vertex count while saving")
uncoupled[:, column] .= tree.uncoupled
coupled[column] = tree.coupled
isdual[:, column] .= tree.isdual
innerlines[:, column] .= tree.innerlines
vertices[:, column] .= tree.vertices
end
return (; uncoupled, coupled, isdual, innerlines, vertices)
end

"""Decode a columnar fusion-tree table and validate its basic representation."""
function _decode_fusiontrees(table, ::Type{I}, numlegs::Int, description::AbstractString) where {I <: Sector}
_require_record_fields(table, _FUSIONTREE_TABLE_FIELDS, description)
table.uncoupled isa Matrix{I} ||
throw(ArgumentError("serialized $description has invalid uncoupled sectors"))
table.coupled isa Vector{I} ||
throw(ArgumentError("serialized $description has invalid coupled sectors"))
table.isdual isa BitMatrix ||
throw(ArgumentError("serialized $description has invalid duality flags"))
table.innerlines isa Matrix{I} ||
throw(ArgumentError("serialized $description has invalid inner lines"))
table.vertices isa Matrix{Int} ||
throw(ArgumentError("serialized $description has invalid vertices"))

numtrees = length(table.coupled)
expected_sizes = (
(numlegs, numtrees),
(numlegs, numtrees),
(max(0, numlegs - 2), numtrees),
(max(0, numlegs - 1), numtrees),
)
actual_sizes = (
size(table.uncoupled), size(table.isdual),
size(table.innerlines), size(table.vertices),
)
actual_sizes == expected_sizes ||
throw(DimensionMismatch("serialized $description has inconsistent table dimensions"))

trees = Vector{FusionTree{I, numlegs}}(undef, numtrees)
for column in 1:numtrees
uncoupled = ntuple(row -> table.uncoupled[row, column], numlegs)
isdual = ntuple(row -> table.isdual[row, column], numlegs)
innerlines = ntuple(row -> table.innerlines[row, column], max(0, numlegs - 2))
vertices = ntuple(row -> table.vertices[row, column], max(0, numlegs - 1))
trees[column] = try
FusionTree{I}(uncoupled, table.coupled[column], isdual, innerlines, vertices)
catch error
message = sprint(showerror, error)
throw(ArgumentError("serialized $description contains an invalid fusion tree: $message"))
end
end
_check_unique_values(trees, description)
return trees
end

"""Decode explicit fusion-tree pair identifiers and validate them against a tensor-map space."""
function _decode_fusiontree_pairs(pair_ids, codomain_trees, domain_trees, tensor_space::TensorMapSpace)
pair_ids isa Matrix{Int} ||
throw(ArgumentError("serialized tensor has invalid fusion-tree pair identifiers"))
size(pair_ids, 1) == 2 ||
throw(DimensionMismatch("serialized fusion-tree pair identifiers must have two rows"))

numpairs = size(pair_ids, 2)
pairs = Vector{Tuple{eltype(codomain_trees), eltype(domain_trees)}}(undef, numpairs)
used_codomain = falses(length(codomain_trees))
used_domain = falses(length(domain_trees))
for column in 1:numpairs
codomain_id = pair_ids[1, column]
domain_id = pair_ids[2, column]
checkbounds(Bool, codomain_trees, codomain_id) ||
throw(ArgumentError("serialized tensor has an out-of-range codomain fusion-tree identifier"))
checkbounds(Bool, domain_trees, domain_id) ||
throw(ArgumentError("serialized tensor has an out-of-range domain fusion-tree identifier"))
pairs[column] = (codomain_trees[codomain_id], domain_trees[domain_id])
used_codomain[codomain_id] = true
used_domain[domain_id] = true
end
_check_unique_values(pairs, "fusion-tree pairs")

all(used_codomain) ||
throw(ArgumentError("serialized tensor contains an unused codomain fusion tree"))
all(used_domain) ||
throw(ArgumentError("serialized tensor contains an unused domain fusion tree"))

expected_pairs = collect(fusiontrees(tensor_space))
length(pairs) == length(expected_pairs) &&
all(pair -> any(==(pair), expected_pairs), pairs) ||
throw(ArgumentError("serialized fusion-tree pairs do not match the tensor-map space"))
return pairs
end
27 changes: 27 additions & 0 deletions ext/TensorKitJLD2Ext/io.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
function save_tensor(path::AbstractString, tensor::AbstractTensorMap)
record = _pack_tensormap(tensor)
destination = abspath(path)
mktemp(dirname(destination)) do temporary, io
close(io)
JLD2.jldsave(
temporary; format = TENSORMAP_FILE_FORMAT,
version = TENSORMAP_FILE_VERSION, tensor = record
)
mv(temporary, destination; force = true)
end
return nothing
end

function load_tensor(path::AbstractString)
record = JLD2.jldopen(path, "r") do file
all(key -> haskey(file, key), ("format", "version", "tensor")) ||
throw(ArgumentError("file is not a TensorKit tensor-map file"))
file["format"] == TENSORMAP_FILE_FORMAT ||
throw(ArgumentError("file has an invalid TensorKit tensor-map format marker"))
version = file["version"]
version == TENSORMAP_FILE_VERSION ||
throw(ArgumentError("unsupported TensorKit tensor-map file version $version"))
return file["tensor"]
end
return _unpack_tensormap(record)
end
Loading
Loading