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
2 changes: 1 addition & 1 deletion docs/source/ENVI.rst
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
ENVI
=========

.. automodule:: scenvi.ENVI
.. automodule:: scenvi._envi
:members:
:undoc-members:
:show-inheritance:
28 changes: 27 additions & 1 deletion scenvi/__init__.py
Original file line number Diff line number Diff line change
@@ -1,2 +1,28 @@
from scenvi.ENVI import ENVI # noqa: F401
"""scENVI — ENVI and COVET.

``ENVI`` is resolved lazily via ``__getattr__`` (PEP 562), so importing scenvi does
not import jax, flax, optax, clu or tensorflow_probability. COVET is pure
numpy/sklearn/scanpy and uses none of them, and resolving ENVI on first use keeps a
breakage anywhere in that stack from taking ``compute_covet`` down with it — which
is what #9 was, and what the tensorflow_probability pin does today.

``from scenvi import ENVI`` behaves exactly as before.
"""

from scenvi.utils import compute_covet # noqa: F401

__all__ = ["ENVI", "compute_covet"]


def __getattr__(name):
"""Resolve ``ENVI`` on first access, so importing scenvi stays free of jax."""
if name == "ENVI":
from scenvi._envi import ENVI

return ENVI
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")


def __dir__():
"""Keep ``ENVI`` discoverable despite the lazy import."""
return sorted(__all__)
3 changes: 2 additions & 1 deletion scenvi/ENVI.py → scenvi/_envi.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,8 @@
log_zinb_pdf,
)

from scenvi.utils import CVAE, Metrics, TrainState, compute_covet, niche_cell_type
from scenvi._nn import CVAE, Metrics, TrainState
from scenvi.utils import compute_covet, niche_cell_type


class ENVI:
Expand Down
142 changes: 142 additions & 0 deletions scenvi/_nn.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,142 @@
"""Flax/CLU model components for the ENVI CVAE.

Kept apart from ``utils.py`` so that COVET — which is pure numpy/sklearn/scanpy
— can be imported without jax, flax, clu or tensorflow_probability. Importing
this module requires the ``envi`` extra.
"""

import jax
import jax.numpy as jnp
from clu import metrics
from flax import linen as nn
from flax import struct
from flax.training import train_state
from jax import random

class FeedForward(nn.Module):
"""
:meta private:
"""

n_layers: int
n_neurons: int
n_output: int

@nn.compact
def __call__(self, x):
"""
:meta private:
"""

n_layers = self.n_layers
n_neurons = self.n_neurons
n_output = self.n_output

x = nn.Dense(
features=n_neurons,
dtype=jnp.float32,
kernel_init=nn.initializers.glorot_uniform(),
bias_init=nn.initializers.zeros_init(),
)(x)
x = nn.leaky_relu(x)
x = nn.LayerNorm(dtype=jnp.float32)(x)

for _ in range(n_layers - 1):

x = nn.Dense(
features=n_neurons,
dtype=jnp.float32,
kernel_init=nn.initializers.glorot_uniform(),
bias_init=nn.initializers.zeros_init(),
)(x)
x = nn.leaky_relu(x) + x
x = nn.LayerNorm(dtype=jnp.float32)(x)

output = nn.Dense(
features=n_output,
dtype=jnp.float32,
kernel_init=nn.initializers.glorot_uniform(),
bias_init=nn.initializers.zeros_init(),
)(x)

return output


class CVAE(nn.Module):
"""
:meta private:
"""

n_layers: int
n_neurons: int
n_latent: int
n_output_exp: int
n_output_cov: int

def setup(self):
"""
:meta private:
"""

n_layers = self.n_layers
n_neurons = self.n_neurons
n_latent = self.n_latent
n_output_exp = self.n_output_exp
n_output_cov = self.n_output_cov

self.encoder = FeedForward(
n_layers=n_layers, n_neurons=n_neurons, n_output=n_latent * 2
)

self.decoder_exp = FeedForward(
n_layers=n_layers, n_neurons=n_neurons, n_output=n_output_exp
)

self.decoder_cov = FeedForward(
n_layers=n_layers, n_neurons=n_neurons, n_output=n_output_cov
)

def __call__(self, x, mode="spatial", key=random.key(0)):
"""
:meta private:
"""

conf_const = 0 if mode == "spatial" else 1
conf_neurons = jax.nn.one_hot(
conf_const * jnp.ones(x.shape[0], dtype=jnp.int8), 2, dtype=jnp.float32
)

x_conf = jnp.concatenate([x, conf_neurons], axis=-1)
enc_mu, enc_logstd = jnp.split(self.encoder(x_conf), 2, axis=-1)

key, subkey = random.split(key)
z = enc_mu + random.normal(key=subkey, shape=enc_logstd.shape) * jnp.exp(
enc_logstd
)
z_conf = jnp.concatenate([z, conf_neurons], axis=-1)

