Instructions to use dn6/RFDiffusion-3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use dn6/RFDiffusion-3 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("dn6/RFDiffusion-3", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
[RFD3] Feed MPNN real backbone atoms and the schema it validates
#2
by dn6 HF Staff - opened
- README.md +8 -15
- before_denoise.py +133 -54
- denoise.py +65 -45
- modular_blocks.py +24 -21
- modular_model_index.json +1 -2
- mpnn/model_mpnn.py +34 -3
- mpnn_ligand/model_mpnn.py +34 -3
- mpnn_soluble/model_mpnn.py +34 -3
- scheduler/model.py +15 -1
- transformer/model_rfdiffusion.py +30 -57
README.md
CHANGED
|
@@ -20,7 +20,7 @@ import torch
|
|
| 20 |
from diffusers import ModularPipeline
|
| 21 |
|
| 22 |
pipe = ModularPipeline.from_pretrained("dn6/RFDiffusion-3", trust_remote_code=True)
|
| 23 |
-
pipe.load_components(device_map="cuda", torch_dtype=torch.
|
| 24 |
|
| 25 |
state = pipe(contigs="100")
|
| 26 |
print(state.output.xyz.shape) # [1, 100, 3]
|
|
@@ -34,7 +34,6 @@ The active workflow is selected automatically based on which inputs you provide:
|
|
| 34 |
|----------|---------------|-----------|
|
| 35 |
| `structure_only` | `contigs` | RFdiffusion3 |
|
| 36 |
| `structure_and_sequence` | `contigs`, `temperature` | RFdiffusion3 → MPNN |
|
| 37 |
-
| `motif_structure_and_sequence` | `contigs`, `input_xyz`, `temperature` | Motif-conditioned RFdiffusion3 → MPNN |
|
| 38 |
|
| 39 |
### Structure Only
|
| 40 |
|
|
@@ -67,18 +66,12 @@ Three MPNN variants are available:
|
|
| 67 |
|
| 68 |
### Motif-Conditioned Design
|
| 69 |
|
| 70 |
-
|
|
|
|
| 71 |
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
motif_coords = torch.randn(16, 3) # [N_motif, 3]
|
| 76 |
-
state = pipe(
|
| 77 |
-
contigs="A10-25/50",
|
| 78 |
-
input_xyz=motif_coords,
|
| 79 |
-
temperature=0.1,
|
| 80 |
-
)
|
| 81 |
-
```
|
| 82 |
|
| 83 |
### Full Design Pipeline
|
| 84 |
|
|
@@ -94,7 +87,7 @@ from diffusers import AutoModel, ModularPipeline
|
|
| 94 |
|
| 95 |
# 1. Design a backbone + sequence
|
| 96 |
design_pipe = ModularPipeline.from_pretrained("dn6/RFDiffusion-3", trust_remote_code=True)
|
| 97 |
-
design_pipe.load_components(device_map="cuda", torch_dtype=torch.
|
| 98 |
|
| 99 |
mpnn = AutoModel.from_pretrained("dn6/RFDiffusion-3", subfolder="mpnn", trust_remote_code=True)
|
| 100 |
design_pipe.update_components(mpnn=mpnn)
|
|
@@ -104,7 +97,7 @@ designed_sequence = state.mpnn_output.designed_sequence
|
|
| 104 |
|
| 105 |
# 2. Validate the fold with RF3
|
| 106 |
fold_pipe = ModularPipeline.from_pretrained("dn6/RosettaFold-3", trust_remote_code=True)
|
| 107 |
-
fold_pipe.load_components(device_map="cuda", torch_dtype=torch.
|
| 108 |
|
| 109 |
state = fold_pipe(sequence=designed_sequence, output_type="cif.gz", output_path="prediction")
|
| 110 |
```
|
|
|
|
| 20 |
from diffusers import ModularPipeline
|
| 21 |
|
| 22 |
pipe = ModularPipeline.from_pretrained("dn6/RFDiffusion-3", trust_remote_code=True)
|
| 23 |
+
pipe.load_components(device_map="cuda", torch_dtype=torch.float32, trust_remote_code=True)
|
| 24 |
|
| 25 |
state = pipe(contigs="100")
|
| 26 |
print(state.output.xyz.shape) # [1, 100, 3]
|
|
|
|
| 34 |
|----------|---------------|-----------|
|
| 35 |
| `structure_only` | `contigs` | RFdiffusion3 |
|
| 36 |
| `structure_and_sequence` | `contigs`, `temperature` | RFdiffusion3 → MPNN |
|
|
|
|
| 37 |
|
| 38 |
### Structure Only
|
| 39 |
|
|
|
|
| 66 |
|
| 67 |
### Motif-Conditioned Design
|
| 68 |
|
| 69 |
+
Not supported yet. Motif contigs such as `"A10-25/50"` and the `input_xyz` argument raise a
|
| 70 |
+
`ValueError`.
|
| 71 |
|
| 72 |
+
Conditioning on a motif requires per-atom element, atom-name and occupancy annotations for the
|
| 73 |
+
fixed residues, which the feature pipeline derives from a reference structure. A bare coordinate
|
| 74 |
+
tensor cannot supply them, so this needs a structure input rather than `input_xyz`.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 75 |
|
| 76 |
### Full Design Pipeline
|
| 77 |
|
|
|
|
| 87 |
|
| 88 |
# 1. Design a backbone + sequence
|
| 89 |
design_pipe = ModularPipeline.from_pretrained("dn6/RFDiffusion-3", trust_remote_code=True)
|
| 90 |
+
design_pipe.load_components(device_map="cuda", torch_dtype=torch.float32, trust_remote_code=True)
|
| 91 |
|
| 92 |
mpnn = AutoModel.from_pretrained("dn6/RFDiffusion-3", subfolder="mpnn", trust_remote_code=True)
|
| 93 |
design_pipe.update_components(mpnn=mpnn)
|
|
|
|
| 97 |
|
| 98 |
# 2. Validate the fold with RF3
|
| 99 |
fold_pipe = ModularPipeline.from_pretrained("dn6/RosettaFold-3", trust_remote_code=True)
|
| 100 |
+
fold_pipe.load_components(device_map="cuda", torch_dtype=torch.float32, trust_remote_code=True)
|
| 101 |
|
| 102 |
state = fold_pipe(sequence=designed_sequence, output_type="cif.gz", output_path="prediction")
|
| 103 |
```
|
before_denoise.py
CHANGED
|
@@ -23,6 +23,65 @@ from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, In
|
|
| 23 |
|
| 24 |
logger = logging.get_logger(__name__)
|
| 25 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
|
| 27 |
def parse_contig_string(contig_str: str) -> Tuple[int, List[Tuple[int, int]]]:
|
| 28 |
"""
|
|
@@ -100,34 +159,45 @@ class RFDiffusionInputStep(ModularPipelineBlocks):
|
|
| 100 |
),
|
| 101 |
]
|
| 102 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 103 |
@property
|
| 104 |
def intermediate_outputs(self) -> List[OutputParam]:
|
| 105 |
return [
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 106 |
OutputParam(
|
| 107 |
"motif_mask",
|
| 108 |
type_hint=torch.Tensor,
|
| 109 |
-
description="
|
| 110 |
),
|
| 111 |
OutputParam(
|
| 112 |
-
"
|
| 113 |
type_hint=torch.Tensor,
|
| 114 |
-
description="
|
| 115 |
),
|
| 116 |
OutputParam(
|
| 117 |
"L",
|
| 118 |
type_hint=int,
|
| 119 |
-
description="Total length of the protein being designed",
|
| 120 |
),
|
| 121 |
OutputParam(
|
| 122 |
"batch_size",
|
| 123 |
type_hint=int,
|
| 124 |
description="Batch size (typically 1 for RFDiffusion)",
|
| 125 |
),
|
| 126 |
-
OutputParam(
|
| 127 |
-
"dtype",
|
| 128 |
-
type_hint=torch.dtype,
|
| 129 |
-
description="Data type for tensors",
|
| 130 |
-
),
|
| 131 |
]
|
| 132 |
|
| 133 |
def check_inputs(self, components, block_state):
|
|
@@ -149,20 +219,33 @@ class RFDiffusionInputStep(ModularPipelineBlocks):
|
|
| 149 |
|
| 150 |
L, motif_ranges = parse_contig_string(contig_str)
|
| 151 |
|
| 152 |
-
|
| 153 |
-
for
|
| 154 |
-
|
| 155 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 156 |
if input_xyz is not None:
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 160 |
|
| 161 |
-
block_state.
|
| 162 |
-
block_state.
|
|
|
|
|
|
|
| 163 |
block_state.L = L
|
| 164 |
-
block_state.batch_size =
|
| 165 |
-
block_state.dtype = torch.float32
|
| 166 |
|
| 167 |
self.set_block_state(state, block_state)
|
| 168 |
return components, state
|
|
@@ -216,11 +299,15 @@ class RFDiffusionSetTimestepsStep(ModularPipelineBlocks):
|
|
| 216 |
def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
|
| 217 |
block_state = self.get_block_state(state)
|
| 218 |
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 224 |
|
| 225 |
block_state.noise_schedule = noise_schedule
|
| 226 |
block_state.num_inference_steps = len(noise_schedule)
|
|
@@ -261,52 +348,44 @@ class RFDiffusionPrepareLatentsStep(ModularPipelineBlocks):
|
|
| 261 |
InputParam("generator", type_hint=torch.Generator, description="Random generator for reproducibility"),
|
| 262 |
InputParam("diffusion_batch_size", default=1, type_hint=int, description="Number of samples to generate in parallel"),
|
| 263 |
InputParam("L", required=True, type_hint=int, description="Protein length"),
|
|
|
|
|
|
|
| 264 |
InputParam("motif_mask", required=True, type_hint=torch.Tensor),
|
| 265 |
-
InputParam("motif_xyz", type_hint=torch.Tensor),
|
| 266 |
InputParam("noise_schedule", required=True, type_hint=torch.Tensor),
|
| 267 |
-
InputParam("dtype", type_hint=torch.dtype),
|
| 268 |
]
|
| 269 |
|
| 270 |
@property
|
| 271 |
def intermediate_outputs(self) -> List[OutputParam]:
|
| 272 |
return [
|
| 273 |
-
OutputParam("xyz", type_hint=torch.Tensor, description="Initial noised coordinates [D,
|
|
|
|
| 274 |
]
|
| 275 |
|
| 276 |
@torch.no_grad()
|
| 277 |
def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
|
| 278 |
block_state = self.get_block_state(state)
|
| 279 |
|
| 280 |
-
L = block_state.L
|
| 281 |
-
motif_mask = block_state.motif_mask
|
| 282 |
-
motif_xyz = block_state.motif_xyz
|
| 283 |
noise_schedule = block_state.noise_schedule
|
| 284 |
-
dtype = block_state.dtype or torch.float32
|
| 285 |
generator = block_state.generator
|
| 286 |
D = block_state.diffusion_batch_size or 1
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
|
| 294 |
-
|
| 295 |
-
#
|
| 296 |
-
|
| 297 |
-
|
| 298 |
-
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
|
| 302 |
-
|
| 303 |
-
for i, idx in enumerate(motif_indices):
|
| 304 |
-
if i < motif_xyz.shape[0]:
|
| 305 |
-
coord_motif[:, idx] = motif_xyz[i].to(dtype=dtype, device=device)
|
| 306 |
-
|
| 307 |
-
xyz = noise + coord_motif
|
| 308 |
-
|
| 309 |
block_state.xyz = xyz
|
|
|
|
| 310 |
|
| 311 |
self.set_block_state(state, block_state)
|
| 312 |
return components, state
|
|
|
|
| 23 |
|
| 24 |
logger = logging.get_logger(__name__)
|
| 25 |
|
| 26 |
+
# Feature widths the checkpoint's token initializer was trained with, from
|
| 27 |
+
# rfd3/configs/model/components/rfd3_net.yaml. They set the input width of the embedding layers,
|
| 28 |
+
# so a mismatch raises on layer shape rather than silently degrading the design.
|
| 29 |
+
_TOKEN_1D_FEATURES = {"ref_motif_token_type": 3, "restype": 32, "ref_plddt": 1, "is_non_loopy": 1}
|
| 30 |
+
_ATOM_1D_FEATURES = {
|
| 31 |
+
"ref_atom_name_chars": 256,
|
| 32 |
+
"ref_element": 128,
|
| 33 |
+
"ref_charge": 1,
|
| 34 |
+
"ref_mask": 1,
|
| 35 |
+
"ref_is_motif_atom_with_fixed_coord": 1,
|
| 36 |
+
"ref_is_motif_atom_unindexed": 1,
|
| 37 |
+
"has_zero_occupancy": 1,
|
| 38 |
+
"ref_pos": 3,
|
| 39 |
+
"ref_atomwise_rasa": 3,
|
| 40 |
+
"active_donor": 1,
|
| 41 |
+
"active_acceptor": 1,
|
| 42 |
+
"is_atom_level_hotspot": 1,
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
# Every token is padded to this many atom slots, so the coordinate tensors the model sees are
|
| 46 |
+
# atom level with L = _N_ATOMS_PER_TOKEN * n_residues.
|
| 47 |
+
_N_ATOMS_PER_TOKEN = 14
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def build_design_features(length: int, diffusion_batch_size: int, sigma_data: float) -> dict:
|
| 51 |
+
"""
|
| 52 |
+
Build the foundry feature dict for an unconditional design of `length` residues.
|
| 53 |
+
|
| 54 |
+
The token initializer embeds real chemical and positional features; there is no meaningful
|
| 55 |
+
zero substitute for them, and no API in foundry that turns a length into features in one
|
| 56 |
+
call. This mirrors what `rfd3.engine.RFD3InferenceEngine` does, minus the hydra app and the
|
| 57 |
+
checkpoint, by driving the same specification and transform pipeline directly.
|
| 58 |
+
|
| 59 |
+
Returns:
|
| 60 |
+
The transformed example, carrying `feats` and `coord_atom_lvl_to_be_noised`.
|
| 61 |
+
"""
|
| 62 |
+
from rfd3.inference.input_parsing import DesignInputSpecification
|
| 63 |
+
from rfd3.transforms.pipelines import build_atom14_base_pipeline
|
| 64 |
+
|
| 65 |
+
# `length` and `contig` are mutually exclusive in the specification: passing `contig` without
|
| 66 |
+
# a structure input fails validation, so unconditional designs go through `length`.
|
| 67 |
+
spec = DesignInputSpecification(length=str(length))
|
| 68 |
+
data = spec.to_pipeline_input(example_id=f"rfd3_{length}")
|
| 69 |
+
|
| 70 |
+
pipeline = build_atom14_base_pipeline(
|
| 71 |
+
is_inference=True,
|
| 72 |
+
diffusion_batch_size=diffusion_batch_size,
|
| 73 |
+
sigma_data=sigma_data,
|
| 74 |
+
central_atom="CB",
|
| 75 |
+
n_atoms_per_token=_N_ATOMS_PER_TOKEN,
|
| 76 |
+
generate_conformers=True,
|
| 77 |
+
provide_reference_conformer_when_unmasked=True,
|
| 78 |
+
ground_truth_conformer_policy="IGNORE",
|
| 79 |
+
use_element_for_atom_names_of_atomized_tokens=True,
|
| 80 |
+
token_1d_features=_TOKEN_1D_FEATURES,
|
| 81 |
+
atom_1d_features=_ATOM_1D_FEATURES,
|
| 82 |
+
)
|
| 83 |
+
return pipeline(data)
|
| 84 |
+
|
| 85 |
|
| 86 |
def parse_contig_string(contig_str: str) -> Tuple[int, List[Tuple[int, int]]]:
|
| 87 |
"""
|
|
|
|
| 159 |
),
|
| 160 |
]
|
| 161 |
|
| 162 |
+
@property
|
| 163 |
+
def expected_components(self) -> List[ComponentSpec]:
|
| 164 |
+
return [
|
| 165 |
+
ComponentSpec("scheduler", description="RFDiffusion3 EDM scheduler"),
|
| 166 |
+
]
|
| 167 |
+
|
| 168 |
@property
|
| 169 |
def intermediate_outputs(self) -> List[OutputParam]:
|
| 170 |
return [
|
| 171 |
+
OutputParam(
|
| 172 |
+
"f",
|
| 173 |
+
type_hint=dict,
|
| 174 |
+
description="Foundry feature dict consumed by the token initializer",
|
| 175 |
+
),
|
| 176 |
+
OutputParam(
|
| 177 |
+
"coord_atom_lvl_to_be_noised",
|
| 178 |
+
type_hint=torch.Tensor,
|
| 179 |
+
description="Reference atom-level coordinates [D, L_atom, 3]",
|
| 180 |
+
),
|
| 181 |
OutputParam(
|
| 182 |
"motif_mask",
|
| 183 |
type_hint=torch.Tensor,
|
| 184 |
+
description="Atom-level boolean mask for motif (fixed) positions [L_atom]",
|
| 185 |
),
|
| 186 |
OutputParam(
|
| 187 |
+
"motif_token_mask",
|
| 188 |
type_hint=torch.Tensor,
|
| 189 |
+
description="Residue-level boolean mask for motif (fixed) positions [L]",
|
| 190 |
),
|
| 191 |
OutputParam(
|
| 192 |
"L",
|
| 193 |
type_hint=int,
|
| 194 |
+
description="Total length of the protein being designed, in residues",
|
| 195 |
),
|
| 196 |
OutputParam(
|
| 197 |
"batch_size",
|
| 198 |
type_hint=int,
|
| 199 |
description="Batch size (typically 1 for RFDiffusion)",
|
| 200 |
),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 201 |
]
|
| 202 |
|
| 203 |
def check_inputs(self, components, block_state):
|
|
|
|
| 219 |
|
| 220 |
L, motif_ranges = parse_contig_string(contig_str)
|
| 221 |
|
| 222 |
+
# Motif conditioning needs a reference structure so the transform pipeline can build
|
| 223 |
+
# per-atom features for the fixed residues. A coordinate tensor alone cannot supply the
|
| 224 |
+
# element, atom-name and occupancy annotations those features are derived from.
|
| 225 |
+
if motif_ranges:
|
| 226 |
+
raise ValueError(
|
| 227 |
+
f"Motif-conditioned contigs are not supported yet, got `contigs={contig_str!r}`. "
|
| 228 |
+
"Pass a plain design length such as `contigs='100'`."
|
| 229 |
+
)
|
| 230 |
if input_xyz is not None:
|
| 231 |
+
raise ValueError(
|
| 232 |
+
"`input_xyz` is not supported yet. Pass a plain design length such as "
|
| 233 |
+
"`contigs='100'`."
|
| 234 |
+
)
|
| 235 |
+
|
| 236 |
+
batch_size = 1
|
| 237 |
+
example = build_design_features(
|
| 238 |
+
length=L,
|
| 239 |
+
diffusion_batch_size=batch_size,
|
| 240 |
+
sigma_data=components.scheduler.config.sigma_data,
|
| 241 |
+
)
|
| 242 |
|
| 243 |
+
block_state.f = example["feats"]
|
| 244 |
+
block_state.coord_atom_lvl_to_be_noised = example["coord_atom_lvl_to_be_noised"]
|
| 245 |
+
block_state.motif_mask = example["feats"]["is_motif_atom_with_fixed_coord"]
|
| 246 |
+
block_state.motif_token_mask = example["feats"]["is_motif_token_with_fully_fixed_coord"]
|
| 247 |
block_state.L = L
|
| 248 |
+
block_state.batch_size = batch_size
|
|
|
|
| 249 |
|
| 250 |
self.set_block_state(state, block_state)
|
| 251 |
return components, state
|
|
|
|
| 299 |
def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
|
| 300 |
block_state = self.get_block_state(state)
|
| 301 |
|
| 302 |
+
# A linear stand-in for the EDM schedule silently changes the sampler, so require
|
| 303 |
+
# the real one rather than degrading.
|
| 304 |
+
if components.scheduler is None:
|
| 305 |
+
raise ValueError(
|
| 306 |
+
"`scheduler` is not loaded. Call `load_components(trust_remote_code=True)` on the "
|
| 307 |
+
"pipeline before calling it."
|
| 308 |
+
)
|
| 309 |
+
|
| 310 |
+
noise_schedule = components.scheduler.get_noise_schedule()
|
| 311 |
|
| 312 |
block_state.noise_schedule = noise_schedule
|
| 313 |
block_state.num_inference_steps = len(noise_schedule)
|
|
|
|
| 348 |
InputParam("generator", type_hint=torch.Generator, description="Random generator for reproducibility"),
|
| 349 |
InputParam("diffusion_batch_size", default=1, type_hint=int, description="Number of samples to generate in parallel"),
|
| 350 |
InputParam("L", required=True, type_hint=int, description="Protein length"),
|
| 351 |
+
InputParam("f", required=True, type_hint=dict),
|
| 352 |
+
InputParam("coord_atom_lvl_to_be_noised", required=True, type_hint=torch.Tensor),
|
| 353 |
InputParam("motif_mask", required=True, type_hint=torch.Tensor),
|
|
|
|
| 354 |
InputParam("noise_schedule", required=True, type_hint=torch.Tensor),
|
|
|
|
| 355 |
]
|
| 356 |
|
| 357 |
@property
|
| 358 |
def intermediate_outputs(self) -> List[OutputParam]:
|
| 359 |
return [
|
| 360 |
+
OutputParam("xyz", type_hint=torch.Tensor, description="Initial noised coordinates [D, L_atom, 3]"),
|
| 361 |
+
OutputParam("initializer_outputs", type_hint=dict, description="Embedded conditioning, reused every step"),
|
| 362 |
]
|
| 363 |
|
| 364 |
@torch.no_grad()
|
| 365 |
def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
|
| 366 |
block_state = self.get_block_state(state)
|
| 367 |
|
|
|
|
|
|
|
|
|
|
| 368 |
noise_schedule = block_state.noise_schedule
|
|
|
|
| 369 |
generator = block_state.generator
|
| 370 |
D = block_state.diffusion_batch_size or 1
|
| 371 |
+
device = components.transformer.device
|
| 372 |
+
|
| 373 |
+
# The feature dict is built on CPU and is read on every denoising step, so move it once.
|
| 374 |
+
f = {k: v.to(device) if torch.is_tensor(v) else v for k, v in block_state.f.items()}
|
| 375 |
+
coord = block_state.coord_atom_lvl_to_be_noised.to(device)
|
| 376 |
+
motif_mask = f["is_motif_atom_with_fixed_coord"]
|
| 377 |
+
|
| 378 |
+
# Matches rfd3.model.inference_sampler._get_initial_structure:
|
| 379 |
+
# noise = c0 * randn(D, L, 3); noise[..., is_motif, :] = 0; X_L = noise + coord
|
| 380 |
+
c0 = noise_schedule[0].to(device)
|
| 381 |
+
L_atom = coord.shape[-2]
|
| 382 |
+
noise = c0 * torch.randn((D, L_atom, 3), device=device, generator=generator)
|
| 383 |
+
noise[..., motif_mask, :] = 0.0
|
| 384 |
+
xyz = noise + coord
|
| 385 |
+
|
| 386 |
+
block_state.f = f
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 387 |
block_state.xyz = xyz
|
| 388 |
+
block_state.initializer_outputs = components.transformer.encode_conditioning(f)
|
| 389 |
|
| 390 |
self.set_block_state(state, block_state)
|
| 391 |
return components, state
|
denoise.py
CHANGED
|
@@ -80,15 +80,18 @@ class RFDiffusionDenoiseStep(ModularPipelineBlocks):
|
|
| 80 |
type_hint=int,
|
| 81 |
description="Frequency of callback invocation",
|
| 82 |
),
|
| 83 |
-
InputParam("xyz", required=True, type_hint=torch.Tensor, description="Initial noised coordinates [D,
|
| 84 |
InputParam("noise_schedule", required=True, type_hint=torch.Tensor, description="EDM noise schedule"),
|
| 85 |
InputParam("motif_mask", required=True, type_hint=torch.Tensor, description="Mask for fixed motif positions"),
|
|
|
|
|
|
|
| 86 |
]
|
| 87 |
|
| 88 |
@property
|
| 89 |
def intermediate_outputs(self) -> List[OutputParam]:
|
| 90 |
return [
|
| 91 |
OutputParam("xyz", type_hint=torch.Tensor, description="Denoised coordinates [D, L, 3]"),
|
|
|
|
| 92 |
OutputParam("single", type_hint=torch.Tensor, description="Single representation"),
|
| 93 |
OutputParam("pair", type_hint=torch.Tensor, description="Pair representation"),
|
| 94 |
OutputParam("sequence_logits", type_hint=torch.Tensor, description="Predicted sequence logits"),
|
|
@@ -98,11 +101,27 @@ class RFDiffusionDenoiseStep(ModularPipelineBlocks):
|
|
| 98 |
|
| 99 |
@torch.no_grad()
|
| 100 |
def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
block_state = self.get_block_state(state)
|
| 102 |
|
| 103 |
xyz = block_state.xyz
|
| 104 |
noise_schedule = block_state.noise_schedule
|
| 105 |
motif_mask = block_state.motif_mask
|
|
|
|
|
|
|
| 106 |
|
| 107 |
n_recycle = block_state.n_recycle
|
| 108 |
callback = block_state.callback
|
|
@@ -123,9 +142,6 @@ class RFDiffusionDenoiseStep(ModularPipelineBlocks):
|
|
| 123 |
sequence_logits = None
|
| 124 |
sequence_indices = None
|
| 125 |
|
| 126 |
-
has_transformer = hasattr(components, "transformer") and components.transformer is not None
|
| 127 |
-
has_scheduler = hasattr(components, "scheduler") and components.scheduler is not None
|
| 128 |
-
|
| 129 |
# Iterate over consecutive pairs (c_t_minus_1, c_t) in the noise schedule
|
| 130 |
# noise_schedule goes from high noise to low noise
|
| 131 |
for step_num in range(len(noise_schedule) - 1):
|
|
@@ -133,61 +149,65 @@ class RFDiffusionDenoiseStep(ModularPipelineBlocks):
|
|
| 133 |
c_t = noise_schedule[step_num + 1]
|
| 134 |
|
| 135 |
# Step 1: Inject stochastic noise (matching original sampler)
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
)
|
| 140 |
-
else:
|
| 141 |
-
X_noisy_L = X_L
|
| 142 |
-
t_hat = c_t_minus_1
|
| 143 |
|
| 144 |
# Step 2: Model forward pass
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
else:
|
| 163 |
-
X_denoised_L = X_noisy_L
|
| 164 |
|
| 165 |
# Step 3: Euler update with step_scale (matching original sampler)
|
| 166 |
-
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
)
|
| 174 |
-
else:
|
| 175 |
-
# Fallback simple Euler step
|
| 176 |
-
delta_L = (X_noisy_L - X_denoised_L) / (t_hat + 1e-8)
|
| 177 |
-
d_t = c_t - t_hat
|
| 178 |
-
X_L = X_noisy_L + d_t * delta_L
|
| 179 |
|
| 180 |
X_denoised_L_traj.append(X_denoised_L.clone())
|
| 181 |
|
| 182 |
if callback is not None and step_num % callback_steps == 0:
|
| 183 |
callback(step_num, c_t_minus_1, X_L)
|
| 184 |
|
| 185 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 186 |
block_state.single = single
|
| 187 |
block_state.pair = pair
|
| 188 |
block_state.sequence_logits = sequence_logits
|
| 189 |
block_state.sequence_indices = sequence_indices
|
| 190 |
-
block_state.trajectory = X_denoised_L_traj
|
| 191 |
|
| 192 |
self.set_block_state(state, block_state)
|
| 193 |
return components, state
|
|
|
|
| 80 |
type_hint=int,
|
| 81 |
description="Frequency of callback invocation",
|
| 82 |
),
|
| 83 |
+
InputParam("xyz", required=True, type_hint=torch.Tensor, description="Initial noised coordinates [D, L_atom, 3]"),
|
| 84 |
InputParam("noise_schedule", required=True, type_hint=torch.Tensor, description="EDM noise schedule"),
|
| 85 |
InputParam("motif_mask", required=True, type_hint=torch.Tensor, description="Mask for fixed motif positions"),
|
| 86 |
+
InputParam("f", required=True, type_hint=dict, description="Foundry feature dict"),
|
| 87 |
+
InputParam("initializer_outputs", required=True, type_hint=dict, description="Embedded conditioning"),
|
| 88 |
]
|
| 89 |
|
| 90 |
@property
|
| 91 |
def intermediate_outputs(self) -> List[OutputParam]:
|
| 92 |
return [
|
| 93 |
OutputParam("xyz", type_hint=torch.Tensor, description="Denoised coordinates [D, L, 3]"),
|
| 94 |
+
OutputParam("xyz_backbone", type_hint=torch.Tensor, description="Backbone N/CA/C/O coordinates [D, L, 4, 3]"),
|
| 95 |
OutputParam("single", type_hint=torch.Tensor, description="Single representation"),
|
| 96 |
OutputParam("pair", type_hint=torch.Tensor, description="Pair representation"),
|
| 97 |
OutputParam("sequence_logits", type_hint=torch.Tensor, description="Predicted sequence logits"),
|
|
|
|
| 101 |
|
| 102 |
@torch.no_grad()
|
| 103 |
def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
|
| 104 |
+
# Both components are required. Falling back to a hand-rolled Euler step when the
|
| 105 |
+
# scheduler is missing silently swaps out the EDM sampler and yields coordinates that
|
| 106 |
+
# look finite but are not a protein backbone.
|
| 107 |
+
if components.transformer is None:
|
| 108 |
+
raise ValueError(
|
| 109 |
+
"`transformer` is not loaded. Call `load_components(trust_remote_code=True)` on the "
|
| 110 |
+
"pipeline before calling it."
|
| 111 |
+
)
|
| 112 |
+
if components.scheduler is None:
|
| 113 |
+
raise ValueError(
|
| 114 |
+
"`scheduler` is not loaded. Call `load_components(trust_remote_code=True)` on the "
|
| 115 |
+
"pipeline before calling it."
|
| 116 |
+
)
|
| 117 |
+
|
| 118 |
block_state = self.get_block_state(state)
|
| 119 |
|
| 120 |
xyz = block_state.xyz
|
| 121 |
noise_schedule = block_state.noise_schedule
|
| 122 |
motif_mask = block_state.motif_mask
|
| 123 |
+
f = block_state.f
|
| 124 |
+
initializer_outputs = block_state.initializer_outputs
|
| 125 |
|
| 126 |
n_recycle = block_state.n_recycle
|
| 127 |
callback = block_state.callback
|
|
|
|
| 142 |
sequence_logits = None
|
| 143 |
sequence_indices = None
|
| 144 |
|
|
|
|
|
|
|
|
|
|
| 145 |
# Iterate over consecutive pairs (c_t_minus_1, c_t) in the noise schedule
|
| 146 |
# noise_schedule goes from high noise to low noise
|
| 147 |
for step_num in range(len(noise_schedule) - 1):
|
|
|
|
| 149 |
c_t = noise_schedule[step_num + 1]
|
| 150 |
|
| 151 |
# Step 1: Inject stochastic noise (matching original sampler)
|
| 152 |
+
X_noisy_L, t_hat = components.scheduler.add_noise(
|
| 153 |
+
X_L, c_t_minus_1, c_t, motif_mask=motif_mask
|
| 154 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 155 |
|
| 156 |
# Step 2: Model forward pass
|
| 157 |
+
# t_hat is a scalar, tile to batch dimension
|
| 158 |
+
t_batch = (t_hat.to(device).expand(D) if isinstance(t_hat, torch.Tensor)
|
| 159 |
+
else torch.full((D,), t_hat, device=device))
|
| 160 |
+
|
| 161 |
+
output = components.transformer(
|
| 162 |
+
xyz_noisy=X_noisy_L,
|
| 163 |
+
t=t_batch,
|
| 164 |
+
f=f,
|
| 165 |
+
initializer_outputs=initializer_outputs,
|
| 166 |
+
n_recycle=n_recycle,
|
| 167 |
+
)
|
| 168 |
+
|
| 169 |
+
X_denoised_L = output.xyz
|
| 170 |
+
single = output.single
|
| 171 |
+
pair = output.pair
|
| 172 |
+
sequence_logits = output.sequence_logits
|
| 173 |
+
sequence_indices = output.sequence_indices
|
|
|
|
|
|
|
| 174 |
|
| 175 |
# Step 3: Euler update with step_scale (matching original sampler)
|
| 176 |
+
X_L = components.scheduler.step(
|
| 177 |
+
xyz_pred=X_denoised_L,
|
| 178 |
+
xyz_noisy=X_noisy_L,
|
| 179 |
+
c_t_minus_1=c_t_minus_1,
|
| 180 |
+
c_t=c_t,
|
| 181 |
+
motif_mask=motif_mask,
|
| 182 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 183 |
|
| 184 |
X_denoised_L_traj.append(X_denoised_L.clone())
|
| 185 |
|
| 186 |
if callback is not None and step_num % callback_steps == 0:
|
| 187 |
callback(step_num, c_t_minus_1, X_L)
|
| 188 |
|
| 189 |
+
# The sampler runs on padded atom-level coordinates (14 slots per residue). Downstream
|
| 190 |
+
# blocks and the documented output are one point per residue, so collapse to CA here.
|
| 191 |
+
is_ca = f["is_ca"]
|
| 192 |
+
block_state.xyz = X_L[:, is_ca]
|
| 193 |
+
|
| 194 |
+
# MPNN wants N, CA, C and O per residue. They are present in the padded tensor, so select
|
| 195 |
+
# them instead of synthesising them from CA: foundry's `BACKBONE_ATOM_NAMES` is
|
| 196 |
+
# ["N", "CA", "C", "O"], which is also the CCD atom order for amino acids.
|
| 197 |
+
is_backbone = f["is_backbone"]
|
| 198 |
+
n_residues = int(is_ca.sum())
|
| 199 |
+
n_backbone = int(is_backbone.sum())
|
| 200 |
+
if n_backbone != 4 * n_residues:
|
| 201 |
+
raise ValueError(
|
| 202 |
+
f"Expected 4 backbone atoms per residue, got {n_backbone} for {n_residues} residues. "
|
| 203 |
+
"The padded atom layout does not match the N/CA/C/O ordering MPNN assumes."
|
| 204 |
+
)
|
| 205 |
+
block_state.xyz_backbone = X_L[:, is_backbone].reshape(X_L.shape[0], n_residues, 4, 3)
|
| 206 |
block_state.single = single
|
| 207 |
block_state.pair = pair
|
| 208 |
block_state.sequence_logits = sequence_logits
|
| 209 |
block_state.sequence_indices = sequence_indices
|
| 210 |
+
block_state.trajectory = [step[:, is_ca] for step in X_denoised_L_traj]
|
| 211 |
|
| 212 |
self.set_block_state(state, block_state)
|
| 213 |
return components, state
|
modular_blocks.py
CHANGED
|
@@ -175,8 +175,12 @@ class MPNNSequenceDesignStep(ModularPipelineBlocks):
|
|
| 175 |
description="Protein backbone coordinates [B, L, 3] (CA atoms)",
|
| 176 |
),
|
| 177 |
InputParam(
|
| 178 |
-
"
|
| 179 |
-
description="
|
|
|
|
|
|
|
|
|
|
|
|
|
| 180 |
),
|
| 181 |
InputParam(
|
| 182 |
"sequence_indices", type_hint=torch.Tensor,
|
|
@@ -214,7 +218,8 @@ class MPNNSequenceDesignStep(ModularPipelineBlocks):
|
|
| 214 |
block_state = self.get_block_state(state)
|
| 215 |
|
| 216 |
xyz = block_state.xyz
|
| 217 |
-
|
|
|
|
| 218 |
known_seq = block_state.sequence_indices
|
| 219 |
temperature = block_state.temperature or 0.1
|
| 220 |
output_type = block_state.output_type or "tensor"
|
|
@@ -227,7 +232,7 @@ class MPNNSequenceDesignStep(ModularPipelineBlocks):
|
|
| 227 |
|
| 228 |
if has_mpnn:
|
| 229 |
sequence_logits, sequence_indices = self._run_mpnn(
|
| 230 |
-
components.mpnn,
|
| 231 |
)
|
| 232 |
else:
|
| 233 |
if known_seq is not None:
|
|
@@ -264,32 +269,30 @@ class MPNNSequenceDesignStep(ModularPipelineBlocks):
|
|
| 264 |
self.set_block_state(state, block_state)
|
| 265 |
return components, state
|
| 266 |
|
| 267 |
-
def _run_mpnn(self, mpnn,
|
| 268 |
-
"""Run the MPNNModel wrapper on
|
| 269 |
-
B
|
| 270 |
-
|
| 271 |
-
dtype = xyz.dtype
|
| 272 |
-
|
| 273 |
-
ca = xyz
|
| 274 |
-
n_offset = torch.tensor([-1.458, 0.0, 0.0], device=device, dtype=dtype)
|
| 275 |
-
c_offset = torch.tensor([0.550, 1.424, 0.0], device=device, dtype=dtype)
|
| 276 |
-
o_offset = torch.tensor([0.550, 2.500, 0.0], device=device, dtype=dtype)
|
| 277 |
-
|
| 278 |
-
X = torch.stack([
|
| 279 |
-
ca + n_offset, ca, ca + c_offset, ca + o_offset,
|
| 280 |
-
], dim=2)
|
| 281 |
|
| 282 |
if motif_mask is not None:
|
| 283 |
designed_mask = ~motif_mask.unsqueeze(0).expand(B, -1)
|
| 284 |
else:
|
| 285 |
designed_mask = None
|
| 286 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 287 |
output = mpnn(
|
| 288 |
-
X=
|
|
|
|
|
|
|
|
|
|
| 289 |
)
|
| 290 |
|
| 291 |
-
logits = output.sequence_logits
|
| 292 |
-
indices = output.sequence_indices
|
| 293 |
|
| 294 |
if motif_mask is not None and known_seq is not None:
|
| 295 |
indices[:, motif_mask] = known_seq[:, motif_mask]
|
|
|
|
| 175 |
description="Protein backbone coordinates [B, L, 3] (CA atoms)",
|
| 176 |
),
|
| 177 |
InputParam(
|
| 178 |
+
"xyz_backbone", required=True, type_hint=torch.Tensor,
|
| 179 |
+
description="Backbone N/CA/C/O coordinates [B, L, 4, 3]",
|
| 180 |
+
),
|
| 181 |
+
InputParam(
|
| 182 |
+
"motif_token_mask", type_hint=torch.Tensor,
|
| 183 |
+
description="Residue-level mask for fixed/motif positions [L]. True = fixed sequence.",
|
| 184 |
),
|
| 185 |
InputParam(
|
| 186 |
"sequence_indices", type_hint=torch.Tensor,
|
|
|
|
| 218 |
block_state = self.get_block_state(state)
|
| 219 |
|
| 220 |
xyz = block_state.xyz
|
| 221 |
+
xyz_backbone = block_state.xyz_backbone
|
| 222 |
+
motif_mask = block_state.motif_token_mask
|
| 223 |
known_seq = block_state.sequence_indices
|
| 224 |
temperature = block_state.temperature or 0.1
|
| 225 |
output_type = block_state.output_type or "tensor"
|
|
|
|
| 232 |
|
| 233 |
if has_mpnn:
|
| 234 |
sequence_logits, sequence_indices = self._run_mpnn(
|
| 235 |
+
components.mpnn, xyz_backbone, motif_mask, known_seq, temperature,
|
| 236 |
)
|
| 237 |
else:
|
| 238 |
if known_seq is not None:
|
|
|
|
| 269 |
self.set_block_state(state, block_state)
|
| 270 |
return components, state
|
| 271 |
|
| 272 |
+
def _run_mpnn(self, mpnn, xyz_backbone, motif_mask, known_seq, temperature):
|
| 273 |
+
"""Run the MPNNModel wrapper on the sampler's N/CA/C/O coordinates."""
|
| 274 |
+
B = xyz_backbone.shape[0]
|
| 275 |
+
out_device = xyz_backbone.device
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 276 |
|
| 277 |
if motif_mask is not None:
|
| 278 |
designed_mask = ~motif_mask.unsqueeze(0).expand(B, -1)
|
| 279 |
else:
|
| 280 |
designed_mask = None
|
| 281 |
|
| 282 |
+
# MPNN is attached with `update_components` and is often left wherever it loaded, which
|
| 283 |
+
# need not be the device the sampler ran on. It is a 1.7M parameter model, so following it
|
| 284 |
+
# is cheaper than making the caller place it. Only the call arguments move; everything used
|
| 285 |
+
# afterwards stays on the sampler's device.
|
| 286 |
+
mpnn_device = next(mpnn.parameters()).device
|
| 287 |
output = mpnn(
|
| 288 |
+
X=xyz_backbone.to(mpnn_device),
|
| 289 |
+
S=None if known_seq is None else known_seq.to(mpnn_device),
|
| 290 |
+
designed_residue_mask=None if designed_mask is None else designed_mask.to(mpnn_device),
|
| 291 |
+
temperature=temperature,
|
| 292 |
)
|
| 293 |
|
| 294 |
+
logits = output.sequence_logits.to(out_device)
|
| 295 |
+
indices = output.sequence_indices.to(out_device)
|
| 296 |
|
| 297 |
if motif_mask is not None and known_seq is not None:
|
| 298 |
indices[:, motif_mask] = known_seq[:, motif_mask]
|
modular_model_index.json
CHANGED
|
@@ -30,8 +30,7 @@
|
|
| 30 |
"AutoModel"
|
| 31 |
],
|
| 32 |
"revision": null,
|
| 33 |
-
"variant": null
|
| 34 |
-
"default_creation_method": "from_config"
|
| 35 |
}
|
| 36 |
]
|
| 37 |
}
|
|
|
|
| 30 |
"AutoModel"
|
| 31 |
],
|
| 32 |
"revision": null,
|
| 33 |
+
"variant": null
|
|
|
|
| 34 |
}
|
| 35 |
]
|
| 36 |
}
|
mpnn/model_mpnn.py
CHANGED
|
@@ -39,6 +39,27 @@ MODEL_CLASSES = {
|
|
| 39 |
}
|
| 40 |
|
| 41 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 42 |
@dataclass
|
| 43 |
class MPNNModelOutput:
|
| 44 |
"""Output from the MPNN model wrapper."""
|
|
@@ -130,7 +151,7 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 130 |
designed_residue_mask: Which residues to design [B, L] (default: all).
|
| 131 |
chain_labels: Chain identifiers [B, L] (default: single chain).
|
| 132 |
R_idx: Residue indices [B, L] (default: 0..L-1).
|
| 133 |
-
temperature: Sampling temperature (default: 0.1).
|
| 134 |
|
| 135 |
Returns:
|
| 136 |
MPNNModelOutput with sequence logits and sampled indices.
|
|
@@ -152,7 +173,12 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 152 |
# Atom mask: mark all atoms as valid based on coordinate presence
|
| 153 |
X_m = (X.abs().sum(dim=-1) > 0).float() # [B, L, num_atoms]
|
| 154 |
|
| 155 |
-
network_input
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 156 |
"X": X,
|
| 157 |
"X_m": X_m,
|
| 158 |
"S": S,
|
|
@@ -160,11 +186,16 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 160 |
"chain_labels": chain_labels,
|
| 161 |
"residue_mask": residue_mask,
|
| 162 |
"designed_residue_mask": designed_residue_mask,
|
|
|
|
| 163 |
"temperature": temperature,
|
|
|
|
| 164 |
**kwargs,
|
| 165 |
}
|
|
|
|
|
|
|
|
|
|
| 166 |
|
| 167 |
-
output = self.model(
|
| 168 |
|
| 169 |
logits = output["decoder_features"]["logits"] # [B, L, n_vocab]
|
| 170 |
S_sampled = output["decoder_features"].get(
|
|
|
|
| 39 |
}
|
| 40 |
|
| 41 |
|
| 42 |
+
# Scalar decoding settings, matching MPNN_PER_INPUT_INFERENCE_DEFAULTS in foundry's
|
| 43 |
+
# mpnn/utils/inference.py. foundry's feature-aggregation transform normally fills these in;
|
| 44 |
+
# calling the network directly means supplying them here.
|
| 45 |
+
_INFERENCE_DECODE_SETTINGS = {
|
| 46 |
+
"structure_noise": 0.0,
|
| 47 |
+
"decode_type": "auto_regressive",
|
| 48 |
+
"causality_pattern": "auto_regressive",
|
| 49 |
+
"initialize_sequence_embedding_with_ground_truth": False,
|
| 50 |
+
"atomize_side_chains": False,
|
| 51 |
+
"features_to_return": None,
|
| 52 |
+
"repeat_sample_num": None,
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
_OPTIONAL_CONDITIONING_KEYS = (
|
| 56 |
+
"bias",
|
| 57 |
+
"pair_bias",
|
| 58 |
+
"symmetry_equivalence_group",
|
| 59 |
+
"symmetry_weight",
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
@dataclass
|
| 64 |
class MPNNModelOutput:
|
| 65 |
"""Output from the MPNN model wrapper."""
|
|
|
|
| 151 |
designed_residue_mask: Which residues to design [B, L] (default: all).
|
| 152 |
chain_labels: Chain identifiers [B, L] (default: single chain).
|
| 153 |
R_idx: Residue indices [B, L] (default: 0..L-1).
|
| 154 |
+
temperature: Sampling temperature, scalar or per-residue [B, L] (default: 0.1).
|
| 155 |
|
| 156 |
Returns:
|
| 157 |
MPNNModelOutput with sequence logits and sampled indices.
|
|
|
|
| 173 |
# Atom mask: mark all atoms as valid based on coordinate presence
|
| 174 |
X_m = (X.abs().sum(dim=-1) > 0).float() # [B, L, num_atoms]
|
| 175 |
|
| 176 |
+
# foundry reads every tensor from `network_input["input_features"]`, and wants a
|
| 177 |
+
# per-residue temperature rather than a scalar.
|
| 178 |
+
if not torch.is_tensor(temperature):
|
| 179 |
+
temperature = torch.full((B, L), float(temperature), device=device)
|
| 180 |
+
|
| 181 |
+
input_features = {
|
| 182 |
"X": X,
|
| 183 |
"X_m": X_m,
|
| 184 |
"S": S,
|
|
|
|
| 186 |
"chain_labels": chain_labels,
|
| 187 |
"residue_mask": residue_mask,
|
| 188 |
"designed_residue_mask": designed_residue_mask,
|
| 189 |
+
"mask_for_loss": residue_mask,
|
| 190 |
"temperature": temperature,
|
| 191 |
+
**_INFERENCE_DECODE_SETTINGS,
|
| 192 |
**kwargs,
|
| 193 |
}
|
| 194 |
+
# The network checks these keys are present but accepts None for all of them.
|
| 195 |
+
for key in _OPTIONAL_CONDITIONING_KEYS:
|
| 196 |
+
input_features.setdefault(key, None)
|
| 197 |
|
| 198 |
+
output = self.model({"input_features": input_features})
|
| 199 |
|
| 200 |
logits = output["decoder_features"]["logits"] # [B, L, n_vocab]
|
| 201 |
S_sampled = output["decoder_features"].get(
|
mpnn_ligand/model_mpnn.py
CHANGED
|
@@ -39,6 +39,27 @@ MODEL_CLASSES = {
|
|
| 39 |
}
|
| 40 |
|
| 41 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 42 |
@dataclass
|
| 43 |
class MPNNModelOutput:
|
| 44 |
"""Output from the MPNN model wrapper."""
|
|
@@ -130,7 +151,7 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 130 |
designed_residue_mask: Which residues to design [B, L] (default: all).
|
| 131 |
chain_labels: Chain identifiers [B, L] (default: single chain).
|
| 132 |
R_idx: Residue indices [B, L] (default: 0..L-1).
|
| 133 |
-
temperature: Sampling temperature (default: 0.1).
|
| 134 |
|
| 135 |
Returns:
|
| 136 |
MPNNModelOutput with sequence logits and sampled indices.
|
|
@@ -152,7 +173,12 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 152 |
# Atom mask: mark all atoms as valid based on coordinate presence
|
| 153 |
X_m = (X.abs().sum(dim=-1) > 0).float() # [B, L, num_atoms]
|
| 154 |
|
| 155 |
-
network_input
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 156 |
"X": X,
|
| 157 |
"X_m": X_m,
|
| 158 |
"S": S,
|
|
@@ -160,11 +186,16 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 160 |
"chain_labels": chain_labels,
|
| 161 |
"residue_mask": residue_mask,
|
| 162 |
"designed_residue_mask": designed_residue_mask,
|
|
|
|
| 163 |
"temperature": temperature,
|
|
|
|
| 164 |
**kwargs,
|
| 165 |
}
|
|
|
|
|
|
|
|
|
|
| 166 |
|
| 167 |
-
output = self.model(
|
| 168 |
|
| 169 |
logits = output["decoder_features"]["logits"] # [B, L, n_vocab]
|
| 170 |
S_sampled = output["decoder_features"].get(
|
|
|
|
| 39 |
}
|
| 40 |
|
| 41 |
|
| 42 |
+
# Scalar decoding settings, matching MPNN_PER_INPUT_INFERENCE_DEFAULTS in foundry's
|
| 43 |
+
# mpnn/utils/inference.py. foundry's feature-aggregation transform normally fills these in;
|
| 44 |
+
# calling the network directly means supplying them here.
|
| 45 |
+
_INFERENCE_DECODE_SETTINGS = {
|
| 46 |
+
"structure_noise": 0.0,
|
| 47 |
+
"decode_type": "auto_regressive",
|
| 48 |
+
"causality_pattern": "auto_regressive",
|
| 49 |
+
"initialize_sequence_embedding_with_ground_truth": False,
|
| 50 |
+
"atomize_side_chains": False,
|
| 51 |
+
"features_to_return": None,
|
| 52 |
+
"repeat_sample_num": None,
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
_OPTIONAL_CONDITIONING_KEYS = (
|
| 56 |
+
"bias",
|
| 57 |
+
"pair_bias",
|
| 58 |
+
"symmetry_equivalence_group",
|
| 59 |
+
"symmetry_weight",
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
@dataclass
|
| 64 |
class MPNNModelOutput:
|
| 65 |
"""Output from the MPNN model wrapper."""
|
|
|
|
| 151 |
designed_residue_mask: Which residues to design [B, L] (default: all).
|
| 152 |
chain_labels: Chain identifiers [B, L] (default: single chain).
|
| 153 |
R_idx: Residue indices [B, L] (default: 0..L-1).
|
| 154 |
+
temperature: Sampling temperature, scalar or per-residue [B, L] (default: 0.1).
|
| 155 |
|
| 156 |
Returns:
|
| 157 |
MPNNModelOutput with sequence logits and sampled indices.
|
|
|
|
| 173 |
# Atom mask: mark all atoms as valid based on coordinate presence
|
| 174 |
X_m = (X.abs().sum(dim=-1) > 0).float() # [B, L, num_atoms]
|
| 175 |
|
| 176 |
+
# foundry reads every tensor from `network_input["input_features"]`, and wants a
|
| 177 |
+
# per-residue temperature rather than a scalar.
|
| 178 |
+
if not torch.is_tensor(temperature):
|
| 179 |
+
temperature = torch.full((B, L), float(temperature), device=device)
|
| 180 |
+
|
| 181 |
+
input_features = {
|
| 182 |
"X": X,
|
| 183 |
"X_m": X_m,
|
| 184 |
"S": S,
|
|
|
|
| 186 |
"chain_labels": chain_labels,
|
| 187 |
"residue_mask": residue_mask,
|
| 188 |
"designed_residue_mask": designed_residue_mask,
|
| 189 |
+
"mask_for_loss": residue_mask,
|
| 190 |
"temperature": temperature,
|
| 191 |
+
**_INFERENCE_DECODE_SETTINGS,
|
| 192 |
**kwargs,
|
| 193 |
}
|
| 194 |
+
# The network checks these keys are present but accepts None for all of them.
|
| 195 |
+
for key in _OPTIONAL_CONDITIONING_KEYS:
|
| 196 |
+
input_features.setdefault(key, None)
|
| 197 |
|
| 198 |
+
output = self.model({"input_features": input_features})
|
| 199 |
|
| 200 |
logits = output["decoder_features"]["logits"] # [B, L, n_vocab]
|
| 201 |
S_sampled = output["decoder_features"].get(
|
mpnn_soluble/model_mpnn.py
CHANGED
|
@@ -39,6 +39,27 @@ MODEL_CLASSES = {
|
|
| 39 |
}
|
| 40 |
|
| 41 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 42 |
@dataclass
|
| 43 |
class MPNNModelOutput:
|
| 44 |
"""Output from the MPNN model wrapper."""
|
|
@@ -130,7 +151,7 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 130 |
designed_residue_mask: Which residues to design [B, L] (default: all).
|
| 131 |
chain_labels: Chain identifiers [B, L] (default: single chain).
|
| 132 |
R_idx: Residue indices [B, L] (default: 0..L-1).
|
| 133 |
-
temperature: Sampling temperature (default: 0.1).
|
| 134 |
|
| 135 |
Returns:
|
| 136 |
MPNNModelOutput with sequence logits and sampled indices.
|
|
@@ -152,7 +173,12 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 152 |
# Atom mask: mark all atoms as valid based on coordinate presence
|
| 153 |
X_m = (X.abs().sum(dim=-1) > 0).float() # [B, L, num_atoms]
|
| 154 |
|
| 155 |
-
network_input
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 156 |
"X": X,
|
| 157 |
"X_m": X_m,
|
| 158 |
"S": S,
|
|
@@ -160,11 +186,16 @@ class MPNNModel(ModelMixin, ConfigMixin):
|
|
| 160 |
"chain_labels": chain_labels,
|
| 161 |
"residue_mask": residue_mask,
|
| 162 |
"designed_residue_mask": designed_residue_mask,
|
|
|
|
| 163 |
"temperature": temperature,
|
|
|
|
| 164 |
**kwargs,
|
| 165 |
}
|
|
|
|
|
|
|
|
|
|
| 166 |
|
| 167 |
-
output = self.model(
|
| 168 |
|
| 169 |
logits = output["decoder_features"]["logits"] # [B, L, n_vocab]
|
| 170 |
S_sampled = output["decoder_features"].get(
|
|
|
|
| 39 |
}
|
| 40 |
|
| 41 |
|
| 42 |
+
# Scalar decoding settings, matching MPNN_PER_INPUT_INFERENCE_DEFAULTS in foundry's
|
| 43 |
+
# mpnn/utils/inference.py. foundry's feature-aggregation transform normally fills these in;
|
| 44 |
+
# calling the network directly means supplying them here.
|
| 45 |
+
_INFERENCE_DECODE_SETTINGS = {
|
| 46 |
+
"structure_noise": 0.0,
|
| 47 |
+
"decode_type": "auto_regressive",
|
| 48 |
+
"causality_pattern": "auto_regressive",
|
| 49 |
+
"initialize_sequence_embedding_with_ground_truth": False,
|
| 50 |
+
"atomize_side_chains": False,
|
| 51 |
+
"features_to_return": None,
|
| 52 |
+
"repeat_sample_num": None,
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
_OPTIONAL_CONDITIONING_KEYS = (
|
| 56 |
+
"bias",
|
| 57 |
+
"pair_bias",
|
| 58 |
+
"symmetry_equivalence_group",
|
| 59 |
+
"symmetry_weight",
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
@dataclass
|
| 64 |
class MPNNModelOutput:
|
| 65 |
"""Output from the MPNN model wrapper."""
|
|
|
|
| 151 |
designed_residue_mask: Which residues to design [B, L] (default: all).
|
| 152 |
chain_labels: Chain identifiers [B, L] (default: single chain).
|
| 153 |
R_idx: Residue indices [B, L] (default: 0..L-1).
|
| 154 |
+
temperature: Sampling temperature, scalar or per-residue [B, L] (default: 0.1).
|
| 155 |
|
| 156 |
Returns:
|
| 157 |
MPNNModelOutput with sequence logits and sampled indices.
|
|
|
|
| 173 |
# Atom mask: mark all atoms as valid based on coordinate presence
|
| 174 |
X_m = (X.abs().sum(dim=-1) > 0).float() # [B, L, num_atoms]
|
| 175 |
|
| 176 |
+
# foundry reads every tensor from `network_input["input_features"]`, and wants a
|
| 177 |
+
# per-residue temperature rather than a scalar.
|
| 178 |
+
if not torch.is_tensor(temperature):
|
| 179 |
+
temperature = torch.full((B, L), float(temperature), device=device)
|
| 180 |
+
|
| 181 |
+
input_features = {
|
| 182 |
"X": X,
|
| 183 |
"X_m": X_m,
|
| 184 |
"S": S,
|
|
|
|
| 186 |
"chain_labels": chain_labels,
|
| 187 |
"residue_mask": residue_mask,
|
| 188 |
"designed_residue_mask": designed_residue_mask,
|
| 189 |
+
"mask_for_loss": residue_mask,
|
| 190 |
"temperature": temperature,
|
| 191 |
+
**_INFERENCE_DECODE_SETTINGS,
|
| 192 |
**kwargs,
|
| 193 |
}
|
| 194 |
+
# The network checks these keys are present but accepts None for all of them.
|
| 195 |
+
for key in _OPTIONAL_CONDITIONING_KEYS:
|
| 196 |
+
input_features.setdefault(key, None)
|
| 197 |
|
| 198 |
+
output = self.model({"input_features": input_features})
|
| 199 |
|
| 200 |
logits = output["decoder_features"]["logits"] # [B, L, n_vocab]
|
| 201 |
S_sampled = output["decoder_features"].get(
|
scheduler/model.py
CHANGED
|
@@ -24,12 +24,19 @@ from typing import Optional
|
|
| 24 |
import torch
|
| 25 |
|
| 26 |
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
|
|
|
| 27 |
|
| 28 |
# Reuse the original noise schedule and sampling config directly
|
| 29 |
from rfd3.model.inference_sampler import SampleDiffusionWithMotif
|
| 30 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
|
| 32 |
-
|
|
|
|
| 33 |
"""
|
| 34 |
Diffusers-compatible scheduler wrapping the foundry EDM sampler.
|
| 35 |
|
|
@@ -65,6 +72,13 @@ class RFDiffusionScheduler(ConfigMixin):
|
|
| 65 |
step_scale=step_scale,
|
| 66 |
)
|
| 67 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 68 |
@property
|
| 69 |
def sampler(self) -> SampleDiffusionWithMotif:
|
| 70 |
return self._sampler
|
|
|
|
| 24 |
import torch
|
| 25 |
|
| 26 |
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
| 27 |
+
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
| 28 |
|
| 29 |
# Reuse the original noise schedule and sampling config directly
|
| 30 |
from rfd3.model.inference_sampler import SampleDiffusionWithMotif
|
| 31 |
|
| 32 |
+
# `ComponentSpec.load` only strips `dtype` for a `type_hint` that is not a torch module, and
|
| 33 |
+
# modular_model_index.json records this scheduler as `AutoModel`, so the loader forwards its
|
| 34 |
+
# weight-placement kwargs here. None of them are constructor arguments, and letting
|
| 35 |
+
# `register_to_config` capture `dtype` would write a torch.dtype into config.json.
|
| 36 |
+
_LOADER_ONLY_KWARGS = ("dtype", "torch_dtype", "device_map", "variant", "trust_remote_code", "low_cpu_mem_usage")
|
| 37 |
|
| 38 |
+
|
| 39 |
+
class RFDiffusionScheduler(SchedulerMixin, ConfigMixin):
|
| 40 |
"""
|
| 41 |
Diffusers-compatible scheduler wrapping the foundry EDM sampler.
|
| 42 |
|
|
|
|
| 72 |
step_scale=step_scale,
|
| 73 |
)
|
| 74 |
|
| 75 |
+
@classmethod
|
| 76 |
+
def from_pretrained(cls, pretrained_model_name_or_path=None, subfolder=None, **kwargs):
|
| 77 |
+
for key in _LOADER_ONLY_KWARGS:
|
| 78 |
+
kwargs.pop(key, None)
|
| 79 |
+
|
| 80 |
+
return super().from_pretrained(pretrained_model_name_or_path, subfolder=subfolder, **kwargs)
|
| 81 |
+
|
| 82 |
@property
|
| 83 |
def sampler(self) -> SampleDiffusionWithMotif:
|
| 84 |
return self._sampler
|
transformer/model_rfdiffusion.py
CHANGED
|
@@ -207,12 +207,25 @@ class RFDiffusionTransformerModel(ModelMixin, ConfigMixin):
|
|
| 207 |
def sigma_data(self) -> float:
|
| 208 |
return self.diffusion_module.sigma_data
|
| 209 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 210 |
def forward(
|
| 211 |
self,
|
| 212 |
xyz_noisy: torch.Tensor,
|
| 213 |
t: torch.Tensor,
|
| 214 |
-
f:
|
| 215 |
-
|
| 216 |
n_recycle: Optional[int] = None,
|
| 217 |
**kwargs,
|
| 218 |
) -> RFDiffusionTransformerOutput:
|
|
@@ -220,77 +233,37 @@ class RFDiffusionTransformerModel(ModelMixin, ConfigMixin):
|
|
| 220 |
Forward pass delegated to the foundry RFD3DiffusionModule.
|
| 221 |
|
| 222 |
Args:
|
| 223 |
-
xyz_noisy: Noisy atom coordinates [
|
| 224 |
-
t: Noise level / timestep [
|
| 225 |
-
f:
|
| 226 |
-
|
| 227 |
-
motif_mask: Mask for fixed motif atoms [L] (used when f is None)
|
| 228 |
n_recycle: Number of recycling iterations
|
| 229 |
|
| 230 |
Returns:
|
| 231 |
RFDiffusionTransformerOutput with denoised coordinates and predictions
|
| 232 |
"""
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
#
|
| 236 |
-
|
| 237 |
-
|
| 238 |
-
|
| 239 |
-
|
| 240 |
-
|
| 241 |
-
|
| 242 |
-
n_recycle=n_recycle,
|
| 243 |
-
**initializer_outputs,
|
| 244 |
)
|
| 245 |
-
return RFDiffusionTransformerOutput(
|
| 246 |
-
xyz=outs["X_L"],
|
| 247 |
-
single=torch.zeros(1), # not directly exposed by foundry
|
| 248 |
-
pair=torch.zeros(1),
|
| 249 |
-
sequence_logits=outs.get("sequence_logits_I"),
|
| 250 |
-
sequence_indices=outs.get("sequence_indices_I"),
|
| 251 |
-
)
|
| 252 |
-
|
| 253 |
-
# Simplified path: construct minimal feature dict and call dm.forward()
|
| 254 |
-
# For unconditional generation, each residue has 1 atom (CA), so L = I
|
| 255 |
-
device = xyz_noisy.device
|
| 256 |
-
dtype = xyz_noisy.dtype
|
| 257 |
-
|
| 258 |
-
if motif_mask is None:
|
| 259 |
-
motif_mask = torch.zeros(L, dtype=torch.bool, device=device)
|
| 260 |
-
else:
|
| 261 |
-
motif_mask = motif_mask.to(device)
|
| 262 |
-
|
| 263 |
-
# Construct minimal feature dict with all keys required by foundry
|
| 264 |
-
f = {
|
| 265 |
-
"atom_to_token_map": torch.arange(L, device=device), # 1:1 atom-to-token
|
| 266 |
-
"unindexing_pair_mask": torch.zeros(L, L, dtype=torch.bool, device=device),
|
| 267 |
-
"is_ca": torch.ones(L, dtype=torch.bool, device=device),
|
| 268 |
-
"is_motif_atom_with_fixed_coord": motif_mask,
|
| 269 |
-
"is_motif_token_with_fully_fixed_coord": motif_mask,
|
| 270 |
-
}
|
| 271 |
-
|
| 272 |
-
# Zero-initialized TokenInitializer outputs (no conditioning features)
|
| 273 |
-
Q_L_init = torch.zeros(L, self.config.c_atom, device=device, dtype=dtype)
|
| 274 |
-
C_L = torch.zeros(L, self.config.c_atom, device=device, dtype=dtype)
|
| 275 |
-
P_LL = torch.zeros(L, L, self.config.c_atompair, device=device, dtype=dtype)
|
| 276 |
-
S_I = torch.zeros(L, self.config.c_s, device=device, dtype=dtype)
|
| 277 |
-
Z_II = torch.zeros(L, L, self.config.c_z, device=device, dtype=dtype)
|
| 278 |
|
| 279 |
outs = self.diffusion_module(
|
| 280 |
X_noisy_L=xyz_noisy,
|
| 281 |
t=t,
|
| 282 |
f=f,
|
| 283 |
-
Q_L_init=Q_L_init,
|
| 284 |
-
C_L=C_L,
|
| 285 |
-
P_LL=P_LL,
|
| 286 |
-
S_I=S_I,
|
| 287 |
-
Z_II=Z_II,
|
| 288 |
n_recycle=n_recycle,
|
|
|
|
| 289 |
)
|
| 290 |
|
| 291 |
return RFDiffusionTransformerOutput(
|
| 292 |
xyz=outs["X_L"],
|
| 293 |
-
single=torch.zeros(1),
|
| 294 |
pair=torch.zeros(1),
|
| 295 |
sequence_logits=outs.get("sequence_logits_I"),
|
| 296 |
sequence_indices=outs.get("sequence_indices_I"),
|
|
|
|
| 207 |
def sigma_data(self) -> float:
|
| 208 |
return self.diffusion_module.sigma_data
|
| 209 |
|
| 210 |
+
def encode_conditioning(self, f: dict) -> dict:
|
| 211 |
+
"""
|
| 212 |
+
Embed the feature dict once, before the denoising loop.
|
| 213 |
+
|
| 214 |
+
The token initializer does not depend on the noise level, so foundry runs it once per
|
| 215 |
+
design and reuses the result for every step. Doing it inside `forward` would repeat the
|
| 216 |
+
pairformer stack on all 200 steps.
|
| 217 |
+
|
| 218 |
+
Returns:
|
| 219 |
+
The `Q_L_init` / `C_L` / `P_LL` / `S_I` / `Z_II` kwargs for `forward`.
|
| 220 |
+
"""
|
| 221 |
+
return self.token_initializer(f)
|
| 222 |
+
|
| 223 |
def forward(
|
| 224 |
self,
|
| 225 |
xyz_noisy: torch.Tensor,
|
| 226 |
t: torch.Tensor,
|
| 227 |
+
f: dict,
|
| 228 |
+
initializer_outputs: dict,
|
| 229 |
n_recycle: Optional[int] = None,
|
| 230 |
**kwargs,
|
| 231 |
) -> RFDiffusionTransformerOutput:
|
|
|
|
| 233 |
Forward pass delegated to the foundry RFD3DiffusionModule.
|
| 234 |
|
| 235 |
Args:
|
| 236 |
+
xyz_noisy: Noisy atom coordinates [D, L, 3], atom level, L = 14 * n_residues
|
| 237 |
+
t: Noise level / timestep [D]
|
| 238 |
+
f: Foundry feature dictionary, from `build_design_features`
|
| 239 |
+
initializer_outputs: Output of `encode_conditioning`
|
|
|
|
| 240 |
n_recycle: Number of recycling iterations
|
| 241 |
|
| 242 |
Returns:
|
| 243 |
RFDiffusionTransformerOutput with denoised coordinates and predictions
|
| 244 |
"""
|
| 245 |
+
# foundry builds the per-atom noise level with a hardcoded `.float()`
|
| 246 |
+
# (rfd3/model/RFD3_diffusion_module.py), so the scaled coordinates it feeds to the first
|
| 247 |
+
# linear layer are float32 no matter what dtype the caller supplies. Half-precision weights
|
| 248 |
+
# cannot consume them. Fail here with the remedy rather than deep inside foundry.
|
| 249 |
+
weight_dtype = next(self.parameters()).dtype
|
| 250 |
+
if weight_dtype != torch.float32:
|
| 251 |
+
raise ValueError(
|
| 252 |
+
f"{self.__class__.__name__} only runs in float32, but its weights are {weight_dtype}. "
|
| 253 |
+
"Reload the pipeline with `torch_dtype=torch.float32`."
|
|
|
|
|
|
|
| 254 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 255 |
|
| 256 |
outs = self.diffusion_module(
|
| 257 |
X_noisy_L=xyz_noisy,
|
| 258 |
t=t,
|
| 259 |
f=f,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 260 |
n_recycle=n_recycle,
|
| 261 |
+
**initializer_outputs,
|
| 262 |
)
|
| 263 |
|
| 264 |
return RFDiffusionTransformerOutput(
|
| 265 |
xyz=outs["X_L"],
|
| 266 |
+
single=torch.zeros(1), # not directly exposed by foundry
|
| 267 |
pair=torch.zeros(1),
|
| 268 |
sequence_logits=outs.get("sequence_logits_I"),
|
| 269 |
sequence_indices=outs.get("sequence_indices_I"),
|