RiboSphere / src /models /cfm.py
zz312's picture
Upload folder using huggingface_hub
cf5d356 verified
Raw
History Blame
5.76 kB
"""Conditional flow-matching probability path utilities.
Adapted from the MIT-licensed conditional-flow-matching implementation by
Alex Tong and Kilian Fatras.
"""
from __future__ import annotations
import torch
from torch import Tensor
def broadcast_time_like(
time: float | Tensor,
reference: Tensor,
) -> float | Tensor:
"""Reshape batch time values so they broadcast over ``reference``."""
if not isinstance(time, Tensor):
return time
if time.ndim > 1:
raise ValueError("time must be a scalar or one-dimensional tensor.")
if time.ndim == 1 and time.shape[0] not in {1, reference.shape[0]}:
raise ValueError("time must contain one value per batch item.")
return time.reshape(-1, *([1] * (reference.ndim - 1)))
class ConditionalFlowMatcher:
"""Independent conditional flow matching with a linear probability path."""
SUPPORTED_TIME_SAMPLING_MODES = {"uniform", "uniform_beta"}
def __init__(
self,
sigma: float = 0.0,
time_sampling_mode: str = "uniform",
) -> None:
if sigma < 0.0:
raise ValueError("sigma must be non-negative.")
if time_sampling_mode not in self.SUPPORTED_TIME_SAMPLING_MODES:
supported = ", ".join(sorted(self.SUPPORTED_TIME_SAMPLING_MODES))
raise ValueError(
f"Unsupported time_sampling_mode={time_sampling_mode!r}. "
f"Expected one of: {supported}."
)
self.sigma = sigma
self.time_sampling_mode = time_sampling_mode
self.beta_distribution: torch.distributions.Beta | None = None
if time_sampling_mode == "uniform_beta":
self.beta_distribution = torch.distributions.Beta(
torch.tensor(1.9),
torch.tensor(1.0),
)
def sample_time(
self,
shape: tuple[int, ...],
*,
device: torch.device | None = None,
) -> Tensor:
"""Sample interpolation times with the configured distribution."""
if self.time_sampling_mode == "uniform":
return torch.rand(shape, device=device)
if self.time_sampling_mode == "uniform_beta":
if self.beta_distribution is None:
raise RuntimeError("Beta distribution was not initialized.")
uniform_times = torch.rand(shape, device=device)
beta_times = self.beta_distribution.sample(shape).to(device=device)
use_uniform = torch.rand(shape, device=device) > 0.98
return torch.where(use_uniform, uniform_times, beta_times)
raise RuntimeError(
f"Unsupported time sampling mode: {self.time_sampling_mode}"
)
def compute_path_mean(
self,
source_samples: Tensor,
target_samples: Tensor,
times: Tensor,
) -> Tensor:
"""Compute t * target + (1 - t) * source."""
times = broadcast_time_like(times, source_samples)
return (
times * target_samples + (1 - times) * source_samples
)
def compute_path_std(self, times: Tensor) -> Tensor:
"""Return the constant probability-path standard deviation."""
return torch.full_like(times, self.sigma)
def sample_intermediate(
self,
source_samples: Tensor,
target_samples: Tensor,
times: Tensor,
noise: Tensor,
) -> Tensor:
"""Sample an intermediate point from the probability path."""
path_mean = self.compute_path_mean(
source_samples, target_samples, times
)
path_std = self.compute_path_std(times)
path_std = broadcast_time_like(path_std, source_samples)
return path_mean + path_std * noise
def compute_conditional_vector_field(
self,
source_samples: Tensor,
target_samples: Tensor,
times: Tensor,
intermediate_samples: Tensor,
) -> Tensor:
"""Compute the conditional vector field target - source."""
del times, intermediate_samples
return target_samples - source_samples
def sample_noise_like(self, reference: Tensor) -> Tensor:
"""Sample standard Gaussian noise matching ``reference``."""
return torch.randn_like(reference)
def sample_flow(
self,
source_samples: Tensor,
target_samples: Tensor,
return_noise: bool = False,
) -> (
tuple[Tensor, Tensor, Tensor]
| tuple[Tensor, Tensor, Tensor, Tensor]
):
"""Sample times, intermediate states, and conditional vector fields."""
if source_samples.shape != target_samples.shape:
raise ValueError("source_samples and target_samples must match.")
times = self.sample_time(
(source_samples.shape[0],),
device=source_samples.device,
).to(dtype=source_samples.dtype)
noise = self.sample_noise_like(source_samples)
intermediate_samples = self.sample_intermediate(
source_samples,
target_samples,
times,
noise,
)
conditional_vector_field = self.compute_conditional_vector_field(
source_samples,
target_samples,
times,
intermediate_samples,
)
if return_noise:
return (
times,
intermediate_samples,
conditional_vector_field,
noise,
)
return times, intermediate_samples, conditional_vector_field
def compute_score_weight(self, times: Tensor) -> Tensor:
"""Compute the score weighting function."""
path_std = self.compute_path_std(times)
return 2 * path_std / (self.sigma ** 2 + 1e-8)