Repository navigation
Expand file tree
/
Copy path2_visualize.py
More file actions
84 lines (65 loc) · 2.25 KB
/
Copy path2_visualize.py
File metadata and controls
84 lines (65 loc) · 2.25 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
import os
from rich.console import Console
from rich.panel import Panel
from allen_brain.TOSICA.train import set_seed
from allen_brain.cell_data.cell_dataset import make_dataset
from allen_brain.cell_data.cell_load import ALL_DATASETS
from allen_brain.cell_data.cell_vis import DatasetVisualizer
SEED = 1
console = Console()
DATASETS = [(info['dir'], name) for name, info in ALL_DATASETS.items()]
def run_visualizations(data_dir: str, tag: str):
set_seed(SEED)
if not (os.path.exists(os.path.join(data_dir, 'X_train.npy'))
or os.path.exists(os.path.join(data_dir, 'X_train.npz'))):
console.print(f" [yellow]No data for {tag}, skipping[/yellow]")
return
console.print(Panel(
f"[bold]{tag.upper()}[/bold] · {data_dir}",
title="Dataset", border_style="cyan", expand=False,
))
train_ds = make_dataset(data_dir, split="train")
if train_ds is None:
console.print(f" [yellow]No train split found for {tag}[/yellow], loading full matrix ...")
train_ds = make_dataset(data_dir)
if train_ds is None:
console.print(f" [bold red]\\[SKIP][/bold red] No data found in {data_dir}")
return
gene_names = train_ds.gene_names
fig_dir = "figures"
vis = DatasetVisualizer(train_ds, fig_dir=fig_dir, seed=SEED)
vis.plot_class_distribution(
save_path=os.path.join(fig_dir, f"{tag}_class_distribution.png"),
)
pca, X_pca = vis.plot_pca(
n_components=20,
save_path=fig_dir,
file_name=f"{tag}_pca.png",
)
vis.plot_umap(
X_pca,
max_cells=6000,
save_path=os.path.join(fig_dir, f"{tag}_umap.png"),
)
vis.plot_heatmap(
gene_names=gene_names,
n_genes=20,
save_path=os.path.join(fig_dir, f"{tag}_heatmap.png"),
)
vis.plot_violin(
gene_names=gene_names,
top_n=6,
save_path=os.path.join(fig_dir, f"{tag}_violin.png"),
)
vis.plot_cv2(
gene_names=gene_names,
n_top=1000,
save_path=os.path.join(fig_dir, f"{tag}_cv2.png"),
)
def main():
set_seed(SEED)
for data_dir, tag in DATASETS:
run_visualizations(data_dir, tag)
console.print("\nAll visualizations complete.")
if __name__ == "__main__":
main()