Commit 15569e6f authored by alat-rights's avatar alat-rights
Browse files

bug fix and cleanup

parent acec200b
Loading
Loading
Loading
Loading
+28 −17
Original line number Diff line number Diff line
@@ -17,6 +17,7 @@ from deepchem.utils.data_utils import load_image_files, load_csv_files, load_jso
from deepchem.feat import UserDefinedFeaturizer, Featurizer
from deepchem.data import Dataset, DiskDataset, NumpyDataset, ImageDataset
from deepchem.feat.molecule_featurizers import OneHotFeaturizer
from deepchem.utils.genomics_utils import encode_bio_sequence

logger = logging.getLogger(__name__)

@@ -910,22 +911,33 @@ class FASTALoader(DataLoader):

    # Process legacy toggle
    if legacy:
      logger.info("Deprecation warning: Legacy mode will soon be deprecated.")
      if not isinstance(featurizer, None) or auto_add_annotations:
        logger.warning(f"featurizer option must be None and 
        auto_add_annotations must be false when legacy mode is enabled. You set
        featurizer to {featurizer} and auto_add_annotations to
        {auto_add_annotations}. So we set legacy = False.")
      logger.info("""
                  Deprecation warning: Legacy mode will soon be deprecated.
                  Disable legacy mode by passing legacy=False during
                  construction of FASTALoader object.
                  """)
      if featurizer is not None or auto_add_annotations:
        logger.warning(f"""
                       featurizer option must be None and
                       auto_add_annotations must be false when legacy mode is
                       enabled. You set featurizer to {featurizer} and
                       auto_add_annotations to {auto_add_annotations}.
                       So we set legacy = False.
                       """)
        legacy = False

    # Set attributes
    self.legacy = legacy
    self.auto_add_annotations = auto_add_annotations

    self.user_specified_features = None

    # Handle special featurizer cases
    if isinstance(featurizer, UserDefinedFeaturizer):  # User defined featurizer
      self.user_specified_features = featurizer.feature_fields
    elif isinstance(featurizer, None): # Default featurizer
      featurizer = featurizer(charset = ("A", "C", "T", "G"),
                              max_length = None)
    elif featurizer is None:  # Default featurizer
      featurizer = OneHotFeaturizer(
          charset=["A", "C", "T", "G"], max_length=None)

    # Set self.featurizer
    self.featurizer = featurizer
@@ -959,13 +971,11 @@ class FASTALoader(DataLoader):
      input_files = [input_files]

    def shard_generator():  # TODO Enable sharding with shard size parameter
      sequences = np.array([])
      for input_file in input_files:
        if self.legacy:
          X = encode_bio_sequence(input_file)
          ids = np.ones(len(X))
        else:
          sequences = np.append(sequences, _read_file(input_file))
          sequences = _read_file(input_file)
          X = self.featurizer(sequences)
        ids = np.ones(len(X))
        # (X, y, w, ids)
@@ -981,6 +991,7 @@ class FASTALoader(DataLoader):
        """
        Uses a fasta_file to create a numpy array of annotated FASTA-format strings
        """
        self.generate_ran = True
        sequences = np.array([])
        sequence = np.array([])
        header_read = False
@@ -995,7 +1006,7 @@ class FASTALoader(DataLoader):
              line = line[0:-1]  # Remove last character
            sequence = np.append(sequence, line)
        sequences = _add_sequence(sequences, sequence)
        yield sequences
        return sequences

      def _add_sequence(sequences: np.array, sequence: np.array) -> np.array:
        # Handle empty sequence
+28 −5
Original line number Diff line number Diff line
@@ -5,6 +5,7 @@ import os
import unittest

import deepchem as dc
from deepchem.feat.molecule_featurizers import OneHotFeaturizer


class TestFASTALoader(unittest.TestCase):
@@ -16,24 +17,46 @@ class TestFASTALoader(unittest.TestCase):
    super(TestFASTALoader, self).setUp()
    self.current_dir = os.path.dirname(os.path.abspath(__file__))

  def test_fasta_one_hot(self):
  def test_legacy_fasta_one_hot(self):
    input_file = os.path.join(self.current_dir,
                              "../../data/tests/example.fasta")
    loader = dc.data.FASTALoader()
    loader = dc.data.FASTALoader(legacy=True)
    sequences = loader.create_dataset(input_file)

    # example.fasta contains 3 sequences each of length 58.
    # The one-hot encoding turns base-pairs into vectors of length 5 (ATCGN).
    # There is one "image channel".

    # Previously expected shape was (3, 5, 58, 1).
    # Due to FASTALoader redesign, expected shape is now (3, 58, 5).
    assert sequences.X.shape == (3, 5, 58, 1)

  def test_fasta_one_hot(self):
    input_file = os.path.join(self.current_dir,
                              "../../data/tests/example.fasta")
    loader = dc.data.FASTALoader(legacy=False)
    sequences = loader.create_dataset(input_file)

    # Due to FASTALoader redesign, expected shape is now (3, 58, 5)

    assert sequences.X.shape == (3, 58, 5)

  def test_fasta_one_hot_big(self):
    protein = ['A', 'B', 'C', 'D', 'E', 'F', 'G', 'H', 'I', 'J', 'K', 'L', 'M',
               'N', 'O', 'P', 'Q', 'R', 'S', 'T', 'U', 'V', 'W', 'X', 'Y', 'Z',
               '*', '-']
    input_file = os.path.join(self.current_dir,
                              "../../data/tests/uniprot_truncated.fasta")
    loader = dc.data.FASTALoader(OneHotFeaturizer(charset=protein, max_length=1000), legacy=False)
    sequences = loader.create_dataset(input_file)

    assert sequences.X.shape

  def test_fasta_legacy_soft_fail(self):
    protein = ['A', 'B', 'C', 'D', 'E', 'F', 'G', 'H', 'I', 'J', 'K', 'L', 'M',
               'N', 'O', 'P', 'Q', 'R', 'S', 'T', 'U', 'V', 'W', 'X', 'Y', 'Z',
               '*', '-']
    input_file = os.path.join(self.current_dir,
                              "../../data/tests/uniprot_truncated.fasta")
    loader = dc.data.FASTALoader(charset="protein", max_length=1000)
    loader = dc.data.FASTALoader(OneHotFeaturizer(charset=protein, max_length=1000))
    sequences = loader.create_dataset(input_file)

    assert sequences.X.shape