jeevan0704/sequential-model-for-sequential_dataset
0
1"""2Data preprocessing for sequence labeling task3"""4import pandas as pd5import numpy as np6import torch7from torch.utils.data import Dataset, DataLoader8from sklearn.model_selection import train_test_split9import config10from utils import pad_sequences, create_padding_mask, set_seed11 12 13class SequenceLabelingDataset(Dataset):14 """PyTorch Dataset for sequence labeling"""15 16 def __init__(self, input_sequences, target_sequences, max_length):17 """18 Args:19 input_sequences: List of input token sequences20 target_sequences: List of target label sequences21 max_length: Maximum sequence length for padding22 """23 self.input_sequences = input_sequences24 self.target_sequences = target_sequences25 self.max_length = max_length26 27 # Remap labels from {-1, 0, 1} to {0, 1, 2} for CrossEntropyLoss28 # -1 -> 0, 0 -> 1, 1 -> 229 remapped_targets = []30 for seq in target_sequences:31 remapped_seq = np.array(seq) + 1 # Shift all labels by 132 remapped_targets.append(remapped_seq)33 34 # Pad sequences35 self.padded_inputs = pad_sequences(input_sequences, max_length, padding_value=0)36 self.padded_targets = pad_sequences(remapped_targets, max_length, padding_value=-100)37 38 # Create masks39 self.masks = create_padding_mask(input_sequences, max_length)40 41 def __len__(self):42 return len(self.input_sequences)43 44 def __getitem__(self, idx):45 return {46 'input_ids': torch.tensor(self.padded_inputs[idx], dtype=torch.long),47 'labels': torch.tensor(self.padded_targets[idx], dtype=torch.long),48 'mask': torch.tensor(self.masks[idx], dtype=torch.bool)49 }50 51 52def load_data(data_path=None):53 """54 Load pickled dataset55 56 Args:57 data_path: Path to pickled data file58 59 Returns:60 Pandas Series containing tuples of (input_seq, target_seq)61 """62 if data_path is None:63 data_path = config.DATA_PATH64 65 print(f"Loading data from {data_path}...")66 data = pd.read_pickle(data_path)67 print(f"Loaded {len(data)} samples")68 69 return data70 71 72def analyze_data(data):73 """74 Analyze dataset characteristics75 76 Args:77 data: Pandas Series of tuples78 79 Returns:80 Dictionary with dataset statistics81 """82 print("\n" + "="*80)83 print("DATASET ANALYSIS")84 print("="*80)85 86 # Extract sequences87 input_sequences = [item[0] for item in data]88 target_sequences = [item[1] for item in data]89 90 # Sequence lengths91 input_lengths = [len(seq) for seq in input_sequences]92 target_lengths = [len(seq) for seq in target_sequences]93 94 print(f"\nTotal samples: {len(data)}")95 print(f"\nSequence length statistics:")96 print(f" Min: {min(input_lengths)}")97 print(f" Max: {max(input_lengths)}")98 print(f" Mean: {np.mean(input_lengths):.2f}")99 print(f" Median: {np.median(input_lengths):.2f}")100 101 # Vocabulary size102 all_tokens = []103 for seq in input_sequences:104 all_tokens.extend(seq)105 unique_tokens = set(all_tokens)106 vocab_size_unique = len(unique_tokens)107 max_token_id = max(unique_tokens)108 109 # Vocabulary size should be max_token_id + 1 to accommodate all tokens110 vocab_size = max_token_id + 1111 112 print(f"\nVocabulary statistics:")113 print(f" Unique tokens: {vocab_size_unique}")114 print(f" Max token ID: {max_token_id}")115 print(f" Vocabulary size (for embedding): {vocab_size}")116 print(f" Token range: [{min(unique_tokens)}, {max_token_id}]")117 118 # Label distribution119 all_labels = []120 for seq in target_sequences:121 all_labels.extend(seq)122 unique_labels = sorted(set(all_labels))123 124 print(f"\nLabel statistics:")125 print(f" Unique labels: {unique_labels}")126 label_counts = {}127 for label in unique_labels:128 count = all_labels.count(label)129 label_counts[label] = count130 print(f" Label {label}: {count} ({count/len(all_labels)*100:.2f}%)")131 132 stats = {133 'num_samples': len(data),134 'vocab_size': vocab_size, # max_token_id + 1135 'vocab_size_unique': vocab_size_unique, # number of unique tokens136 'max_seq_length': max(input_lengths),137 'unique_labels': unique_labels,138 'label_counts': label_counts139 }140 141 return stats142 143 144def prepare_data(data, max_length=None, train_split=None, val_split=None, test_split=None, random_seed=None):145 """146 Prepare data for training147 148 Args:149 data: Pandas Series of tuples150 max_length: Maximum sequence length151 train_split: Training set proportion152 val_split: Validation set proportion153 test_split: Test set proportion154 random_seed: Random seed for reproducibility155 156 Returns:157 Tuple of (train_dataset, val_dataset, test_dataset, stats)158 """159 # Set defaults160 if max_length is None:161 max_length = config.MAX_SEQ_LENGTH162 if train_split is None:163 train_split = config.TRAIN_SPLIT164 if val_split is None:165 val_split = config.VAL_SPLIT166 if test_split is None:167 test_split = config.TEST_SPLIT168 if random_seed is None:169 random_seed = config.RANDOM_SEED170 171 # Set seed172 set_seed(random_seed)173 174 # Analyze data175 stats = analyze_data(data)176 177 # Extract sequences178 input_sequences = [item[0] for item in data]179 target_sequences = [item[1] for item in data]180 181 print(f"\n" + "="*80)182 print("SPLITTING DATA")183 print("="*80)184 print(f"Train: {train_split*100:.0f}%, Val: {val_split*100:.0f}%, Test: {test_split*100:.0f}%")185 186 # First split: train+val vs test187 train_val_inputs, test_inputs, train_val_targets, test_targets = train_test_split(188 input_sequences, target_sequences,189 test_size=test_split,190 random_state=random_seed191 )192 193 # Second split: train vs val194 val_ratio = val_split / (train_split + val_split)195 train_inputs, val_inputs, train_targets, val_targets = train_test_split(196 train_val_inputs, train_val_targets,197 test_size=val_ratio,198 random_state=random_seed199 )200 201 print(f"\nSplit sizes:")202 print(f" Train: {len(train_inputs)}")203 print(f" Val: {len(val_inputs)}")204 print(f" Test: {len(test_inputs)}")205 206 # Create datasets207 print(f"\nCreating datasets with max_length={max_length}...")208 train_dataset = SequenceLabelingDataset(train_inputs, train_targets, max_length)209 val_dataset = SequenceLabelingDataset(val_inputs, val_targets, max_length)210 test_dataset = SequenceLabelingDataset(test_inputs, test_targets, max_length)211 212 print("Datasets created successfully!")213 214 return train_dataset, val_dataset, test_dataset, stats215 216 217def create_data_loaders(train_dataset, val_dataset, test_dataset, batch_size=None):218 """219 Create PyTorch DataLoaders220 221 Args:222 train_dataset: Training dataset223 val_dataset: Validation dataset224 test_dataset: Test dataset225 batch_size: Batch size226 227 Returns:228 Tuple of (train_loader, val_loader, test_loader)229 """230 if batch_size is None:231 batch_size = config.BATCH_SIZE232 233 print(f"\nCreating data loaders with batch_size={batch_size}...")234 235 train_loader = DataLoader(236 train_dataset,237 batch_size=batch_size,238 shuffle=True,239 num_workers=0, # Set to 0 for Windows compatibility240 pin_memory=True if torch.cuda.is_available() else False241 )242 243 val_loader = DataLoader(244 val_dataset,245 batch_size=batch_size,246 shuffle=False,247 num_workers=0,248 pin_memory=True if torch.cuda.is_available() else False249 )250 251 test_loader = DataLoader(252 test_dataset,253 batch_size=batch_size,254 shuffle=False,255 num_workers=0,256 pin_memory=True if torch.cuda.is_available() else False257 )258 259 print(f"Data loaders created:")260 print(f" Train batches: {len(train_loader)}")261 print(f" Val batches: {len(val_loader)}")262 print(f" Test batches: {len(test_loader)}")263 264 return train_loader, val_loader, test_loader265 266 267if __name__ == "__main__":268 # Test data loading and preprocessing269 print("Testing data preprocessing pipeline...")270 271 # Load data272 data = load_data()273 274 # Prepare data275 train_dataset, val_dataset, test_dataset, stats = prepare_data(data)276 277 # Create data loaders278 train_loader, val_loader, test_loader = create_data_loaders(279 train_dataset, val_dataset, test_dataset280 )281 282 # Test a batch283 print("\n" + "="*80)284 print("TESTING BATCH")285 print("="*80)286 batch = next(iter(train_loader))287 print(f"Batch keys: {batch.keys()}")288 print(f"Input shape: {batch['input_ids'].shape}")289 print(f"Labels shape: {batch['labels'].shape}")290 print(f"Mask shape: {batch['mask'].shape}")291 print(f"\nSample input (first sequence, first 10 tokens): {batch['input_ids'][0, :10]}")292 print(f"Sample labels (first sequence, first 10 tokens): {batch['labels'][0, :10]}")293 print(f"Sample mask (first sequence, first 10 tokens): {batch['mask'][0, :10]}")294 295 print("\n" + "="*80)296 print("Data preprocessing pipeline test completed successfully!")297 print("="*80)298 