import numpy as np META = { "name": "conditional_ot_domains", "domain": "domain_generalization", "description": "Three source domains with shared class signal and domain-specific nuisance, plus a shifted target test domain." } def get_dataset(seed, n_train, n_test): rng = np.random.default_rng(seed) domains = (-2.0, 0.0, 2.0) base = n_train // len(domains) sizes = [base + (i < (n_train % len(domains))) for i in range(len(domains))] xs, ys = [], [] for d, ns in zip(domains, sizes): y = rng.integers(0, 2, ns) signal = (2 * y - 1).astype(np.float32) x = np.stack([ signal + rng.normal(0, 0.65, ns), d + rng.normal(0, 0.75, ns), np.full(ns, d / 2.0), ], axis=1) xs.append(x.astype(np.float32)) ys.append(y.astype(np.int64)) xtr, ytr = np.concatenate(xs), np.concatenate(ys) yte = rng.integers(0, 2, n_test) xte = np.stack([ (2 * yte - 1) + rng.normal(0, 0.65, n_test), 3.5 + rng.normal(0, 0.75, n_test), np.full(n_test, 1.75), ], axis=1).astype(np.float32) return { "xtr": xtr, "ytr": ytr, "xte": xte, "yte": yte.astype(np.int64), "task": "classification", "metric": "cross_entropy", "out_dim": 2, }