Skip to content
Draft
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
161 changes: 161 additions & 0 deletions skada/deep/_multi_source.py
Original file line number Diff line number Diff line change
@@ -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)
147 changes: 147 additions & 0 deletions skada/deep/test_multi_source.py
Original file line number Diff line number Diff line change
@@ -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_)