CoolFace
Modelpublic

sentence-transformers/multi-qa-MiniLM-L6-dot-v1

sourceHugging Faceupdated 2y agoView on Hugging Face
17likes90kdownloads
train_script.py382 linesDownload Raw Back to root
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