sentence-transformers/multi-qa-MiniLM-L6-dot-v1
1790k
1"""2Train script for a single file3 4Need to set the TPU address first:5export XRT_TPU_CONFIG="localservice;0;localhost:51011"6"""7 8import torch.multiprocessing as mp9import threading10import time11import random12import sys13import argparse14import gzip15import json16import logging17import tqdm18import torch19from torch import nn20from torch.utils.data import DataLoader21import torch22import torch_xla23import torch_xla.core24import torch_xla.core.functions25import torch_xla.core.xla_model as xm26import torch_xla.distributed.xla_multiprocessing as xmp27import torch_xla.distributed.parallel_loader as pl28import os29from shutil import copyfile30 31 32from transformers import (33 AdamW,34 AutoModel,35 AutoTokenizer,36 get_linear_schedule_with_warmup,37 set_seed,38)39 40class AutoModelForSentenceEmbedding(nn.Module):41 def __init__(self, model_name, tokenizer, args):42 super(AutoModelForSentenceEmbedding, self).__init__()43 44 assert args.pooling in ['mean', 'cls']45 46 self.model = AutoModel.from_pretrained(model_name)47 self.normalize = not args.no_normalize48 self.tokenizer = tokenizer49 self.pooling = args.pooling50 51 def forward(self, **kwargs):52 model_output = self.model(**kwargs)53 if self.pooling == 'mean':54 embeddings = self.mean_pooling(model_output, kwargs['attention_mask'])55 elif self.pooling == 'cls':56 embeddings = self.cls_pooling(model_output, kwargs['attention_mask'])57 58 if self.normalize:59 embeddings = torch.nn.functional.normalize(embeddings, p=2, dim=1)60 61 return embeddings62 63 def mean_pooling(self, model_output, attention_mask):64 token_embeddings = model_output[0] # First element of model_output contains all token embeddings65 input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()66 return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(input_mask_expanded.sum(1), min=1e-9)67 68 def cls_pooling(self, model_output, attention_mask):69 return model_output[0][:,0]70 71 def save_pretrained(self, output_path):72 if xm.is_master_ordinal():73 self.tokenizer.save_pretrained(output_path)74 self.model.config.save_pretrained(output_path)75 76 xm.save(self.model.state_dict(), os.path.join(output_path, "pytorch_model.bin"))77 78 79 80 81def train_function(index, args, queue):82 tokenizer = AutoTokenizer.from_pretrained(args.model)83 model = AutoModelForSentenceEmbedding(args.model, tokenizer, args)84 85 86 ### Train Loop87 device = xm.xla_device()88 model = model.to(device)89 90 # Instantiate optimizer91 optimizer = AdamW(params=model.parameters(), lr=2e-5, correct_bias=True)92 93 lr_scheduler = get_linear_schedule_with_warmup(94 optimizer=optimizer,95 num_warmup_steps=args.warmup_steps,96 num_training_steps=args.steps,97 )98 99 # Now we train the model100 cross_entropy_loss = nn.CrossEntropyLoss()101 max_grad_norm = 1102 103 model.train()104 105 for global_step in tqdm.trange(args.steps, disable=not xm.is_master_ordinal()):106 #### Get the batch data107 batch = queue.get()108 #print(index, "batch {}x{}".format(len(batch), ",".join([str(len(b)) for b in batch])))109 110 111 if len(batch[0]) == 2: #(anchor, positive)112 text1 = tokenizer([b[0] for b in batch], return_tensors="pt", max_length=args.max_length_a, truncation=True, padding="max_length")113 text2 = tokenizer([b[1] for b in batch], return_tensors="pt", max_length=args.max_length_b, truncation=True, padding="max_length")114 115 ### Compute embeddings116 embeddings_a = model(**text1.to(device))117 embeddings_b = model(**text2.to(device))118 119 ### Gather all embedings 120 embeddings_a = torch_xla.core.functions.all_gather(embeddings_a)121 embeddings_b = torch_xla.core.functions.all_gather(embeddings_b)122 123 ### Compute similarity scores 512 x 512124 scores = torch.mm(embeddings_a, embeddings_b.transpose(0, 1)) * args.scale125 126 ### Compute cross-entropy loss127 labels = torch.tensor(range(len(scores)), dtype=torch.long, device=embeddings_a.device) # Example a[i] should match with b[i]128 129 ## Symmetric loss as in CLIP130 loss = (cross_entropy_loss(scores, labels) + cross_entropy_loss(scores.transpose(0, 1), labels)) / 2131 132 else: #(anchor, positive, negative)133 text1 = tokenizer([b[0] for b in batch], return_tensors="pt", max_length=args.max_length_a, truncation=True, padding="max_length")134 text2 = tokenizer([b[1] for b in batch], return_tensors="pt", max_length=args.max_length_b, truncation=True, padding="max_length")135 text3 = tokenizer([b[2] for b in batch], return_tensors="pt", max_length=args.max_length_b, truncation=True, padding="max_length")136 137 embeddings_a = model(**text1.to(device))138 embeddings_b1 = model(**text2.to(device))139 embeddings_b2 = model(**text3.to(device))140 141 embeddings_a = torch_xla.core.functions.all_gather(embeddings_a)142 embeddings_b1 = torch_xla.core.functions.all_gather(embeddings_b1)143 embeddings_b2 = torch_xla.core.functions.all_gather(embeddings_b2)144 145 embeddings_b = torch.cat([embeddings_b1, embeddings_b2])146 147 ### Compute similarity scores 512 x 1024148 scores = torch.mm(embeddings_a, embeddings_b.transpose(0, 1)) * args.scale149 150 ### Compute cross-entropy loss151 labels = torch.tensor(range(len(scores)), dtype=torch.long, device=embeddings_a.device) # Example a[i] should match with b[i]152 153 ## One-way loss154 loss = cross_entropy_loss(scores, labels)155 156 157 # Backward pass158 optimizer.zero_grad()159 loss.backward()160 torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)161 162 xm.optimizer_step(optimizer, barrier=True)163 lr_scheduler.step()164 165 166 #Save model167 if (global_step+1) % args.save_steps == 0:168 output_path = os.path.join(args.output, str(global_step+1))169 xm.master_print("save model: "+output_path)170 model.save_pretrained(output_path)171 172 173 output_path = os.path.join(args.output, "final")174 xm.master_print("save model final: "+ output_path)175 model.save_pretrained(output_path)176 177 178def produce_data(args, queue, filepaths, dataset_indices):179 global_batch_size = args.batch_size*args.nprocs #Global batch size180 num_same_dataset = int(args.nprocs / args.datasets_per_batch)181 print("producer", "global_batch_size", global_batch_size)182 print("producer", "num_same_dataset", num_same_dataset)183 184 datasets = []185 for filepath in filepaths:186 if "reddit_" in filepath: #Special dataset class for Reddit files187 data_obj = RedditDataset(filepath)188 else:189 data_obj = Dataset(filepath, args)190 datasets.append(iter(data_obj)) 191 192 # Store if dataset is in a 2 col or 3 col format193 num_cols = {idx: len(next(dataset)) for idx, dataset in enumerate(datasets)}194 195 while True:196 texts_in_batch = set()197 batch_format = None #2 vs 3 col format for this batch198 199 #Add data from several sub datasets200 for _ in range(args.datasets_per_batch):201 valid_dataset = False #Check that datasets have the same 2/3 col format202 while not valid_dataset:203 data_idx = random.choice(dataset_indices)204 if batch_format is None:205 batch_format = num_cols[data_idx]206 valid_dataset = True207 else: #Check that this dataset has the same format208 valid_dataset = (batch_format == num_cols[data_idx])209 210 #Get data from this dataset211 dataset = datasets[data_idx]212 local_batch_size = args.batch_size213 if batch_format == 3 and args.batch_size_triplets is not None:214 local_batch_size = args.batch_size_triplets215 216 for _ in range(num_same_dataset):217 for _ in range(args.nprocs):218 batch_device = [] #A batch for one device219 while len(batch_device) < local_batch_size:220 sample = next(dataset)221 in_batch = False222 for text in sample:223 if text in texts_in_batch:224 in_batch = True225 break226 227 if not in_batch:228 for text in sample:229 texts_in_batch.add(text)230 batch_device.append(sample)231 232 queue.put(batch_device)233 234 235class RedditDataset:236 """237 A class that handles the reddit data files238 """239 def __init__(self, filepath):240 self.filepath = filepath241 242 def __iter__(self):243 while True:244 with gzip.open(self.filepath, "rt") as fIn:245 for line in fIn:246 data = json.loads(line)247 248 if "response" in data and "context" in data:249 yield [data["response"], data["context"]]250 251class Dataset:252 """253 A class that handles one dataset254 """255 def __init__(self, filepath, args):256 self.filepath = filepath257 self.args = args258 259 def __iter__(self):260 max_dataset_size = 20*1000*1000 #Cache small datasets in memory261 min_dataset_size = 50*1000 # Size for the small chunk of the dataset262 dataset = []263 min_dataset = []264 data_format = None265 266 print(self.filepath, "load")267 while dataset is None or len(dataset) == 0:268 with gzip.open(self.filepath, "rt") as fIn:269 for line in fIn:270 data = json.loads(line)271 if isinstance(data, dict):272 data = data['texts']273 274 if data_format is None:275 data_format = len(data)276 277 #Ensure that all entries are of the same 2/3 col format278 assert len(data) == data_format279 280 if dataset is not None:281 dataset.append(data)282 if len(dataset) >= max_dataset_size and not self.args.no_data_streaming:283 dataset = None284 285 if self.args.no_data_streaming:286 min_dataset.append(data)287 if len(min_dataset) >= min_dataset_size:288 random.shuffle(min_dataset)289 for data in min_dataset:290 yield data291 min_dataset = []292 else:293 yield data294 295 print(self.filepath, "fully loaded")296 297 if len(min_dataset) > 0:298 random.shuffle(min_dataset)299 for data in min_dataset:300 yield data301 302 # Data loaded. Now stream to the queue303 # Shuffle for each epoch304 while True:305 random.shuffle(dataset)306 for data in dataset:307 yield data308 309 310 311if __name__ == "__main__":312 parser = argparse.ArgumentParser()313 parser.add_argument('--model', default='nreimers/MiniLM-L6-H384-uncased')314 parser.add_argument('--steps', type=int, default=2000)315 parser.add_argument('--save_steps', type=int, default=10000)316 parser.add_argument('--warmup_steps', type=int, default=500)317 parser.add_argument('--batch_size', type=int, default=64)318 parser.add_argument('--batch_size_triplets', type=int, default=None)319 parser.add_argument('--max_length_a', type=int, default=128)320 parser.add_argument('--max_length_b', type=int, default=128)321 parser.add_argument('--nprocs', type=int, default=8)322 parser.add_argument('--datasets_per_batch', type=int, default=2, help="Number of datasets per batch")323 parser.add_argument('--scale', type=float, default=20, help="Use 20 for cossim, and 1 when you work with unnormalized embeddings with dot product")324 parser.add_argument('--no_normalize', action="store_true", default=False, help="If set: Embeddings are not normalized")325 parser.add_argument('--pooling', default='mean')326 parser.add_argument('--data_folder', default="/data", help="Folder with your dataset files")327 parser.add_argument('--no_data_streaming', action="store_true", default=False, help="If set: All data will first be loaded in memory")328 parser.add_argument('data_config', help="A data_config.json file")329 parser.add_argument('output')330 args = parser.parse_args()331 332 # Ensure num proc is devisible by datasets_per_batch333 assert (args.nprocs % args.datasets_per_batch) == 0334 335 336 logging.info("Output: "+args.output)337 if os.path.exists(args.output):338 print("Output folder already exists.")339 input("Continue?")340 341 # Write train script to output path342 os.makedirs(args.output, exist_ok=True)343 344 data_config_path = os.path.join(args.output, 'data_config.json')345 copyfile(args.data_config, data_config_path)346 347 train_script_path = os.path.join(args.output, 'train_script.py')348 copyfile(__file__, train_script_path)349 with open(train_script_path, 'a') as fOut:350 fOut.write("\n\n# Script was called via:\n#python " + " ".join(sys.argv))351 352 353 354 #Load data config355 with open(args.data_config) as fIn:356 data_config = json.load(fIn)357 358 queue = mp.Queue(maxsize=100*args.nprocs)359 360 filepaths = []361 dataset_indices = []362 for idx, data in enumerate(data_config):363 filepaths.append(os.path.join(os.path.expanduser(args.data_folder), data['name']))364 dataset_indices.extend([idx]*data['weight'])365 366 # Start producer367 p = mp.Process(target=produce_data, args=(args, queue, filepaths, dataset_indices))368 p.start()369 370 # Run training371 print("Start processes:", args.nprocs)372 xmp.spawn(train_function, args=(args, queue), nprocs=args.nprocs, start_method='fork')373 print("Training done")374 print("It might be that not all processes exit automatically. In that case you must manually kill this process.")375 print("With 'pkill python' you can kill all remaining python processes")376 p.kill()377 exit()378 379 380 381# Script was called via:382#python train_many_data_files_v2.py --steps 200000 --batch_size 128 --model nreimers/MiniLM-L6-H384-uncased --max_length_a 64 --max_length_b 250 --scale 1 --pooling cls --no_normalize train_data_configs/multi-qa_v1.json output/multi-qa_v1-MiniLM-L6-cls_dot