From 57b63d3b8b55aef4aaeab4ef5281bc71549db751 Mon Sep 17 00:00:00 2001 From: 51616 Date: Tue, 3 Jun 2025 14:29:20 +0000 Subject: [PATCH] fix ds probs --- intx_sft.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/intx_sft.py b/intx_sft.py index 30e4c2b..84fa251 100755 --- a/intx_sft.py +++ b/intx_sft.py @@ -2,7 +2,7 @@ import logging import os from copy import deepcopy from functools import partial -from math import ceil +from math import ceil, isclose import numpy as np import torch @@ -70,11 +70,13 @@ def get_ds_prob(train_ds_len: list[int], total_len: int): if ds_len / total_len <= 0.01: probs[i] = 0.01 res_probs = 1 - sum(probs) - res_total_len = sum([l for l in train_ds_len if l / total_len > 0.01]) + res_total_len = sum([l for l in train_ds_len if (l / total_len) > 0.01]) for i, ds_len in enumerate(train_ds_len): - if ds_len / total_len > 0.01: + if (ds_len / total_len) > 0.01: probs[i] = ds_len / res_total_len * res_probs - assert sum(probs) == 1 + assert isclose(sum(probs), 1.0), ( + f"Probs sum to {sum(probs)} ({probs}), expected 1.0" + ) return probs