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
6 changes: 0 additions & 6 deletions src/sc_flow/backends/torch/methods/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -270,12 +270,6 @@ def _train_step_forward(
**kwargs,
)

def extract_state_data(
self,
state_data: StateData | None,
) -> torch.Tensor | None:
return self._extract_state_data(state_data)

def train_step(
self,
matched_distr: MatchedDistributions,
Expand Down
33 changes: 23 additions & 10 deletions src/sc_flow/backends/torch/nn/_vf.py
Original file line number Diff line number Diff line change
Expand Up @@ -552,10 +552,29 @@ def is_conditional(
"""Whether a condition encoder is associated to velocity field."""
return self._condition_encoder_input_layers is not None

@classmethod
def _use_source_encoder(cls, is_paired_setting: bool, generate_from_noise: bool) -> bool:
"""Returns a boolean flag indicating whether to initialize the source encore module.

The source encoder module will be initialized only in the paired settings when generating from noise.

:param is_paired_setting: Boolean flag indicating whether the data is configured in the paired setting.
:type is_paired_setting: class: `bool`

:param generate_from_noise: Boolean flag indicating whether the interpolation is made between tractable
noise distribution and data.
:type generate_from_noise: class: `bool`
"""
if is_paired_setting and generate_from_noise:
return True
return False

@classmethod
def init_from_dims_registry(
cls,
dims_registry: DataDimensionalitiesRegistry,
is_paired_setting: bool,
generate_from_noise: bool = False,
condition_encoder_input_layers: NestedLayersDict | None = None,
source_encoder_mlp_kwargs: LayersDict | None = None,
**kwargs,
Expand Down Expand Up @@ -590,16 +609,10 @@ def init_from_dims_registry(
condition_encoder_covariates_not_pooled.append(cov)

# register source state dimensionality when provided
if source_encoder_mlp_kwargs is not None:
# get source dimension
if dims_registry.source_lin_dim is not None and dims_registry.source_quad_dim is not None:
source_dim = dims_registry.source_lin_dim + dims_registry.source_quad_dim
elif dims_registry.source_lin_dim is not None:
source_dim = dims_registry.source_lin_dim
elif dims_registry.source_quad_dim is not None:
source_dim = dims_registry.source_quad_dim
else:
source_dim = None
if cls._use_source_encoder(is_paired_setting, generate_from_noise):
if source_encoder_mlp_kwargs is None:
source_encoder_mlp_kwargs = {}
source_dim = dims_registry.source_dim
source_encoder_mlp_kwargs["input_dim"] = source_dim

# promote arguments passed from kwargs
Expand Down
5 changes: 5 additions & 0 deletions src/sc_flow/data/_dims_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,11 @@ class DataDimensionalitiesRegistry:
target_lin_dim: int | None
target_quad_dim: int | None

@property
def source_dim(self) -> int:
to_concat = [self.source_lin_dim, self.source_quad_dim]
return sum([e for e in to_concat if e is not None])

@classmethod
def _get_dims_from_continuous_data(cls, data: BatchMixin) -> dict[str, int]:
return {cov_name: cov_data.shape[-1] for cov_name, cov_data in data.mapping.items()}
Expand Down
33 changes: 16 additions & 17 deletions src/sc_flow/methods/_methods.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@

from sc_flow.data._dims_registry import DataDimensionalitiesRegistry
from sc_flow.data._manager import DataManager
from sc_flow.data.containers._state import StateData

if TYPE_CHECKING:
# jax backend
Expand Down Expand Up @@ -36,32 +35,34 @@ def __init__(
dm: DataManager,
is_paired_setting: bool,
*args,
generate_from_noise: bool = False,
**kwargs,
) -> None:
# initialize attributes
self._dims_registry = dims_registry
self._dm = dm
self._is_paired_setting = is_paired_setting

# automatically fall back to noise generation when
# no control values are provided
if not self._is_paired_setting:
generate_from_noise = True
self._generate_from_noise = generate_from_noise

# check module is passed
if self._module_cls is None:
raise NotImplementedError(f"{self.__class__.__name__} must define a `_module_cls` class attribute.")

# initialize module with dimensionality registry
self._module = self._module_cls.init_from_dims_registry(self._dims_registry, *args, **kwargs)
self._module = self._module_cls.init_from_dims_registry(
self._dims_registry, self._is_paired_setting, *args, generate_from_noise=self._generate_from_noise, **kwargs
)

@abc.abstractmethod
def set_train_mode(self, mode: bool) -> None:
"""Set the underlying module to training (True) or evaluation (False) mode."""
pass

@abc.abstractmethod
def extract_state_data(
self,
state_data: StateData | None,
) -> Any | None:
pass

@abc.abstractmethod
def train_step(self, *args: Any, **kwargs: Any) -> tuple[Any, dict[str, Any]]:
pass
Expand All @@ -86,6 +87,10 @@ def dims_registry(self) -> DataDimensionalitiesRegistry | None:
def is_paired_setting(self) -> bool:
return self._is_paired_setting

@property
def generate_from_noise(self) -> bool:
return self._generate_from_noise


class BaseGenerativeFlow(BaseMethod):
_default_solver_cls: type[JaxSolver | TorchSolver] | None = None
Expand All @@ -96,28 +101,22 @@ def __init__(
dm: DataManager,
is_paired_setting: bool,
*args,
generate_from_noise: bool = False,
probability_path: JaxProbabilityPath | TorchProbabilityPath | None = None,
match_fn: JaxMatchFn | TorchMatchFn | None = None,
noise_sampler: JaxNoiseSampler | TorchNoiseSampler | None = None,
time_sampler: JaxTimeSampler | TorchTimeSampler | None = None,
generate_from_noise: bool = False,
**kwargs,
) -> None:
# initialize parent class
super().__init__(dims_registry, dm, is_paired_setting, *args, **kwargs)
super().__init__(dims_registry, dm, is_paired_setting, *args, generate_from_noise=generate_from_noise, **kwargs)

# set attributes
self._probability_path = probability_path
self._match_fn = match_fn
self._noise_sampler = noise_sampler
self._time_sampler = time_sampler

# automatically fall back to noise generation when
# no control values are provided
if not self._is_paired_setting:
generate_from_noise = True
self._generate_from_noise = generate_from_noise

@property
def generate_from_noise(self) -> bool:
return self._generate_from_noise
Expand Down
Loading
Loading