Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions RELEASES.md
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,8 @@
- Fix quantized (F)GW solvers that ordered OT based on clusters and not initial node ordering (PR #857, Issue #786)
- Update CircleCI config to use `version: 2.1` instead of the deprecated `version: 2.` (PR #882)

- Support integer class labels on CUDA in domain-adaptation estimators by computing the label-presence mask with broadcast multiplication.

## 0.9.7.post1

This release is identical to 0.9.7 but will allow the upload of a source distribution to PyPI and release on conda-forge (that requires a source distribution).
Expand Down
2 changes: 1 addition & 1 deletion ot/da.py
Original file line number Diff line number Diff line change
Expand Up @@ -606,7 +606,7 @@ class label
# (i.e. neither is masked with MISSING_LABEL)
present_ys = (ys != MISSING_LABEL) + nx.zeros(ys.shape, type_as=ys)
present_yt = (yt != MISSING_LABEL) + nx.zeros(yt.shape, type_as=yt)
present_labels = present_ys[:, None] @ present_yt[None, :]
present_labels = present_ys[:, None] * present_yt[None, :]
# label_mismatch is a (ns, nt) matrix of {True, False} such that
# the cell (i, j) is True if ys[i] != yt[j]
label_mismatch = (ys[:, None] - yt[None, :]) != 0
Expand Down
39 changes: 39 additions & 0 deletions test/test_da.py
Original file line number Diff line number Diff line change
Expand Up @@ -647,6 +647,45 @@ def test_semisupervised_cost_correction(nx):
assert np.all(cost2[:, unlabeled] < limit2), "unlabeled targets wrongly forbidden"


@pytest.mark.parametrize("device", ["cpu", "cuda"])
@pytest.mark.parametrize("dtype_name", ["int32", "int64"])
@pytest.mark.parametrize("missing_labels", [False, True])
def test_semisupervised_integer_labels(device, dtype_name, missing_labels):
torch = pytest.importorskip("torch")
if device == "cuda" and not torch.cuda.is_available():
pytest.skip("CUDA is not available")

Xs = np.array([[0.0], [0.1], [0.3], [0.4]])
Xt = np.array([[0.35], [0.05], [0.45], [0.15]])
ys = np.array([0, 0, 1, 1])
yt = np.array([1, 0, 1, 0])
if missing_labels:
ys[0] = -1
yt[-1] = -1

expected_cost = ((Xs[:, None, :] - Xt[None, :, :]) ** 2).sum(axis=2)
penalty = 10 * expected_cost.max()
for i, source_label in enumerate(ys):
for j, target_label in enumerate(yt):
if (
source_label != -1
and target_label != -1
and source_label != target_label
):
expected_cost[i, j] = penalty

dtype = getattr(torch, dtype_name)
transport = ot.da.EMDTransport().fit(
Xs=torch.tensor(Xs, device=device),
ys=torch.tensor(ys, dtype=dtype, device=device),
Xt=torch.tensor(Xt, device=device),
yt=torch.tensor(yt, dtype=dtype, device=device),
)
assert_allclose(transport.cost_.cpu().numpy(), expected_cost, atol=1e-15)
assert_allclose(transport.coupling_.sum(0).cpu().numpy(), np.full(4, 0.25))
assert_allclose(transport.coupling_.sum(1).cpu().numpy(), np.full(4, 0.25))


@pytest.skip_backend("jax")
@pytest.skip_backend("tf")
@pytest.mark.parametrize("kernel", ["linear", "gaussian"])
Expand Down
Loading