Skip to content

Commit ea57913

Browse files
authored
Merge pull request #59 from JuliaQUBO/fix/issue-58-return-embeddings
Expose returned embeddings for supplied samplers
2 parents 27376a6 + cc1e14f commit ea57913

2 files changed

Lines changed: 141 additions & 5 deletions

File tree

src/sampler.jl

Lines changed: 79 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -104,6 +104,81 @@ function _dwave_chip_info(dwave_sampler)
104104
return chip_info
105105
end
106106

107+
function _normalise_embedding_index(value)
108+
value isa Integer && return Int(value)
109+
110+
if value isa AbstractString
111+
index = tryparse(Int, value)
112+
index === nothing || return index
113+
end
114+
115+
return nothing
116+
end
117+
118+
function _normalise_embedding_chain(chain)
119+
chain isa AbstractVector || chain isa Tuple || return nothing
120+
121+
normal_chain = Int[]
122+
sizehint!(normal_chain, length(chain))
123+
124+
for qubit in chain
125+
index = _normalise_embedding_index(qubit)
126+
index === nothing && return nothing
127+
128+
push!(normal_chain, index)
129+
end
130+
131+
return normal_chain
132+
end
133+
134+
function _normalise_embedding(embedding)
135+
embedding isa AbstractDict || return nothing
136+
137+
normal_embedding = Dict{Int,Vector{Int}}()
138+
139+
for (variable, chain) in pairs(embedding)
140+
normal_variable = _normalise_embedding_index(variable)
141+
normal_variable === nothing && return nothing
142+
normal_chain = _normalise_embedding_chain(chain)
143+
normal_chain === nothing && return nothing
144+
145+
normal_embedding[normal_variable] = normal_chain
146+
end
147+
148+
return normal_embedding
149+
end
150+
151+
function _normalise_dwave_embedding!(dwave_info::AbstractDict)
152+
context = get(dwave_info, "embedding_context", nothing)
153+
context isa AbstractDict || return dwave_info
154+
155+
normal_embedding = _normalise_embedding(get(context, "embedding", nothing))
156+
normal_embedding === nothing || (context["embedding"] = normal_embedding)
157+
158+
return dwave_info
159+
end
160+
161+
@doc raw"""
162+
DWave.embedding(sampleset_or_metadata)
163+
164+
Return the minor embedding recorded in D-Wave sample-set metadata, or `nothing`
165+
when no embedding was returned. Embeddings are normalized as
166+
`Dict{Int,Vector{Int}}`.
167+
"""
168+
function embedding(metadata::AbstractDict)
169+
dwave_info = haskey(metadata, "dwave_info") ? metadata["dwave_info"] : metadata
170+
dwave_info isa AbstractDict || return nothing
171+
172+
context = get(dwave_info, "embedding_context", nothing)
173+
context isa AbstractDict || return nothing
174+
175+
return _normalise_embedding(get(context, "embedding", nothing))
176+
end
177+
178+
function embedding(sampleset::QUBOTools.SampleSet)
179+
return embedding(QUBOTools.metadata(sampleset))
180+
end
181+
107182
function QUBODrivers.sample(sampler::Optimizer{T}) where {T}
108183
# Ising Model
109184
n, h, J, α, β = QUBOTools.ising(sampler, :dict; sense = :min)
@@ -112,8 +187,9 @@ function QUBODrivers.sample(sampler::Optimizer{T}) where {T}
112187
num_reads = MOI.get(sampler, DWave.NumberOfReads())
113188
final_num_reads = MOI.get(sampler, QUBODrivers.FinalNumberOfReads())
114189
sample_params = Dict{Symbol,Any}(
115-
:num_reads => final_num_reads,
116-
:annealing_time => MOI.get(sampler, DWave.AnnealingTime()),
190+
:num_reads => final_num_reads,
191+
:annealing_time => MOI.get(sampler, DWave.AnnealingTime()),
192+
:return_embedding => MOI.get(sampler, DWave.ReturnEmbedding()),
117193
)
118194
dwave_sampler = MOI.get(sampler, DWave.Sampler())
119195

