(sents, embeddings, eos_emb, threshold)
| 66 | |
| 67 | @torch.no_grad() |
| 68 | def predict_next_sentence(sents, embeddings, eos_emb, threshold): |
| 69 | if len(embeddings) == 0: |
| 70 | embeddings = t2vec_model.predict(sents, source_lang="eng_Latn") |
| 71 | embeddings = embeddings.to(device) |
| 72 | |
| 73 | next_sentence, next_embedding = inference_model.inference_step(embeddings) |
| 74 | embedding_re = t2vec_model.predict([next_sentence], source_lang="eng_Latn").to(device) |
| 75 | sim = F.cosine_similarity(embedding_re, eos_emb, dim=1).item() |
| 76 | stop = sim >= threshold |
| 77 | embeddings = torch.cat([embeddings, embedding_re]) |
| 78 | return next_sentence, embeddings, stop |
| 79 | |
| 80 | |
| 81 | @torch.no_grad() |
no test coverage detected