r"""This function gives the regularization path of l2-penalized UOT problem The problem to optimize is the Lasso reformulation of the l2-penalized UOT: .. math:: \min_t \gamma \mathbf{c}^T \mathbf{t} + 0.5 * \|{H} \mathbf{t} - \mathbf{y}\|_2^2 s.t.
(a: np.array, b: np.array, C: np.array, reg=1e-4, itmax=50000)
| 540 | |
| 541 | |
| 542 | def fully_relaxed_path(a: np.array, b: np.array, C: np.array, reg=1e-4, itmax=50000): |
| 543 | r"""This function gives the regularization path of l2-penalized UOT problem |
| 544 | |
| 545 | The problem to optimize is the Lasso reformulation of the l2-penalized UOT: |
| 546 | |
| 547 | .. math:: |
| 548 | \min_t \gamma \mathbf{c}^T \mathbf{t} |
| 549 | + 0.5 * \|{H} \mathbf{t} - \mathbf{y}\|_2^2 |
| 550 | |
| 551 | s.t. |
| 552 | \mathbf{t} \geq 0 |
| 553 | |
| 554 | where : |
| 555 | |
| 556 | - :math:`\mathbf{c}` is the flattened version of the cost matrix \ |
| 557 | :math:`{C}` |
| 558 | - :math:`\gamma = 1/\lambda` is the l2-regularization coefficient |
| 559 | - :math:`\mathbf{y}` is the concatenation of vectors :math:`\mathbf{a}` \ |
| 560 | and :math:`\mathbf{b}`, defined as \ |
| 561 | :math:`\mathbf{y}^T = [\mathbf{a}^T \mathbf{b}^T]` |
| 562 | - :math:`{H}` is a design matrix, see :ref:`[41] <references-regpath>` \ |
| 563 | for the design of :math:`{H}`. The matrix product :math:`H\mathbf{t}` \ |
| 564 | computes both the source marginal and the target marginals. |
| 565 | - :math:`\mathbf{t}` is the flattened version of the transport matrix |
| 566 | |
| 567 | Parameters |
| 568 | ---------- |
| 569 | a : np.ndarray (dim_a,) |
| 570 | Histogram of dimension dim_a |
| 571 | b : np.ndarray (dim_b,) |
| 572 | Histogram of dimension dim_b |
| 573 | C : np.ndarray, shape (dim_a, dim_b) |
| 574 | Cost matrix |
| 575 | reg: float |
| 576 | l2-regularization coefficient |
| 577 | itmax: int |
| 578 | Maximum number of iteration |
| 579 | Returns |
| 580 | ------- |
| 581 | t : np.ndarray (dim_a*dim_b, ) |
| 582 | Flattened vector of the optimal transport matrix |
| 583 | t_list : list |
| 584 | List of solutions in the regularization path |
| 585 | gamma_list : list |
| 586 | List of regularization coefficients in the regularization path |
| 587 | |
| 588 | Examples |
| 589 | -------- |
| 590 | >>> import ot |
| 591 | >>> import numpy as np |
| 592 | >>> n = 3 |
| 593 | >>> xs = np.array([1., 2., 3.]).reshape((n, 1)) |
| 594 | >>> xt = np.array([5., 6., 7.]).reshape((n, 1)) |
| 595 | >>> C = ot.dist(xs, xt) |
| 596 | >>> C /= C.max() |
| 597 | >>> a = np.array([0.2, 0.5, 0.3]) |
| 598 | >>> b = np.array([0.2, 0.5, 0.3]) |
| 599 | >>> t, _, _ = ot.regpath.fully_relaxed_path(a, b, C, 1e-4) |
no test coverage detected