saarantras1 commited on
Commit
9bb7ca8
·
verified ·
1 Parent(s): 4c2f63d

Add sliding-window and allelic-skew inference

Browse files
Files changed (1) hide show
  1. 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):