24 #include "absl/strings/str_format.h"
25 #include "absl/strings/str_join.h"
34 "If true, caching for IntElement is disabled.");
39 void LinkVarExpr(Solver*
const s, IntExpr*
const expr, IntVar*
const var);
46 explicit VectorLess(
const std::vector<T>* values) : values_(values) {}
47 bool operator()(
const T& x,
const T& y)
const {
48 return (*values_)[x] < (*values_)[y];
52 const std::vector<T>* values_;
58 explicit VectorGreater(
const std::vector<T>* values) : values_(values) {}
59 bool operator()(
const T& x,
const T& y)
const {
60 return (*values_)[x] > (*values_)[y];
64 const std::vector<T>* values_;
69 class BaseIntExprElement :
public BaseIntExpr {
71 BaseIntExprElement(Solver*
const s, IntVar*
const e);
72 ~BaseIntExprElement()
override {}
73 int64_t Min()
const override;
74 int64_t Max()
const override;
75 void Range(int64_t* mi, int64_t* ma)
override;
76 void SetMin(int64_t m)
override;
77 void SetMax(int64_t m)
override;
78 void SetRange(int64_t mi, int64_t ma)
override;
79 bool Bound()
const override {
return (
expr_->Bound()); }
81 void WhenRange(Demon* d)
override {
expr_->WhenRange(d); }
84 virtual int64_t ElementValue(
int index)
const = 0;
85 virtual int64_t ExprMin()
const = 0;
86 virtual int64_t ExprMax()
const = 0;
91 void UpdateSupports()
const;
93 void UpdateElementIndexBounds(T check_value) {
94 const int64_t emin = ExprMin();
95 const int64_t emax = ExprMax();
97 int64_t
value = ElementValue(nmin);
98 while (nmin < emax && check_value(
value)) {
100 value = ElementValue(nmin);
102 if (nmin == emax && check_value(
value)) {
106 value = ElementValue(nmax);
107 while (nmax >= nmin && check_value(
value)) {
109 value = ElementValue(nmax);
111 expr_->SetRange(nmin, nmax);
114 mutable int64_t min_;
115 mutable int min_support_;
116 mutable int64_t max_;
117 mutable int max_support_;
118 mutable bool initial_update_;
119 IntVarIterator*
const expr_iterator_;
122 BaseIntExprElement::BaseIntExprElement(Solver*
const s, IntVar*
const e)
129 initial_update_(true),
130 expr_iterator_(
expr_->MakeDomainIterator(true)) {
135 int64_t BaseIntExprElement::Min()
const {
140 int64_t BaseIntExprElement::Max()
const {
151 void BaseIntExprElement::SetMin(int64_t m) {
152 UpdateElementIndexBounds([m](int64_t
value) {
return value < m; });
155 void BaseIntExprElement::SetMax(int64_t m) {
156 UpdateElementIndexBounds([m](int64_t
value) {
return value > m; });
159 void BaseIntExprElement::SetRange(int64_t mi, int64_t ma) {
163 UpdateElementIndexBounds(
164 [mi, ma](int64_t
value) {
return value < mi || value > ma; });
167 void BaseIntExprElement::UpdateSupports()
const {
168 if (initial_update_ || !
expr_->Contains(min_support_) ||
169 !
expr_->Contains(max_support_)) {
170 const int64_t emin = ExprMin();
171 const int64_t emax = ExprMax();
172 int64_t min_value = ElementValue(emax);
173 int64_t max_value = min_value;
174 int min_support = emax;
175 int max_support = emax;
176 const uint64_t expr_size =
expr_->Size();
178 if (expr_size == emax - emin + 1) {
182 if (
value > max_value) {
185 }
else if (
value < min_value) {
191 for (
const int64_t
index : InitAndGetValues(expr_iterator_)) {
194 if (
value > max_value) {
197 }
else if (
value < min_value) {
205 Solver* s = solver();
206 s->SaveAndSetValue(&min_, min_value);
207 s->SaveAndSetValue(&min_support_, min_support);
208 s->SaveAndSetValue(&max_, max_value);
209 s->SaveAndSetValue(&max_support_, max_support);
210 s->SaveAndSetValue(&initial_update_,
false);
219 class IntElementConstraint :
public CastConstraint {
221 IntElementConstraint(Solver*
const s,
const std::vector<int64_t>& values,
222 IntVar*
const index, IntVar*
const elem)
223 : CastConstraint(s, elem),
226 index_iterator_(index_->MakeDomainIterator(true)) {
227 CHECK(
index !=
nullptr);
230 void Post()
override {
232 solver()->MakeDelayedConstraintInitialPropagateCallback(
this);
233 index_->WhenDomain(d);
237 void InitialPropagate()
override {
238 index_->SetRange(0, values_.size() - 1);
241 int64_t new_min = target_var_max;
242 int64_t new_max = target_var_min;
244 for (
const int64_t
index : InitAndGetValues(index_iterator_)) {
246 if (value < target_var_min || value > target_var_max) {
249 if (
value < new_min) {
252 if (
value > new_max) {
263 std::string DebugString()
const override {
264 return absl::StrFormat(
"IntElementConstraint(%s, %s, %s)",
265 absl::StrJoin(values_,
", "), index_->DebugString(),
269 void Accept(ModelVisitor*
const visitor)
const override {
270 visitor->BeginVisitConstraint(ModelVisitor::kElementEqual,
this);
271 visitor->VisitIntegerArrayArgument(ModelVisitor::kValuesArgument, values_);
272 visitor->VisitIntegerExpressionArgument(ModelVisitor::kIndexArgument,
274 visitor->VisitIntegerExpressionArgument(ModelVisitor::kTargetArgument,
276 visitor->EndVisitConstraint(ModelVisitor::kElementEqual,
this);
280 const std::vector<int64_t> values_;
281 IntVar*
const index_;
282 IntVarIterator*
const index_iterator_;
288 IntVar* BuildDomainIntVar(Solver*
const solver, std::vector<int64_t>* values);
290 class IntExprElement :
public BaseIntExprElement {
292 IntExprElement(Solver*
const s,
const std::vector<int64_t>& vals,
294 : BaseIntExprElement(s, expr), values_(vals) {}
296 ~IntExprElement()
override {}
298 std::string
name()
const override {
299 const int size = values_.size();
301 return absl::StrFormat(
"IntElement(array of size %d, %s)", size,
304 return absl::StrFormat(
"IntElement(%s, %s)", absl::StrJoin(values_,
", "),
309 std::string DebugString()
const override {
310 const int size = values_.size();
312 return absl::StrFormat(
"IntElement(array of size %d, %s)", size,
313 expr_->DebugString());
315 return absl::StrFormat(
"IntElement(%s, %s)", absl::StrJoin(values_,
", "),
316 expr_->DebugString());
320 IntVar* CastToVar()
override {
321 Solver*
const s = solver();
322 IntVar*
const var = s->MakeIntVar(values_);
323 s->AddCastConstraint(
324 s->RevAlloc(
new IntElementConstraint(s, values_,
expr_,
var)),
var,
329 void Accept(ModelVisitor*
const visitor)
const override {
330 visitor->BeginVisitIntegerExpression(ModelVisitor::kElement,
this);
331 visitor->VisitIntegerArrayArgument(ModelVisitor::kValuesArgument, values_);
332 visitor->VisitIntegerExpressionArgument(ModelVisitor::kIndexArgument,
334 visitor->EndVisitIntegerExpression(ModelVisitor::kElement,
this);
338 int64_t ElementValue(
int index)
const override {
339 DCHECK_LT(
index, values_.size());
340 return values_[
index];
342 int64_t ExprMin()
const override {
343 return std::max<int64_t>(0,
expr_->Min());
345 int64_t ExprMax()
const override {
346 return values_.empty()
348 : std::min<int64_t>(values_.size() - 1,
expr_->Max());
352 const std::vector<int64_t> values_;
357 class RangeMinimumQueryExprElement :
public BaseIntExpr {
359 RangeMinimumQueryExprElement(Solver* solver,
360 const std::vector<int64_t>& values,
362 ~RangeMinimumQueryExprElement()
override {}
363 int64_t Min()
const override;
364 int64_t Max()
const override;
365 void Range(int64_t* mi, int64_t* ma)
override;
366 void SetMin(int64_t m)
override;
367 void SetMax(int64_t m)
override;
368 void SetRange(int64_t mi, int64_t ma)
override;
369 bool Bound()
const override {
return (index_->Bound()); }
371 void WhenRange(Demon* d)
override { index_->WhenRange(d); }
372 IntVar* CastToVar()
override {
376 IntVar*
const var = solver()->MakeIntVar(min_rmq_.array());
377 solver()->AddCastConstraint(solver()->RevAlloc(
new IntElementConstraint(
378 solver(), min_rmq_.array(), index_,
var)),
382 void Accept(ModelVisitor*
const visitor)
const override {
383 visitor->BeginVisitIntegerExpression(ModelVisitor::kElement,
this);
384 visitor->VisitIntegerArrayArgument(ModelVisitor::kValuesArgument,
386 visitor->VisitIntegerExpressionArgument(ModelVisitor::kIndexArgument,
388 visitor->EndVisitIntegerExpression(ModelVisitor::kElement,
this);
392 int64_t IndexMin()
const {
return std::max<int64_t>(0, index_->Min()); }
393 int64_t IndexMax()
const {
394 return std::min<int64_t>(min_rmq_.array().size() - 1, index_->Max());
397 IntVar*
const index_;
398 const RangeMinimumQuery<int64_t, std::less<int64_t>> min_rmq_;
399 const RangeMinimumQuery<int64_t, std::greater<int64_t>> max_rmq_;
402 RangeMinimumQueryExprElement::RangeMinimumQueryExprElement(
403 Solver* solver,
const std::vector<int64_t>& values, IntVar*
index)
404 : BaseIntExpr(solver), index_(
index), min_rmq_(values), max_rmq_(values) {
405 CHECK(solver !=
nullptr);
406 CHECK(
index !=
nullptr);
409 int64_t RangeMinimumQueryExprElement::Min()
const {
410 return min_rmq_.GetMinimumFromRange(IndexMin(), IndexMax() + 1);
413 int64_t RangeMinimumQueryExprElement::Max()
const {
414 return max_rmq_.GetMinimumFromRange(IndexMin(), IndexMax() + 1);
418 const int64_t range_min = IndexMin();
419 const int64_t range_max = IndexMax() + 1;
420 *mi = min_rmq_.GetMinimumFromRange(range_min, range_max);
421 *ma = max_rmq_.GetMinimumFromRange(range_min, range_max);
424 #define UPDATE_RMQ_BASE_ELEMENT_INDEX_BOUNDS(test) \
425 const std::vector<int64_t>& values = min_rmq_.array(); \
426 int64_t index_min = IndexMin(); \
427 int64_t index_max = IndexMax(); \
428 int64_t value = values[index_min]; \
429 while (index_min < index_max && (test)) { \
431 value = values[index_min]; \
433 if (index_min == index_max && (test)) { \
436 value = values[index_max]; \
437 while (index_max >= index_min && (test)) { \
439 value = values[index_max]; \
441 index_->SetRange(index_min, index_max);
443 void RangeMinimumQueryExprElement::SetMin(int64_t m) {
447 void RangeMinimumQueryExprElement::SetMax(int64_t m) {
451 void RangeMinimumQueryExprElement::SetRange(int64_t mi, int64_t ma) {
458 #undef UPDATE_RMQ_BASE_ELEMENT_INDEX_BOUNDS
462 class IncreasingIntExprElement :
public BaseIntExpr {
464 IncreasingIntExprElement(Solver*
const s,
const std::vector<int64_t>& values,
465 IntVar*
const index);
466 ~IncreasingIntExprElement()
override {}
468 int64_t Min()
const override;
469 void SetMin(int64_t m)
override;
470 int64_t Max()
const override;
471 void SetMax(int64_t m)
override;
472 void SetRange(int64_t mi, int64_t ma)
override;
473 bool Bound()
const override {
return (index_->Bound()); }
475 std::string
name()
const override {
476 return absl::StrFormat(
"IntElement(%s, %s)", absl::StrJoin(values_,
", "),
479 std::string DebugString()
const override {
480 return absl::StrFormat(
"IntElement(%s, %s)", absl::StrJoin(values_,
", "),
481 index_->DebugString());
484 void Accept(ModelVisitor*
const visitor)
const override {
485 visitor->BeginVisitIntegerExpression(ModelVisitor::kElement,
this);
486 visitor->VisitIntegerArrayArgument(ModelVisitor::kValuesArgument, values_);
487 visitor->VisitIntegerExpressionArgument(ModelVisitor::kIndexArgument,
489 visitor->EndVisitIntegerExpression(ModelVisitor::kElement,
this);
492 void WhenRange(Demon* d)
override { index_->WhenRange(d); }
494 IntVar* CastToVar()
override {
495 Solver*
const s = solver();
496 IntVar*
const var = s->MakeIntVar(values_);
502 const std::vector<int64_t> values_;
503 IntVar*
const index_;
506 IncreasingIntExprElement::IncreasingIntExprElement(
507 Solver*
const s,
const std::vector<int64_t>& values, IntVar*
const index)
508 : BaseIntExpr(s), values_(values), index_(
index) {
513 int64_t IncreasingIntExprElement::Min()
const {
514 const int64_t expression_min = std::max<int64_t>(0, index_->Min());
515 return (expression_min < values_.size()
516 ? values_[expression_min]
520 void IncreasingIntExprElement::SetMin(int64_t m) {
521 const int64_t index_min = std::max<int64_t>(0, index_->Min());
522 const int64_t index_max =
523 std::min<int64_t>(values_.size() - 1, index_->Max());
525 if (index_min > index_max || m > values_[index_max]) {
529 const std::vector<int64_t>::const_iterator first =
531 const int64_t new_index_min = first - values_.begin();
532 index_->SetMin(new_index_min);
535 int64_t IncreasingIntExprElement::Max()
const {
536 const int64_t expression_max =
537 std::min<int64_t>(values_.size() - 1, index_->Max());
538 return (expression_max >= 0 ? values_[expression_max]
542 void IncreasingIntExprElement::SetMax(int64_t m) {
543 int64_t index_min = std::max<int64_t>(0, index_->Min());
544 if (m < values_[index_min]) {
548 const std::vector<int64_t>::const_iterator last_after =
550 const int64_t new_index_max = (last_after - values_.begin()) - 1;
551 index_->SetRange(0, new_index_max);
554 void IncreasingIntExprElement::SetRange(int64_t mi, int64_t ma) {
558 const int64_t index_min = std::max<int64_t>(0, index_->Min());
559 const int64_t index_max =
560 std::min<int64_t>(values_.size() - 1, index_->Max());
562 if (mi > ma || ma < values_[index_min] || mi > values_[index_max]) {
566 const std::vector<int64_t>::const_iterator first =
568 const int64_t new_index_min = first - values_.begin();
570 const std::vector<int64_t>::const_iterator last_after =
572 const int64_t new_index_max = (last_after - values_.begin()) - 1;
575 index_->SetRange(new_index_min, new_index_max);
579 IntExpr* BuildElement(Solver*
const solver,
const std::vector<int64_t>& values,
580 IntVar*
const index) {
584 solver->AddConstraint(solver->MakeBetweenCt(
index, 0, values.size() - 1));
585 return solver->MakeIntConst(values[0]);
590 std::vector<int64_t> ones;
592 for (
int i = 0; i < values.size(); ++i) {
593 if (values[i] == 1) {
599 if (ones.size() == 1) {
600 DCHECK_EQ(int64_t{1}, values[ones.back()]);
601 solver->AddConstraint(solver->MakeBetweenCt(
index, 0, values.size() - 1));
602 return solver->MakeIsEqualCstVar(
index, ones.back());
603 }
else if (ones.size() == values.size() - 1) {
604 solver->AddConstraint(solver->MakeBetweenCt(
index, 0, values.size() - 1));
605 return solver->MakeIsDifferentCstVar(
index, first_zero);
606 }
else if (ones.size() == ones.back() - ones.front() + 1) {
607 solver->AddConstraint(solver->MakeBetweenCt(
index, 0, values.size() - 1));
608 IntVar*
const b = solver->MakeBoolVar(
"ContiguousBooleanElementVar");
609 solver->AddConstraint(
610 solver->MakeIsBetweenCt(
index, ones.front(), ones.back(),
b));
613 IntVar*
const b = solver->MakeBoolVar(
"NonContiguousBooleanElementVar");
614 solver->AddConstraint(solver->MakeBetweenCt(
index, 0, values.size() - 1));
615 solver->AddConstraint(solver->MakeIsMemberCt(
index, ones,
b));
619 IntExpr* cache =
nullptr;
620 if (!absl::GetFlag(FLAGS_cp_disable_element_cache)) {
621 cache = solver->Cache()->FindVarConstantArrayExpression(
622 index, values, ModelCache::VAR_CONSTANT_ARRAY_ELEMENT);
624 if (cache !=
nullptr) {
627 IntExpr* result =
nullptr;
628 if (values.size() >= 2 &&
index->Min() == 0 &&
index->Max() == 1) {
629 result = solver->MakeSum(solver->MakeProd(
index, values[1] - values[0]),
631 }
else if (values.size() == 2 &&
index->Contains(0) &&
index->Contains(1)) {
632 solver->AddConstraint(solver->MakeBetweenCt(
index, 0, 1));
633 result = solver->MakeSum(solver->MakeProd(
index, values[1] - values[0]),
636 result = solver->MakeSum(
index, values[0]);
638 result = solver->RegisterIntExpr(solver->RevAlloc(
639 new IncreasingIntExprElement(solver, values,
index)));
641 if (solver->parameters().use_element_rmq()) {
642 result = solver->RegisterIntExpr(solver->RevAlloc(
643 new RangeMinimumQueryExprElement(solver, values,
index)));
645 result = solver->RegisterIntExpr(
646 solver->RevAlloc(
new IntExprElement(solver, values,
index)));
649 if (!absl::GetFlag(FLAGS_cp_disable_element_cache)) {
650 solver->Cache()->InsertVarConstantArrayExpression(
651 result,
index, values, ModelCache::VAR_CONSTANT_ARRAY_ELEMENT);
658 IntExpr* Solver::MakeElement(
const std::vector<int64_t>& values,
661 DCHECK_EQ(
this,
index->solver());
662 if (
index->Bound()) {
663 return MakeIntConst(values[
index->Min()]);
665 return BuildElement(
this, values,
index);
668 IntExpr* Solver::MakeElement(
const std::vector<int>& values,
671 DCHECK_EQ(
this,
index->solver());
672 if (
index->Bound()) {
673 return MakeIntConst(values[
index->Min()]);
681 class IntExprFunctionElement :
public BaseIntExprElement {
685 ~IntExprFunctionElement()
override;
687 std::string
name()
const override {
688 return absl::StrFormat(
"IntFunctionElement(%s)",
expr_->name());
691 std::string DebugString()
const override {
692 return absl::StrFormat(
"IntFunctionElement(%s)",
expr_->DebugString());
695 void Accept(ModelVisitor*
const visitor)
const override {
697 visitor->BeginVisitIntegerExpression(ModelVisitor::kElement,
this);
698 visitor->VisitIntegerExpressionArgument(ModelVisitor::kIndexArgument,
700 visitor->VisitInt64ToInt64Extension(values_,
expr_->Min(),
expr_->Max());
701 visitor->EndVisitIntegerExpression(ModelVisitor::kElement,
this);
705 int64_t ElementValue(
int index)
const override {
return values_(
index); }
706 int64_t ExprMin()
const override {
return expr_->Min(); }
707 int64_t ExprMax()
const override {
return expr_->Max(); }
710 Solver::IndexEvaluator1 values_;
713 IntExprFunctionElement::IntExprFunctionElement(Solver*
const s,
714 Solver::IndexEvaluator1 values,
716 : BaseIntExprElement(s, e), values_(std::move(values)) {
717 CHECK(values_ !=
nullptr);
720 IntExprFunctionElement::~IntExprFunctionElement() {}
724 class IncreasingIntExprFunctionElement :
public BaseIntExpr {
726 IncreasingIntExprFunctionElement(Solver*
const s,
729 : BaseIntExpr(s), values_(std::move(values)), index_(
index) {
730 DCHECK(values_ !=
nullptr);
735 ~IncreasingIntExprFunctionElement()
override {}
737 int64_t Min()
const override {
return values_(index_->Min()); }
739 void SetMin(int64_t m)
override {
740 const int64_t index_min = index_->Min();
741 const int64_t index_max = index_->Max();
742 if (m > values_(index_max)) {
745 const int64_t new_index_min = FindNewIndexMin(index_min, index_max, m);
746 index_->SetMin(new_index_min);
749 int64_t Max()
const override {
return values_(index_->Max()); }
751 void SetMax(int64_t m)
override {
752 int64_t index_min = index_->Min();
753 int64_t index_max = index_->Max();
754 if (m < values_(index_min)) {
757 const int64_t new_index_max = FindNewIndexMax(index_min, index_max, m);
758 index_->SetMax(new_index_max);
761 void SetRange(int64_t mi, int64_t ma)
override {
762 const int64_t index_min = index_->Min();
763 const int64_t index_max = index_->Max();
764 const int64_t value_min = values_(index_min);
765 const int64_t value_max = values_(index_max);
766 if (mi > ma || ma < value_min || mi > value_max) {
769 if (mi <= value_min && ma >= value_max) {
774 const int64_t new_index_min = FindNewIndexMin(index_min, index_max, mi);
775 const int64_t new_index_max = FindNewIndexMax(new_index_min, index_max, ma);
777 index_->SetRange(new_index_min, new_index_max);
780 std::string
name()
const override {
781 return absl::StrFormat(
"IncreasingIntExprFunctionElement(values, %s)",
785 std::string DebugString()
const override {
786 return absl::StrFormat(
"IncreasingIntExprFunctionElement(values, %s)",
787 index_->DebugString());
790 void WhenRange(Demon* d)
override { index_->WhenRange(d); }
792 void Accept(ModelVisitor*
const visitor)
const override {
797 if (index_->Min() == 0) {
801 visitor->VisitInt64ToInt64Extension(values_, index_->Min(),
808 int64_t FindNewIndexMin(int64_t index_min, int64_t index_max, int64_t m) {
809 if (m <= values_(index_min)) {
813 DCHECK_LT(values_(index_min), m);
814 DCHECK_GE(values_(index_max), m);
816 int64_t index_lower_bound = index_min;
817 int64_t index_upper_bound = index_max;
818 while (index_upper_bound - index_lower_bound > 1) {
819 DCHECK_LT(values_(index_lower_bound), m);
820 DCHECK_GE(values_(index_upper_bound), m);
821 const int64_t pivot = (index_lower_bound + index_upper_bound) / 2;
822 const int64_t pivot_value = values_(pivot);
823 if (pivot_value < m) {
824 index_lower_bound = pivot;
826 index_upper_bound = pivot;
829 DCHECK(values_(index_upper_bound) >= m);
830 return index_upper_bound;
833 int64_t FindNewIndexMax(int64_t index_min, int64_t index_max, int64_t m) {
834 if (m >= values_(index_max)) {
838 DCHECK_LE(values_(index_min), m);
839 DCHECK_GT(values_(index_max), m);
841 int64_t index_lower_bound = index_min;
842 int64_t index_upper_bound = index_max;
843 while (index_upper_bound - index_lower_bound > 1) {
844 DCHECK_LE(values_(index_lower_bound), m);
845 DCHECK_GT(values_(index_upper_bound), m);
846 const int64_t pivot = (index_lower_bound + index_upper_bound) / 2;
847 const int64_t pivot_value = values_(pivot);
848 if (pivot_value > m) {
849 index_upper_bound = pivot;
851 index_lower_bound = pivot;
854 DCHECK(values_(index_lower_bound) <= m);
855 return index_lower_bound;
859 IntVar*
const index_;
865 CHECK_EQ(
this,
index->solver());
867 RevAlloc(
new IntExprFunctionElement(
this, std::move(values),
index)));
872 CHECK_EQ(
this,
index->solver());
875 RevAlloc(
new IncreasingIntExprFunctionElement(
this, values,
index)));
883 new IncreasingIntExprFunctionElement(
this, opposite_values,
index))));
890 class IntIntExprFunctionElement :
public BaseIntExpr {
894 ~IntIntExprFunctionElement()
override;
895 std::string DebugString()
const override {
896 return absl::StrFormat(
"IntIntFunctionElement(%s,%s)",
897 expr1_->DebugString(), expr2_->DebugString());
899 int64_t Min()
const override;
900 int64_t Max()
const override;
905 bool Bound()
const override {
return expr1_->Bound() && expr2_->Bound(); }
907 void WhenRange(Demon* d)
override {
908 expr1_->WhenRange(d);
909 expr2_->WhenRange(d);
912 void Accept(ModelVisitor*
const visitor)
const override {
919 const int64_t expr1_min = expr1_->Min();
920 const int64_t expr1_max = expr1_->Max();
923 for (
int i = expr1_min; i <= expr1_max; ++i) {
924 visitor->VisitInt64ToInt64Extension(
925 [
this, i](int64_t j) {
return values_(i, j); }, expr2_->Min(),
932 int64_t ElementValue(
int index1,
int index2)
const {
933 return values_(index1, index2);
935 void UpdateSupports()
const;
937 IntVar*
const expr1_;
938 IntVar*
const expr2_;
939 mutable int64_t min_;
940 mutable int min_support1_;
941 mutable int min_support2_;
942 mutable int64_t max_;
943 mutable int max_support1_;
944 mutable int max_support2_;
945 mutable bool initial_update_;
947 IntVarIterator*
const expr1_iterator_;
948 IntVarIterator*
const expr2_iterator_;
951 IntIntExprFunctionElement::IntIntExprFunctionElement(
963 initial_update_(true),
964 values_(std::move(values)),
965 expr1_iterator_(expr1_->MakeDomainIterator(true)),
966 expr2_iterator_(expr2_->MakeDomainIterator(true)) {
967 CHECK(values_ !=
nullptr);
970 IntIntExprFunctionElement::~IntIntExprFunctionElement() {}
972 int64_t IntIntExprFunctionElement::Min()
const {
977 int64_t IntIntExprFunctionElement::Max()
const {
989 #define UPDATE_ELEMENT_INDEX_BOUNDS(test) \
990 const int64_t emin1 = expr1_->Min(); \
991 const int64_t emax1 = expr1_->Max(); \
992 const int64_t emin2 = expr2_->Min(); \
993 const int64_t emax2 = expr2_->Max(); \
994 int64_t nmin1 = emin1; \
995 bool found = false; \
996 while (nmin1 <= emax1 && !found) { \
997 for (int i = emin2; i <= emax2; ++i) { \
998 int64_t value = ElementValue(nmin1, i); \
1008 if (nmin1 > emax1) { \
1011 int64_t nmin2 = emin2; \
1013 while (nmin2 <= emax2 && !found) { \
1014 for (int i = emin1; i <= emax1; ++i) { \
1015 int64_t value = ElementValue(i, nmin2); \
1025 if (nmin2 > emax2) { \
1028 int64_t nmax1 = emax1; \
1030 while (nmax1 >= nmin1 && !found) { \
1031 for (int i = emin2; i <= emax2; ++i) { \
1032 int64_t value = ElementValue(nmax1, i); \
1042 int64_t nmax2 = emax2; \
1044 while (nmax2 >= nmin2 && !found) { \
1045 for (int i = emin1; i <= emax1; ++i) { \
1046 int64_t value = ElementValue(i, nmax2); \
1056 expr1_->SetRange(nmin1, nmax1); \
1057 expr2_->SetRange(nmin2, nmax2);
1059 void IntIntExprFunctionElement::SetMin(int64_t
lower_bound) {
1063 void IntIntExprFunctionElement::SetMax(int64_t
upper_bound) {
1067 void IntIntExprFunctionElement::SetRange(int64_t
lower_bound,
1075 #undef UPDATE_ELEMENT_INDEX_BOUNDS
1077 void IntIntExprFunctionElement::UpdateSupports()
const {
1078 if (initial_update_ || !expr1_->
Contains(min_support1_) ||
1080 !expr2_->
Contains(max_support2_)) {
1081 const int64_t emax1 = expr1_->
Max();
1082 const int64_t emax2 = expr2_->
Max();
1083 int64_t min_value = ElementValue(emax1, emax2);
1084 int64_t max_value = min_value;
1085 int min_support1 = emax1;
1086 int max_support1 = emax1;
1087 int min_support2 = emax2;
1088 int max_support2 = emax2;
1089 for (
const int64_t index1 : InitAndGetValues(expr1_iterator_)) {
1090 for (
const int64_t index2 : InitAndGetValues(expr2_iterator_)) {
1091 const int64_t
value = ElementValue(index1, index2);
1092 if (
value > max_value) {
1094 max_support1 = index1;
1095 max_support2 = index2;
1096 }
else if (
value < min_value) {
1098 min_support1 = index1;
1099 min_support2 = index2;
1103 Solver* s = solver();
1104 s->SaveAndSetValue(&min_, min_value);
1105 s->SaveAndSetValue(&min_support1_, min_support1);
1106 s->SaveAndSetValue(&min_support2_, min_support2);
1107 s->SaveAndSetValue(&max_, max_value);
1108 s->SaveAndSetValue(&max_support1_, max_support1);
1109 s->SaveAndSetValue(&max_support2_, max_support2);
1110 s->SaveAndSetValue(&initial_update_,
false);
1117 CHECK_EQ(
this, index1->
solver());
1118 CHECK_EQ(
this, index2->
solver());
1120 new IntIntExprFunctionElement(
this, std::move(values), index1, index2)));
1132 condition_(condition),
1152 if (condition_->
Max() == 0) {
1153 zero_->
SetRange(target_var_min, target_var_max);
1154 zero_->
Range(&new_min, &new_max);
1155 }
else if (condition_->
Min() == 1) {
1156 one_->
SetRange(target_var_min, target_var_max);
1157 one_->
Range(&new_min, &new_max);
1159 if (target_var_max < zero_->Min() || target_var_min > zero_->
Max()) {
1161 one_->
SetRange(target_var_min, target_var_max);
1162 one_->
Range(&new_min, &new_max);
1163 }
else if (target_var_max < one_->Min() || target_var_min > one_->
Max()) {
1165 zero_->
SetRange(target_var_min, target_var_max);
1166 zero_->
Range(&new_min, &new_max);
1172 zero_->
Range(&zl, &zu);
1173 one_->
Range(&ol, &ou);
1182 return absl::StrFormat(
"(%s ? %s : %s) == %s", condition_->
DebugString(),
1190 IntVar*
const condition_;
1202 class IntExprEvaluatorElementCt :
public CastConstraint {
1205 int64_t range_start, int64_t range_end,
1206 IntVar*
const index, IntVar*
const target_var);
1207 ~IntExprEvaluatorElementCt()
override {}
1209 void Post()
override;
1210 void InitialPropagate()
override;
1213 void Update(
int index);
1216 std::string DebugString()
const override;
1217 void Accept(ModelVisitor*
const visitor)
const override;
1220 IntVar*
const index_;
1224 const int64_t range_start_;
1225 const int64_t range_end_;
1230 IntExprEvaluatorElementCt::IntExprEvaluatorElementCt(
1232 int64_t range_end, IntVar*
const index, IntVar*
const target_var)
1233 : CastConstraint(s, target_var),
1236 range_start_(range_start),
1237 range_end_(range_end),
1241 void IntExprEvaluatorElementCt::Post() {
1243 solver(),
this, &IntExprEvaluatorElementCt::Propagate,
"Propagate");
1244 for (
int i = range_start_; i < range_end_; ++i) {
1246 current_var->WhenRange(delayed_propagate_demon);
1248 solver(),
this, &IntExprEvaluatorElementCt::Update,
"Update", i);
1249 current_var->WhenRange(update_demon);
1251 index_->
WhenRange(delayed_propagate_demon);
1253 solver(),
this, &IntExprEvaluatorElementCt::UpdateExpr,
"UpdateExpr");
1256 solver(),
this, &IntExprEvaluatorElementCt::Propagate,
"UpdateVar");
1261 void IntExprEvaluatorElementCt::InitialPropagate() { Propagate(); }
1263 void IntExprEvaluatorElementCt::Propagate() {
1264 const int64_t emin =
std::max(range_start_, index_->
Min());
1265 const int64_t emax = std::min<int64_t>(range_end_ - 1, index_->
Max());
1272 int64_t nmin = emin;
1273 for (; nmin <= emax; nmin++) {
1278 if (nmin_var->Min() <= vmax && nmin_var->Max() >= vmin)
break;
1280 int64_t nmax = emax;
1281 for (; nmin <= nmax; nmax--) {
1286 if (nmax_var->Min() <= vmax && nmax_var->Max() >= vmin)
break;
1293 if (min_support_ == -1 || max_support_ == -1) {
1294 int min_support = -1;
1295 int max_support = -1;
1298 for (
int i = index_->
Min(); i <= index_->Max(); ++i) {
1300 const int64_t vmin = var_i->Min();
1304 const int64_t vmax = var_i->Max();
1309 solver()->SaveAndSetValue(&min_support_, min_support);
1310 solver()->SaveAndSetValue(&max_support_, max_support);
1315 void IntExprEvaluatorElementCt::Update(
int index) {
1316 if (
index == min_support_ ||
index == max_support_) {
1317 solver()->SaveAndSetValue(&min_support_, -1);
1318 solver()->SaveAndSetValue(&max_support_, -1);
1322 void IntExprEvaluatorElementCt::UpdateExpr() {
1324 solver()->SaveAndSetValue(&min_support_, -1);
1325 solver()->SaveAndSetValue(&max_support_, -1);
1331 int64_t range_start, int64_t range_end) {
1333 for (int64_t i = range_start; i < range_end; ++i) {
1334 if (i != range_start) {
1337 out += absl::StrFormat(
"%d -> %s", i, evaluator(i)->DebugString());
1343 int64_t range_begin, int64_t range_end) {
1345 if (range_end - range_begin > 10) {
1346 out = absl::StrFormat(
1347 "IntToIntVar(%s, ...%s)",
1348 StringifyEvaluatorBare(evaluator, range_begin, range_begin + 5),
1349 StringifyEvaluatorBare(evaluator, range_end - 5, range_end));
1351 out = absl::StrFormat(
1353 StringifyEvaluatorBare(evaluator, range_begin, range_end));
1359 std::string IntExprEvaluatorElementCt::DebugString()
const {
1360 return StringifyInt64ToIntVar(
evaluator_, range_start_, range_end_);
1363 void IntExprEvaluatorElementCt::Accept(ModelVisitor*
const visitor)
const {
1365 visitor->VisitIntegerVariableEvaluatorArgument(
1378 class IntExprArrayElementCt :
public IntExprEvaluatorElementCt {
1380 IntExprArrayElementCt(Solver*
const s, std::vector<IntVar*> vars,
1381 IntVar*
const index, IntVar*
const target_var);
1383 std::string DebugString()
const override;
1384 void Accept(ModelVisitor*
const visitor)
const override;
1387 const std::vector<IntVar*>
vars_;
1390 IntExprArrayElementCt::IntExprArrayElementCt(Solver*
const s,
1391 std::vector<IntVar*> vars,
1392 IntVar*
const index,
1393 IntVar*
const target_var)
1394 : IntExprEvaluatorElementCt(
1395 s, [this](int64_t idx) {
return vars_[idx]; }, 0, vars.size(),
index,
1397 vars_(std::move(vars)) {}
1399 std::string IntExprArrayElementCt::DebugString()
const {
1400 int64_t size =
vars_.size();
1402 return absl::StrFormat(
1403 "IntExprArrayElement(var array of size %d, %s) == %s", size,
1406 return absl::StrFormat(
"IntExprArrayElement([%s], %s) == %s",
1412 void IntExprArrayElementCt::Accept(ModelVisitor*
const visitor)
const {
1426 class IntExprArrayElementCstCt :
public Constraint {
1428 IntExprArrayElementCstCt(Solver*
const s,
const std::vector<IntVar*>& vars,
1429 IntVar*
const index, int64_t target)
1434 demons_(vars.size()) {}
1436 ~IntExprArrayElementCstCt()
override {}
1438 void Post()
override {
1439 for (
int i = 0; i <
vars_.size(); ++i) {
1441 solver(),
this, &IntExprArrayElementCstCt::Propagate,
"Propagate", i);
1442 vars_[i]->WhenDomain(demons_[i]);
1445 solver(),
this, &IntExprArrayElementCstCt::PropagateIndex,
1447 index_->WhenBound(index_demon);
1450 void InitialPropagate()
override {
1451 for (
int i = 0; i <
vars_.size(); ++i) {
1457 void Propagate(
int index) {
1458 if (!vars_[
index]->Contains(target_)) {
1459 index_->RemoveValue(
index);
1460 demons_[
index]->inhibit(solver());
1464 void PropagateIndex() {
1465 if (index_->Bound()) {
1466 vars_[index_->Min()]->SetValue(target_);
1470 std::string DebugString()
const override {
1471 return absl::StrFormat(
"IntExprArrayElement([%s], %s) == %d",
1473 index_->DebugString(), target_);
1476 void Accept(ModelVisitor*
const visitor)
const override {
1487 const std::vector<IntVar*>
vars_;
1488 IntVar*
const index_;
1489 const int64_t target_;
1490 std::vector<Demon*> demons_;
1495 class IntExprIndexOfCt :
public Constraint {
1497 IntExprIndexOfCt(Solver*
const s,
const std::vector<IntVar*>& vars,
1498 IntVar*
const index, int64_t target)
1503 demons_(
vars_.size()),
1504 index_iterator_(
index->MakeHoleIterator(true)) {}
1506 ~IntExprIndexOfCt()
override {}
1508 void Post()
override {
1509 for (
int i = 0; i <
vars_.size(); ++i) {
1511 solver(),
this, &IntExprIndexOfCt::Propagate,
"Propagate", i);
1512 vars_[i]->WhenDomain(demons_[i]);
1515 solver(),
this, &IntExprIndexOfCt::PropagateIndex,
"PropagateIndex");
1516 index_->WhenDomain(index_demon);
1519 void InitialPropagate()
override {
1520 for (
int i = 0; i <
vars_.size(); ++i) {
1521 if (!index_->Contains(i)) {
1522 vars_[i]->RemoveValue(target_);
1523 }
else if (!vars_[i]->Contains(target_)) {
1524 index_->RemoveValue(i);
1525 demons_[i]->inhibit(solver());
1526 }
else if (vars_[i]->Bound()) {
1527 index_->SetValue(i);
1528 demons_[i]->inhibit(solver());
1533 void Propagate(
int index) {
1534 if (!vars_[
index]->Contains(target_)) {
1535 index_->RemoveValue(
index);
1536 demons_[
index]->inhibit(solver());
1537 }
else if (vars_[
index]->Bound()) {
1538 index_->SetValue(
index);
1542 void PropagateIndex() {
1543 const int64_t oldmax = index_->OldMax();
1544 const int64_t vmin = index_->Min();
1545 const int64_t vmax = index_->Max();
1548 demons_[
value]->inhibit(solver());
1550 for (
const int64_t
value : InitAndGetValues(index_iterator_)) {
1552 demons_[
value]->inhibit(solver());
1556 demons_[
value]->inhibit(solver());
1558 if (index_->Bound()) {
1559 vars_[index_->Min()]->SetValue(target_);
1563 std::string DebugString()
const override {
1564 return absl::StrFormat(
"IntExprIndexOf([%s], %s) == %d",
1566 index_->DebugString(), target_);
1569 void Accept(ModelVisitor*
const visitor)
const override {
1580 const std::vector<IntVar*>
vars_;
1581 IntVar*
const index_;
1582 const int64_t target_;
1583 std::vector<Demon*> demons_;
1584 IntVarIterator*
const index_iterator_;
1589 Constraint* MakeElementEqualityFunc(Solver*
const solver,
1590 const std::vector<int64_t>& vals,
1591 IntVar*
const index, IntVar*
const target) {
1592 if (
index->Bound()) {
1593 const int64_t val =
index->Min();
1594 if (val < 0 || val >= vals.size()) {
1595 return solver->MakeFalseConstraint();
1597 return solver->MakeEquality(target, vals[val]);
1601 return solver->MakeEquality(target, solver->MakeSum(
index, vals[0]));
1603 return solver->RevAlloc(
1604 new IntElementConstraint(solver, vals,
index, target));
1613 IntVar*
const target_var) {
1615 new IfThenElseCt(
this, condition, then_expr, else_expr, target_var));
1620 if (
index->Bound()) {
1621 return vars[
index->Min()];
1623 const int size = vars.size();
1625 std::vector<int64_t> values(size);
1626 for (
int i = 0; i < size; ++i) {
1627 values[i] = vars[i]->Value();
1632 index->Min() >= 0 &&
index->Max() < vars.size()) {
1637 const std::string
name = absl::StrFormat(
1647 std::unique_ptr<IntVarIterator> iterator(
index->MakeDomainIterator(
false));
1649 if (index_value >= 0 && index_value < size) {
1650 emin =
std::min(emin, vars[index_value]->Min());
1651 emax =
std::max(emax, vars[index_value]->Max());
1654 const std::string vname =
1655 size > 10 ? absl::StrFormat(
"ElementVar(var array of size %d, %s)", size,
1656 index->DebugString())
1657 : absl::StrFormat(
"ElementVar([%s], %s)",
1661 RevAlloc(
new IntExprArrayElementCt(
this, vars,
index, element_var)));
1666 int64_t range_end,
IntVar* argument) {
1667 const std::string index_name =
1669 const std::string vname = absl::StrFormat(
1670 "ElementVar(%s, %s)",
1671 StringifyInt64ToIntVar(vars, range_start, range_end), index_name);
1672 IntVar*
const element_var =
1675 IntExprEvaluatorElementCt* evaluation_ct =
new IntExprEvaluatorElementCt(
1676 this, std::move(vars), range_start, range_end, argument, element_var);
1678 evaluation_ct->Propagate();
1685 return MakeElementEqualityFunc(
this, vals,
index, target);
1698 std::vector<int64_t> values(vars.size());
1699 for (
int i = 0; i < vars.size(); ++i) {
1700 values[i] = vars[i]->Value();
1704 if (
index->Bound()) {
1705 const int64_t val =
index->Min();
1706 if (val < 0 || val >= vars.size()) {
1712 if (target->
Bound()) {
1714 new IntExprArrayElementCstCt(
this, vars,
index, target->
Min()));
1716 return RevAlloc(
new IntExprArrayElementCt(
this, vars,
index, target));
1724 std::vector<int> valid_indices;
1725 for (
int i = 0; i < vars.size(); ++i) {
1726 if (vars[i]->
Value() == target) {
1727 valid_indices.push_back(i);
1732 if (
index->Bound()) {
1733 const int64_t pos =
index->Min();
1734 if (pos >= 0 && pos < vars.size()) {
1741 return RevAlloc(
new IntExprArrayElementCstCt(
this, vars,
index, target));
1747 if (
index->Bound()) {
1748 const int64_t pos =
index->Min();
1749 if (pos >= 0 && pos < vars.size()) {
1756 return RevAlloc(
new IntExprIndexOfCt(
this, vars,
index, target));
1762 IntExpr*
const cache = model_cache_->FindVarArrayConstantExpression(
1764 if (cache !=
nullptr) {
1765 return cache->
Var();
1767 const std::string
name =
1771 model_cache_->InsertVarArrayConstantExpression(
const std::vector< IntVar * > vars_
Cast constraints are special channeling constraints designed to keep a variable in sync with an expre...
IntVar *const target_var_
A constraint is the main modeling object.
A Demon is the base element of a propagation queue.
void Post() override
This method is called when the constraint is processed by the solver.
void InitialPropagate() override
This method performs the initial propagation of the constraint.
IfThenElseCt(Solver *const solver, IntVar *const condition, IntExpr *const one, IntExpr *const zero, IntVar *const target)
void Accept(ModelVisitor *const visitor) const override
Accepts the given visitor.
std::string DebugString() const override
Utility class to encapsulate an IntVarIterator and use it in a range-based loop.
The class IntExpr is the base of all integer expressions in constraint programming.
virtual IntVar * Var()=0
Creates a variable from the expression.
virtual void SetRange(int64_t l, int64_t u)
This method sets both the min and the max of the expression.
virtual bool Bound() const
Returns true if the min and the max of the expression are equal.
virtual void SetValue(int64_t v)
This method sets the value of the expression.
virtual int64_t Min() const =0
virtual int64_t Max() const =0
virtual void Range(int64_t *l, int64_t *u)
By default calls Min() and Max(), but can be redefined when Min and Max code can be factorized.
virtual void WhenRange(Demon *d)=0
Attach a demon that will watch the min or the max of the expression.
The class IntVar is a subset of IntExpr.
virtual bool Contains(int64_t v) const =0
This method returns whether the value 'v' is in the domain of the variable.
virtual void WhenBound(Demon *d)=0
This method attaches a demon that will be awakened when the variable is bound.
@ VAR_ARRAY_CONSTANT_INDEX
static const char kIndex2Argument[]
static const char kMinArgument[]
static const char kElementEqual[]
static const char kTargetArgument[]
static const char kMaxArgument[]
static const char kEvaluatorArgument[]
static const char kVarsArgument[]
static const char kIndexOf[]
static const char kElement[]
static const char kValuesArgument[]
static const char kIndexArgument[]
virtual std::string name() const
Object naming.
std::string DebugString() const override
IntExpr * MakeIndexExpression(const std::vector< IntVar * > &vars, int64_t value)
Returns the expression expr such that vars[expr] == value.
IntExpr * RegisterIntExpr(IntExpr *const expr)
Registers a new IntExpr and wraps it inside a TraceIntExpr if necessary.
Constraint * MakeFalseConstraint()
This constraint always fails.
Constraint * MakeEquality(IntExpr *const left, IntExpr *const right)
left == right
Constraint * MakeElementEquality(const std::vector< int64_t > &vals, IntVar *const index, IntVar *const target)
IntVar * MakeIntVar(int64_t min, int64_t max, const std::string &name)
MakeIntVar will create the best range based int var for the bounds given.
Constraint * MakeMemberCt(IntExpr *const expr, const std::vector< int64_t > &values)
expr in set.
std::function< int64_t(int64_t, int64_t)> IndexEvaluator2
void AddConstraint(Constraint *const c)
Adds the constraint 'c' to the model.
IntExpr * MakeOpposite(IntExpr *const expr)
-expr
Constraint * MakeIfThenElseCt(IntVar *const condition, IntExpr *const then_expr, IntExpr *const else_expr, IntVar *const target_var)
Special cases with arrays of size two.
Demon * MakeConstraintInitialPropagateCallback(Constraint *const ct)
This method is a specialized case of the MakeConstraintDemon method to call the InitiatePropagate of ...
Constraint * MakeIndexOfConstraint(const std::vector< IntVar * > &vars, IntVar *const index, int64_t target)
This constraint is a special case of the element constraint with an array of integer variables,...
IntExpr * MakeElement(const std::vector< int64_t > &values, IntVar *const index)
values[index]
T * RevAlloc(T *object)
Registers the given object as being reversible.
IntExpr * MakeSum(IntExpr *const left, IntExpr *const right)
left + right.
std::function< int64_t(int64_t)> IndexEvaluator1
Callback typedefs.
std::function< IntVar *(int64_t)> Int64ToIntVar
IntExpr * MakeMonotonicElement(IndexEvaluator1 values, bool increasing, IntVar *const index)
Function based element.
std::vector< int64_t > to_remove_
#define UPDATE_ELEMENT_INDEX_BOUNDS(test)
ABSL_FLAG(bool, cp_disable_element_cache, true, "If true, caching for IntElement is disabled.")
#define UPDATE_RMQ_BASE_ELEMENT_INDEX_BOUNDS(test)
std::pair< double, double > Range
std::function< int64_t(const Model &)> Value(IntegerVariable v)
Collection of objects used to extend the Constraint Solver library.
bool IsArrayConstant(const std::vector< T > &values, const T &value)
bool IsIncreasing(const std::vector< T > &values)
Demon * MakeConstraintDemon0(Solver *const s, T *const ct, void(T::*method)(), const std::string &name)
bool IsArrayBoolean(const std::vector< T > &values)
Demon * MakeConstraintDemon1(Solver *const s, T *const ct, void(T::*method)(P), const std::string &name, P param1)
Demon * MakeDelayedConstraintDemon0(Solver *const s, T *const ct, void(T::*method)(), const std::string &name)
std::string JoinDebugStringPtr(const std::vector< T > &v, const std::string &separator)
bool IsIncreasingContiguous(const std::vector< T > &values)
std::vector< int64_t > ToInt64Vector(const std::vector< int > &input)
void LinkVarExpr(Solver *const s, IntExpr *const expr, IntVar *const var)
bool AreAllBound(const std::vector< IntVar * > &vars)
std::string JoinNamePtr(const std::vector< T > &v, const std::string &separator)
IntervalVar *const target_var_
std::function< int64_t(int64_t, int64_t)> evaluator_