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
1,987 changes: 1,987 additions & 0 deletions notebooks/script_demo.ipynb

Large diffs are not rendered by default.

70 changes: 70 additions & 0 deletions src/eval/metrics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
import sklearn


# def macro_average_precision(ground_truth, prediction_probability):
# """
# Calculates macro-averaged precision for multi-class classification.

# Args:
# ground_truth (array-like): Array of true labels.
# prediction_probability (array-like): Array of predicted class probabilities.

# Returns:
# float: Macro-averaged precision score.
# """
# num_classes = len(ground_truth.unique())

# precision_per_class = []

# # Calculate precision score for each class
# for class_label in range(num_classes):
# class_mask = ground_truth == class_label
# ground_truth_filtered = ground_truth[class_mask]
# prediction_probability_filtered = prediction_probability[class_mask]
# # Calculate precision for this class
# precision = sklearn.metrics.precision_score(ground_truth_filtered, prediction_probability_filtered[:, class_label], average='binary', zero_division=0)
# precision_per_class.append(precision)

# # Macro-average the precision scores
# macro_average_precision = np.mean(precision_per_class)
# return macro_average_precision


def calculate_metrics(
ground_truth, prediction, prediction_probability, classes, cross_validation_metrics
):
f1_score_per_cell_type = sklearn.metrics.f1_score(
ground_truth, prediction, labels=classes, average=None
)
f1_score = sklearn.metrics.f1_score(
ground_truth, prediction, labels=classes, average="macro"
)
accuracy = sklearn.metrics.accuracy_score(ground_truth, prediction)
if prediction_probability is not None:
average_precision_per_cell_type = sklearn.metrics.average_precision_score(
ground_truth, prediction_probability, average=None
)
roc_auc_per_cell_type = sklearn.metrics.roc_auc_score(
ground_truth,
prediction_probability,
multi_class="ovr",
average=None,
labels=classes,
)
else:
average_precision_per_cell_type = None
roc_auc_per_cell_type = None
confusion_matrix = sklearn.metrics.confusion_matrix(
ground_truth, prediction, labels=classes
)

metrics = {
"f1_score_per_cell_type": f1_score_per_cell_type,
"f1_score": f1_score,
"accuracy": accuracy,
"average_precision_per_cell_type": average_precision_per_cell_type,
"roc_auc_per_cell_type": roc_auc_per_cell_type,
"confusion_matrix": confusion_matrix,
}

cross_validation_metrics.loc[len(cross_validation_metrics.index)] = metrics
63 changes: 63 additions & 0 deletions src/eval/utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,63 @@
import pandas as pd
import anndata
import torch
from sklearn.model_selection import KFold
from eval.metrics import calculate_metrics
from models.ModelBase import ModelBase


def cross_validation(
data: anndata.AnnData, model: ModelBase, random_state: int = 42, n_folds: int = 5
) -> pd.DataFrame:

metrics_names = [
"f1_score_per_cell_type",
"f1_score",
"accuracy",
"average_precision_per_cell_type",
"roc_auc_per_cell_type",
"confusion_matrix",
]

cross_validation_metrics = pd.DataFrame(columns=metrics_names)

for i, (train_data, test_data) in enumerate(k_folds(data, n_folds, random_state)):
model.train(train_data)
prediction = model.predict(test_data)
prediction_probability = model.predict_proba(test_data)
ground_truth = test_data.obs["cell_labels"]

calculate_metrics(
ground_truth,
prediction,
prediction_probability,
test_data.obs["cell_labels"].cat.categories,
cross_validation_metrics,
)

print(
f"Validation accuracy of {i} fold:",
cross_validation_metrics.loc[i]["accuracy"],
)

average_metrics = {
metric_name: cross_validation_metrics[metric_name].mean()
for metric_name in metrics_names
}
cross_validation_metrics.loc[len(cross_validation_metrics.index)] = average_metrics

return cross_validation_metrics


def k_folds(data: anndata.AnnData, n_folds: int, random_state: int):
sample_ids = data.obs["sample_id"].cat.remove_unused_categories()
sample_ids_unique = sample_ids.cat.categories

kfold = KFold(n_splits=n_folds, shuffle=True, random_state=random_state)
split = kfold.split(sample_ids_unique.tolist())

for train, test in split:
train_mask = data.obs["sample_id"].isin(sample_ids_unique[train])
test_mask = data.obs["sample_id"].isin(sample_ids_unique[test])

