PeptEdgeV2 / data_utils.py
devansh0703's picture
Initial release: PeptEdgeV2 (3.43M params) w/ trained weights, config, source, model card
9f16c4c verified
Raw
History Blame Contribute Delete
2.68 kB
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