File size: 2,853 Bytes
35cdf53
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
// Copyright 2024 DeepMind Technologies Limited
//
// AlphaFold 3 source code is licensed under CC BY-NC-SA 4.0. To view a copy of
// this license, visit https://creativecommons.org/licenses/by-nc-sa/4.0/
//
// To request access to the AlphaFold 3 model parameters, follow the process set
// out at https://github.com/google-deepmind/alphafold3. You may only use these
// if received directly from Google. Use is subject to terms of use available at
// https://github.com/google-deepmind/alphafold3/blob/main/WEIGHTS_TERMS_OF_USE.md

#include <algorithm>

#include "absl/strings/str_cat.h"
#include "pybind11/cast.h"
#include "pybind11/numpy.h"
#include "pybind11/pybind11.h"

namespace {

namespace py = pybind11;

py::array_t<double> ComputeMsaProfile(
    const py::array_t<int, py::array::c_style>& msa, int num_residue_types) {
  if (msa.size() == 0) {
    throw py::value_error("The MSA must be non-empty.");
  }
  if (msa.ndim() != 2) {
    throw py::value_error(absl::StrCat("The MSA must be rectangular, got ",
                                       msa.ndim(), "-dimensional MSA array."));
  }
  const int msa_depth = msa.shape()[0];
  const int sequence_length = msa.shape()[1];

  py::array_t<double> profile({sequence_length, num_residue_types});
  std::fill(profile.mutable_data(), profile.mutable_data() + profile.size(),
            0.0f);
  auto profile_unchecked = profile.mutable_unchecked<2>();

  const double normalized_count = 1.0 / msa_depth;
  const int* msa_it = msa.data();
  for (int row_index = 0; row_index < msa_depth; ++row_index) {
    for (int column_index = 0; column_index < sequence_length; ++column_index) {
      const int residue_code = *(msa_it++);
      if (residue_code < 0 || residue_code >= num_residue_types) {
        throw py::value_error(
            absl::StrCat("All residue codes must be positive and smaller than "
                         "num_residue_types ",
                         num_residue_types, ", got ", residue_code));
      }
      profile_unchecked(column_index, residue_code) += normalized_count;
    }
  }
  return profile;
}

constexpr char kComputeMsaProfileDoc[] = R"(
Computes MSA profile for the given encoded MSA.

Args:
  msa: A Numpy array of shape (num_msa, num_res) with the integer coded MSA.
  num_residue_types: Integer that determines the number of unique residue types.
    This will determine the shape of the output profile.

Returns:
  A float Numpy array of shape (num_res, num_residue_types) with residue
  frequency (residue type count normalized by MSA depth) for every column of the
  MSA.
)";

}  // namespace

namespace alphafold3 {

void RegisterModuleMsaProfile(pybind11::module m) {
  m.def("compute_msa_profile", &ComputeMsaProfile, py::arg("msa"),
        py::arg("num_residue_types"), py::doc(kComputeMsaProfileDoc + 1));
}

}  // namespace alphafold3