MCPcopy Create free account
hub / github.com/THUDM/AgentTuning / sbert_search

Function sbert_search

eval_heldout/science-world/eval_utils.py:306–338  ·  view source on GitHub ↗
(action_list, validActions, sbert_model, logger, k=1, N=1, return_scores=False)

Source from the content-addressed store, hash-verified

304 return validActions
305
306def sbert_search(action_list, validActions, sbert_model, logger, k=1, N=1, return_scores=False):
307 validActions = list(validActions)
308 pred_vectors = sbert_model.encode(action_list[:k], batch_size=5, show_progress_bar=False)
309 valid_action_vectors = sbert_model.encode(validActions, batch_size=min(len(validActions), 128), show_progress_bar=False)
310
311 # Calculate cosine similarity between each vector in pred_vectors and all vectors in valid_action_vectors
312 similarity_matrix = cosine_similarity(pred_vectors, valid_action_vectors)
313
314 # Take the sum of cosine similarities for each vector in valid_action_vectors
315 sum_similarities = similarity_matrix.sum(axis=0)
316
317 N = min(N, len(validActions))
318 # Find the indices of the k vectors with the highest sum of cosine similarities
319 # N = 10 # Change this to the number of top vectors you want to retrieve
320 top_indices = np.argpartition(sum_similarities, -N)[-N:]
321
322 # Print the indices of the top vectors
323 # print(f"The indices of the top {k} vectors in valid_action_vectors are: {top_indices}")
324 # logger.info("The most similar valid actions to the predictions:")
325 # for ti in top_indices:
326 # logger.info("\t\t - "+validActions[ti])
327 if N == 1:
328 action = validActions[top_indices[0]]
329 score = sum_similarities[top_indices[0]]
330 if return_scores:
331 return action, score
332 return action
333 else:
334 action_list = []
335 for i in range(N):
336 action = validActions[top_indices[i]]
337 action_list.append(action)
338 return action_list
339
340
341

Callers 1

Calls 1

encodeMethod · 0.80

Tested by

no test coverage detected