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
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
name = "ITensorCPD"
uuid = "8ca0d870-8743-11ef-3aad-a3ebf6911431"
authors = ["Karl Pierce <kpierce@flatironinstitute.org>"]
version = "0.0.82"
version = "0.0.83"

[deps]
Adapt = "79e6a3ab-5dfb-504d-930d-738a2a938a0e"
Expand Down
11 changes: 6 additions & 5 deletions src/algebra/ldiv_solve.jl
Original file line number Diff line number Diff line change
Expand Up @@ -11,19 +11,20 @@ end

## For now just call with QR CP if factorize. Later this will be more complex.
function ldiv_solve!(A, B; factorizeA = false)
sol = nothing
if factorizeA
szA = size(A)
if (szA[1] == szA[2])
try
return cholesky(Hermitian(A), RowMaximum(), check=true, tol=cholesky_epsilon) \ B
sol = cholesky(Hermitian(A), RowMaximum(), check=true, tol=cholesky_epsilon) \ B
catch
# println("Warning: Cholesky based solver failed.")
return qr(A, ColumnNorm()) \ B
sol = qr(A, ColumnNorm()) \ B
end
else
return qr(A, ColumnNorm()) \ B
sol = qr(A, ColumnNorm()) \ B
end
else
return A \ B
sol = A \ B
end
return sol
end
Original file line number Diff line number Diff line change
Expand Up @@ -55,14 +55,13 @@ abstract type ProjectionAlgorithm end

## Default algorithm uses the pivoted QR to solve LS problem.
function solve_ls_problem(::ProjectionAlgorithm, als, projected_KRP, project_target, rank)
# direction = qr(array(projected_KRP * prime(projected_KRP, tags=tags(rank))), ColumnNorm()) \ transpose(array(project_target * projected_KRP))
#direction = qr(array(projected_KRP), ColumnNorm()) \ transpose(array(project_target))
projected_KRP, normalizers = row_norm(projected_KRP, ind(projected_KRP, 1))
direction = nothing
if als.additional_items[:normal]
direction = ldiv_solve!!(expose(array(dag(projected_KRP * prime(dag(projected_KRP); tags=tags(rank))))), expose(transpose(array(project_target * prime(dag(projected_KRP); tags=tags(rank)))));factorizeA=true)
else
direction = ldiv_solve!!(expose(array(projected_KRP)), expose(transpose(array(project_target)));factorizeA=true)
end
i = ind(project_target, 1)
return itensor(copy(transpose(direction)), i,rank)
return itensor(copy(transpose(direction)), i,rank), normalizers
end
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,7 @@ end
end

function matricize_tensor(::LevScoreSampled, ::Val{true}, als, factors, cp, rank::Index, fact::Int)
if als.check.iter ≤ als.additional_items[:stop_resample]
if als.check.iter < als.additional_items[:stop_resample] || !isassigned(als.additional_items[:sampled_targets], fact)
als.additional_items[:sampled_targets][fact] = fused_flatten_sample(als.target, fact, als.additional_items[:projects_tensors][fact])
end

Expand All @@ -54,7 +54,7 @@ end

function post_solve(::LevScoreSampled, als, factors, λ, cp, rank::Index, fact::Integer)
## update the factor weights.
@inbounds als.additional_items[:factor_weights][fact] = compute_leverage_score_probabilitiy(factors[fact], ind(cp, fact))
@inbounds als.additional_items[:factor_weights][fact] = compute_leverage_score_probabilitiy(factors[fact], ind(cp, fact); use_variance = als.additional_items[:variance_truncation])
end

### With this solver we are going to compute sampling projectors for LS decomposition
Expand Down Expand Up @@ -104,6 +104,6 @@ end

function post_solve(::BlockLevScoreSampled, als, factors, λ, cp, rank::Index, fact::Integer)
## update the factor weights.
als.additional_items[:factor_weights][fact] = compute_leverage_score_probabilitiy(factors[fact], ind(cp, fact))
als.additional_items[:factor_weights][fact] = compute_leverage_score_probabilitiy(factors[fact], ind(cp, fact); use_variance = als.additional_items[:variance_truncation])
end

2 changes: 1 addition & 1 deletion src/algorithms/als_algorithms/standard/MttkrpAlgorithm.jl
Original file line number Diff line number Diff line change
Expand Up @@ -37,5 +37,5 @@ abstract type MttkrpAlgorithm end
#solution = array(dag(krp)) \ transpose(array(mtkrp))
solution = ldiv_solve!!(expose(array(dag(krp))), expose(transpose(array(mtkrp))); factorizeA = true)
i = ind(mtkrp, 1)
return itensor(copy(transpose(solution)), i,rank)
return itensor(copy(transpose(solution)), i,rank), ones(eltype(solution), dim(rank))
end
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ struct InvKRP <: ProjectionAlgorithm end

function solve_ls_problem(::InvKRP, _, projected_KRP, projected_target, rank)
U, S, V = svd(dag(projected_KRP), rank; use_absolute_cutoff = true, cutoff = 0)
return prime(projected_target; tags = tags(rank)) * V * (1 ./ S) * U
return prime(projected_target; tags = tags(rank)) * V * (1 ./ S) * U, ones(eltype(U), dim(rank))
end

