Branch data Line data Source code
1 : : /******************************************************************************
2 : : * This file is part of the cvc5 project.
3 : : *
4 : : * Copyright (c) 2009-2026 by the authors listed in the file AUTHORS
5 : : * in the top-level source directory and their institutional affiliations.
6 : : * All rights reserved. See the file COPYING in the top-level source
7 : : * directory for licensing information.
8 : : * ****************************************************************************
9 : : *
10 : : * Model object for the non-linear extension class.
11 : : */
12 : :
13 : : #include "theory/arith/nl/nl_model.h"
14 : :
15 : : #include "expr/node_algorithm.h"
16 : : #include "options/arith_options.h"
17 : : #include "options/smt_options.h"
18 : : #include "options/theory_options.h"
19 : : #include "theory/arith/arith_msum.h"
20 : : #include "theory/arith/arith_utilities.h"
21 : : #include "theory/arith/nl/nl_lemma_utils.h"
22 : : #include "theory/rewriter.h"
23 : : #include "theory/theory_model.h"
24 : :
25 : : using namespace cvc5::internal::kind;
26 : :
27 : : namespace cvc5::internal {
28 : : namespace theory {
29 : : namespace arith {
30 : : namespace nl {
31 : :
32 : 13917 : NlModel::NlModel(Env& env) : EnvObj(env), d_used_approx(false)
33 : : {
34 : 13917 : d_true = nodeManager()->mkConst(true);
35 : 13917 : d_false = nodeManager()->mkConst(false);
36 : 13917 : d_zero = nodeManager()->mkConstReal(Rational(0));
37 : 13917 : d_one = nodeManager()->mkConstReal(Rational(1));
38 : 13917 : d_two = nodeManager()->mkConstReal(Rational(2));
39 : 13917 : }
40 : :
41 : 20287 : NlModel::~NlModel() {}
42 : :
43 : 12663 : void NlModel::reset(const std::map<Node, Node>& arithModel)
44 : : {
45 : 12663 : d_concreteModelCache.clear();
46 : 12663 : d_abstractModelCache.clear();
47 : 12663 : d_arithVal = arithModel;
48 : 12663 : }
49 : :
50 : 11387 : void NlModel::resetCheck()
51 : : {
52 : 11387 : d_used_approx = false;
53 : 11387 : d_check_model_solved.clear();
54 : 11387 : d_check_model_bounds.clear();
55 : 11387 : d_substitutions.clear();
56 : 11387 : }
57 : :
58 : 1364707 : Node NlModel::computeConcreteModelValue(TNode n)
59 : : {
60 : 1364707 : return computeModelValue(n, true);
61 : : }
62 : :
63 : 882640 : Node NlModel::computeAbstractModelValue(TNode n)
64 : : {
65 : 882640 : return computeModelValue(n, false);
66 : : }
67 : :
68 : 6871380 : Node NlModel::computeModelValue(TNode n, bool isConcrete)
69 : : {
70 [ + + ]: 6871380 : auto& cache = isConcrete ? d_concreteModelCache : d_abstractModelCache;
71 [ + + ]: 6871380 : if (auto it = cache.find(n); it != cache.end())
72 : : {
73 : 4204722 : return it->second;
74 : : }
75 [ + - ]: 5333316 : Trace("nl-ext-mv-debug") << "computeModelValue " << n
76 : 2666658 : << ", isConcrete=" << isConcrete << std::endl;
77 : 2666658 : Node ret;
78 [ + + ]: 2666658 : if (n.isConst())
79 : : {
80 : 202533 : ret = n;
81 : : }
82 [ + + ][ + + ]: 2464125 : else if (!isConcrete && hasLinearModelValue(n, ret))
[ + + ][ + + ]
[ - - ]
83 : : {
84 : : // use model value for abstraction
85 : : }
86 [ + + ]: 2312940 : else if (n.getNumChildren() == 0)
87 : : {
88 : : // we are interested in the exact value of PI, which cannot be computed.
89 : : // hence, we return PI itself when asked for the concrete value.
90 [ + + ]: 83346 : if (n.getKind() == Kind::PI)
91 : : {
92 : 1397 : ret = n;
93 : : }
94 : : else
95 : : {
96 : 81949 : ret = getValueInternal(n);
97 : : }
98 : : }
99 : : else
100 : : {
101 : : // otherwise, compute true value
102 : 2229594 : TheoryId ctid = theory::kindToTheoryId(n.getKind());
103 [ + + ][ + + ]: 2229594 : if (ctid != THEORY_ARITH && ctid != THEORY_BOOL && ctid != THEORY_BUILTIN)
[ + + ]
104 : : {
105 : : // we directly look up terms not belonging to arithmetic
106 : 33108 : ret = getValueInternal(n);
107 : : }
108 : : else
109 : : {
110 : 2196486 : std::vector<Node> children;
111 [ + + ]: 2196486 : if (n.getMetaKind() == metakind::PARAMETERIZED)
112 : : {
113 : 727 : children.emplace_back(n.getOperator());
114 : : }
115 [ + + ]: 6219479 : for (size_t i = 0, nchild = n.getNumChildren(); i < nchild; i++)
116 : : {
117 : 4022993 : children.emplace_back(computeModelValue(n[i], isConcrete));
118 : : }
119 : 2196486 : ret = nodeManager()->mkNode(n.getKind(), children);
120 : 2196486 : ret = rewrite(ret);
121 : 2196486 : }
122 : : }
123 [ + - ][ - - ]: 5333316 : Trace("nl-ext-mv-debug") << "computed " << (isConcrete ? "M" : "M_A") << "["
124 : 2666658 : << n << "] = " << ret << std::endl;
125 [ - + ][ - + ]: 7999974 : AssertEqual(n.getType(), ret.getType());
[ - - ]
126 : 2666658 : cache[n] = ret;
127 : 2666658 : return ret;
128 : 2666658 : }
129 : :
130 : 281356 : int NlModel::compare(TNode i, TNode j, bool isConcrete, bool isAbsolute)
131 : : {
132 [ - + ]: 281356 : if (i == j)
133 : : {
134 : 0 : return 0;
135 : : }
136 : 281356 : Node ci = computeModelValue(i, isConcrete);
137 : 281356 : Node cj = computeModelValue(j, isConcrete);
138 [ + - ]: 281356 : if (ci.isConst())
139 : : {
140 [ + - ]: 281356 : if (cj.isConst())
141 : : {
142 : 281356 : return compareValue(ci, cj, isAbsolute);
143 : : }
144 : 0 : return 1;
145 : : }
146 [ - - ]: 0 : return cj.isConst() ? -1 : 0;
147 : 281356 : }
148 : :
149 : 319684 : int NlModel::compareValue(TNode i, TNode j, bool isAbsolute) const
150 : : {
151 [ + - ][ + - ]: 319684 : Assert(i.isConst() && j.isConst());
[ - + ][ - + ]
[ - - ]
152 [ + + ]: 319684 : if (i == j)
153 : : {
154 : 38086 : return 0;
155 : : }
156 [ + + ]: 281598 : if (!isAbsolute)
157 : : {
158 [ + + ]: 599 : return i.getConst<Rational>() < j.getConst<Rational>() ? -1 : 1;
159 : : }
160 : 280999 : Rational iabs = i.getConst<Rational>().abs();
161 : 280999 : Rational jabs = j.getConst<Rational>().abs();
162 [ + + ]: 280999 : if (iabs == jabs)
163 : : {
164 : 21685 : return 0;
165 : : }
166 [ + + ]: 259314 : return iabs < jabs ? -1 : 1;
167 : 280999 : }
168 : :
169 : 663 : bool NlModel::checkModel(const std::vector<Node>& assertions,
170 : : unsigned d,
171 : : std::vector<NlLemma>& lemmas)
172 : : {
173 [ + - ]: 1326 : Trace("nl-ext-cm-debug") << "NlModel::checkModel: solve for equalities..."
174 : 663 : << std::endl;
175 [ + + ]: 16533 : for (const Node& atom : assertions)
176 : : {
177 [ + - ]: 15870 : Trace("nl-ext-cm-debug") << "- assertion: " << atom << std::endl;
178 : : // see if it corresponds to a univariate polynomial equation of degree two
179 [ + + ]: 15870 : if (atom.getKind() == Kind::EQUAL)
180 : : {
181 : : // we substitute inside of solve equality simple
182 [ + + ]: 2468 : if (!solveEqualitySimple(atom, d, lemmas))
183 : : {
184 : : // no chance we will satisfy this equality
185 [ + - ]: 1920 : Trace("nl-ext-cm") << "...check-model : failed to solve equality : "
186 : 960 : << atom << std::endl;
187 : : }
188 : : }
189 : : }
190 : :
191 : : // all remaining variables are constrained to their exact model values
192 [ + - ]: 1326 : Trace("nl-ext-cm-debug") << " set exact bounds for remaining variables..."
193 : 663 : << std::endl;
194 : 663 : std::unordered_set<TNode> visited;
195 : 663 : std::vector<TNode> visit;
196 : 663 : TNode cur;
197 [ + + ]: 16533 : for (const Node& a : assertions)
198 : : {
199 : 15870 : visit.push_back(a);
200 : : do
201 : : {
202 : 77070 : cur = visit.back();
203 : 77070 : visit.pop_back();
204 [ + + ]: 77070 : if (visited.find(cur) == visited.end())
205 : : {
206 : 41453 : visited.insert(cur);
207 [ + + ][ + + ]: 41453 : if (cur.getType().isRealOrInt() && !cur.isConst())
[ + - ][ + + ]
[ - - ]
208 : : {
209 : 12608 : Kind k = cur.getKind();
210 [ + + ][ + + ]: 9378 : if (k != Kind::MULT && k != Kind::ADD && k != Kind::NONLINEAR_MULT
211 [ + + ][ + + ]: 2549 : && k != Kind::TO_REAL && !isTranscendentalKind(k)
212 [ + + ][ + - ]: 21986 : && k != Kind::IAND && k != Kind::PIAND && k != Kind::POW2)
[ + - ][ + - ]
[ + + ]
213 : : {
214 : : // if we have not set an approximate bound for it
215 [ + + ]: 1709 : if (!hasAssignment(cur))
216 : : {
217 : : // set its exact model value in the substitution, if we compute
218 : : // a constant value
219 : 905 : Node curv = computeConcreteModelValue(cur);
220 [ + + ]: 905 : if (curv.isConst())
221 : : {
222 [ - + ]: 899 : if (TraceIsOn("nl-ext-cm"))
223 : : {
224 [ - - ]: 0 : Trace("nl-ext-cm")
225 : 0 : << "check-model-bound : exact : " << cur << " = ";
226 : 0 : printRationalApprox("nl-ext-cm", curv);
227 [ - - ]: 0 : Trace("nl-ext-cm") << std::endl;
228 : : }
229 : 899 : bool ret = addSubstitution(cur, curv);
230 [ - + ][ - + ]: 899 : AlwaysAssert(ret);
[ - - ]
231 : : }
232 : 905 : }
233 : : }
234 : : }
235 : 41453 : visit.insert(visit.end(), cur.begin(), cur.end());
236 : : }
237 [ + + ]: 77070 : } while (!visit.empty());
238 : : }
239 : :
240 [ + - ]: 663 : Trace("nl-ext-cm-debug") << " check assertions..." << std::endl;
241 : 663 : std::vector<Node> check_assertions;
242 [ + + ]: 16533 : for (const Node& a : assertions)
243 : : {
244 [ + + ]: 15870 : if (d_check_model_solved.find(a) == d_check_model_solved.end())
245 : : {
246 : : // apply the substitution to a
247 : 14362 : Node av = getSubstitutedForm(a);
248 [ + - ]: 28724 : Trace("nl-ext-cm") << "simpleCheckModelLit " << av << " (from " << a
249 : 14362 : << ")" << std::endl;
250 : : // simple check literal
251 [ + + ]: 14362 : if (!simpleCheckModelLit(av))
252 : : {
253 [ + - ]: 2646 : Trace("nl-ext-cm") << "...check-model : assertion failed : " << a
254 : 1323 : << std::endl;
255 : 1323 : check_assertions.push_back(av);
256 [ + - ]: 2646 : Trace("nl-ext-cm-debug")
257 : 1323 : << "...check-model : failed assertion, value : " << av << std::endl;
258 : : }
259 : 14362 : }
260 : : }
261 : :
262 [ + + ]: 663 : if (!check_assertions.empty())
263 : : {
264 [ + - ]: 430 : Trace("nl-ext-cm") << "...simple check failed." << std::endl;
265 : : // TODO (#1450) check model for general case
266 : 430 : return false;
267 : : }
268 [ + - ]: 233 : Trace("nl-ext-cm") << "...simple check succeeded!" << std::endl;
269 : 233 : return true;
270 : 663 : }
271 : :
272 : 2645 : bool NlModel::addSubstitution(TNode v, TNode s)
273 : : {
274 [ - + ][ - + ]: 2645 : Assert(v.getKind() != Kind::TO_REAL);
[ - - ]
275 [ + - ]: 5290 : Trace("nl-ext-model") << "* check model substitution : " << v << " -> " << s
276 : 2645 : << std::endl;
277 : 2645 : Assert(getSubstitutedForm(s) == s)
278 : 0 : << "Added a substitution whose range is not in substituted form " << s;
279 : : // cannot substitute real for integer
280 : 2645 : Assert(v.getType().isReal() || s.getType().isInteger());
281 : : // should not substitute the same variable twice
282 : : // should not set exact bound more than once
283 [ - + ]: 2645 : if (d_substitutions.contains(v))
284 : : {
285 : 0 : Node cur = d_substitutions.getSubs(v);
286 [ - - ]: 0 : if (cur != s)
287 : : {
288 [ - - ]: 0 : Trace("nl-ext-model")
289 : 0 : << "...warning: already has value: " << cur << std::endl;
290 : : // We set two different substitutions for a variable v. If both are
291 : : // constant, then we throw an error. Otherwise, we ignore the newer
292 : : // substitution and return false here.
293 : 0 : Assert(!cur.isConst() || !s.isConst())
294 : 0 : << "Conflicting exact bounds given for a variable (" << cur << " and "
295 : 0 : << s << ") for " << v;
296 : 0 : return false;
297 : : }
298 [ - - ]: 0 : }
299 : : // Check if the substitution is cyclic, considering arithmetic subterms.
300 : : // This prevents an assignment like x -> (* 2 x) but allows an assignment
301 : : // like x -> (f x) where f is an uninterpreted function.
302 : 2645 : Node subsFull = d_substitutions.applyArith(s);
303 [ - + ]: 2645 : if (ArithSubs::hasArithSubterm(subsFull, v))
304 : : {
305 [ - - ]: 0 : Trace("nl-ext-model") << "ERROR: has subterm " << subsFull << std::endl;
306 : 0 : return false;
307 : : }
308 : :
309 : : // if we previously had an approximate bound, the exact bound should be in its
310 : : // range
311 : : std::map<Node, std::pair<Node, Node>>::iterator itb =
312 : 2645 : d_check_model_bounds.find(v);
313 [ - + ]: 2645 : if (itb != d_check_model_bounds.end())
314 : : {
315 : 0 : Assert(s.isConst());
316 : 0 : if (s.getConst<Rational>() <= itb->second.first.getConst<Rational>()
317 : 0 : || s.getConst<Rational>() >= itb->second.second.getConst<Rational>())
318 : : {
319 [ - - ]: 0 : Trace("nl-ext-model")
320 : 0 : << "...ERROR: already has bound which is out of range." << std::endl;
321 : 0 : DebugUnhandled()
322 : : << "Out of bounds exact bound given for a variable with an "
323 : 0 : "approximate bound";
324 : : return false;
325 : : }
326 : : }
327 : 2645 : ArithSubs tmp;
328 : 2645 : tmp.addArith(v, s);
329 [ + + ]: 19646 : for (auto& sub : d_substitutions.d_subs)
330 : : {
331 : 17001 : Node ms = tmp.applyArith(sub);
332 [ + + ]: 17001 : if (ms != sub)
333 : : {
334 : 472 : sub = rewrite(ms);
335 : : }
336 : 17001 : }
337 : 2645 : d_substitutions.addArith(v, s);
338 : 2645 : return true;
339 : 2645 : }
340 : :
341 : 1280 : bool NlModel::addBound(TNode v, TNode l, TNode u)
342 : : {
343 [ - + ][ - + ]: 3840 : AssertEqual(l.getType(), v.getType());
[ - - ]
344 [ - + ][ - + ]: 3840 : AssertEqual(u.getType(), v.getType());
[ - - ]
345 [ + - ]: 2560 : Trace("nl-ext-model") << "* check model bound : " << v << " -> [" << l << " "
346 : 1280 : << u << "]" << std::endl;
347 [ + + ]: 1280 : if (l == u)
348 : : {
349 : : // bound is exact, can add as substitution
350 : 378 : return addSubstitution(v, l);
351 : : }
352 : : // should not set a bound for a value that is exact
353 [ - + ]: 902 : if (d_substitutions.contains(v))
354 : : {
355 [ - - ]: 0 : Trace("nl-ext-model")
356 : 0 : << "...ERROR: setting bound for variable that already has exact value."
357 : 0 : << std::endl;
358 : 0 : DebugUnhandled()
359 : 0 : << "Setting bound for variable that already has exact value.";
360 : : return false;
361 : : }
362 [ - + ][ - + ]: 902 : Assert(l.isConst());
[ - - ]
363 [ - + ][ - + ]: 902 : Assert(u.isConst());
[ - - ]
364 [ - + ][ - + ]: 902 : Assert(l.getConst<Rational>() <= u.getConst<Rational>());
[ - - ]
365 : 902 : d_check_model_bounds[v] = std::pair<Node, Node>(l, u);
366 [ - + ]: 902 : if (TraceIsOn("nl-ext-cm"))
367 : : {
368 [ - - ]: 0 : Trace("nl-ext-cm") << "check-model-bound : approximate : ";
369 : 0 : printRationalApprox("nl-ext-cm", l);
370 [ - - ]: 0 : Trace("nl-ext-cm") << " <= " << v << " <= ";
371 : 0 : printRationalApprox("nl-ext-cm", u);
372 [ - - ]: 0 : Trace("nl-ext-cm") << std::endl;
373 : : }
374 : 902 : return true;
375 : : }
376 : :
377 : 522 : void NlModel::setUsedApproximate() { d_used_approx = true; }
378 : :
379 : 65 : bool NlModel::usedApproximate() const { return d_used_approx; }
380 : :
381 : 2730 : bool NlModel::solveEqualitySimple(Node eq,
382 : : unsigned d,
383 : : std::vector<NlLemma>& lemmas)
384 : : {
385 : 2730 : Node seq = eq;
386 [ + + ]: 2730 : if (!d_substitutions.empty())
387 : : {
388 : 2450 : seq = getSubstitutedForm(eq);
389 [ + + ]: 2450 : if (seq.isConst())
390 : : {
391 [ + + ]: 1068 : if (seq.getConst<bool>())
392 : : {
393 : : // already true
394 : 1003 : d_check_model_solved[eq] = Node::null();
395 : 1003 : return true;
396 : : }
397 : 65 : return false;
398 : : }
399 : : }
400 [ + - ]: 1662 : Trace("nl-ext-cms") << "simple solve equality " << seq << "..." << std::endl;
401 [ - + ][ - + ]: 1662 : Assert(seq.getKind() == Kind::EQUAL);
[ - - ]
402 : 1662 : std::map<Node, Node> msum;
403 [ - + ]: 1662 : if (!ArithMSum::getMonomialSumLit(seq, msum))
404 : : {
405 [ - - ]: 0 : Trace("nl-ext-cms") << "...fail, could not determine monomial sum."
406 : 0 : << std::endl;
407 : 0 : return false;
408 : : }
409 : 1662 : bool is_valid = true;
410 : : // the variable we will solve a quadratic equation for
411 : 1662 : Node var;
412 : 1662 : Node b = d_zero;
413 : 1662 : Node c = d_zero;
414 : 1662 : NodeManager* nm = nodeManager();
415 : : // the list of variables that occur as a monomial in msum, and whose value
416 : : // is so far unconstrained in the model.
417 : 1662 : std::unordered_set<Node> unc_vars;
418 : : // the list of variables that occur as a factor in a monomial, and whose
419 : : // value is so far unconstrained in the model.
420 : 1662 : std::unordered_set<Node> unc_vars_factor;
421 [ + + ]: 5655 : for (std::pair<const Node, Node>& m : msum)
422 : : {
423 : 3993 : Node v = m.first;
424 [ + + ]: 3993 : Node coeff = m.second.isNull() ? d_one : m.second;
425 [ + + ]: 3993 : if (v.isNull())
426 : : {
427 : 863 : c = coeff;
428 : : }
429 [ + + ]: 3130 : else if (v.getKind() == Kind::NONLINEAR_MULT)
430 : : {
431 : 537 : is_valid = false;
432 [ + - ]: 1074 : Trace("nl-ext-cms-debug")
433 : 537 : << "...invalid due to non-linear monomial " << v << std::endl;
434 : : // may wish to set an exact bound for a factor and repeat
435 [ + + ]: 1646 : for (const Node& vc : v)
436 : : {
437 : 1109 : unc_vars_factor.insert(vc);
438 : 1109 : }
439 : : }
440 [ + + ][ + + ]: 2593 : else if (!v.isVar() || (!var.isNull() && var != v))
[ + - ][ + + ]
441 : : {
442 [ + - ]: 3138 : Trace("nl-ext-cms-debug")
443 : 1569 : << "...invalid due to factor " << v << std::endl;
444 : : // cannot solve multivariate
445 [ + + ]: 1569 : if (is_valid)
446 : : {
447 : 1074 : is_valid = false;
448 : : // if b is non-zero, then var is also an unconstrained variable
449 [ + + ]: 1074 : if (b != d_zero)
450 : : {
451 : 276 : unc_vars.insert(var);
452 : 276 : unc_vars_factor.insert(var);
453 : : }
454 : : }
455 : : // if v is unconstrained, we may turn this equality into a substitution
456 : 1569 : unc_vars.insert(v);
457 : 1569 : unc_vars_factor.insert(v);
458 : : }
459 : : else
460 : : {
461 : : // set the variable to solve for
462 : 1024 : b = coeff;
463 : 1024 : var = v;
464 : : }
465 : 3993 : }
466 [ + + ]: 1662 : if (!is_valid)
467 : : {
468 : : // see if we can solve for a variable?
469 [ + + ]: 2521 : for (const Node& uv : unc_vars)
470 : : {
471 [ + - ]: 1511 : Trace("nl-ext-cm-debug") << "check subs var : " << uv << std::endl;
472 : : // cannot already have a bound
473 [ + + ][ + - ]: 1511 : if (uv.isVar() && !hasAssignment(uv))
[ + + ][ + + ]
[ - - ]
474 : : {
475 : 462 : Node slv;
476 : 462 : Node veqc;
477 [ + - ]: 462 : if (ArithMSum::isolate(uv, msum, veqc, slv, Kind::EQUAL) != 0)
478 : : {
479 [ - + ][ - + ]: 462 : Assert(!slv.isNull());
[ - - ]
480 : : // must rewrite here to be in substituted form
481 : 462 : slv = rewrite(slv);
482 : : // Currently do not support substitution-with-coefficients.
483 : : // We also ensure types are correct here, which avoids substituting
484 : : // a term of non-integer type for a variable of integer type.
485 : 745 : if (veqc.isNull() && !expr::hasSubterm(slv, uv)
486 [ + + ][ + - ]: 1287 : && CVC5_EQUAL(slv.getType(), uv.getType()))
[ + + ]
487 : : {
488 [ + - ]: 542 : Trace("nl-ext-cm")
489 : 271 : << "check-model-subs : " << uv << " -> " << slv << std::endl;
490 : 271 : bool ret = addSubstitution(uv, slv);
491 [ + - ]: 271 : if (ret)
492 : : {
493 [ + - ]: 542 : Trace("nl-ext-cms") << "...success, model substitution " << uv
494 : 271 : << " -> " << slv << std::endl;
495 : 271 : d_check_model_solved[eq] = uv;
496 : : }
497 : 271 : return ret;
498 : : }
499 : : }
500 [ + + ][ + + ]: 733 : }
501 : : }
502 : : // see if we can assign a variable to a constant
503 [ + + ]: 1950 : for (const Node& uvf : unc_vars_factor)
504 : : {
505 [ + - ]: 1202 : Trace("nl-ext-cm-debug") << "check set var : " << uvf << std::endl;
506 : : // cannot already have a bound
507 [ + + ][ + - ]: 1202 : if (uvf.isVar() && !hasAssignment(uvf))
[ + + ][ + + ]
[ - - ]
508 : : {
509 : 262 : Node uvfv = computeConcreteModelValue(uvf);
510 : : // fail if model value is non-constant
511 [ - + ]: 262 : if (!uvfv.isConst())
512 : : {
513 : 0 : return false;
514 : : }
515 [ - + ]: 262 : if (TraceIsOn("nl-ext-cm"))
516 : : {
517 [ - - ]: 0 : Trace("nl-ext-cm") << "check-model-bound : exact : " << uvf << " = ";
518 : 0 : printRationalApprox("nl-ext-cm", uvfv);
519 [ - - ]: 0 : Trace("nl-ext-cm") << std::endl;
520 : : }
521 : 262 : bool ret = addSubstitution(uvf, uvfv);
522 : : // recurse
523 [ + - ][ + - ]: 262 : return ret ? solveEqualitySimple(eq, d, lemmas) : false;
[ - - ]
524 : 262 : }
525 : : }
526 [ + - ]: 1496 : Trace("nl-ext-cms") << "...fail due to constrained invalid terms."
527 : 748 : << std::endl;
528 : 748 : return false;
529 : : }
530 [ + - ][ + + ]: 381 : else if (var.isNull() || var.getType().isInteger())
[ + - ][ + + ]
[ - - ]
531 : : {
532 : : // cannot solve quadratic equations for integer variables
533 [ + - ]: 147 : Trace("nl-ext-cms") << "...fail due to variable to solve for." << std::endl;
534 : 147 : return false;
535 : : }
536 : :
537 : : // we are linear, it is simple
538 [ - + ]: 234 : if (b == d_zero)
539 : : {
540 [ - - ]: 0 : Trace("nl-ext-cms") << "...fail due to zero a/b." << std::endl;
541 : 0 : DebugUnhandled();
542 : : return false;
543 : : }
544 : 468 : Node val = nm->mkConstReal(-c.getConst<Rational>() / b.getConst<Rational>());
545 [ - + ]: 234 : if (TraceIsOn("nl-ext-cm"))
546 : : {
547 [ - - ]: 0 : Trace("nl-ext-cm") << "check-model-bound : exact : " << var << " = ";
548 : 0 : printRationalApprox("nl-ext-cm", val);
549 [ - - ]: 0 : Trace("nl-ext-cm") << std::endl;
550 : : }
551 : 234 : bool ret = addSubstitution(var, val);
552 [ + - ]: 234 : if (ret)
553 : : {
554 [ + - ]: 234 : Trace("nl-ext-cms") << "...success, solved linear." << std::endl;
555 : 234 : d_check_model_solved[eq] = var;
556 : : }
557 : 234 : return ret;
558 : 2730 : }
559 : :
560 : 15710 : bool NlModel::simpleCheckModelLit(Node lit)
561 : : {
562 [ + - ]: 31420 : Trace("nl-ext-cms") << "*** Simple check-model lit for " << lit << "..."
563 : 15710 : << std::endl;
564 [ + + ]: 15710 : if (lit.isConst())
565 : : {
566 [ + - ]: 10297 : Trace("nl-ext-cms") << " return constant." << std::endl;
567 : 10297 : return lit.getConst<bool>();
568 : : }
569 : 5413 : NodeManager* nm = nodeManager();
570 : 5413 : bool pol = lit.getKind() != Kind::NOT;
571 [ + + ]: 5413 : Node atom = lit.getKind() == Kind::NOT ? lit[0] : lit;
572 : :
573 [ + + ]: 5413 : if (atom.getKind() == Kind::EQUAL)
574 : : {
575 : : // x = a is ( x >= a ^ x <= a )
576 [ + - ]: 1348 : for (unsigned i = 0; i < 2; i++)
577 : : {
578 : 2696 : Node lit2 = nm->mkNode(Kind::GEQ, atom[i], atom[1 - i]);
579 [ + + ]: 1348 : if (!pol)
580 : : {
581 : 1036 : lit2 = lit2.negate();
582 : : }
583 : 1348 : lit2 = rewrite(lit2);
584 : 1348 : bool success = simpleCheckModelLit(lit2);
585 [ + + ]: 1348 : if (success != pol)
586 : : {
587 : : // false != true -> one conjunct of equality is false, we fail
588 : : // true != false -> one disjunct of disequality is true, we succeed
589 : 759 : return success;
590 : : }
591 [ + + ]: 1348 : }
592 : : // both checks passed and polarity is true, or both checks failed and
593 : : // polarity is false
594 : 0 : return pol;
595 : : }
596 [ - + ]: 4654 : else if (atom.getKind() != Kind::GEQ)
597 : : {
598 [ - - ]: 0 : Trace("nl-ext-cms") << " failed due to unknown literal." << std::endl;
599 : 0 : return false;
600 : : }
601 : : // get the monomial sum
602 : 4654 : std::map<Node, Node> msum;
603 [ - + ]: 4654 : if (!ArithMSum::getMonomialSumLit(atom, msum))
604 : : {
605 [ - - ]: 0 : Trace("nl-ext-cms") << " failed due to get msum." << std::endl;
606 : 0 : return false;
607 : : }
608 : : // simple interval analysis
609 [ + + ]: 4654 : if (simpleCheckModelMsum(msum, pol))
610 : : {
611 : 3448 : return true;
612 : : }
613 : : // can also try reasoning about univariate quadratic equations
614 [ + - ]: 2412 : Trace("nl-ext-cms-debug")
615 : 1206 : << "* Try univariate quadratic analysis..." << std::endl;
616 : 1206 : std::vector<Node> vs_invalid;
617 : 1206 : std::unordered_set<Node> vs;
618 : 1206 : std::map<Node, Node> v_a;
619 : 1206 : std::map<Node, Node> v_b;
620 : : // get coefficients...
621 [ + + ]: 3542 : for (std::pair<const Node, Node>& m : msum)
622 : : {
623 : 2336 : Node v = m.first;
624 [ + + ]: 2336 : if (!v.isNull())
625 : : {
626 [ - + ]: 1239 : if (v.isVar())
627 : : {
628 [ - - ]: 0 : v_b[v] = m.second.isNull() ? d_one : m.second;
629 : 0 : vs.insert(v);
630 : : }
631 [ - - ]: 0 : else if (v.getKind() == Kind::NONLINEAR_MULT && v.getNumChildren() == 2
632 : 1239 : && v[0] == v[1] && v[0].isVar())
633 : : {
634 [ - - ]: 0 : v_a[v[0]] = m.second.isNull() ? d_one : m.second;
635 : 0 : vs.insert(v[0]);
636 : : }
637 : : else
638 : : {
639 : 1239 : vs_invalid.push_back(v);
640 : : }
641 : : }
642 : 2336 : }
643 : : // solve the valid variables...
644 : : Node invalid_vsum =
645 : 1206 : vs_invalid.empty()
646 : 0 : ? d_zero
647 : 2379 : : (vs_invalid.size() == 1 ? vs_invalid[0]
648 [ - + ][ + + ]: 3585 : : nm->mkNode(Kind::ADD, vs_invalid));
649 : : // substitution to try
650 : 1206 : ArithSubs qsub;
651 [ - + ]: 1206 : for (const Node& v : vs)
652 : : {
653 : : // is it a valid variable?
654 : : std::map<Node, std::pair<Node, Node>>::iterator bit =
655 : 0 : d_check_model_bounds.find(v);
656 : 0 : if (!expr::hasSubterm(invalid_vsum, v) && bit != d_check_model_bounds.end())
657 : : {
658 : 0 : std::map<Node, Node>::iterator it = v_a.find(v);
659 [ - - ]: 0 : if (it != v_a.end())
660 : : {
661 : 0 : Node a = it->second;
662 : 0 : Assert(a.isConst());
663 : 0 : int asgn = a.getConst<Rational>().sgn();
664 : 0 : Assert(asgn != 0);
665 : 0 : Node t = nm->mkNode(Kind::MULT, a, v, v);
666 : 0 : Node b = d_zero;
667 : 0 : it = v_b.find(v);
668 [ - - ]: 0 : if (it != v_b.end())
669 : : {
670 : 0 : b = it->second;
671 : 0 : t = nm->mkNode(Kind::ADD, t, nm->mkNode(Kind::MULT, b, v));
672 : : }
673 : 0 : t = rewrite(t);
674 [ - - ]: 0 : Trace("nl-ext-cms-debug") << "Trying to find min/max for quadratic "
675 : 0 : << t << "..." << std::endl;
676 [ - - ]: 0 : Trace("nl-ext-cms-debug") << " a = " << a << std::endl;
677 [ - - ]: 0 : Trace("nl-ext-cms-debug") << " b = " << b << std::endl;
678 : : // find maximal/minimal value on the interval
679 : 0 : Node apex = nm->mkNode(
680 : : Kind::DIVISION,
681 : 0 : {nm->mkNode(Kind::NEG, b), nm->mkNode(Kind::MULT, d_two, a)});
682 : 0 : apex = rewrite(apex);
683 : 0 : Assert(apex.isConst());
684 : : // for lower, upper, whether we are greater than the apex
685 : : bool cmp[2];
686 [ - - ]: 0 : Node boundn[2];
687 [ - - ]: 0 : for (unsigned r = 0; r < 2; r++)
688 : : {
689 [ - - ]: 0 : boundn[r] = r == 0 ? bit->second.first : bit->second.second;
690 : 0 : Node cmpn = nm->mkNode(Kind::GT, boundn[r], apex);
691 : 0 : cmpn = rewrite(cmpn);
692 : 0 : Assert(cmpn.isConst());
693 : 0 : cmp[r] = cmpn.getConst<bool>();
694 : 0 : }
695 [ - - ]: 0 : Trace("nl-ext-cms-debug") << " apex " << apex << std::endl;
696 [ - - ]: 0 : Trace("nl-ext-cms-debug")
697 : 0 : << " lower " << boundn[0] << ", cmp: " << cmp[0] << std::endl;
698 [ - - ]: 0 : Trace("nl-ext-cms-debug")
699 : 0 : << " upper " << boundn[1] << ", cmp: " << cmp[1] << std::endl;
700 : 0 : Assert(boundn[0].getConst<Rational>()
701 : : <= boundn[1].getConst<Rational>());
702 : 0 : Node s;
703 : 0 : qsub.addArith(v, Node());
704 [ - - ]: 0 : if (cmp[0] != cmp[1])
705 : : {
706 : 0 : Assert(!cmp[0] && cmp[1]);
707 : : // does the sign match the bound?
708 [ - - ]: 0 : if ((asgn == 1) == pol)
709 : : {
710 : : // the apex is the max/min value
711 : 0 : s = apex;
712 [ - - ]: 0 : Trace("nl-ext-cms-debug") << " ...set to apex." << std::endl;
713 : : }
714 : : else
715 : : {
716 : : // it is one of the endpoints, plug in and compare
717 [ - - ]: 0 : Node tcmpn[2];
718 [ - - ]: 0 : for (unsigned r = 0; r < 2; r++)
719 : : {
720 : 0 : qsub.d_subs.back() = boundn[r];
721 : 0 : Node ts = qsub.applyArith(t);
722 : 0 : tcmpn[r] = rewrite(ts);
723 : 0 : }
724 : 0 : Node tcmp = nm->mkNode(Kind::LT, tcmpn[0], tcmpn[1]);
725 [ - - ]: 0 : Trace("nl-ext-cms-debug")
726 : 0 : << " ...both sides of apex, compare " << tcmp << std::endl;
727 : 0 : tcmp = rewrite(tcmp);
728 : 0 : Assert(tcmp.isConst());
729 [ - - ]: 0 : unsigned bindex_use = (tcmp.getConst<bool>() == pol) ? 1 : 0;
730 [ - - ]: 0 : Trace("nl-ext-cms-debug")
731 [ - - ]: 0 : << " ...set to " << (bindex_use == 1 ? "upper" : "lower")
732 : 0 : << std::endl;
733 : 0 : s = boundn[bindex_use];
734 [ - - ][ - - ]: 0 : }
735 : : }
736 : : else
737 : : {
738 : : // both to one side of the apex
739 : : // we figure out which bound to use (lower or upper) based on
740 : : // three factors:
741 : : // (1) whether a's sign is positive,
742 : : // (2) whether we are greater than the apex of the parabola,
743 : : // (3) the polarity of the constraint, i.e. >= or <=.
744 : : // there are 8 cases of these factors, which we test here.
745 : 0 : unsigned bindex_use = (((asgn == 1) == cmp[0]) == pol) ? 0 : 1;
746 [ - - ]: 0 : Trace("nl-ext-cms-debug")
747 [ - - ]: 0 : << " ...set to " << (bindex_use == 1 ? "upper" : "lower")
748 : 0 : << std::endl;
749 : 0 : s = boundn[bindex_use];
750 : : }
751 : 0 : Assert(!s.isNull());
752 : 0 : qsub.d_subs.back() = s;
753 [ - - ]: 0 : Trace("nl-ext-cms") << "* set bound based on quadratic : " << v
754 : 0 : << " -> " << s << std::endl;
755 [ - - ][ - - ]: 0 : }
756 : : }
757 : : }
758 [ - + ]: 1206 : if (!qsub.empty())
759 : : {
760 : 0 : Node slit = qsub.applyArith(lit);
761 : 0 : slit = rewrite(slit);
762 : 0 : return simpleCheckModelLit(slit);
763 : 0 : }
764 : 1206 : return false;
765 : 5413 : }
766 : :
767 : 4654 : bool NlModel::simpleCheckModelMsum(const std::map<Node, Node>& msum, bool pol)
768 : : {
769 [ + - ]: 4654 : Trace("nl-ext-cms-debug") << "* Try simple interval analysis..." << std::endl;
770 : 4654 : NodeManager* nm = nodeManager();
771 : : // map from transcendental functions to whether they were set to lower
772 : : // bound
773 : 4654 : bool simpleSuccess = true;
774 : 4654 : std::map<Node, bool> set_bound;
775 : 4654 : std::vector<Node> sum_bound;
776 [ + + ]: 13440 : for (const std::pair<const Node, Node>& m : msum)
777 : : {
778 : 8786 : Node v = m.first;
779 [ + + ]: 8786 : if (v.isNull())
780 : : {
781 [ - + ]: 3844 : sum_bound.push_back(m.second.isNull() ? d_one : m.second);
782 : : }
783 : : else
784 : : {
785 [ + - ]: 4942 : Trace("nl-ext-cms-debug") << "- monomial : " << v << std::endl;
786 : : // --- whether we should set a lower bound for this monomial
787 : : bool set_lower =
788 [ + + ][ - + ]: 4942 : (m.second.isNull() || m.second.getConst<Rational>().sgn() == 1)
789 : 4942 : == pol;
790 [ + - ]: 9884 : Trace("nl-ext-cms-debug")
791 [ - - ]: 4942 : << "set bound to " << (set_lower ? "lower" : "upper") << std::endl;
792 : :
793 : : // --- Collect variables and factors in v
794 : 4942 : std::vector<Node> vars;
795 : 4942 : std::vector<unsigned> factors;
796 [ - + ]: 4942 : if (v.getKind() == Kind::NONLINEAR_MULT)
797 : : {
798 : 0 : unsigned last_start = 0;
799 [ - - ]: 0 : for (unsigned i = 0, nchildren = v.getNumChildren(); i < nchildren; i++)
800 : : {
801 : : // are we at the end?
802 : 0 : if (i + 1 == nchildren || v[i + 1] != v[i])
803 : : {
804 : 0 : unsigned vfact = 1 + (i - last_start);
805 : 0 : last_start = (i + 1);
806 : 0 : vars.push_back(v[i]);
807 : 0 : factors.push_back(vfact);
808 : : }
809 : : }
810 : : }
811 : : else
812 : : {
813 : 4942 : vars.push_back(v);
814 : 4942 : factors.push_back(1);
815 : : }
816 : :
817 : : // --- Get the lower and upper bounds and sign information.
818 : : // Whether we have an (odd) number of negative factors in vars, apart
819 : : // from the variable at choose_index.
820 : 4942 : bool has_neg_factor = false;
821 : 4942 : int choose_index = -1;
822 : 4942 : std::vector<Node> ls;
823 : 4942 : std::vector<Node> us;
824 : : // the relevant sign information for variables with odd exponents:
825 : : // 1: both signs of the interval of this variable are positive,
826 : : // -1: both signs of the interval of this variable are negative.
827 : 4942 : std::vector<int> signs;
828 [ + - ]: 4942 : Trace("nl-ext-cms-debug") << "get sign information..." << std::endl;
829 [ + + ]: 9884 : for (unsigned i = 0, size = vars.size(); i < size; i++)
830 : : {
831 : 4942 : Node vc = vars[i];
832 : 4942 : unsigned vcfact = factors[i];
833 [ - + ]: 4942 : if (TraceIsOn("nl-ext-cms-debug"))
834 : : {
835 [ - - ]: 0 : Trace("nl-ext-cms-debug") << "-- " << vc;
836 [ - - ]: 0 : if (vcfact > 1)
837 : : {
838 [ - - ]: 0 : Trace("nl-ext-cms-debug") << "^" << vcfact;
839 : : }
840 [ - - ]: 0 : Trace("nl-ext-cms-debug") << " ";
841 : : }
842 : : std::map<Node, std::pair<Node, Node>>::iterator bit =
843 : 4942 : d_check_model_bounds.find(vc);
844 : : // if there is a model bound for this term
845 [ + - ]: 4942 : if (bit != d_check_model_bounds.end())
846 : : {
847 : 4942 : Node l = bit->second.first;
848 : 4942 : Node u = bit->second.second;
849 : 4942 : ls.push_back(l);
850 : 4942 : us.push_back(u);
851 : 4942 : int vsign = 0;
852 [ + - ]: 4942 : if (vcfact % 2 == 1)
853 : : {
854 : 4942 : vsign = 1;
855 : 4942 : int lsgn = l.getConst<Rational>().sgn();
856 : 4942 : int usgn = u.getConst<Rational>().sgn();
857 [ + - ]: 9884 : Trace("nl-ext-cms-debug")
858 : 4942 : << "bound_sign(" << lsgn << "," << usgn << ") ";
859 [ + + ]: 4942 : if (lsgn == -1)
860 : : {
861 [ + - ]: 60 : if (usgn < 1)
862 : : {
863 : : // must have a negative factor
864 : 60 : has_neg_factor = !has_neg_factor;
865 : 60 : vsign = -1;
866 : : }
867 [ - - ]: 0 : else if (choose_index == -1)
868 : : {
869 : : // set the choose index to this
870 : 0 : choose_index = i;
871 : 0 : vsign = 0;
872 : : }
873 : : else
874 : : {
875 : : // ambiguous, can't determine the bound
876 [ - - ]: 0 : Trace("nl-ext-cms")
877 : 0 : << " failed due to ambiguious monomial." << std::endl;
878 : 0 : return false;
879 : : }
880 : : }
881 : : }
882 [ + - ]: 4942 : Trace("nl-ext-cms-debug") << " -> " << vsign << std::endl;
883 : 4942 : signs.push_back(vsign);
884 [ + - ][ + - ]: 4942 : }
885 : : else
886 : : {
887 [ - - ]: 0 : Trace("nl-ext-cms-debug") << std::endl;
888 [ - - ]: 0 : Trace("nl-ext-cms")
889 : 0 : << " failed due to unknown bound for " << vc << std::endl;
890 : : // should either assign a model bound or eliminate the variable
891 : : // via substitution
892 : 0 : DebugUnhandled() << "A variable " << vc
893 : 0 : << " is missing a bound/value in the model";
894 : : return false;
895 : : }
896 [ + - ]: 4942 : }
897 : : // whether we will try to minimize/maximize (-1/1) the absolute value
898 [ + + ]: 4942 : int setAbs = (set_lower == has_neg_factor) ? 1 : -1;
899 [ + - ]: 9884 : Trace("nl-ext-cms-debug")
900 [ - - ]: 0 : << "set absolute value to " << (setAbs == 1 ? "maximal" : "minimal")
901 : 4942 : << std::endl;
902 : :
903 : 4942 : std::vector<Node> vbs;
904 [ + - ]: 4942 : Trace("nl-ext-cms-debug") << "set bounds..." << std::endl;
905 [ + + ]: 9884 : for (unsigned i = 0, size = vars.size(); i < size; i++)
906 : : {
907 : 4942 : Node vc = vars[i];
908 : 4942 : unsigned vcfact = factors[i];
909 : 4942 : Node l = ls[i];
910 : 4942 : Node u = us[i];
911 : : bool vc_set_lower;
912 : 4942 : int vcsign = signs[i];
913 [ + - ]: 9884 : Trace("nl-ext-cms-debug")
914 : 0 : << "Bounds for " << vc << " : " << l << ", " << u
915 : 4942 : << ", sign : " << vcsign << ", factor : " << vcfact << std::endl;
916 [ - + ]: 4942 : if (l == u)
917 : : {
918 : : // by convention, always say it is lower if they are the same
919 : 0 : vc_set_lower = true;
920 [ - - ]: 0 : Trace("nl-ext-cms-debug")
921 : 0 : << "..." << vc << " equal bound, set to lower" << std::endl;
922 : : }
923 : : else
924 : : {
925 [ - + ]: 4942 : if (vcfact % 2 == 0)
926 : : {
927 : : // minimize or maximize its absolute value
928 : 0 : Rational la = l.getConst<Rational>().abs();
929 : 0 : Rational ua = u.getConst<Rational>().abs();
930 [ - - ]: 0 : if (la == ua)
931 : : {
932 : : // by convention, always say it is lower if abs are the same
933 : 0 : vc_set_lower = true;
934 [ - - ]: 0 : Trace("nl-ext-cms-debug")
935 : 0 : << "..." << vc << " equal abs, set to lower" << std::endl;
936 : : }
937 : : else
938 : : {
939 : 0 : vc_set_lower = (la > ua) == (setAbs == 1);
940 : : }
941 : 0 : }
942 [ - + ]: 4942 : else if (signs[i] == 0)
943 : : {
944 : : // we choose this index to match the overall set_lower
945 : 0 : vc_set_lower = set_lower;
946 : : }
947 : : else
948 : : {
949 : 4942 : vc_set_lower = (signs[i] != setAbs);
950 : : }
951 [ + - ]: 9884 : Trace("nl-ext-cms-debug")
952 [ - - ]: 0 : << "..." << vc << " set to " << (vc_set_lower ? "lower" : "upper")
953 : 4942 : << std::endl;
954 : : }
955 : : // check whether this is a conflicting bound
956 : 4942 : std::map<Node, bool>::iterator itsb = set_bound.find(vc);
957 [ + - ]: 4942 : if (itsb == set_bound.end())
958 : : {
959 : 4942 : set_bound[vc] = vc_set_lower;
960 : : }
961 [ - - ]: 0 : else if (itsb->second != vc_set_lower)
962 : : {
963 [ - - ]: 0 : Trace("nl-ext-cms")
964 : 0 : << " failed due to conflicting bound for " << vc << std::endl;
965 : 0 : return false;
966 : : }
967 : : // must over/under approximate based on vc_set_lower, computed above
968 [ + + ]: 4942 : Node vb = vc_set_lower ? l : u;
969 [ + + ]: 9884 : for (unsigned i2 = 0; i2 < vcfact; i2++)
970 : : {
971 : 4942 : vbs.push_back(vb);
972 : : }
973 [ + - ][ + - ]: 4942 : }
[ + - ]
974 [ - + ]: 4942 : if (!simpleSuccess)
975 : : {
976 : 0 : break;
977 : : }
978 [ + - ]: 4942 : Node vbound = vbs.size() == 1 ? vbs[0] : nm->mkNode(Kind::MULT, vbs);
979 : 4942 : sum_bound.push_back(ArithMSum::mkCoeffTerm(m.second, vbound));
980 [ + - ][ - + ]: 4942 : }
[ - - ][ + - ]
[ - + ][ - - ]
[ + - ][ - + ]
[ - - ]
981 [ + - ][ - ]: 8786 : }
982 : : // if the exact bound was computed via simple analysis above
983 : : // make the bound
984 : 4654 : Node bound;
985 [ + + ]: 4654 : if (sum_bound.size() > 1)
986 : : {
987 : 3916 : bound = nm->mkNode(Kind::ADD, sum_bound);
988 : : }
989 [ + - ]: 738 : else if (sum_bound.size() == 1)
990 : : {
991 : 738 : bound = sum_bound[0];
992 : : }
993 : : else
994 : : {
995 : 0 : bound = d_zero;
996 : : }
997 : : // make the comparison
998 : 9308 : Node comp = nm->mkNode(Kind::GEQ, bound, d_zero);
999 [ + + ]: 4654 : if (!pol)
1000 : : {
1001 : 2722 : comp = comp.negate();
1002 : : }
1003 [ + - ]: 4654 : Trace("nl-ext-cms") << " comparison is : " << comp << std::endl;
1004 : 4654 : comp = rewrite(comp);
1005 [ - + ][ - + ]: 4654 : Assert(comp.isConst());
[ - - ]
1006 [ + - ]: 4654 : Trace("nl-ext-cms") << " returned : " << comp << std::endl;
1007 : 4654 : return comp == d_true;
1008 : 4654 : }
1009 : :
1010 : 96429 : void NlModel::printModelValue(const char* c, Node n, unsigned prec) const
1011 : : {
1012 [ - + ]: 96429 : if (TraceIsOn(c))
1013 : : {
1014 [ - - ]: 0 : Trace(c) << " " << n << " -> ";
1015 : 0 : const Node& aval = d_abstractModelCache.at(n);
1016 [ - - ]: 0 : if (aval.isConst())
1017 : : {
1018 : 0 : printRationalApprox(c, aval, prec);
1019 : : }
1020 : : else
1021 : : {
1022 [ - - ]: 0 : Trace(c) << "?";
1023 : : }
1024 [ - - ]: 0 : Trace(c) << " [actual: ";
1025 : 0 : const Node& cval = d_concreteModelCache.at(n);
1026 [ - - ]: 0 : if (cval.isConst())
1027 : : {
1028 : 0 : printRationalApprox(c, cval, prec);
1029 : : }
1030 : : else
1031 : : {
1032 [ - - ]: 0 : Trace(c) << "?";
1033 : : }
1034 [ - - ]: 0 : Trace(c) << " ]" << std::endl;
1035 : : }
1036 : 96429 : }
1037 : :
1038 : 1334 : void NlModel::getModelValueRepair(std::map<Node, Node>& arithModel)
1039 : : {
1040 : 1334 : NodeManager* nm = nodeManager();
1041 [ + - ]: 1334 : Trace("nl-model") << "NlModel::getModelValueRepair:" << std::endl;
1042 : : // If we extended the model with entries x -> 0 for unconstrained values,
1043 : : // we first update the map to the extended one.
1044 [ + + ]: 1334 : if (d_arithVal.size() > arithModel.size())
1045 : : {
1046 : 148 : arithModel = d_arithVal;
1047 : : }
1048 : : // Record the approximations we used. This code calls the
1049 : : // recordApproximation method of the model, which overrides the model
1050 : : // values for variables that we solved for, using techniques specific to
1051 : : // this class.
1052 : 1334 : for (const std::pair<const Node, std::pair<Node, Node>>& cb :
1053 [ + + ]: 2803 : d_check_model_bounds)
1054 : : {
1055 : 135 : Node l = cb.second.first;
1056 : 135 : Node u = cb.second.second;
1057 : 135 : Node v = cb.first;
1058 [ + - ]: 135 : if (l != u)
1059 : : {
1060 [ + - ]: 270 : Trace("nl-model") << v << " is in interval " << l << "..." << u
1061 : 135 : << std::endl;
1062 : : }
1063 : : else
1064 : : {
1065 : : // overwrite, ensure the type is correct
1066 : 0 : Assert(l.isConst());
1067 : 0 : Node ll = nm->mkConstRealOrInt(v.getType(), l.getConst<Rational>());
1068 : 0 : arithModel[v] = ll;
1069 [ - - ]: 0 : Trace("nl-model") << v << " exact approximation is " << ll << std::endl;
1070 : 0 : }
1071 : 135 : }
1072 : : // Also record the exact values we used. An exact value can be seen as a
1073 : : // special kind approximation of the form (witness x. x = exact_value).
1074 : : // Notice that the above term gets rewritten such that the choice function
1075 : : // is eliminated.
1076 [ + + ]: 2434 : for (size_t i = 0; i < d_substitutions.size(); ++i)
1077 : : {
1078 : : // overwrite, ensure the type is correct
1079 : 1100 : Node v = d_substitutions.d_vars[i];
1080 : 1100 : Node s = d_substitutions.d_subs[i];
1081 : 1100 : Node ss = s;
1082 : : // If its a rational constant, ensure it has the proper type now. It
1083 : : // also may be a RAN, in which case v should be a real.
1084 [ + + ]: 1100 : if (s.isConst())
1085 : : {
1086 : 1009 : ss = nm->mkConstRealOrInt(v.getType(), s.getConst<Rational>());
1087 : : }
1088 : 1100 : arithModel[v] = ss;
1089 [ + - ]: 1100 : Trace("nl-model") << v << " solved is " << ss << std::endl;
1090 : 1100 : }
1091 : :
1092 : : // multiplication terms should not be given values; their values are
1093 : : // implied by the monomials that they consist of
1094 : 1334 : std::vector<Node> amErase;
1095 [ + + ]: 20311 : for (const std::pair<const Node, Node>& am : arithModel)
1096 : : {
1097 [ + + ]: 18977 : if (am.first.getKind() == Kind::NONLINEAR_MULT)
1098 : : {
1099 : 4204 : amErase.push_back(am.first);
1100 : : }
1101 : : }
1102 [ + + ]: 5538 : for (const Node& ae : amErase)
1103 : : {
1104 : 4204 : arithModel.erase(ae);
1105 : : }
1106 : 1334 : }
1107 : :
1108 : 115057 : Node NlModel::getValueInternal(TNode n)
1109 : : {
1110 [ - + ]: 115057 : if (n.isConst())
1111 : : {
1112 : 0 : return n;
1113 : : }
1114 [ + + ]: 115057 : if (auto it = d_arithVal.find(n); it != d_arithVal.end())
1115 : : {
1116 [ - + ][ - + ]: 113709 : AlwaysAssert(it->second.isConst());
[ - - ]
1117 : 113709 : return it->second;
1118 : : }
1119 : : // It is unconstrained in the model, return 0. We additionally add it
1120 : : // to mapping from the linear solver. This ensures that if the nonlinear
1121 : : // solver assumes that n = 0, then this assumption is recorded in the overall
1122 : : // model.
1123 : 1348 : Node zero = mkZero(n.getType());
1124 : 1348 : d_arithVal[n] = zero;
1125 : 1348 : return zero;
1126 : 1348 : }
1127 : :
1128 : 2433 : bool NlModel::hasAssignment(Node v) const
1129 : : {
1130 [ - + ]: 2433 : if (d_check_model_bounds.find(v) != d_check_model_bounds.end())
1131 : : {
1132 : 0 : return true;
1133 : : }
1134 : 2433 : return (d_substitutions.contains(v));
1135 : : }
1136 : :
1137 : 762497 : bool NlModel::hasLinearModelValue(TNode v, Node& val) const
1138 : : {
1139 : 762497 : auto it = d_arithVal.find(v);
1140 [ + + ]: 762497 : if (it != d_arithVal.end())
1141 : : {
1142 : 151185 : val = it->second;
1143 : 151185 : return true;
1144 : : }
1145 : 611312 : return false;
1146 : : }
1147 : :
1148 : 19893 : Node NlModel::getSubstitutedForm(TNode s) const
1149 : : {
1150 [ + + ]: 19893 : if (d_substitutions.empty())
1151 : : {
1152 : : // no substitutions, just return s
1153 : 1850 : return s;
1154 : : }
1155 : 18043 : return rewrite(d_substitutions.applyArith(s));
1156 : : }
1157 : :
1158 : : } // namespace nl
1159 : : } // namespace arith
1160 : : } // namespace theory
1161 : : } // namespace cvc5::internal
|