"""Top-level RiboSphere model and Hugging Face serialization helpers.""" from __future__ import annotations import json from collections.abc import Mapping from dataclasses import asdict, dataclass, fields from math import prod from os import PathLike from pathlib import Path from typing import Any, Literal import torch from numpy.typing import ArrayLike from torch import Tensor, nn from .attention import TransformerStack from .cfm import ConditionalFlowMatcher from .dit import DiffusionTransformer from .fsq import FiniteScalarQuantizer from .layers import FeedForward AtomRepresentation = Literal["a1", "a6", "a10", "a11"] SamplingSchedule = Literal["us", "tan", "1/t"] WeightValue = Tensor | ArrayLike ATOM_COUNTS: dict[str, int] = { "a1": 1, "a6": 6, "a10": 10, "a11": 11, } def _get_variant_metadata( manifest: Mapping[str, Any], requested_variant: str, ) -> Mapping[str, Any]: """Return metadata for an exact canonical variant name.""" if requested_variant in manifest: return manifest[requested_variant] available_variants = ", ".join(sorted(manifest)) raise ValueError( f"Unknown variant {requested_variant!r}. Available variants: " f"{available_variants}." ) @dataclass class RiboSphereConfig: """Serializable architecture configuration for :class:`RiboSphere`.""" n_channels_encoder: int = 256 n_channels_decoder: int = 512 n_layers_encoder: int = 2 n_layers_decoder: int = 8 n_heads: int = 8 mlp_factor: int = 4 use_qknorm: bool = False sigma: float = 0.0 levels: tuple[int, ...] = (7, 5, 5, 5, 5) drop_cond_p: float = 0.0 conditioning_type: str = "cat" n_channels_pair: int = 64 encoder_type: str = "xformer" attention_backend: str = "sdpa" window_size: int = 8 share_adaln: bool = False atoms: AtomRepresentation = "a11" def __post_init__(self) -> None: self.levels = tuple(self.levels) positive_values = { "n_channels_encoder": self.n_channels_encoder, "n_channels_decoder": self.n_channels_decoder, "n_layers_encoder": self.n_layers_encoder, "n_layers_decoder": self.n_layers_decoder, "n_heads": self.n_heads, "mlp_factor": self.mlp_factor, "window_size": self.window_size, } invalid_names = [ name for name, value in positive_values.items() if value <= 0 ] if invalid_names: raise ValueError( f"Configuration values must be positive: {', '.join(invalid_names)}" ) if self.n_channels_pair < 0: raise ValueError("n_channels_pair must be non-negative.") if self.atoms not in ATOM_COUNTS: raise ValueError( f"atoms must be one of {', '.join(ATOM_COUNTS)}, " f"received {self.atoms!r}." ) if not self.levels or any(level < 2 for level in self.levels): raise ValueError("levels must contain integers greater than one.") if not 0.0 <= self.drop_cond_p <= 1.0: raise ValueError("drop_cond_p must be in [0, 1].") if self.sigma < 0: raise ValueError("sigma must be non-negative.") if self.attention_backend not in {"sdpa", "flex"}: raise ValueError("attention_backend must be 'sdpa' or 'flex'.") if self.conditioning_type != "cat": raise ValueError("Only 'cat' conditioning is currently supported.") @classmethod def from_dict(cls, values: Mapping[str, Any]) -> RiboSphereConfig: """Build a config while ignoring Hugging Face metadata fields.""" valid_keys = {field.name for field in fields(cls)} normalized = dict(values) if "share_adaLN" in normalized and "share_adaln" not in normalized: normalized["share_adaln"] = normalized.pop("share_adaLN") normalized.pop("gpt_prior", None) normalized.pop("gpt_weight", None) config_values = { key: value for key, value in normalized.items() if key in valid_keys } return cls(**config_values) class RiboSphere(nn.Module): """Finite-scalar RNA tokenizer and flow-matching coordinate decoder.""" def __init__(self, config: RiboSphereConfig) -> None: super().__init__() self.config = config levels = list(config.levels) self.quantizer = FiniteScalarQuantizer( levels=levels, output_dimension=config.n_channels_decoder, input_dimension=config.n_channels_encoder, jitter_spread=0.0, ) self.codebook_size = prod(levels) self.flow_matcher = ConditionalFlowMatcher(config.sigma, "uniform") self.condition_dropout_probability = config.drop_cond_p self.use_pairwise_bias = config.n_channels_pair > 0 if self.use_pairwise_bias: self.pairwise_feature_embedder = PairwiseFeatureEmbedder( config.n_channels_pair, 100, ) self.num_atoms = ATOM_COUNTS[config.atoms] self.coordinate_encoder = nn.Sequential( nn.Linear(self.num_atoms * 3, config.n_channels_encoder), nn.SiLU(), nn.Linear(config.n_channels_encoder, config.n_channels_encoder), nn.LayerNorm(config.n_channels_encoder), ) self.encoder = TransformerStack( num_channels=config.n_channels_encoder, num_heads=config.n_heads, mlp_factor=config.mlp_factor, window_size=config.window_size, attention_backend=config.attention_backend, num_layers=config.n_layers_encoder, pairwise_channels=config.n_channels_pair, is_causal=False, ) self.decoder = DiffusionTransformer( num_channels=config.n_channels_decoder, input_channels=self.num_atoms * 3, num_layers=config.n_layers_decoder, num_heads=config.n_heads, mlp_factor=config.mlp_factor, normalize_queries_and_keys=config.use_qknorm, conditioning_type=config.conditioning_type, share_adaln=config.share_adaln, attention_backend=config.attention_backend, ) @classmethod def from_pretrained( cls, model_path: str | PathLike[str], *, variant: str | None = None, subfolder: str | PathLike[str] | None = None, ) -> RiboSphere: """Load a variant from a local path or Hugging Face Hub repository.""" if variant is not None and subfolder is not None: raise ValueError("Specify either variant or subfolder, not both.") if variant is not None: selected_variant = variant elif subfolder is not None: selected_variant = str(subfolder) else: selected_variant = None if selected_variant is not None and ( not selected_variant or Path(selected_variant).name != selected_variant ): raise ValueError("variant must be a single directory-safe name.") repository_path = Path(model_path) if not repository_path.exists() and selected_variant is not None: from huggingface_hub import hf_hub_download repository_id = str(model_path) manifest_path = Path( hf_hub_download( repo_id=repository_id, filename="variants.json", ) ) with manifest_path.open(encoding="utf-8") as handle: manifest = json.load(handle) variant_metadata = _get_variant_metadata( manifest, selected_variant, ) config_path = Path( hf_hub_download( repo_id=repository_id, filename=variant_metadata["config"], ) ) weights_path = Path( hf_hub_download( repo_id=repository_id, filename=variant_metadata["weights"], ) ) else: if not repository_path.exists(): from huggingface_hub import snapshot_download repository_path = Path( snapshot_download(repo_id=str(model_path)) ) if not repository_path.is_dir(): raise ValueError( "from_pretrained expects a Hugging Face model directory " "or repository ID." ) manifest_path = repository_path / "variants.json" manifest: dict[str, Any] = {} if manifest_path.is_file(): with manifest_path.open(encoding="utf-8") as handle: manifest = json.load(handle) if selected_variant is not None and manifest: variant_metadata = _get_variant_metadata( manifest, selected_variant, ) config_path = repository_path / variant_metadata.get( "config", f"configs/{selected_variant}.json", ) weights_path = repository_path / variant_metadata.get( "weights", f"weights/{selected_variant}.safetensors", ) elif selected_variant is not None: variant_directory = repository_path / selected_variant config_path = variant_directory / "config.json" weights_path = variant_directory / "model.safetensors" else: config_path = repository_path / "config.json" weights_path = repository_path / "model.safetensors" if ( selected_variant is None and manifest and ( not config_path.is_file() or not weights_path.is_file() ) ): available_variants = ", ".join(sorted(manifest)) raise ValueError( "This repository contains multiple variants. Pass " "variant=. Available variants: " f"{available_variants}." ) if not config_path.is_file() or not weights_path.is_file(): raise FileNotFoundError( "Checkpoint files are missing: " f"{config_path} and {weights_path}." ) from safetensors.torch import load_file with config_path.open(encoding="utf-8") as handle: config = RiboSphereConfig.from_dict(json.load(handle)) model = cls(config) model.load_state_dict( load_file(weights_path, device="cpu"), strict=True, ) return model def save_pretrained( self, output_directory: str | PathLike[str], *, variant: str | None = None, ) -> None: """Save a standalone checkpoint or a named repository variant.""" from safetensors.torch import save_file resolved_directory = Path(output_directory) resolved_directory.mkdir(parents=True, exist_ok=True) if variant is not None and ( not variant or Path(variant).name != variant ): raise ValueError("variant must be a single directory-safe name.") config_values = asdict(self.config) config_values.update( { "architectures": ["RiboSphere"], "model_type": "ribosphere", } ) if variant is None: config_path = resolved_directory / "config.json" weights_path = resolved_directory / "model.safetensors" else: config_directory = resolved_directory / "configs" weights_directory = resolved_directory / "weights" config_directory.mkdir(exist_ok=True) weights_directory.mkdir(exist_ok=True) config_path = config_directory / f"{variant}.json" weights_path = weights_directory / f"{variant}.safetensors" with config_path.open( "w", encoding="utf-8", ) as handle: json.dump(config_values, handle, indent=2) handle.write("\n") save_file(self.state_dict(), weights_path) if variant is not None: manifest_path = resolved_directory / "variants.json" manifest: dict[str, Any] = {} if manifest_path.is_file(): with manifest_path.open(encoding="utf-8") as handle: manifest = json.load(handle) manifest[variant] = { "atoms": self.config.atoms, "levels": list(self.config.levels), "codebook_size": self.codebook_size, "config": config_path.relative_to( resolved_directory ).as_posix(), "weights": weights_path.relative_to( resolved_directory ).as_posix(), } with manifest_path.open("w", encoding="utf-8") as handle: json.dump( dict(sorted(manifest.items())), handle, indent=2, ) handle.write("\n") def num_parameters(self, *, trainable_only: bool = False) -> int: """Return the total or trainable parameter count.""" return sum( parameter.numel() for parameter in self.parameters() if not trainable_only or parameter.requires_grad ) def _validate_coordinates(self, coordinates: Tensor) -> None: if coordinates.ndim != 4 or coordinates.shape[-1] != 3: raise ValueError("coordinates must have shape [B, L, A, 3].") if coordinates.shape[2] != self.num_atoms: raise ValueError( f"Expected {self.num_atoms} atoms per residue, " f"received {coordinates.shape[2]}." ) if not coordinates.is_floating_point(): raise TypeError("coordinates must use a floating-point dtype.") def encode( self, coordinates: Tensor, *, preprocess: bool = False, ) -> tuple[Tensor, Tensor, Tensor]: """Encode coordinates into continuous, quantized, and token states. Args: coordinates: Tensor shaped ``[B, L, A, 3]``. preprocess: If true, center Angstrom coordinates and convert to nm. Returns: ``(encoder_states, quantized_states, token_ids)`` with shapes ``[B, L, E]``, ``[B, L, D]``, and ``[B, L]``. """ self._validate_coordinates(coordinates) if preprocess: coordinates = coordinates - coordinates.mean( dim=(1, 2), keepdim=True, ) coordinates = coordinates / 10.0 batch_size, sequence_length, num_atoms, _ = coordinates.shape centered_coordinates = coordinates - coordinates.mean( dim=(1, 2), keepdim=True, ) pairwise_features = None if self.use_pairwise_bias: pairwise_features = self.pairwise_feature_embedder( centered_coordinates ) flattened_coordinates = centered_coordinates.reshape( batch_size, sequence_length, num_atoms * 3, ) encoder_states = self.coordinate_encoder(flattened_coordinates) encoder_states = self.encoder( encoder_states, pairwise_features=pairwise_features, ) quantized_states, token_ids = self.quantizer(encoder_states) return encoder_states, quantized_states, token_ids @staticmethod def _sampling_weights( noise_weight: WeightValue, score_weight: WeightValue, guidance_weight: WeightValue, *, device: torch.device, dtype: torch.dtype, ) -> tuple[Tensor, Tensor, Tensor]: tensors = [ torch.as_tensor(weight, device=device, dtype=dtype).flatten() for weight in (noise_weight, score_weight, guidance_weight) ] setting_count = max(tensor.numel() for tensor in tensors) if setting_count == 0: raise ValueError("Sampling weights cannot be empty.") normalized: list[Tensor] = [] for tensor in tensors: if tensor.numel() == 1: tensor = tensor.expand(setting_count) elif tensor.numel() != setting_count: raise ValueError( "Non-scalar sampling weights must have equal lengths." ) normalized.append(tensor.reshape(setting_count, 1, 1, 1)) return normalized[0], normalized[1], normalized[2] @torch.no_grad() def decode( self, token_ids: Tensor, *, num_steps: int = 200, noise_weight: WeightValue = 0.2, score_weight: WeightValue = 1.0, guidance_weight: WeightValue = 1.0, ) -> Tensor: """Generate centered nm coordinates from token IDs. One-dimensional weight inputs evaluate multiple sampling settings and return setting-major batches with shape ``[S * B, L, A, 3]``. """ if token_ids.ndim != 2: raise ValueError("token_ids must have shape [B, L].") if num_steps < 2: raise ValueError("num_steps must be at least 2.") if torch.any(token_ids < 0) or torch.any(token_ids >= self.codebook_size): raise ValueError( f"token_ids must be in [0, {self.codebook_size})." ) conditioning_states = self.quantizer.indices_to_codes(token_ids) device = conditioning_states.device dtype = conditioning_states.dtype original_batch_size, sequence_length, _ = conditioning_states.shape noise_weights, score_weights, guidance_weights = ( self._sampling_weights( noise_weight, score_weight, guidance_weight, device=device, dtype=dtype, ) ) setting_count = noise_weights.shape[0] if setting_count > 1: conditioning_states = conditioning_states.repeat( setting_count, 1, 1, ) batch_size = original_batch_size * setting_count coordinates = torch.randn( batch_size, sequence_length, self.num_atoms, 3, device=device, dtype=dtype, ) coordinates = coordinates - coordinates.mean( dim=(1, 2), keepdim=True, ) time_steps = torch.linspace( 0, 1, num_steps, device=device, dtype=dtype, ) sampling_schedule = self.compute_sampling_schedule(time_steps) step_size = time_steps[1] - time_steps[0] if setting_count > 1 and original_batch_size > 1: noise_weights = noise_weights.repeat_interleave( original_batch_size, dim=0, ) score_weights = score_weights.repeat_interleave( original_batch_size, dim=0, ) guidance_weights = guidance_weights.repeat_interleave( original_batch_size, dim=0, ) for step_index, current_time in enumerate(time_steps): flattened_coordinates = coordinates.reshape( batch_size, sequence_length, self.num_atoms * 3, ) batch_times = current_time.expand(batch_size) conditional_vector_field = self.decoder( flattened_coordinates, batch_times, conditioning_states=conditioning_states, ).view(batch_size, sequence_length, self.num_atoms, 3) conditional_vector_field = ( conditional_vector_field - conditional_vector_field.mean( dim=(1, 2), keepdim=True, ) ) unconditional_vector_field = self.decoder( flattened_coordinates, batch_times, conditioning_states=torch.zeros_like(conditioning_states), ).view(batch_size, sequence_length, self.num_atoms, 3) unconditional_vector_field = ( unconditional_vector_field - unconditional_vector_field.mean( dim=(1, 2), keepdim=True, ) ) guided_vector_field = ( unconditional_vector_field + guidance_weights * ( conditional_vector_field - unconditional_vector_field ) ) if current_time.item() >= 0.99: coordinates = coordinates + guided_vector_field * step_size continue score_times = current_time.expand(coordinates.shape[:-1]) conditional_score = self.vector_field_to_score( coordinates, conditional_vector_field, score_times, ) unconditional_score = self.vector_field_to_score( coordinates, unconditional_vector_field, score_times, ) guided_score = ( unconditional_score + guidance_weights * (conditional_score - unconditional_score) ) noise = torch.randn_like(coordinates) noise = noise - noise.mean(dim=(1, 2), keepdim=True) noise_std = torch.sqrt( 2 * sampling_schedule[step_index] * noise_weights * step_size ) coordinate_delta = ( guided_vector_field + sampling_schedule[step_index] * guided_score * score_weights ) * step_size + noise_std * noise coordinates = coordinates + coordinate_delta return coordinates def forward( self, coordinates: Tensor, ) -> tuple[Tensor, dict[str, Tensor]]: """Compute token IDs and flow-matching training loss.""" self._validate_coordinates(coordinates) batch_size, sequence_length, num_atoms, _ = coordinates.shape centered_coordinates = coordinates - coordinates.mean( dim=(1, 2), keepdim=True, ) _, conditioning_states, token_ids = self.encode( centered_coordinates ) source_coordinates = torch.randn_like(centered_coordinates) source_coordinates = ( source_coordinates - source_coordinates.mean(dim=(1, 2), keepdim=True) ) times, intermediate_coordinates, target_vector_field = ( self.flow_matcher.sample_flow( source_coordinates, centered_coordinates, ) ) condition_mask = ( torch.rand( (batch_size,), device=centered_coordinates.device, ) > self.condition_dropout_probability )[:, None, None] conditioning_states = conditioning_states * condition_mask flattened_intermediate_coordinates = ( intermediate_coordinates.reshape( batch_size, sequence_length, num_atoms * 3, ) ) predicted_vector_field = self.decoder( flattened_intermediate_coordinates, times, conditioning_states=conditioning_states, ).reshape(batch_size, sequence_length, num_atoms, 3) flow_loss = ( (target_vector_field - predicted_vector_field) ** 2 ).mean() return token_ids, {"flow_loss": flow_loss} @staticmethod def compute_sampling_schedule( times: Tensor, mode: SamplingSchedule = "us", exponent: float = 1.0, maximum: float | None = None, epsilon: float = 1e-2, ) -> Tensor: """Compute a reverse-time sampling schedule.""" if times.ndim != 1: raise ValueError("times must be one-dimensional.") if exponent <= 0: raise ValueError("exponent must be positive.") if maximum is not None and maximum < 0: raise ValueError("maximum must be non-negative or None.") if epsilon <= 0: raise ValueError("epsilon must be positive.") def transform_schedule(schedule: Tensor, power: float) -> Tensor: if power == 1.0: return schedule log_schedule = torch.log(schedule) mean_log_schedule = torch.mean(log_schedule) centered_log_schedule = log_schedule - mean_log_schedule normalized = torch.sigmoid(centered_log_schedule).pow(power) reconstructed_centered_log = torch.logit( normalized, eps=1e-6, ) return torch.exp( reconstructed_centered_log + mean_log_schedule ) clamped_times = torch.clamp(times, 0, 1 - 1e-5) if mode == "us": schedule = ( (1.0 - clamped_times) / (clamped_times + epsilon) ) elif mode == "tan": angle = (1.0 - clamped_times) * torch.pi / 2.0 schedule = ( (torch.pi / 2.0) * torch.sin(angle) / (torch.cos(angle) + epsilon) ) elif mode == "1/t": schedule = 1.0 / (clamped_times + epsilon) else: raise ValueError(f"Unsupported sampling schedule mode: {mode}") schedule = transform_schedule(schedule, exponent) if maximum is not None: schedule = torch.clamp_max(schedule, maximum) return torch.clamp_min(schedule, 0) @staticmethod def vector_field_to_score( noisy_coordinates: Tensor, vector_field: Tensor, times: Tensor, reference_scale: float = 1.0, ) -> Tensor: """Convert a learned vector field into a noisy-density score.""" if noisy_coordinates.shape != vector_field.shape: raise ValueError( "noisy_coordinates and vector_field must have identical shapes." ) if reference_scale <= 0: raise ValueError("reference_scale must be positive.") if torch.any(times >= 1.0): raise ValueError("times must be strictly less than 1.") numerator = times[..., None] * vector_field - noisy_coordinates denominator = ( (1.0 - times)[..., None] * reference_scale**2 ) return numerator / denominator class PairwiseFeatureEmbedder(nn.Module): """Embed residue distances and relative sequence positions.""" def __init__( self, num_channels: int, num_distance_buckets: int, ) -> None: super().__init__() if num_channels <= 0 or num_distance_buckets < 2: raise ValueError( "num_channels must be positive and " "num_distance_buckets must be at least 2." ) self.distance_embedding = nn.Embedding( num_distance_buckets, num_channels, ) self.relative_position_embedding = nn.Embedding(128, num_channels) self.register_buffer( "bins", torch.linspace(0, 4**2, num_distance_buckets - 1), ) self.projection = FeedForward( num_channels, 4 * num_channels, num_channels, activation=nn.GELU, ) self.norm = nn.LayerNorm(num_channels) self.num_channels = num_channels def forward(self, coordinates: Tensor) -> Tensor: """Return pair features shaped ``[B, L, L, C]``.""" if coordinates.ndim != 4 or coordinates.shape[-1] != 3: raise ValueError("coordinates must have shape [B, L, A, 3].") sequence_length = coordinates.shape[1] residue_centers = coordinates.mean(dim=2) squared_distances = ( ( residue_centers[:, :, None] - residue_centers[:, None, :] ) ** 2 ).sum(dim=-1) residue_indices = torch.arange( sequence_length, device=residue_centers.device, ) relative_indices = ( (residue_indices[:, None] - residue_indices[None, :]) .clip(min=-64, max=63) + 64 ) relative_position_features = self.relative_position_embedding( relative_indices ) distance_buckets = torch.bucketize( squared_distances, self.bins, ) pairwise_features = ( self.distance_embedding(distance_buckets) + relative_position_features ) return self.projection(self.norm(pairwise_features))