sparse-encoder-testing/splade-bert-tiny-nq
092k
1from datasets import load_dataset2from sentence_transformers import (3 SparseEncoder,4 SparseEncoderTrainer,5 SparseEncoderTrainingArguments,6 SparseEncoderModelCardData,7)8from sentence_transformers.sparse_encoder.losses import SpladeLoss, SparseMultipleNegativesRankingLoss9from sentence_transformers.training_args import BatchSamplers10from sentence_transformers.sparse_encoder.evaluation import SparseNanoBEIREvaluator11from sentence_transformers.sparse_encoder.models import SpladePooling, MLMTransformer12 13# 1. Load a model to finetune with 2. (Optional) model card data14mlm_transformer = MLMTransformer("prajjwal1/bert-tiny")15splade_pooling = SpladePooling(pooling_strategy="max", word_embedding_dimension=mlm_transformer.get_sentence_embedding_dimension())16 17model = SparseEncoder(18 modules=[mlm_transformer, splade_pooling],19 model_card_data=SparseEncoderModelCardData(20 language="en",21 license="apache-2.0",22 model_name="SPLADE BERT-tiny trained on Natural-Questions tuples",23 )24)25 26# 3. Load a dataset to finetune on27full_dataset = load_dataset("sentence-transformers/natural-questions", split="train").select(range(100_000))28dataset_dict = full_dataset.train_test_split(test_size=1_000, seed=12)29train_dataset = dataset_dict["train"]30eval_dataset = dataset_dict["test"]31 32# 4. Define a loss function33loss = SpladeLoss(34 model=model,35 loss=SparseMultipleNegativesRankingLoss(model=model),36 lambda_query=5e-5,37 lambda_corpus=3e-5,38)39 40# 5. (Optional) Specify training arguments41args = SparseEncoderTrainingArguments(42 # Required parameter:43 output_dir="models/splade-bert-tiny-nq",44 # Optional training parameters:45 num_train_epochs=1,46 per_device_train_batch_size=64,47 per_device_eval_batch_size=64,48 learning_rate=2e-5,49 warmup_ratio=0.1,50 fp16=True, # Set to False if you get an error that your GPU can't run on FP1651 bf16=False, # Set to True if you have a GPU that supports BF1652 batch_sampler=BatchSamplers.NO_DUPLICATES, # MultipleNegativesRankingLoss benefits from no duplicate samples in a batch53 # Optional tracking/debugging parameters:54 eval_strategy="steps",55 eval_steps=200,56 save_strategy="steps",57 save_steps=200,58 save_total_limit=2,59 logging_steps=20,60 run_name="splade-bert-tiny-nq", # Will be used in W&B if `wandb` is installed61)62 63# 6. (Optional) Create an evaluator & evaluate the base model64dev_evaluator = SparseNanoBEIREvaluator(dataset_names=["msmarco", "nfcorpus", "nq"], batch_size=16)65 66# 7. Create a trainer & train67trainer = SparseEncoderTrainer(68 model=model,69 args=args,70 train_dataset=train_dataset,71 eval_dataset=eval_dataset,72 loss=loss,73 evaluator=dev_evaluator,74)75trainer.train()76 77# 8. Evaluate the model performance again after training78dev_evaluator(model)79 80# 9. Save the trained model81model.save_pretrained("models/splade-bert-tiny-nq/final")82 83# 10. (Optional) Push it to the Hugging Face Hub84model.push_to_hub("splade-bert-tiny-nq")