MCPcopy Create free account
hub / github.com/BIT-DataLab/LakeBench / train_function

Function train_function

join/Deepjoin/all-mpnet-base-v2/train_script.py:71–165  ·  view source on GitHub ↗
(index, args, queue)

Source from the content-addressed store, hash-verified

69
70
71def train_function(index, args, queue):
72 tokenizer = AutoTokenizer.from_pretrained(args.model)
73 model = AutoModelForSentenceEmbedding(args.model, tokenizer)
74
75
76 ### Train Loop
77 device = xm.xla_device()
78 model = model.to(device)
79
80 # Instantiate optimizer
81 optimizer = AdamW(params=model.parameters(), lr=2e-5, correct_bias=True)
82
83 lr_scheduler = get_linear_schedule_with_warmup(
84 optimizer=optimizer,
85 num_warmup_steps=500,
86 num_training_steps=args.steps,
87 )
88
89 # Now we train the model
90 cross_entropy_loss = nn.CrossEntropyLoss()
91 max_grad_norm = 1
92
93 model.train()
94
95 for global_step in tqdm.trange(args.steps, disable=not xm.is_master_ordinal()):
96 #### Get the batch data
97 batch = queue.get()
98 #print(index, "batch {}x{}".format(len(batch), ",".join([str(len(b)) for b in batch])))
99
100
101 if len(batch[0]) == 2: #(anchor, positive)
102 text1 = tokenizer([b[0] for b in batch], return_tensors="pt", max_length=args.max_length, truncation=True, padding="max_length")
103 text2 = tokenizer([b[1] for b in batch], return_tensors="pt", max_length=args.max_length, truncation=True, padding="max_length")
104
105 ### Compute embeddings
106 embeddings_a = model(**text1.to(device))
107 embeddings_b = model(**text2.to(device))
108
109 ### Gather all embedings
110 embeddings_a = torch_xla.core.functions.all_gather(embeddings_a)
111 embeddings_b = torch_xla.core.functions.all_gather(embeddings_b)
112
113 ### Compute similarity scores 512 x 512
114 scores = torch.mm(embeddings_a, embeddings_b.transpose(0, 1)) * args.scale
115
116 ### Compute cross-entropy loss
117 labels = torch.tensor(range(len(scores)), dtype=torch.long, device=embeddings_a.device) # Example a[i] should match with b[i]
118
119 ## Symmetric loss as in CLIP
120 loss = (cross_entropy_loss(scores, labels) + cross_entropy_loss(scores.transpose(0, 1), labels)) / 2
121
122 else: #(anchor, positive, negative)
123 text1 = tokenizer([b[0] for b in batch], return_tensors="pt", max_length=args.max_length, truncation=True, padding="max_length")
124 text2 = tokenizer([b[1] for b in batch], return_tensors="pt", max_length=args.max_length, truncation=True, padding="max_length")
125 text3 = tokenizer([b[2] for b in batch], return_tensors="pt", max_length=args.max_length, truncation=True, padding="max_length")
126
127 embeddings_a = model(**text1.to(device))
128 embeddings_b1 = model(**text2.to(device))

Callers

nothing calls this directly

Calls 3

save_pretrainedMethod · 0.95
getMethod · 0.45

Tested by

no test coverage detected