sabr-kit 0.4.2__tar.gz → 0.4.3.dev1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (28) hide show
  1. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/PKG-INFO +10 -6
  2. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/README.md +8 -4
  3. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/alignment.py +36 -59
  4. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/constants.py +0 -7
  5. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/corrections.py +70 -10
  6. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/model.py +2 -22
  7. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/structure.py +0 -5
  8. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/.gitignore +0 -0
  9. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/LICENSE +0 -0
  10. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/pyproject.toml +0 -0
  11. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/__init__.py +0 -0
  12. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/_anarci/LICENSE +0 -0
  13. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/_anarci/__init__.py +0 -0
  14. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/_anarci/schemes.py +0 -0
  15. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/api.py +0 -0
  16. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/__init__.py +0 -0
  17. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_0.0.npz +0 -0
  18. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_0.2.npz +0 -0
  19. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_0.5.npz +0 -0
  20. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_1.0.npz +0 -0
  21. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_2.0.npz +0 -0
  22. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/modified_residues.json +0 -0
  23. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/mpnn_encoder.npz +0 -0
  24. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/softalign_embeddings.npz +0 -0
  25. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/softalign_encoder.npz +0 -0
  26. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/softalign_gap.npz +0 -0
  27. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/cli.py +0 -0
  28. {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/numbering.py +0 -0
@@ -1,6 +1,6 @@
1
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.5
2
2
  Name: sabr-kit
3
- Version: 0.4.2
3
+ Version: 0.4.3.dev1
4
4
  Summary: Structure-based Antibody Renumbering
5
5
  Project-URL: Homepage, https://github.com/delalamo/sabr
6
6
  Project-URL: Issues, https://github.com/delalamo/sabr/issues
@@ -174,18 +174,22 @@ the CLI with mmCIF output.
174
174
  - In `sabr` mode, gap extension is `-0.175027` and gap opening is `-2.525591`.
175
175
  In `softalign` mode, they are `0.1942468136548996` and
176
176
  `-2.5441808700561523`, respectively, as stored in the repository asset.
177
- - CDR gap distribution is always applied.
178
- - No deterministic light-chain DE-loop or C-terminal correction is applied.
177
+ - Deterministic CDR gap distribution and DE-loop correction are always
178
+ applied. Between IMGT anchors 79 and 85, DE-loop residues fill 80 first,
179
+ then 84 back through 81; additional residues are inserted after 82.
180
+ - No deterministic C-terminal correction is applied.
179
181
  - Automatic chain selection aligns against H, K, and L references and uses
180
182
  the highest score, with deterministic H/K/L tie order.
181
183
  - scFv mode appends H:K, H:L, K:H, and L:H reference candidates in that order.
184
+ - Composite candidates do not apply gap-open or gap-extension costs to query
185
+ linker residues aligned at the boundary between their two references.
182
186
  - Composite candidates receive normal affine gap-open and gap-extension costs
183
187
  for unaligned query and reference termini when their selection scores are
184
188
  compared; the underlying alignments and raw alignment scores are unchanged.
185
189
 
186
190
  A structural gap is detected when the C–N distance between consecutive
187
- residues exceeds 2.66 Å. A gap skips only the affected CDR correction and
188
- emits a warning; other regions continue normally.
191
+ residues exceeds 2.66 Å. A gap skips only the affected CDR or DE-loop
192
+ correction and emits a warning; other regions continue normally.
189
193
 
190
194
  T-cell receptors are not an officially supported SAbR target. For
191
195
  experimental low-level use, align a TCR against the K reference because that
@@ -130,18 +130,22 @@ the CLI with mmCIF output.
130
130
  - In `sabr` mode, gap extension is `-0.175027` and gap opening is `-2.525591`.
131
131
  In `softalign` mode, they are `0.1942468136548996` and
132
132
  `-2.5441808700561523`, respectively, as stored in the repository asset.
133
- - CDR gap distribution is always applied.
134
- - No deterministic light-chain DE-loop or C-terminal correction is applied.
133
+ - Deterministic CDR gap distribution and DE-loop correction are always
134
+ applied. Between IMGT anchors 79 and 85, DE-loop residues fill 80 first,
135
+ then 84 back through 81; additional residues are inserted after 82.
136
+ - No deterministic C-terminal correction is applied.
135
137
  - Automatic chain selection aligns against H, K, and L references and uses
136
138
  the highest score, with deterministic H/K/L tie order.
137
139
  - scFv mode appends H:K, H:L, K:H, and L:H reference candidates in that order.
140
+ - Composite candidates do not apply gap-open or gap-extension costs to query
141
+ linker residues aligned at the boundary between their two references.
138
142
  - Composite candidates receive normal affine gap-open and gap-extension costs
139
143
  for unaligned query and reference termini when their selection scores are
140
144
  compared; the underlying alignments and raw alignment scores are unchanged.
141
145
 
142
146
  A structural gap is detected when the C–N distance between consecutive
143
- residues exceeds 2.66 Å. A gap skips only the affected CDR correction and
144
- emits a warning; other regions continue normally.
147
+ residues exceeds 2.66 Å. A gap skips only the affected CDR or DE-loop
148
+ correction and emits a warning; other regions continue normally.
145
149
 
146
150
  T-cell receptors are not an officially supported SAbR target. For
147
151
  experimental low-level use, align a TCR against the K reference because that
@@ -1,4 +1,3 @@
1
- #!/usr/bin/env python3
2
1
  """Reference alignment using the validated affine Smith-Waterman method."""
3
2
 
4
3
  import functools
@@ -15,16 +14,18 @@ from sabr.corrections import apply_corrections
15
14
  LOGGER = logging.getLogger(__name__)
16
15
 
17
16
 
18
- def _rotate_for_dp(x, NINF):
17
+ def _rotate_for_dp(x, NINF, free_gap_boundary=-1):
19
18
  """Rotate a matrix for striped dynamic programming."""
20
19
  a, b = x.shape
21
20
  ar = jnp.arange(a)[::-1, None]
22
21
  br = jnp.arange(b)[None, :]
23
22
  i, j = (br - ar) + (a - 1), (ar + br) // 2
24
23
  n, m = (a + b - 1), (a + b) // 2
24
+ free_gap = jnp.broadcast_to(br == free_gap_boundary, (a, b))
25
25
  output = {
26
26
  "x": jnp.full([n, m], NINF).at[i, j].set(x),
27
27
  "o": (jnp.arange(n) + a % 2) % 2,
28
+ "free_gap": jnp.full([n, m], False).at[i, j].set(free_gap),
28
29
  }
29
30
  prev = (jnp.full((m, 3), NINF), jnp.full((m, 3), NINF))
30
31
  return output, prev, (i, j)
@@ -71,9 +72,10 @@ def _affine_score(
71
72
  temperature,
72
73
  gap_extend,
73
74
  gap_open,
75
+ free_gap_boundary=-1,
74
76
  NINF=-1e30,
75
77
  ):
76
- """Return the smooth affine Smith-Waterman score for one matrix."""
78
+ """Return a score, optionally making one reference gap boundary free."""
77
79
  right_penalties = jnp.asarray(
78
80
  [
79
81
  gap_open,
@@ -105,7 +107,11 @@ def _affine_score(
105
107
  _pad(h1[1:], ([0, 1], [0, 0]), NINF),
106
108
  )
107
109
  right += right_penalties
108
- down += down_penalties
110
+ down += jnp.where(
111
+ stripe["free_gap"][:, None],
112
+ 0,
113
+ down_penalties,
114
+ )
109
115
  right = right[:, :2]
110
116
 
111
117
  h0 = jnp.stack(
@@ -119,7 +125,9 @@ def _affine_score(
119
125
  return (h1, h0), h0
120
126
 
121
127
  similarities, mask = _apply_length_mask(similarities, lengths, NINF)
122
- stripes, previous, indices = _rotate_for_dp(similarities[:-1, :-1], NINF)
128
+ stripes, previous, indices = _rotate_for_dp(
129
+ similarities[:-1, :-1], NINF, free_gap_boundary
130
+ )
123
131
  scores = jax.lax.scan(step, previous, stripes, unroll=2)[-1][indices]
124
132
  return _soft_maximum(
125
133
  scores + similarities[1:, 1:, None],
@@ -130,7 +138,10 @@ def _affine_score(
130
138
 
131
139
 
132
140
  _AFFINE_ALIGNMENT = jax.jit(
133
- jax.vmap(jax.value_and_grad(_affine_score), (0, 0, None, None, None))
141
+ jax.vmap(
142
+ jax.value_and_grad(_affine_score),
143
+ (0, 0, None, None, None, None),
144
+ )
134
145
  )
135
146
 
136
147
 
@@ -218,48 +229,6 @@ def _alignment_path(alignment: np.ndarray) -> np.ndarray:
218
229
  return path
219
230
 
220
231
 
221
- def _validate_alignment(alignment: np.ndarray, chain_type: str) -> None:
222
- """Reject ambiguous paths before they are converted to numbering states."""
223
- path = _alignment_path(alignment)
224
- rows = path[:, 0]
225
- columns = path[:, 1]
226
-
227
- first_row = int(rows[0])
228
- first_position = int(columns[0]) + 1
229
- if first_row and first_position - first_row < 1:
230
- raise ValueError(
231
- "N-terminal residues would require non-positive numbering; "
232
- "use residue_range to select the antibody domain."
233
- )
234
-
235
- regions = list(constants.IMGT_LOOPS.values())
236
- if chain_type in ("K", "L"):
237
- regions.append((79, 84))
238
- for index, (left_row, right_row) in enumerate(zip(rows, rows[1:])):
239
- if right_row == left_row + 1:
240
- continue
241
- left_position = int(columns[index]) + 1
242
- right_position = int(columns[index + 1]) + 1
243
- if not any(
244
- start <= left_position <= end and start <= right_position <= end
245
- for start, end in regions
246
- ):
247
- raise ValueError(
248
- f"Unassigned query rows {left_row + 1}-{right_row - 1} "
249
- f"are bracketed by IMGT {left_position} and "
250
- f"{right_position}; use residue_range to select one "
251
- "antibody domain."
252
- )
253
-
254
- last_row = int(rows[-1])
255
- last_position = int(columns[-1]) + 1
256
- if last_row < alignment.shape[0] - 1 and last_position < 125:
257
- raise ValueError(
258
- f"Unassigned trailing query rows follow IMGT {last_position}; "
259
- "use residue_range to select the antibody domain."
260
- )
261
-
262
-
263
232
  def _validate_scfv_alignment(
264
233
  alignment: np.ndarray, representation: str
265
234
  ) -> None:
@@ -271,7 +240,6 @@ def _validate_scfv_alignment(
271
240
  )
272
241
  _alignment_path(alignment)
273
242
 
274
- domain_rows = []
275
243
  for domain_index, chain_type in enumerate(representation.split(":")):
276
244
  start = domain_index * constants.IMGT_MAX_POSITION
277
245
  end = start + constants.IMGT_MAX_POSITION
@@ -281,21 +249,15 @@ def _validate_scfv_alignment(
281
249
  raise ValueError(
282
250
  f"scFv {chain_type} domain contains no assigned residues."
283
251
  )
284
- domain_rows.append((int(assigned_rows[0]), int(assigned_rows[-1])))
285
- if domain_index == 0:
286
- _validate_alignment(domain[: assigned_rows[-1] + 1], chain_type)
287
- else:
288
- _validate_alignment(domain[assigned_rows[0] :], chain_type)
289
-
290
- if domain_rows[0][1] >= domain_rows[1][0]:
291
- raise ValueError("scFv domain assignments overlap or are out of order.")
292
252
 
293
253
 
294
254
  def _align_reference(
295
255
  query: np.ndarray,
296
256
  reference: np.ndarray,
297
257
  mode: str = "sabr",
258
+ free_gap_boundary: int = -1,
298
259
  ):
260
+ """Align one reference, optionally allowing a free query insertion."""
299
261
  anchor = np.zeros((1, reference.shape[1]), dtype=reference.dtype)
300
262
  augmented_reference = np.concatenate((anchor, reference, anchor), axis=0)
301
263
  query_batch = jnp.asarray(query[None, :])
@@ -309,6 +271,7 @@ def _align_reference(
309
271
  constants.DEFAULT_TEMPERATURE,
310
272
  gap_extend,
311
273
  gap_open,
274
+ free_gap_boundary,
312
275
  )
313
276
  return (
314
277
  np.asarray(soft_alignment[0])[:, 1:-1],
@@ -367,7 +330,21 @@ def align(
367
330
  best = None
368
331
  for candidate in candidates:
369
332
  reference, positions = references[candidate]
370
- reduced, similarity, score = _align_reference(query, reference, mode)
333
+ if ":" in candidate:
334
+ free_gap_boundary = sum(
335
+ position <= constants.IMGT_MAX_POSITION
336
+ for position in positions
337
+ )
338
+ reduced, similarity, score = _align_reference(
339
+ query,
340
+ reference,
341
+ mode,
342
+ free_gap_boundary=free_gap_boundary,
343
+ )
344
+ else:
345
+ reduced, similarity, score = _align_reference(
346
+ query, reference, mode
347
+ )
371
348
  if (
372
349
  not np.isfinite(score)
373
350
  or not np.isfinite(reduced).all()
@@ -419,5 +396,5 @@ def align(
419
396
  _validate_scfv_alignment(corrected, selected_type)
420
397
  else:
421
398
  corrected = apply_corrections(full_alignment, gap_indices=gap_indices)
422
- _validate_alignment(corrected, selected_type)
399
+ _alignment_path(corrected)
423
400
  return corrected, selected_type, score
@@ -1,4 +1,3 @@
1
- #!/usr/bin/env python3
2
1
  """Constants and configuration values for SAbR.
3
2
 
4
3
  This module defines constants used throughout the SAbR package including:
@@ -25,12 +24,6 @@ PEPTIDE_BOND_LENGTH = 1.33
25
24
  PEPTIDE_BOND_MAX_DISTANCE = 2 * PEPTIDE_BOND_LENGTH
26
25
  MAX_SELECTED_RESIDUES = 1024
27
26
 
28
- # Backbone indices for MPNN
29
- BACKBONE_N_IDX = 0
30
- BACKBONE_CA_IDX = 1
31
- BACKBONE_C_IDX = 2
32
- BACKBONE_CB_IDX = 3
33
-
34
27
  # Gap scores (from NPZ)
35
28
  SW_GAP_EXTEND = -0.175027
36
29
  SW_GAP_OPEN = -2.525591
@@ -1,4 +1,4 @@
1
- """Deterministic CDR corrections."""
1
+ """Deterministic CDR and DE-loop corrections."""
2
2
 
3
3
  import logging
4
4
  import warnings
@@ -10,12 +10,6 @@ from sabr import constants
10
10
  LOGGER = logging.getLogger(__name__)
11
11
 
12
12
 
13
- def _has_gap_in_region(
14
- gap_indices: frozenset[int], start_row: int, end_row: int
15
- ) -> bool:
16
- return any(index in gap_indices for index in range(start_row, end_row))
17
-
18
-
19
13
  def _aligned_row_near(aln: np.ndarray, target_col: int) -> int | None:
20
14
  """Return the row aligned at or within two columns of a target."""
21
15
  for offset in (0, -1, 1, -2, 2):
@@ -49,7 +43,9 @@ def _skip_for_structural_gap(
49
43
  region_name: str,
50
44
  ) -> bool:
51
45
  """Warn and return true when a regional correction crosses a gap."""
52
- if gap_indices and _has_gap_in_region(gap_indices, start_row, end_row):
46
+ if gap_indices and any(
47
+ index in gap_indices for index in range(start_row, end_row)
48
+ ):
53
49
  message = (
54
50
  f"Skipping {region_name} deterministic correction: structural "
55
51
  f"gap detected between rows {start_row} and {end_row}; using "
@@ -143,14 +139,78 @@ def correct_cdr_loop(
143
139
  return aln
144
140
 
145
141
 
142
+ def de_loop_positions(n_residues: int) -> list[int]:
143
+ """Return IMGT positions for residues between anchors 79 and 85.
144
+
145
+ Position 80 is filled first, followed by positions 84 through 81 from
146
+ right to left. Once positions 80 through 84 are occupied, additional
147
+ residues become insertions on position 82.
148
+ """
149
+ if n_residues <= 0:
150
+ return []
151
+ if n_residues <= 5:
152
+ return [80, *range(86 - n_residues, 85)]
153
+ return [80, 81, 82, *([82] * (n_residues - 5)), 83, 84]
154
+
155
+
156
+ def correct_de_loop(
157
+ aln: np.ndarray,
158
+ gap_indices: frozenset[int] | None = None,
159
+ ) -> np.ndarray:
160
+ """Assign the residues between IMGT 79 and 85 by loop length."""
161
+ anchor_79_row = _aligned_row_near(aln, 78)
162
+ anchor_85_row = _aligned_row_near(aln, 84)
163
+
164
+ if anchor_79_row is None or anchor_85_row is None:
165
+ LOGGER.warning(
166
+ "Skipping DE loop correction; missing anchor near IMGT 79 or 85."
167
+ )
168
+ return aln
169
+ if anchor_79_row >= anchor_85_row:
170
+ LOGGER.warning(
171
+ "Skipping DE loop correction; anchor 79 row (%d) is not before "
172
+ "anchor 85 row (%d).",
173
+ anchor_79_row,
174
+ anchor_85_row,
175
+ )
176
+ return aln
177
+ if _skip_for_structural_gap(
178
+ gap_indices, anchor_79_row, anchor_85_row, "DE loop (79-85)"
179
+ ):
180
+ return aln
181
+
182
+ intermediate_rows = list(range(anchor_79_row + 1, anchor_85_row))
183
+ positions = de_loop_positions(len(intermediate_rows))
184
+
185
+ # Clear the learned assignments for this region before rebuilding it.
186
+ # Additional occurrences of 82 remain unassigned here; alignment_to_states
187
+ # converts those orphan rows to 82A, 82B, and later insertion states.
188
+ aln[intermediate_rows, :] = 0
189
+ assigned_positions = set()
190
+ for row, position in zip(intermediate_rows, positions):
191
+ if position in assigned_positions:
192
+ continue
193
+ aln[row, position - 1] = 1
194
+ assigned_positions.add(position)
195
+
196
+ if intermediate_rows:
197
+ LOGGER.info(
198
+ "DE loop correction: assigned %d residue(s) to %s.",
199
+ len(intermediate_rows),
200
+ positions,
201
+ )
202
+ return aln
203
+
204
+
146
205
  def apply_corrections(
147
206
  aln: np.ndarray,
148
207
  gap_indices: frozenset[int] | None = None,
149
208
  ) -> np.ndarray:
150
- """Apply deterministic CDR corrections."""
209
+ """Apply deterministic CDR and DE-loop corrections."""
151
210
  for loop_name, (cdr_start, cdr_end) in constants.IMGT_LOOPS.items():
152
- correct_cdr_loop(
211
+ aln = correct_cdr_loop(
153
212
  aln, loop_name, cdr_start, cdr_end, gap_indices=gap_indices
154
213
  )
155
214
 
215
+ correct_de_loop(aln, gap_indices=gap_indices)
156
216
  return aln
@@ -1,4 +1,3 @@
1
- #!/usr/bin/env python3
2
1
  """MPNN (Message Passing Neural Network) encoder for protein structures.
3
2
 
4
3
  This module provides the ENC class which encodes protein backbone structures
@@ -130,20 +129,15 @@ class ProteinFeatures(hk.Module):
130
129
  def __init__(
131
130
  self,
132
131
  edge_features: int,
133
- node_features: int,
134
132
  num_positional_embeddings: int = 16,
135
133
  num_rbf: int = 16,
136
134
  top_k: int = 30,
137
135
  augment_eps: float = 0.0,
138
- num_chain_embeddings: int = 16,
139
136
  ):
140
137
  super(ProteinFeatures, self).__init__()
141
- self.edge_features = edge_features
142
- self.node_features = node_features
143
138
  self.top_k = top_k
144
139
  self.augment_eps = augment_eps
145
140
  self.num_rbf = num_rbf
146
- self.num_positional_embeddings = num_positional_embeddings
147
141
 
148
142
  self.embeddings = PositionalEncodings(num_positional_embeddings)
149
143
  # edge_in = num_positional_embeddings + num_rbf * 25 (for reference)
@@ -270,7 +264,7 @@ class PositionWiseFeedForward(hk.Module):
270
264
  self.act = Gelu
271
265
 
272
266
  def __call__(self, h_V):
273
- h = self.act(self.W_in(h_V), approximate=False)
267
+ h = self.act(self.W_in(h_V))
274
268
  h = self.W_out(h)
275
269
  return h
276
270
 
@@ -294,19 +288,14 @@ class EncLayer(hk.Module):
294
288
  def __init__(
295
289
  self,
296
290
  num_hidden: int,
297
- num_in: int,
298
291
  dropout: float = 0.1,
299
- num_heads: int = None,
300
292
  scale: int = 30,
301
293
  name: str = None,
302
294
  ):
303
295
  super(EncLayer, self).__init__()
304
296
  self.num_hidden = num_hidden
305
- self.num_in = num_in
306
297
  self.scale = scale
307
298
 
308
- self.safe_key = SafeKey(hk.next_rng_key())
309
-
310
299
  self.dropout1 = DropoutCust(dropout)
311
300
  self.dropout2 = DropoutCust(dropout)
312
301
  self.dropout3 = DropoutCust(dropout)
@@ -387,7 +376,6 @@ class ENC:
387
376
 
388
377
  def __init__(
389
378
  self,
390
- node_features: int,
391
379
  edge_features: int,
392
380
  hidden_dim: int,
393
381
  num_encoder_layers: int = 1,
@@ -398,7 +386,6 @@ class ENC:
398
386
  """Initialize the MPNN encoder.
399
387
 
400
388
  Args:
401
- node_features: Dimension of node features.
402
389
  edge_features: Dimension of edge features.
403
390
  hidden_dim: Hidden dimension for layers.
404
391
  num_encoder_layers: Number of encoder layers.
@@ -407,12 +394,8 @@ class ENC:
407
394
  dropout: Dropout rate.
408
395
  """
409
396
  super(ENC, self).__init__()
410
- self.node_features = node_features
411
- self.edge_features = edge_features
412
- self.hidden_dim = hidden_dim
413
397
 
414
398
  self.features = ProteinFeatures(
415
- node_features,
416
399
  edge_features,
417
400
  top_k=k_neighbors,
418
401
  augment_eps=augment_eps,
@@ -420,9 +403,7 @@ class ENC:
420
403
 
421
404
  self.W_e = hk.Linear(hidden_dim, with_bias=True, name="W_e")
422
405
  self.encoder_layers = [
423
- EncLayer(
424
- hidden_dim, hidden_dim * 2, dropout=dropout, name="enc" + str(i)
425
- )
406
+ EncLayer(hidden_dim, dropout=dropout, name="enc" + str(i))
426
407
  for i in range(num_encoder_layers)
427
408
  ]
428
409
 
@@ -481,7 +462,6 @@ def load_parameters(mode: str = "sabr") -> dict:
481
462
 
482
463
  def _encode(coords, mask, chain_ids, residue_indices):
483
464
  encoder = ENC(
484
- constants.EMBED_DIM,
485
465
  constants.EMBED_DIM,
486
466
  constants.EMBED_DIM,
487
467
  constants.N_MPNN_LAYERS,
@@ -352,11 +352,6 @@ def _new_residue_ids(data: _ChainData, numbered: list) -> dict:
352
352
  mapping = {}
353
353
  first_row, first_number = numbered[0][:2]
354
354
  first_assigned_number = first_number - first_row
355
- if first_assigned_number < 1:
356
- raise ValueError(
357
- "N-terminal residues would require non-positive numbering; "
358
- "use residue_range to select the antibody domain."
359
- )
360
355
  for query_index in range(first_row):
361
356
  mapping[data.residue_indices[query_index]] = (
362
357
  first_assigned_number + query_index,
File without changes
File without changes
File without changes
File without changes
File without changes