from tqdm import tqdm import torch from multiprocessing import Pool, cpu_count #defining the device the data ends up living on device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') #number of examples in the batch batch_size = 128 # should be divisible by 2 #sequence length of model max_input_length = 64 #defining parallelizable function for processign batches def process_batch(batch_index): #establishing bounds of the batch start_index = batch_index * batch_size end_index = start_index + batch_size if end_index > len(positive_pairs): return None, None, None #getting the sentence pairs of the batch, and if they're pos or neg sentence_pairs = [] is_positives = [] # Creating positive pairs sentence_pairs.extend(positive_pairs[start_index:start_index + int(batch_size / 2)]) is_positives.extend([1] * int(batch_size / 2)) # Creating negative pairs sentence_pairs.extend(negative_pairs[start_index + int(batch_size / 2):end_index]) is_positives.extend([0] * int(batch_size / 2)) # Defining outputs # At the end of the day we need to know three things: # - the tokens for the sequences in a batch # - which sentence the tokens belong to, for positional encoding # - if the examples in the batch are positive or negative # these keep track of the first two batch_sentence_location_tokens = [] batch_sequence_tokens = [] # Tokenizing pairs for sentence_pair in sentence_pairs: sentence1 = sentence_pair[0] sentence2 = sentence_pair[1] # Tokenizing both sentences tokens = tokenizer([sentence1, sentence2]) sentence1_tokens = tokens['input_ids'][0] sentence2_tokens = tokens['input_ids'][1] # Trimming down tokens if len(sentence1_tokens) + len(sentence2_tokens) > max_input_length: sentence1_tokens = [101] + sentence1_tokens[-int(max_input_length / 2) + 1:] sentence2_tokens = sentence2_tokens[:int(max_input_length / 2) - 1] + [102] # Creating sentence tokens sentence_tokens = [0] * len(sentence1_tokens) + [1] * len(sentence2_tokens) # Combining and padding pad_num = max_input_length - (len(sentence1_tokens) + len(sentence2_tokens)) sequence_tokens = sentence1_tokens + sentence2_tokens + [0] * pad_num sentence_location_tokens = sentence_tokens + [1] * pad_num # Adding to batch batch_sequence_tokens.append(sequence_tokens) batch_sentence_location_tokens.append(sentence_location_tokens) return torch.tensor(batch_sentence_location_tokens), torch.tensor(batch_sequence_tokens), torch.tensor(is_positives) # Determine the number of batches num_batches = len(positive_pairs) // batch_size # Use a Pool of workers equal to the number of CPU cores with Pool(processes=cpu_count()) as pool: results = list(tqdm(pool.imap(process_batch, range(num_batches)), total=num_batches)) # Filter out None results from the process_batch function results = [result for result in results if result[0] is not None] # Unpack results into batches sentence_location_batches, sequence_tokens_batches, is_positives_batches = zip(*results) # Stack tensors into final batches sentence_location_batches = torch.stack(sentence_location_batches).to(device) sequence_tokens_batches = torch.stack(sequence_tokens_batches).to(device) is_positives_batches = torch.stack(is_positives_batches).to(device)