Add sliding-window and allelic-skew inference
Browse files- modeling_malinois.py +48 -0
modeling_malinois.py
CHANGED
|
@@ -50,6 +50,7 @@ except ImportError: # keeps the file usable as a plain torch module offline
|
|
| 50 |
__all__ = [
|
| 51 |
'STANDARD_NT', 'MPRA_UPSTREAM', 'MPRA_DOWNSTREAM', 'CELL_TYPES',
|
| 52 |
'dna2tensor', 'MPACModel', 'MalinoisModel', 'MPACEnsemble', 'fold_for_chromosome',
|
|
|
|
| 53 |
]
|
| 54 |
|
| 55 |
# -----------------------------------------------------------------------------
|
|
@@ -66,6 +67,13 @@ MPRA_DOWNSTREAM = 'CACTGCGGCTCCTGCGATCTAACTGGCCGGTACCTGAGCTCGCTAGCCTCGAGGATATCAA
|
|
| 66 |
|
| 67 |
CELL_TYPES = ['K562', 'HepG2', 'SKNSH']
|
| 68 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 69 |
|
| 70 |
def dna2tensor(sequence_str, vocab_list=STANDARD_NT):
|
| 71 |
"""One-hot encode a DNA string as a (4, len) float tensor."""
|
|
@@ -425,6 +433,44 @@ class MPACModel(
|
|
| 425 |
|
| 426 |
return torch.cat(results, dim=0)
|
| 427 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 428 |
|
| 429 |
class MPACEnsemble(nn.Module):
|
| 430 |
"""Mean prediction over a set of architecturally identical `MPACModel`s.
|
|
@@ -474,6 +520,8 @@ class MPACEnsemble(nn.Module):
|
|
| 474 |
return self._template.add_flanks(x)
|
| 475 |
|
| 476 |
predict = MPACModel.predict
|
|
|
|
|
|
|
| 477 |
|
| 478 |
@classmethod
|
| 479 |
def from_pretrained(cls, repo_id, chromosome, device='cpu', **kwargs):
|
|
|
|
| 50 |
__all__ = [
|
| 51 |
'STANDARD_NT', 'MPRA_UPSTREAM', 'MPRA_DOWNSTREAM', 'CELL_TYPES',
|
| 52 |
'dna2tensor', 'MPACModel', 'MalinoisModel', 'MPACEnsemble', 'fold_for_chromosome',
|
| 53 |
+
'MPAC_CONTEXT_UPSTREAM', 'MPAC_CONTEXT_DOWNSTREAM', 'MPAC_STEP_SIZE',
|
| 54 |
]
|
| 55 |
|
| 56 |
# -----------------------------------------------------------------------------
|
|
|
|
| 67 |
|
| 68 |
CELL_TYPES = ['K562', 'HepG2', 'SKNSH']
|
| 69 |
|
| 70 |
+
# Genomic context reproducing `vcf_predict.py --relative_start 9 --relative_end 181
|
| 71 |
+
# --step_size 10`: 371 bp running from 180 bp before the variant to 190 bp after it,
|
| 72 |
+
# sliced into eighteen 200 bp windows at stride 10.
|
| 73 |
+
MPAC_CONTEXT_UPSTREAM = 180
|
| 74 |
+
MPAC_CONTEXT_DOWNSTREAM = 190
|
| 75 |
+
MPAC_STEP_SIZE = 10
|
| 76 |
+
|
| 77 |
|
| 78 |
def dna2tensor(sequence_str, vocab_list=STANDARD_NT):
|
| 79 |
"""One-hot encode a DNA string as a (4, len) float tensor."""
|
|
|
|
| 433 |
|
| 434 |
return torch.cat(results, dim=0)
|
| 435 |
|
| 436 |
+
def predict_windows(self, sequences, step_size=MPAC_STEP_SIZE, **kwargs):
|
| 437 |
+
"""Average predictions over the tiled windows of longer sequences.
|
| 438 |
+
|
| 439 |
+
Each sequence is cut into `variable_region_len` windows at `step_size`
|
| 440 |
+
stride and every window is scored by `predict` (flanks attached, strands
|
| 441 |
+
averaged), then averaged. Passing `MPAC_CONTEXT_UPSTREAM + 1 +
|
| 442 |
+
MPAC_CONTEXT_DOWNSTREAM` bp around a variant reproduces the sliding-window
|
| 443 |
+
scheme used for the published MPAC predictions.
|
| 444 |
+
"""
|
| 445 |
+
width = self.variable_region_len
|
| 446 |
+
offsets = [range(0, len(s) - width + 1, step_size) for s in sequences]
|
| 447 |
+
assert all(len(o) for o in offsets), \
|
| 448 |
+
f"every sequence must be at least {width} bp"
|
| 449 |
+
|
| 450 |
+
flat = [s[i:i + width] for s, o in zip(sequences, offsets) for i in o]
|
| 451 |
+
preds = self.predict(flat, **kwargs)
|
| 452 |
+
assert preds.shape[0] == len(flat), \
|
| 453 |
+
f"got {preds.shape[0]} predictions for {len(flat)} windows"
|
| 454 |
+
|
| 455 |
+
out, cursor = [], 0
|
| 456 |
+
for o in offsets:
|
| 457 |
+
out.append(preds[cursor:cursor + len(o)].mean(dim=0))
|
| 458 |
+
cursor += len(o)
|
| 459 |
+
assert cursor == preds.shape[0], f"consumed {cursor} of {preds.shape[0]}"
|
| 460 |
+
return torch.stack(out)
|
| 461 |
+
|
| 462 |
+
def predict_skew(self, ref_sequences, alt_sequences, **kwargs):
|
| 463 |
+
"""Allelic skew for matched reference/alternate contexts.
|
| 464 |
+
|
| 465 |
+
Returns a dict of (n, n_outputs) tensors: `ref`, `alt`, and `skew`, the
|
| 466 |
+
latter being alt minus ref.
|
| 467 |
+
"""
|
| 468 |
+
assert len(ref_sequences) == len(alt_sequences), \
|
| 469 |
+
f"{len(ref_sequences)} ref vs {len(alt_sequences)} alt sequences"
|
| 470 |
+
ref = self.predict_windows(ref_sequences, **kwargs)
|
| 471 |
+
alt = self.predict_windows(alt_sequences, **kwargs)
|
| 472 |
+
return {'ref': ref, 'alt': alt, 'skew': alt - ref}
|
| 473 |
+
|
| 474 |
|
| 475 |
class MPACEnsemble(nn.Module):
|
| 476 |
"""Mean prediction over a set of architecturally identical `MPACModel`s.
|
|
|
|
| 520 |
return self._template.add_flanks(x)
|
| 521 |
|
| 522 |
predict = MPACModel.predict
|
| 523 |
+
predict_windows = MPACModel.predict_windows
|
| 524 |
+
predict_skew = MPACModel.predict_skew
|
| 525 |
|
| 526 |
@classmethod
|
| 527 |
def from_pretrained(cls, repo_id, chromosome, device='cpu', **kwargs):
|