dec_exp = self.decoder_exp(z_conf)

if mode == "spatial":
dec_cov = self.decoder_cov(z)
return (enc_mu, enc_logstd, dec_exp, dec_cov)
return (enc_mu, enc_logstd, dec_exp)


@struct.dataclass
class Metrics(metrics.Collection):
"""
:meta private:
"""

enc_loss: metrics.Average
dec_loss: metrics.Average
enc_corr: metrics.Average


class TrainState(train_state.TrainState):
"""
:meta private:
"""

metrics: Metrics
148 changes: 11 additions & 137 deletions scenvi/utils.py
Original file line number Diff line number Diff line change
@@ -1,148 +1,22 @@
"""COVET: covariance-environment niche representation (Haviv et al., Nat Biotechnol 2024).

COVET's maths is numpy + sklearn only — ``sklearn.neighbors`` for the spatial
kNN, ``np.matmul`` for the shifted covariance, ``np.linalg.eigh`` for the matrix
square root — so this module imports nothing from the deep-learning stack and
``compute_covet`` is usable without it. The ENVI CVAE components live in
``scenvi/_nn.py``.
"""

import warnings

import jax
import jax.numpy as jnp
import numpy as np
import pandas as pd
import scanpy as sc
import scipy.sparse
import sklearn.neighbors
from clu import metrics
from flax import linen as nn
from flax import struct
from flax.training import train_state
from jax import random

from sklearn.preprocessing import OneHotEncoder
from tqdm import tqdm
import scipy.sparse

class FeedForward(nn.Module):
"""
:meta private:
"""

n_layers: int
n_neurons: int
n_output: int

@nn.compact
def __call__(self, x):
"""
:meta private:
"""

n_layers = self.n_layers
n_neurons = self.n_neurons
n_output = self.n_output

x = nn.Dense(
features=n_neurons,
dtype=jnp.float32,
kernel_init=nn.initializers.glorot_uniform(),
bias_init=nn.initializers.zeros_init(),
)(x)
x = nn.leaky_relu(x)
x = nn.LayerNorm(dtype=jnp.float32)(x)

for _ in range(n_layers - 1):

x = nn.Dense(
features=n_neurons,
dtype=jnp.float32,
kernel_init=nn.initializers.glorot_uniform(),
bias_init=nn.initializers.zeros_init(),
)(x)
x = nn.leaky_relu(x) + x
x = nn.LayerNorm(dtype=jnp.float32)(x)

output = nn.Dense(
features=n_output,
dtype=jnp.float32,
kernel_init=nn.initializers.glorot_uniform(),
bias_init=nn.initializers.zeros_init(),
)(x)

return output


class CVAE(nn.Module):
"""
:meta private:
"""

n_layers: int
n_neurons: int
n_latent: int
n_output_exp: int
n_output_cov: int

def setup(self):
"""
:meta private:
"""

n_layers = self.n_layers
n_neurons = self.n_neurons
n_latent = self.n_latent
n_output_exp = self.n_output_exp
n_output_cov = self.n_output_cov

self.encoder = FeedForward(
n_layers=n_layers, n_neurons=n_neurons, n_output=n_latent * 2
)

self.decoder_exp = FeedForward(
n_layers=n_layers, n_neurons=n_neurons, n_output=n_output_exp
)

self.decoder_cov = FeedForward(
n_layers=n_layers, n_neurons=n_neurons, n_output=n_output_cov
)

def __call__(self, x, mode="spatial", key=random.key(0)):
"""
:meta private:
"""

conf_const = 0 if mode == "spatial" else 1
conf_neurons = jax.nn.one_hot(
conf_const * jnp.ones(x.shape[0], dtype=jnp.int8), 2, dtype=jnp.float32
)

x_conf = jnp.concatenate([x, conf_neurons], axis=-1)
enc_mu, enc_logstd = jnp.split(self.encoder(x_conf), 2, axis=-1)

key, subkey = random.split(key)
z = enc_mu + random.normal(key=subkey, shape=enc_logstd.shape) * jnp.exp(
enc_logstd
)
z_conf = jnp.concatenate([z, conf_neurons], axis=-1)

dec_exp = self.decoder_exp(z_conf)

if mode == "spatial":
dec_cov = self.decoder_cov(z)
return (enc_mu, enc_logstd, dec_exp, dec_cov)
return (enc_mu, enc_logstd, dec_exp)


@struct.dataclass
class Metrics(metrics.Collection):
"""
:meta private:
"""

enc_loss: metrics.Average
dec_loss: metrics.Average
enc_corr: metrics.Average


class TrainState(train_state.TrainState):
"""
:meta private:
"""
from tqdm import tqdm

metrics: Metrics

def batch_matrix_sqrt(Mats):
"""
Expand Down
Loading
Loading