@@ -123,15 +199,14 @@ function QUBODrivers.sample(sampler::Optimizer{T}) where {T}
123199
token = get(ENV, "DWAVE_API_TOKEN", nothing)
124200
)
125201
)
126-
127-
sample_params[:return_embedding] = MOI.get(sampler, DWave.ReturnEmbedding())
128202
end
129203

130204
# Results
131205
samples = QUBOTools.Sample{T,Int}[]
132206
results = @timed dwave_sampler.sample_ising(h, J; sample_params...)
133207
var_map = pyconvert.(Int, [var for var in results.value.variables])
134208
dw_info = jl_object(results.value.info)
209+
_normalise_dwave_embedding!(dw_info)
135210
chip_info = _dwave_chip_info(dwave_sampler)
136211

137212
if !isempty(chip_info)

test/dwave_metadata.jl

Lines changed: 62 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,54 @@ sampler = FakeEmbeddingSampler()
4545
return ans.sampler
4646
end
4747

48-
function _wrapper_dwave_sampleset(sampler)
48+
function _fake_return_embedding_sampler()
49+
ans = DWave.PythonCall.pyexec(
50+
@NamedTuple{sampler::DWave.PythonCall.Py,calls::DWave.PythonCall.Py},
51+
"""
52+
class FakeEmbeddingSampler:
53+
def __init__(self):
54+
self.calls = []
55+
self.child = type("FakeChildSampler", (), {})()
56+
self.child.properties = {
57+
"chip_id": "mock-chip",
58+
"topology": {"type": "pegasus", "shape": [16]},
59+
"category": "qpu",
60+
"qubits": [3, 4, 5],
61+
"couplers": [[3, 4], [4, 5]],
62+
"num_qubits": 3,
63+
}
64+
self.child.solver = type("FakeSolver", (), {"name": "mock-solver"})()
65+
66+
def sample_ising(self, h, J, **params):
67+
self.calls.append(params)
68+
info = {
69+
"problem_id": "mock-problem",
70+
"timing": {"qpu_access_time": 42},
71+
}
72+
if params.get("return_embedding"):
73+
info["embedding_context"] = {"embedding": {1: (3, 4), 2: (5,)}}
74+
75+
return type(
76+
"FakeSampleSet",
77+
(),
78+
{
79+
"variables": [1, 2],
80+
"record": [([1, -1], -1.25, 3)],
81+
"info": info,
82+
},
83+
)()
84+
85+
sampler = FakeEmbeddingSampler()
86+
calls = sampler.calls
87+
""",
88+
@__MODULE__,
89+
(),
90+
)
91+
92+
return ans
93+
end
94+
95+
function _wrapper_dwave_sampleset(sampler; return_embedding::Bool = false)
4996
model = MOI.instantiate(DWave.Optimizer; with_bridge_type = Float64)
5097
variables, _ = MOI.add_constrained_variables(model, fill(QUBODrivers.Spin(), 2))
5198

@@ -65,6 +112,7 @@ function _wrapper_dwave_sampleset(sampler)
65112
),
66113
)
67114
MOI.set(model, MOI.RawOptimizerAttribute("sampler"), sampler)
115+
MOI.set(model, MOI.RawOptimizerAttribute("return_embedding"), return_embedding)
68116

69117
MOI.optimize!(model)
70118

@@ -100,3 +148,16 @@ Test.@testset "DWave metadata includes chip info" begin
100148
Test.@test chip_info["couplers"] == Any[Any[0, 1], Any[1, 4]]
101149
Test.@test chip_info["num_qubits"] == 3
102150
end
151+
152+
Test.@testset "DWave metadata exposes returned embeddings for supplied samplers" begin
153+
ans = _fake_return_embedding_sampler()
154+
sampleset = _wrapper_dwave_sampleset(ans.sampler; return_embedding = true)
155+
metadata = QUBOTools.metadata(sampleset)
156+
calls = DWave.jl_object(ans.calls)
157+
158+
Test.@test calls[1]["return_embedding"] == true
159+
Test.@test DWave.embedding(sampleset) == Dict(1 => [3, 4], 2 => [5])
160+
Test.@test DWave.embedding(metadata) == Dict(1 => [3, 4], 2 => [5])
161+
Test.@test metadata["dwave_info"]["embedding_context"]["embedding"] ==
162+
Dict(1 => [3, 4], 2 => [5])
163+
end

0 commit comments

Comments
 (0)