yield data[train_mask], data[test_mask]
127 changes: 99 additions & 28 deletions src/models/custom_stellar.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,29 +12,75 @@
from torch import Tensor
from torch_geometric.data import Batch, Data
from torch_geometric.loader import DataLoader, RandomNodeLoader, NeighborLoader
from torch_geometric.nn import (APPNP, AGNNConv, AntiSymmetricConv, ARMAConv,
CGConv, ChebConv, ClusterGCNConv,
CuGraphGATConv, CuGraphRGCNConv,
CuGraphSAGEConv, DirGNNConv, DNAConv,
DynamicEdgeConv, EdgeConv, EGConv, FAConv,
FastRGCNConv, FeaStConv, FiLMConv,
FusedGATConv, GATConv, GatedGraphConv,
GATv2Conv, GCN2Conv, GCNConv, GENConv,
GeneralConv, GINConv, GINEConv, GMMConv,
GPSConv, GraphConv, GravNetConv, HANConv,
HEATConv, HeteroConv, HGTConv, HypergraphConv,
LEConv, LGConv, MFConv, MixHopConv, NNConv,
PANConv, PDNConv, PNAConv, PointGNNConv,
PointNetConv, PointTransformerConv, PPFConv,
ResGatedGraphConv, RGATConv, RGCNConv,
SAGEConv, SGConv, SignedConv, SplineConv,
SSGConv, SuperGATConv, TAGConv,
TransformerConv, WLConv, WLConvContinuous,
XConv)
from torch_geometric.nn import (
APPNP,
AGNNConv,
AntiSymmetricConv,
ARMAConv,
CGConv,
ChebConv,
ClusterGCNConv,
CuGraphGATConv,
CuGraphRGCNConv,
CuGraphSAGEConv,
DirGNNConv,
DNAConv,
DynamicEdgeConv,
EdgeConv,
EGConv,
FAConv,
FastRGCNConv,
FeaStConv,
FiLMConv,
FusedGATConv,
GATConv,
GatedGraphConv,
GATv2Conv,
GCN2Conv,
GCNConv,
GENConv,
GeneralConv,
GINConv,
GINEConv,
GMMConv,
GPSConv,
GraphConv,
GravNetConv,
HANConv,
HEATConv,
HeteroConv,
HGTConv,
HypergraphConv,
LEConv,
LGConv,
MFConv,
MixHopConv,
NNConv,
PANConv,
PDNConv,
PNAConv,
PointGNNConv,
PointNetConv,
PointTransformerConv,
PPFConv,
ResGatedGraphConv,
RGATConv,
RGCNConv,
SAGEConv,
SGConv,
SignedConv,
SplineConv,
SSGConv,
SuperGATConv,
TAGConv,
TransformerConv,
WLConv,
WLConvContinuous,
XConv,
)
from tqdm import tqdm

