Team Ai
Apppublic

Kashtan/Detect_Edits_in_AI-Generated_Text

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
fit_survival_function.py95 linesDownload Raw Back to src
1"""2Script to read log-loss data of many sentences and characterize the empirical distribution.3We also report the mean log-loss as a function of sentence length4"""5from scipy.interpolate import RectBivariateSpline, interp1d6import numpy as np7 8def fit_survival_func(xx, log_space=True):9    """10    Returns an estimated survival function to the data in :xx: using11    interpolation.12 13    Args:14        :xx:  data15        :log_space:  indicates whether fitting is in log space or not.16 17    Returns:18         univariate function19    """20    assert len(xx) > 021 22    eps = 1 / len(xx)23    inf = 1 / eps24 25    sxx = np.sort(xx)26    qq = np.mean(np.expand_dims(sxx,1) >= sxx, 0)27 28    if log_space:29        qq = -np.log(qq)30 31 32    if log_space:33        return interp1d(sxx, qq, fill_value=(0 , np.log(inf)), bounds_error=False)34    else:35        return interp1d(sxx, qq, fill_value=(1 , 0), bounds_error=False)36 37 38def fit_per_length_survival_function(lengths, xx, G=501, log_space=True):39    """40    Returns a survival function for every sentence length in tokens.41    Use 2D interpolation over the empirical survival function of the pairs (length, x)42    43    Args:44        :lengths:, :xx:, 1-D arrays45        :G:  number of grid points to use in the interpolation in the xx dimension46        :log_space:  indicates whether result is in log space or not.47 48    Returns:49        bivariate function (length, x) -> [0,1]50    """51 52    assert len(lengths) == len(xx)53 54    min_tokens_per_sentence = lengths.min()55    max_tokens_per_sentence = lengths.max()56    ll = np.arange(min_tokens_per_sentence, max_tokens_per_sentence)57 58    ppx_min_val = xx.min()59    ppx_max_val = xx.max()60    xx0 = np.linspace(ppx_min_val, ppx_max_val, G)61 62    ll_valid = []63    zz = []64    for l in ll:65        xx1 = xx[lengths == l]66        if len(xx1) > 1:67            univariate_survival_func = fit_survival_func(xx1, log_space=log_space)68            ll_valid.append(l)69            zz.append(univariate_survival_func(xx0))70 71    func = RectBivariateSpline(np.array(ll_valid), xx0, np.vstack(zz))72    if log_space:73        def func2d(x, y):74            return np.exp(-func(x,y))75        return func2d76    else:77        return func78    79 80# import pickle81# import pandas as pd82# df = pd.read_csv('D:\\.Idan\\תואר שני\\תזה\\detectLM\\article_null.csv')83# LOGLOSS_PVAL_FUNC_FILE = 'D:\.Idan\תואר שני\תזה\detectLM\example\logloss_pval_function.pkl'84# LOGLOSS_PVAL_FUNC_FILE_TEST = 'D:\.Idan\תואר שני\תזה\detectLM\example\logloss_pval_function_test.pkl'85# with open(LOGLOSS_PVAL_FUNC_FILE, 'wb') as handle:86#     pickle.dump(fit_per_length_survival_function(df['length'].values, df['response'].values), handle, protocol=pickle.HIGHEST_PROTOCOL)87 88# with open(LOGLOSS_PVAL_FUNC_FILE, 'rb') as f:89#     data = pickle.load(f)90#     print(data)91 92# with open(LOGLOSS_PVAL_FUNC_FILE_TEST, 'rb') as f:93#     data = pickle.load(f)94#     print(data)95