INNER CODE UNIT · Python
similarity
yixinL7/BRIO · main.py:238
similarity = similarity * args.scale
gold_similarity = gold_similarity * args.scale
similarity = similarity.cpu().numpy()
probs = output["probs"] # [bz, seq_len, word_num]
probs = output["probs"][:, :-1] # truncate last token
gold = batch["candidate_ids"][:, 0, 1:] # shift right
mle_loss += mle_fn(probs.transpose(1, 2), gold)
if i % 1000 == 0:
print(f"test similarity: {similarity[0]}")
max_ids = similarity.argmax(1)
for j in range(similarity.shape[0]):
cnt += 1
sample = samples[j]
sents = sample["candidates"][max_ids[j]][0]
score = rouge_scorer.score("\n".join(sample["abstract"]), "\n".join(sents))
rouge1 += score["rouge1"].fmeasure
rouge2 += score["rouge2"].fmeasure
rougeLsum += score["rougeLsum"].fmeasure