Team Ai
Apppublic

Kashtan/Detect_Edits_in_AI-Generated_Text

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
HC_survival_function.py67 linesDownload Raw Back to src
1"""
2This script computes the survival function of the HC statistic for a given sample size n.
3The survival function is computed using a simulation of the null distribution of the HC statistic.
4We use the simulation results to fit a bivariate function of the form Pr[HC >= x | n] = f(n, x).
5The simulation results are saved in a file named HC_null_sim_results.csv.
6use function get_HC_survival_function to load the bivariate function or simulate the distribution. 
7"""
8
9import numpy as np
10import pandas as pd
11from multitest import MultiTest
12from tqdm import tqdm
13from scipy.interpolate import RectBivariateSpline
14from src.fit_survival_function import fit_survival_func
15import logging
16
17HC_NULL_SIM_FILE = "HC_null_sim_results.csv"
18STBL = True
19NN = [25, 50, 75, 100, 125, 150, 200, 250, 300, 400, 500]  # values of n to simulate
20
21def get_HC_survival_function(HC_null_sim_file, log_space=True, nMonte=10000, STBL=True):
22
23    xx = {}
24    if HC_null_sim_file is None:            
25            logging.info("Simulated HC null values file was not provided.")
26            for n in tqdm(NN):
27                logging.info(f"Simulating HC null values for n={n}...")
28                yy = np.zeros(nMonte)
29                for j in range(nMonte):
30                    uu = np.random.rand(n)
31                    mt = MultiTest(uu, stbl=STBL)
32                    yy[j] = mt.hc()[0]
33                xx[n] = yy
34            nn = NN # Idan
35    else:
36        logging.info(f"Loading HC null values from {HC_null_sim_file}...")
37        df = pd.read_csv(HC_null_sim_file, index_col=0)
38        for n in df.index:
39            xx[n] = df.loc[n]
40        nn = df.index.tolist()
41
42    xx0 = np.linspace(-1, 10, 57)
43    zz = []
44    for n in nn:
45        univariate_survival_func = fit_survival_func(xx[n], log_space=log_space)
46        zz.append(univariate_survival_func(xx0))
47        
48    func_log = RectBivariateSpline(np.array(nn), xx0, np.vstack(zz))
49
50    if log_space:
51        def func(x, y):
52            return np.exp(-func_log(x,y))
53        return func
54    else:
55        return func_log
56    
57
58def main():
59    func = get_HC_survival_function(HC_null_sim_file=HC_NULL_SIM_FILE, STBL=STBL)
60    print("Pr[HC >= 3 |n=50] = ", func(50, 3)[0][0]) # 9.680113e-05
61    print("Pr[HC >= 3 |n=100] = ", func(100, 3)[0][0]) # 0.0002335
62    print("Pr[HC >= 3 |n=200] = ", func(200, 3)[0][0]) # 0.00103771
63    
64
65if __name__ == '__main__':
66    main()
67