From a711f6824f13b38d0031d8f036abbe2684789bba Mon Sep 17 00:00:00 2001 From: CEP-24 Date: Fri, 27 Jun 2025 17:40:40 +0200 Subject: [PATCH] [issue_314][multi_source]Added draft version of M3SDA for multi source --- skada/deep/_multi_source.py | 161 ++++++++++++++++++++++++++++++++ skada/deep/test_multi_source.py | 147 +++++++++++++++++++++++++++++ 2 files changed, 308 insertions(+) create mode 100644 skada/deep/_multi_source.py create mode 100644 skada/deep/test_multi_source.py diff --git a/skada/deep/_multi_source.py b/skada/deep/_multi_source.py new file mode 100644 index 00000000..04f5f479 --- /dev/null +++ b/skada/deep/_multi_source.py @@ -0,0 +1,161 @@ +import numpy as np +import torch +from torch import nn +from torch.utils.data import DataLoader, TensorDataset +from sklearn.base import BaseEstimator, ClassifierMixin +from skada import BaseAdapter + +# === Feature Extractor avec domain embedding === +class FeatureExtractor(nn.Module): + def __init__(self, input_dim, hidden_dim, domain_embedding_dim, num_domains): + super().__init__() + self.domain_embedding = nn.Embedding(num_domains, domain_embedding_dim) + self.net = nn.Sequential( + nn.Linear(input_dim + domain_embedding_dim, hidden_dim), + nn.ReLU(), + nn.Linear(hidden_dim, hidden_dim), + nn.ReLU() + ) + + def forward(self, x, domain_ids): + domain_embed = self.domain_embedding(domain_ids) + x_cat = torch.cat([x, domain_embed], dim=1) + return self.net(x_cat) + +# === Classifier simple (partagé) === +class Classifier(nn.Module): + def __init__(self, input_dim, num_classes): + super().__init__() + self.classifier = nn.Linear(input_dim, num_classes) + + def forward(self, x): + return self.classifier(x) + +# === M3SDA Adapter complet === +class M3SDAAdapter(BaseAdapter, BaseEstimator, ClassifierMixin): + def __init__(self, input_dim, hidden_dim=128, num_classes=2, domain_embedding_dim=8, epochs=10, batch_size=64, lr=1e-3, device=None): + self.input_dim = input_dim + self.hidden_dim = hidden_dim + self.num_classes = num_classes + self.domain_embedding_dim = domain_embedding_dim + self.epochs = epochs + self.batch_size = batch_size + self.lr = lr + self.device = device or ('cuda' if torch.cuda.is_available() else 'cpu') + self.is_fitted_ = False + + def fit(self, X, y=None, sample_domain=None): + X = np.asarray(X) + y = np.asarray(y) + sample_domain = np.asarray(sample_domain) + + # Trouver les domaines source et cible + all_domains = np.unique(sample_domain) + self.source_domains_ = [d for d in all_domains if d != 'tgt'] + self.target_domain_ = 'tgt' + + self.domain_to_idx_ = {d: i for i, d in enumerate(self.source_domains_ + [self.target_domain_])} + num_domains = len(self.domain_to_idx_) + + # Séparer par domaine + X_sources, y_sources = [], [] + for d in self.source_domains_: + mask = sample_domain == d + X_sources.append(X[mask]) + y_sources.append(y[mask]) + + X_target = X[sample_domain == self.target_domain_] + + # Loaders + loaders_source = [self._to_loader(Xs, ys, self.domain_to_idx_[d]) for Xs, ys, d in zip(X_sources, y_sources, self.source_domains_)] + loader_target = self._to_loader(X_target, domain_label=self.domain_to_idx_[self.target_domain_]) + + # Init modèle + self.feature_extractor_ = FeatureExtractor(self.input_dim, self.hidden_dim, self.domain_embedding_dim, num_domains).to(self.device) + self.classifier_ = Classifier(self.hidden_dim, self.num_classes).to(self.device) + + self._train_model(loaders_source, loader_target) + self.is_fitted_ = True + return self + + def _train_model(self, loaders_source, loader_target): + self.feature_extractor_.train() + self.classifier_.train() + optimizer = torch.optim.Adam( + list(self.feature_extractor_.parameters()) + list(self.classifier_.parameters()), + lr=self.lr + ) + criterion = nn.CrossEntropyLoss() + + for epoch in range(self.epochs): + for batches in zip(*loaders_source): + optimizer.zero_grad() + losses = [] + + # Target batch + try: + batch_target = next(self.target_iter) + except: + self.target_iter = iter(loader_target) + batch_target = next(self.target_iter) + + x_t, d_t = batch_target[0].to(self.device), batch_target[1].to(self.device) + f_t = self.feature_extractor_(x_t, d_t) + + for x_s, y_s, d_s in batches: + x_s, y_s, d_s = x_s.to(self.device), y_s.to(self.device), d_s.to(self.device) + + # Forward + f_s = self.feature_extractor_(x_s, d_s) + y_pred = self.classifier_(f_s) + + # Classification loss + loss_cls = criterion(y_pred, y_s) + + # Moment matching loss + loss_mm = self._moment_loss(f_s, f_t) + losses.append(loss_cls + loss_mm) + + loss = sum(losses) / len(losses) + loss.backward() + optimizer.step() + + def _moment_loss(self, f_s, f_t): + mu_s = f_s.mean(0) + mu_t = f_t.mean(0) + return torch.norm(mu_s - mu_t, p=2) + + def _to_loader(self, X, y=None, domain_label=None): + X_tensor = torch.tensor(X, dtype=torch.float32) + d_tensor = torch.full((X.shape[0],), domain_label, dtype=torch.long) + if y is not None: + y_tensor = torch.tensor(y, dtype=torch.long) + dataset = TensorDataset(X_tensor, y_tensor, d_tensor) + else: + dataset = TensorDataset(X_tensor, d_tensor) + return DataLoader(dataset, batch_size=self.batch_size, shuffle=True) + + def transform(self, X, sample_domain=None): + if not self.is_fitted_: + raise ValueError("M3SDAAdapter must be fitted before calling transform.") + + self.feature_extractor_.eval() + with torch.no_grad(): + X_tensor = torch.tensor(X, dtype=torch.float32).to(self.device) + d_idx = self.domain_to_idx_.get(sample_domain, self.domain_to_idx_[self.target_domain_]) + d_tensor = torch.full((X.shape[0],), d_idx, dtype=torch.long).to(self.device) + features = self.feature_extractor_(X_tensor, d_tensor) + return features.cpu().numpy() + + def predict(self, X, sample_domain=None): + features = self.transform(X, sample_domain=sample_domain) + self.classifier_.eval() + with torch.no_grad(): + X_tensor = torch.tensor(features, dtype=torch.float32).to(self.device) + preds = self.classifier_(X_tensor).argmax(1) + return preds.cpu().numpy() + + def score(self, X, y, sample_domain=None): + from sklearn.metrics import accuracy_score + y_pred = self.predict(X, sample_domain) + return accuracy_score(y, y_pred) \ No newline at end of file diff --git a/skada/deep/test_multi_source.py b/skada/deep/test_multi_source.py new file mode 100644 index 00000000..574d0bab --- /dev/null +++ b/skada/deep/test_multi_source.py @@ -0,0 +1,147 @@ +import numpy as np +from _multi_source import M3SDAAdapter +import torch.nn as nn +import torch.optim as optim +from torch.utils.data import DataLoader, TensorDataset + + +def generate_2d_gaussian_domains(n_samples=100, seed=42): + np.random.seed(seed) + + def make_domain(mu0, mu1): + X0 = np.random.normal(loc=mu0, scale=0.5, size=(n_samples, 2)) + X1 = np.random.normal(loc=mu1, scale=0.5, size=(n_samples, 2)) + X = np.vstack([X0, X1]) + y = np.array([0] * n_samples + [1] * n_samples) + return X, y + + X_src1, y_src1 = make_domain(mu0=[-2, 0], mu1=[-2, 2]) + X_src2, y_src2 = make_domain(mu0=[0, -2], mu1=[2, -2]) + X_tgt, y_tgt = make_domain(mu0=[1, 1], mu1=[3, 3]) + + X_all = np.vstack([X_src1, X_src2, X_tgt]) + y_all = np.hstack([y_src1, y_src2, y_tgt]) + domains = np.array(['src1'] * len(X_src1) + ['src2'] * len(X_src2) + ['tgt'] * len(X_tgt)) + + return X_all, y_all, domains, X_tgt, y_tgt, X_src1, y_src1, X_src2, y_src2 + + +#Training + + +X_all, y_all, domains, X_tgt, y_tgt, X_src1, y_src1, X_src2, y_src2 = generate_2d_gaussian_domains() + +adapter = M3SDAAdapter(input_dim=2, hidden_dim=32, domain_embedding_dim=4, epochs=50) +adapter.fit(X_all, y_all, sample_domain=domains) + +#Display + +import matplotlib.pyplot as plt +import torch + +def plot_decision_boundary(model, domain_id, domain_name, color, ax, xrange=(-4, 5), yrange=(-4, 5), steps=200): + xx, yy = np.meshgrid(np.linspace(*xrange, steps), np.linspace(*yrange, steps)) + grid = np.c_[xx.ravel(), yy.ravel()] + with torch.no_grad(): + inputs = torch.tensor(grid, dtype=torch.float32) + domains = torch.full((inputs.shape[0],), domain_id, dtype=torch.long) + feats = model.feature_extractor_(inputs, domains) + preds = model.classifier_(feats).argmax(1).numpy() + zz = preds.reshape(xx.shape) + ax.contourf(xx, yy, zz, levels=1, alpha=0.15, colors=[color]) + +def show_class_separation(adapter, X_src1, y_src1, X_src2, y_src2, X_tgt, y_tgt, domain_to_idx): + fig, ax = plt.subplots(figsize=(8, 8)) + + ax.scatter(*X_src1[y_src1==0].T, color='blue', label='src1 - class 0') + ax.scatter(*X_src1[y_src1==1].T, color='navy', label='src1 - class 1') + + ax.scatter(*X_src2[y_src2==0].T, color='green', label='src2 - class 0') + ax.scatter(*X_src2[y_src2==1].T, color='darkgreen', label='src2 - class 1') + + ax.scatter(*X_tgt[y_tgt==0].T, color='red', label='tgt - class 0') + ax.scatter(*X_tgt[y_tgt==1].T, color='darkred', label='tgt - class 1') + + # Décision par domaine + plot_decision_boundary(adapter, domain_to_idx['src1'], 'src1', 'blue', ax) + plot_decision_boundary(adapter, domain_to_idx['src2'], 'src2', 'green', ax) + plot_decision_boundary(adapter, domain_to_idx['tgt'], 'tgt', 'red', ax) + + ax.set_xlim(-4, 5) + ax.set_ylim(-4, 5) + ax.set_title("Séparation des classes + frontières de décision") + ax.legend() + ax.grid(True) + plt.show() + + +show_class_separation(adapter, + X_src1, y_src1, X_src2, y_src2, X_tgt, y_tgt, adapter.domain_to_idx_) + +#Naive classifier + + + +class SimpleClassifier(nn.Module): + def __init__(self, input_dim=2, hidden_dim=32, num_classes=2): + super().__init__() + self.net = nn.Sequential( + nn.Linear(input_dim, hidden_dim), + nn.ReLU(), + nn.Linear(hidden_dim, num_classes) + ) + + def forward(self, x): + return self.net(x) + +def train_naive_classifier(X, y, epochs=50, batch_size=64, lr=1e-3): + X_tensor = torch.tensor(X, dtype=torch.float32) + y_tensor = torch.tensor(y, dtype=torch.long) + dataset = TensorDataset(X_tensor, y_tensor) + loader = DataLoader(dataset, batch_size=batch_size, shuffle=True) + + model = SimpleClassifier() + model.train() + optimizer = optim.Adam(model.parameters(), lr=lr) + criterion = nn.CrossEntropyLoss() + + for epoch in range(epochs): + for x_batch, y_batch in loader: + optimizer.zero_grad() + logits = model(x_batch) + loss = criterion(logits, y_batch) + loss.backward() + optimizer.step() + + return model + + +# Utilise uniquement les sources (comme en adaptation classique) +X_sources = np.vstack([X_src1, X_src2]) +y_sources = np.hstack([y_src1, y_src2]) +baseline_model = train_naive_classifier(X_sources, y_sources) + + +def plot_baseline_decision(model, X_tgt, y_tgt): + xx, yy = np.meshgrid(np.linspace(-4, 5, 200), np.linspace(-4, 5, 200)) + grid = np.c_[xx.ravel(), yy.ravel()] + with torch.no_grad(): + inputs = torch.tensor(grid, dtype=torch.float32) + preds = model(inputs).argmax(1).numpy() + zz = preds.reshape(xx.shape) + + plt.figure(figsize=(8, 6)) + plt.contourf(xx, yy, zz, levels=1, alpha=0.2, colors=["gray", "black"]) + plt.scatter(*X_tgt[y_tgt==0].T, color="red", label="Classe 0 (tgt)", alpha=0.6) + plt.scatter(*X_tgt[y_tgt==1].T, color="darkred", label="Classe 1 (tgt)", alpha=0.6) + plt.title("Frontière du classifieur NON adapté sur la cible") + plt.legend() + plt.grid(True) + plt.show() + + +# Classifieur non adapté +plot_baseline_decision(baseline_model, X_tgt, y_tgt) + +# Classifieur M3SDA adapté (avec séparation + plans de décision) +show_class_separation(adapter, X_src1, y_src1, X_src2, y_src2, X_tgt, y_tgt, adapter.domain_to_idx_) \ No newline at end of file