Download src/models/cfm.py from zz312/RiboSphere: direct link, hf CLI and curl.
- Browser
- Download file 5.76 kB
-
https://huggingface.co/zz312/RiboSphere/resolve/main/src/models/cfm.py
- Command line
-
hf download hf://zz312/RiboSphere/src/models/cfm.py
-
curl -L -o cfm.py https://huggingface.co/zz312/RiboSphere/resolve/main/src/models/cfm.py
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) | |