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.
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/PKG-INFO +10 -6
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/README.md +8 -4
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/alignment.py +36 -59
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/constants.py +0 -7
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/corrections.py +70 -10
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/model.py +2 -22
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/structure.py +0 -5
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/.gitignore +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/LICENSE +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/pyproject.toml +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/__init__.py +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/_anarci/LICENSE +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/_anarci/__init__.py +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/_anarci/schemes.py +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/api.py +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/__init__.py +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_0.0.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_0.2.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_0.5.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_1.0.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/embeddings_noise_2.0.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/modified_residues.json +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/mpnn_encoder.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/softalign_embeddings.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/softalign_encoder.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/assets/softalign_gap.npz +0 -0
- {sabr_kit-0.4.2 → sabr_kit-0.4.3.dev1}/src/sabr/cli.py +0 -0
- {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.
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
2
|
Name: sabr-kit
|
|
3
|
-
Version: 0.4.
|
|
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
|
|
178
|
-
|
|
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
|
|
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
|
|
134
|
-
|
|
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
|
|
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
|
|
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 +=
|
|
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(
|
|
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(
|
|
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
|
-
|
|
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
|
-
|
|
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
|
|
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)
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|