Kashtan/Detect_Edits_in_AI-Generated_Text
0
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 