AR traceback == MILP.
(self, e: dict[str, float])
| 374 | self.register([(a, 1), (b, -1), (c, -1), (d, 1)], dep) |
| 375 | |
| 376 | def why(self, e: dict[str, float]) -> list[Any]: |
| 377 | """AR traceback == MILP.""" |
| 378 | if not self.do_why: |
| 379 | return [] |
| 380 | # why expr == 0? |
| 381 | # Solve min(c^Tx) s.t. A_eq * x = b_eq, x >= 0 |
| 382 | e = strip(e) |
| 383 | if not e: |
| 384 | return [] |
| 385 | |
| 386 | b_eq = [0] * len(self.v2i) |
| 387 | for v, c in e.items(): |
| 388 | b_eq[self.v2i[v]] += float(c) |
| 389 | |
| 390 | try: |
| 391 | x = optimize.linprog(c=self.c, A_eq=self.A, b_eq=b_eq, method='highs')[ |
| 392 | 'x' |
| 393 | ] |
| 394 | except: # pylint: disable=bare-except |
| 395 | x = optimize.linprog( |
| 396 | c=self.c, |
| 397 | A_eq=self.A, |
| 398 | b_eq=b_eq, |
| 399 | )['x'] |
| 400 | |
| 401 | deps = [] |
| 402 | for i, dep in enumerate(self.deps): |
| 403 | if x[2 * i] > 1e-12 or x[2 * i + 1] > 1e-12: |
| 404 | if dep not in deps: |
| 405 | deps.append(dep) |
| 406 | return deps |
| 407 | |
| 408 | def record_eq(self, v1: str, v2: str, v3: str, v4: str) -> None: |
| 409 | self.eqs.add((v1, v2, v3, v4)) |
no test coverage detected