from datasets.stellar_data import (StellarDataloader,
make_graph_list_from_anndata)
from datasets.stellar_data import StellarDataloader, make_graph_list_from_anndata
from models.ModelBase import ModelBase
from models.vanilla_stellar import VanillaStellarClassifficationHead
from utils import calculate_batch_accuracy
Expand All @@ -60,7 +106,12 @@ def __init__(
self.input_dim = input_dim
self.hid_dim = hid_dim
self.input_linear = nn.Linear(input_dim, hid_dim)
self.hidden_linear = nn.Sequential(*[nn.Sequential(nn.Linear(hid_dim, hid_dim), nn.ReLU() ) for _ in range(n_hidden_layers)])
self.hidden_linear = nn.Sequential(
*[
nn.Sequential(nn.Linear(hid_dim, hid_dim), nn.ReLU())
for _ in range(n_hidden_layers)
]
)
self.graph_convs = nn.ModuleList()
self.batch_norm = batch_norm
self.batch_norms = nn.ModuleList()
Expand Down Expand Up @@ -128,6 +179,7 @@ def forward(self, x: Tensor) -> Tensor:
out = self.linear(x)
return out * self.temperature


class CustomSimpleStellarClassifficationHead(nn.Module):
r"""
A classification head that uses a linear layer to make predictions.
Expand All @@ -143,10 +195,11 @@ def forward(self, x: Tensor) -> Tensor:
out = self.linear(x)
return out


CLASSIFICATION_HEAD_IMPLEMENTATIONS = [
VanillaStellarClassifficationHead,
CustomStellarClassifficationHead,
CustomSimpleStellarClassifficationHead
CustomSimpleStellarClassifficationHead,
]


Expand All @@ -171,7 +224,12 @@ def __init__(
):
super(CustomStellarModel, self).__init__()
self.encoder = CustomStellarEncoder(
input_dim, hid_dim, graph_conv_constructor, n_hidden_layers, n_graph_layers, batch_norm
input_dim,
hid_dim,
graph_conv_constructor,
n_hidden_layers,
n_graph_layers,
batch_norm,
)
self.fc_net = fc_net_constructor(hid_dim, num_classes, temperature=temperature)

Expand Down Expand Up @@ -219,16 +277,29 @@ def train(self, data: anndata.AnnData) -> None:
batched_graphs = Batch.from_data_list(graphs)

if self.cfg.batch_type == "graph":
train_data_loader = StellarDataloader(graphs, batch_size=self.cfg.batch_size)
train_data_loader = StellarDataloader(
graphs, batch_size=self.cfg.batch_size
)
self._train_graph_batch(train_data_loader, self.cfg.epochs)
elif self.cfg.batch_type == "neighbors":
train_data_loader = NeighborLoader(batched_graphs, num_neighbors=[5], batch_size=self.cfg.node_batch_size, shuffle=True)
train_data_loader = NeighborLoader(
batched_graphs,
num_neighbors=[5],
batch_size=self.cfg.node_batch_size,
shuffle=True,
)
self._train_graph_batch(train_data_loader, self.cfg.epochs)
elif self.cfg.batch_type == "nodes_in_graph":
train_data_loader = StellarDataloader(graphs, batch_size=self.cfg.batch_size)
train_data_loader = StellarDataloader(
graphs, batch_size=self.cfg.batch_size
)
self._train_node_batch(train_data_loader, self.cfg.epochs)
elif self.cfg.batch_type == "random_nodes":
train_data_loader = RandomNodeLoader(batched_graphs, num_parts=batched_graphs.x.shape[0] // self.cfg.node_batch_size + 1, shuffle=True)
train_data_loader = RandomNodeLoader(
batched_graphs,
num_parts=batched_graphs.x.shape[0] // self.cfg.node_batch_size + 1,
shuffle=True,
)
self._train_graph_batch(train_data_loader, self.cfg.epochs)

def predict(self, data: anndata.AnnData) -> np.ndarray:
Expand Down
4 changes: 2 additions & 2 deletions src/models/sklearn_mlp.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ def __init__(self, config):
random_state=1,
early_stopping=True,
verbose=True,
max_iter=self.config.max_inter
max_iter=self.config.max_inter,
)

def train(self, data: anndata.AnnData) -> None:
Expand Down Expand Up @@ -44,4 +44,4 @@ def save(self, file_path: str) -> str:
return file_path + ".joblib"

def load(self, file_path: str) -> None:
self.mlp_classifier = load(file_path)
self.mlp_classifier = load(file_path + ".joblib")
3 changes: 2 additions & 1 deletion src/models/sklearn_svm.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ def train(self, data: anndata.AnnData) -> None:

def predict(self, data: anndata.AnnData) -> np.ndarray:
X = data.layers["exprs"]
self.scaler.fit(X)
X_scaled = self.scaler.transform(X)

prediction = self.svm.predict(X_scaled)
Expand All @@ -49,7 +50,7 @@ def save(self, file_path: str) -> str:
return path_with_ext

def load(self, file_path: str) -> None:
self.svm = load(file_path)
self.svm = load(file_path + ".joblib")


class SVMSklearnSVC(SVMSklearnModel):
Expand Down
2 changes: 1 addition & 1 deletion src/models/torch_mlp.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,7 @@ def save(self, file_path: str) -> str:
return file_path + ".pt"

def load(self, file_path: str) -> None:
self.mlp.load_state_dict(torch.load(file_path))
self.mlp.load_state_dict(torch.load(file_path + ".pt"))

def _train(self, X_train, y_train, log=False, early_stopping=False):
# train config. Probably it should be moved to config directory
Expand Down
3 changes: 2 additions & 1 deletion src/models/xgboost.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ def train(self, data: anndata.AnnData) -> None:

def predict(self, data: anndata.AnnData) -> np.ndarray:
X = data.layers["exprs"]
self.scaler.fit(X)
X_scaled = self.scaler.transform(X)

prediction = self.xgboost.predict(X_scaled)
Expand All @@ -48,4 +49,4 @@ def save(self, file_path: str) -> str:
return path_with_ext

def load(self, file_path: str) -> None:
self.xgboost.load_model(file_path)
self.xgboost.load_model(file_path + ".json")
Loading