import pandas as pd import numpy as np import torch from torch.utils.data import Dataset, DataLoader from sklearn.model_selection import train_test_split from collections import Counter AA_VOCAB = {aa: i+1 for i, aa in enumerate('ARNDCQEGHILKMFPSTWYV')} AA_VOCAB['PAD'] = 0 AA_VOCAB['X'] = 20 def encode_sequence(seq, max_len=200): ids = [AA_VOCAB.get(c.upper(), AA_VOCAB['X']) for c in seq[:max_len]] ids += [AA_VOCAB['PAD']] * (max_len - len(ids)) return ids class PeptideDataset(Dataset): def __init__(self, sequences, labels, max_len=200): self.sequences = sequences self.labels = labels self.max_len = max_len def __len__(self): return len(self.sequences) def __getitem__(self, idx): seq = self.sequences[idx] label = self.labels[idx] ids = encode_sequence(seq, self.max_len) return torch.tensor(ids, dtype=torch.long), torch.tensor(label, dtype=torch.long) def load_genpept_data(csv_path='GenPept-Curated-2025/data/balanced_11000.csv'): df = pd.read_csv(csv_path) sequences = df['sequence'].values labels = (df['label'].values == 'AMP').astype(np.int64) return sequences, labels def create_splits(sequences, labels, test_size=0.21, val_size=0.09, random_state=42): X_temp, X_test, y_temp, y_test = train_test_split( sequences, labels, test_size=test_size, stratify=labels, random_state=random_state ) val_ratio = val_size / (1 - test_size) X_train, X_val, y_train, y_val = train_test_split( X_temp, y_temp, test_size=val_ratio, stratify=y_temp, random_state=random_state ) return (X_train, y_train), (X_val, y_val), (X_test, y_test) def get_dataloaders(sequences, labels, batch_size=64, max_len=200, num_workers=2): (X_train, y_train), (X_val, y_val), (X_test, y_test) = create_splits(sequences, labels) train_ds = PeptideDataset(X_train, y_train, max_len) val_ds = PeptideDataset(X_val, y_val, max_len) test_ds = PeptideDataset(X_test, y_test, max_len) train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=num_workers) val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=num_workers) test_loader = DataLoader(test_ds, batch_size=batch_size, shuffle=False, num_workers=num_workers) return train_loader, val_loader, test_loader if __name__ == '__main__': seqs, labs = load_genpept_data() print(f'Loaded {len(seqs)} sequences, {Counter(labs)}') train_l, val_l, test_l = get_dataloaders(seqs, labs, batch_size=4) for x, y in train_l: print(f'Batch: x={x.shape}, y={y.shape}') break