function post_solve(::InvKRP, als, factors, λ, cp, rank::Index, fact::Integer) end
24 changes: 19 additions & 5 deletions src/math_tools/probability.jl
Original file line number Diff line number Diff line change
@@ -1,12 +1,26 @@
using LinearAlgebra, StatsBase
using ITensors: Index
function compute_leverage_score_probabilitiy(A, row::Index)
function compute_leverage_score_probabilitiy(A, row::Index; use_variance=true)
## This only works on matrices for now.
@assert ndims(A) == 2
q, _ = qr(A, row)
ITensors.hadamard_product!(q, q, dag(q))
ni = dim(q, 1)
return [real(sum(array(q)[i,:])) for i in 1:ni] ./ minimum(dims(A))
posrow = findall(i-> i!=row, inds(A))
col = ind(A, posrow[1])
Am = array(A, (row, col))
q, r, p = qr(Am, ColumnNorm())

rd = diag(r).^2
rd = rd ./ sum(rd)

q = copy(q)
q = q .* conj(q)
ni = size(q)[1]
mn = minimum(dims(A))

if use_variance
return [real(sum(array(q)[i,1:mn] .* rd)) for i in 1:ni]
end

return [real(sum(array(q)[i,1:mn])) for i in 1:ni] ./ mn
end

function samples_from_probability_vector(PW::Vector, samples)
Expand Down
2 changes: 1 addition & 1 deletion src/math_tools/row_norm.jl
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ function row_norm(t::ITensor, i...)
dataT = NDTensors.datatype(t)
λ = hadamard_product(t, t)
for is in tuple(i...)
d = itensor(NDTensors.Diag(dataT(ones(Float32, dim(is)))), is)
d = itensor(NDTensors.Diag(dataT(ones(elt, dim(is)))), is)
λ = λ * d
end
map!(i -> sqrt(i), data(λ), data(λ))
Expand Down
13 changes: 13 additions & 0 deletions src/optimizers/als_optimizers/als_optimizer.jl
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,19 @@ include("randomized/qr_lev_score_sampled.jl")
include("randomized/krp_lev_score_sampled.jl")
include("randomized/sketched_ls.jl")

### Default ALS constructor algorithm for Arrays which converts to tensors.
### This will develop the "optimization sequence" variable
### and then pass along to more specialized constructors
function compute_als(
target::AbstractArray,
cp::CPD{<:ITensor};
kwargs...
)
@assert ndims(target) == ndims(cp)

return compute_als(itensor(target, inds(cp)), cp; kwargs...)
end

### Default ALS constructor algorithm for Tensors (versus tensor networks).
### This will develop the "optimization sequence" variable
### and then pass along to more specialized constructors
Expand Down
4 changes: 3 additions & 1 deletion src/optimizers/als_optimizers/optimize.jl
Original file line number Diff line number Diff line change
Expand Up @@ -20,10 +20,12 @@ function optimize(cp::CPD, als::ALS; verbose = false)

mtkrp = matricize_tensor(als.mttkrp_alg, als, factors, cp, rank, fact)

solution = solve_ls_problem(als.mttkrp_alg, als, krp, mtkrp, rank)
solution, normalizers = solve_ls_problem(als.mttkrp_alg, als, krp, mtkrp, rank)

factors[fact], λ = row_norm(solution, target_ind)

array(λ) .*= 1 ./ array(normalizers)

post_solve(als.mttkrp_alg, als, factors, λ, cp, rank, fact)
end

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,11 @@ function compute_als(
normal=false,
stop_resample=-1,
cache_sampled_targets=true,
variance_truncation = true,
kwargs...
)
## For each factor matrix compute its weights
extra_args[:factor_weights] = [compute_leverage_score_probabilitiy(cp[i], ind(cp, i)) for i in 1:length(cp)]
extra_args[:factor_weights] = [compute_leverage_score_probabilitiy(cp[i], ind(cp, i); use_variance = variance_truncation) for i in 1:length(cp)]
projects_tensors = Vector{ITensor}()
cache_sampled_targets = (stop_resample == -1 ? false : cache_sampled_targets)
for fact in 1:length(cp)
Expand All @@ -36,6 +37,7 @@ function compute_als(
extra_args[:stop_resample] = stop_resample
extra_args[:sampled_targets] = Vector{ITensor}(undef, ndims(target))
extra_args[:cache_sampled_targets] = cache_sampled_targets
extra_args[:variance_truncation] = variance_truncation
return ALS(target, alg, extra_args, check)
end

Expand All @@ -47,11 +49,14 @@ function compute_als(
check = nothing,
normal=false,
stop_resample=-1,
cache_sampled_targets=true,
variance_truncation = true,
kwargs...
)
## For each factor matrix compute its weights
extra_args[:factor_weights] = [compute_leverage_score_probabilitiy(cp[i], ind(cp, i)) for i in 1:length(cp)]
extra_args[:factor_weights] = [compute_leverage_score_probabilitiy(cp[i], ind(cp, i); use_variance = variance_truncation) for i in 1:length(cp)]
projects_tensors = Vector{ITensor}()
cache_sampled_targets = (stop_resample == -1 ? false : cache_sampled_targets)
for fact in 1:length(cp)
## grab the tensor indices for all other factors but fact
Ris = inds(cp)[1:end .!= fact]
Expand Down Expand Up @@ -79,5 +84,7 @@ function compute_als(
extra_args[:projects_tensors] = projects_tensors
extra_args[:normal] = normal
extra_args[:stop_resample] = stop_resample
extra_args[:cache_sampled_targets] = cache_sampled_targets
extra_args[:variance_truncation] = variance_truncation
return ALS(target, alg, extra_args, check)
end
Loading