[RFD3] Feed MPNN real backbone atoms and the schema it validates

#2
by dn6 HF Staff - opened
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.bfloat16, trust_remote_code=True)
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
- Passing `input_xyz` enables motif conditioning — fix specific residues in place while designing the rest:
 
71
 
72
- ```python
73
- import torch
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.bfloat16, trust_remote_code=True)
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.bfloat16, trust_remote_code=True)
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="Boolean mask for motif (fixed) positions",
110
  ),
111
  OutputParam(
112
- "motif_xyz",
113
  type_hint=torch.Tensor,
114
- description="Coordinates for motif residues",
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
- motif_mask = torch.zeros(L, dtype=torch.bool)
153
- for start, end in motif_ranges:
154
- motif_mask[start:end] = True
155
-
 
 
 
 
156
  if input_xyz is not None:
157
- motif_xyz = input_xyz
158
- else:
159
- motif_xyz = None
 
 
 
 
 
 
 
 
160
 
161
- block_state.motif_mask = motif_mask
162
- block_state.motif_xyz = motif_xyz
 
 
163
  block_state.L = L
164
- block_state.batch_size = 1
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
- if hasattr(components, "scheduler") and components.scheduler is not None:
220
- noise_schedule = components.scheduler.get_noise_schedule()
221
- else:
222
- # Fallback: simple linear schedule
223
- noise_schedule = torch.linspace(160.0 * 16.0, 4e-4 * 16.0, 200)
 
 
 
 
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, L, 3]"),
 
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
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
289
-
290
- # Initial noise scaled by first noise level (c0), matching original:
291
- # noise = c0 * randn(D, L, 3)
292
- c0 = noise_schedule[0]
293
- noise = c0 * torch.randn((D, L, 3), dtype=dtype, device=device, generator=generator)
294
-
295
- # Zero out noise for motif atoms
296
- if motif_mask is not None:
297
- noise[:, motif_mask] = 0.0
298
-
299
- # Build initial coordinates: motif coords + noise
300
- coord_motif = torch.zeros((D, L, 3), dtype=dtype, device=device)
301
- if motif_xyz is not None and motif_mask is not None:
302
- motif_indices = motif_mask.nonzero(as_tuple=True)[0]
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, L, 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
  ]
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
- if has_scheduler:
137
- X_noisy_L, t_hat = components.scheduler.add_noise(
138
- X_L, c_t_minus_1, c_t, motif_mask=motif_mask
139
- )
140
- else:
141
- X_noisy_L = X_L
142
- t_hat = c_t_minus_1
143
 
144
  # Step 2: Model forward pass
145
- if has_transformer:
146
- # t_hat is a scalar, tile to batch dimension
147
- t_batch = (t_hat.to(device).expand(D) if isinstance(t_hat, torch.Tensor)
148
- else torch.full((D,), t_hat, device=device))
149
-
150
- output = components.transformer(
151
- xyz_noisy=X_noisy_L,
152
- t=t_batch,
153
- motif_mask=motif_mask,
154
- n_recycle=n_recycle,
155
- )
156
-
157
- X_denoised_L = output.xyz
158
- single = output.single
159
- pair = output.pair
160
- sequence_logits = output.sequence_logits
161
- sequence_indices = output.sequence_indices
162
- else:
163
- X_denoised_L = X_noisy_L
164
 
165
  # Step 3: Euler update with step_scale (matching original sampler)
166
- if has_scheduler:
167
- X_L = components.scheduler.step(
168
- xyz_pred=X_denoised_L,
169
- xyz_noisy=X_noisy_L,
170
- c_t_minus_1=c_t_minus_1,
171
- c_t=c_t,
172
- motif_mask=motif_mask,
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
- block_state.xyz = X_L
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- "motif_mask", type_hint=torch.Tensor,
179
- description="Mask for fixed/motif positions [L]. True = fixed sequence.",
 
 
 
 
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
- motif_mask = block_state.motif_mask
 
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, xyz, motif_mask, known_seq, temperature,
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, xyz, motif_mask, known_seq, temperature):
268
- """Run the MPNNModel wrapper on backbone coordinates."""
269
- B, L, _ = xyz.shape
270
- device = xyz.device
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=X, S=known_seq, designed_residue_mask=designed_mask, temperature=temperature,
 
 
 
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(network_input)
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(network_input)
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(network_input)
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
- class RFDiffusionScheduler(ConfigMixin):
 
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: Optional[dict] = None,
215
- motif_mask: Optional[torch.Tensor] = None,
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 [B, L, 3]
224
- t: Noise level / timestep [B]
225
- f: Feature dictionary (as expected by foundry). If None, a minimal
226
- feature dict is constructed from xyz_noisy and motif_mask.
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
- B, L, _ = xyz_noisy.shape
234
-
235
- # If caller provides a full feature dict, use the native foundry path
236
- if f is not None:
237
- initializer_outputs = self.token_initializer(f)
238
- outs = self.diffusion_module(
239
- X_noisy_L=xyz_noisy,
240
- t=t,
241
- f=f,
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"),