[RFD3] Build real conditioning features instead of zero tensors

#1
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,40 @@ 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 +214,32 @@ 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 +293,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 +342,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
  "L",
188
  type_hint=int,
189
+ description="Total length of the protein being designed, in residues",
190
  ),
191
  OutputParam(
192
  "batch_size",
193
  type_hint=int,
194
  description="Batch size (typically 1 for RFDiffusion)",
195
  ),
 
 
 
 
 
196
  ]
197
 
198
  def check_inputs(self, components, block_state):
 
214
 
215
  L, motif_ranges = parse_contig_string(contig_str)
216
 
217
+ # Motif conditioning needs a reference structure so the transform pipeline can build
218
+ # per-atom features for the fixed residues. A coordinate tensor alone cannot supply the
219
+ # element, atom-name and occupancy annotations those features are derived from.
220
+ if motif_ranges:
221
+ raise ValueError(
222
+ f"Motif-conditioned contigs are not supported yet, got `contigs={contig_str!r}`. "
223
+ "Pass a plain design length such as `contigs='100'`."
224
+ )
225
  if input_xyz is not None:
226
+ raise ValueError(
227
+ "`input_xyz` is not supported yet. Pass a plain design length such as "
228
+ "`contigs='100'`."
229
+ )
230
+
231
+ batch_size = 1
232
+ example = build_design_features(
233
+ length=L,
234
+ diffusion_batch_size=batch_size,
235
+ sigma_data=components.scheduler.config.sigma_data,
236
+ )
237
 
238
+ block_state.f = example["feats"]
239
+ block_state.coord_atom_lvl_to_be_noised = example["coord_atom_lvl_to_be_noised"]
240
+ block_state.motif_mask = example["feats"]["is_motif_atom_with_fixed_coord"]
241
  block_state.L = L
242
+ block_state.batch_size = batch_size
 
243
 
244
  self.set_block_state(state, block_state)
245
  return components, state
 
293
  def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
294
  block_state = self.get_block_state(state)
295
 
296
+ # A linear stand-in for the EDM schedule silently changes the sampler, so require
297
+ # the real one rather than degrading.
298
+ if components.scheduler is None:
299
+ raise ValueError(
300
+ "`scheduler` is not loaded. Call `load_components(trust_remote_code=True)` on the "
301
+ "pipeline before calling it."
302
+ )
303
+
304
+ noise_schedule = components.scheduler.get_noise_schedule()
305
 
306
  block_state.noise_schedule = noise_schedule
307
  block_state.num_inference_steps = len(noise_schedule)
 
342
  InputParam("generator", type_hint=torch.Generator, description="Random generator for reproducibility"),
343
  InputParam("diffusion_batch_size", default=1, type_hint=int, description="Number of samples to generate in parallel"),
344
  InputParam("L", required=True, type_hint=int, description="Protein length"),
345
+ InputParam("f", required=True, type_hint=dict),
346
+ InputParam("coord_atom_lvl_to_be_noised", required=True, type_hint=torch.Tensor),
347
  InputParam("motif_mask", required=True, type_hint=torch.Tensor),
 
348
  InputParam("noise_schedule", required=True, type_hint=torch.Tensor),
 
349
  ]
350
 
351
  @property
352
  def intermediate_outputs(self) -> List[OutputParam]:
353
  return [
354
+ OutputParam("xyz", type_hint=torch.Tensor, description="Initial noised coordinates [D, L_atom, 3]"),
355
+ OutputParam("initializer_outputs", type_hint=dict, description="Embedded conditioning, reused every step"),
356
  ]
357
 
358
  @torch.no_grad()
359
  def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
360
  block_state = self.get_block_state(state)
361
 
 
 
 
362
  noise_schedule = block_state.noise_schedule
 
363
  generator = block_state.generator
364
  D = block_state.diffusion_batch_size or 1
365
+ device = components.transformer.device
366
+
367
+ # The feature dict is built on CPU and is read on every denoising step, so move it once.
368
+ f = {k: v.to(device) if torch.is_tensor(v) else v for k, v in block_state.f.items()}
369
+ coord = block_state.coord_atom_lvl_to_be_noised.to(device)
370
+ motif_mask = f["is_motif_atom_with_fixed_coord"]
371
+
372
+ # Matches rfd3.model.inference_sampler._get_initial_structure:
373
+ # noise = c0 * randn(D, L, 3); noise[..., is_motif, :] = 0; X_L = noise + coord
374
+ c0 = noise_schedule[0].to(device)
375
+ L_atom = coord.shape[-2]
376
+ noise = c0 * torch.randn((D, L_atom, 3), device=device, generator=generator)
377
+ noise[..., motif_mask, :] = 0.0
378
+ xyz = noise + coord
379
+
380
+ block_state.f = f
 
 
 
 
 
 
381
  block_state.xyz = xyz
382
+ block_state.initializer_outputs = components.transformer.encode_conditioning(f)
383
 
384
  self.set_block_state(state, block_state)
385
  return components, state
denoise.py CHANGED
@@ -80,9 +80,11 @@ 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
@@ -98,11 +100,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 +141,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 +148,52 @@ 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
 
100
 
101
  @torch.no_grad()
102
  def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:
103
+ # Both components are required. Falling back to a hand-rolled Euler step when the
104
+ # scheduler is missing silently swaps out the EDM sampler and yields coordinates that
105
+ # look finite but are not a protein backbone.
106
+ if components.transformer is None:
107
+ raise ValueError(
108
+ "`transformer` is not loaded. Call `load_components(trust_remote_code=True)` on the "
109
+ "pipeline before calling it."
110
+ )
111
+ if components.scheduler is None:
112
+ raise ValueError(
113
+ "`scheduler` is not loaded. Call `load_components(trust_remote_code=True)` on the "
114
+ "pipeline before calling it."
115
+ )
116
+
117
  block_state = self.get_block_state(state)
118
 
119
  xyz = block_state.xyz
120
  noise_schedule = block_state.noise_schedule
121
  motif_mask = block_state.motif_mask
122
+ f = block_state.f
123
+ initializer_outputs = block_state.initializer_outputs
124
 
125
  n_recycle = block_state.n_recycle
126
  callback = block_state.callback
 
141
  sequence_logits = None
142
  sequence_indices = None
143
 
 
 
 
144
  # Iterate over consecutive pairs (c_t_minus_1, c_t) in the noise schedule
145
  # noise_schedule goes from high noise to low noise
146
  for step_num in range(len(noise_schedule) - 1):
 
148
  c_t = noise_schedule[step_num + 1]
149
 
150
  # Step 1: Inject stochastic noise (matching original sampler)
151
+ X_noisy_L, t_hat = components.scheduler.add_noise(
152
+ X_L, c_t_minus_1, c_t, motif_mask=motif_mask
153
+ )
 
 
 
 
154
 
155
  # Step 2: Model forward pass
156
+ # t_hat is a scalar, tile to batch dimension
157
+ t_batch = (t_hat.to(device).expand(D) if isinstance(t_hat, torch.Tensor)
158
+ else torch.full((D,), t_hat, device=device))
159
+
160
+ output = components.transformer(
161
+ xyz_noisy=X_noisy_L,
162
+ t=t_batch,
163
+ f=f,
164
+ initializer_outputs=initializer_outputs,
165
+ n_recycle=n_recycle,
166
+ )
167
+
168
+ X_denoised_L = output.xyz
169
+ single = output.single
170
+ pair = output.pair
171
+ sequence_logits = output.sequence_logits
172
+ sequence_indices = output.sequence_indices
 
 
173
 
174
  # Step 3: Euler update with step_scale (matching original sampler)
175
+ X_L = components.scheduler.step(
176
+ xyz_pred=X_denoised_L,
177
+ xyz_noisy=X_noisy_L,
178
+ c_t_minus_1=c_t_minus_1,
179
+ c_t=c_t,
180
+ motif_mask=motif_mask,
181
+ )
 
 
 
 
 
 
182
 
183
  X_denoised_L_traj.append(X_denoised_L.clone())
184
 
185
  if callback is not None and step_num % callback_steps == 0:
186
  callback(step_num, c_t_minus_1, X_L)
187
 
188
+ # The sampler runs on padded atom-level coordinates (14 slots per residue). Downstream
189
+ # blocks and the documented output are one point per residue, so collapse to CA here.
190
+ is_ca = f["is_ca"]
191
+ block_state.xyz = X_L[:, is_ca]
192
  block_state.single = single
193
  block_state.pair = pair
194
  block_state.sequence_logits = sequence_logits
195
  block_state.sequence_indices = sequence_indices
196
+ block_state.trajectory = [step[:, is_ca] for step in X_denoised_L_traj]
197
 
198
  self.set_block_state(state, block_state)
199
  return components, state
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
  }
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"),