core.py 58 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117111811191120112111221123112411251126112711281129113011311132113311341135113611371138113911401141114211431144114511461147114811491150115111521153115411551156115711581159116011611162116311641165116611671168116911701171117211731174117511761177117811791180118111821183118411851186118711881189119011911192119311941195119611971198119912001201120212031204120512061207120812091210121112121213121412151216121712181219122012211222122312241225122612271228122912301231123212331234123512361237123812391240124112421243124412451246124712481249125012511252125312541255125612571258125912601261126212631264126512661267126812691270127112721273127412751276127712781279128012811282128312841285128612871288128912901291129212931294129512961297129812991300130113021303130413051306130713081309131013111312131313141315131613171318131913201321132213231324132513261327132813291330133113321333133413351336133713381339134013411342134313441345134613471348134913501351135213531354135513561357135813591360136113621363136413651366136713681369137013711372137313741375137613771378137913801381138213831384138513861387138813891390139113921393139413951396139713981399140014011402140314041405140614071408140914101411141214131414141514161417141814191420142114221423142414251426142714281429143014311432143314341435143614371438143914401441144214431444144514461447144814491450145114521453145414551456145714581459146014611462146314641465146614671468146914701471147214731474147514761477147814791480148114821483148414851486148714881489149014911492149314941495149614971498149915001501150215031504150515061507150815091510151115121513151415151516151715181519152015211522152315241525152615271528152915301531153215331534153515361537153815391540154115421543154415451546154715481549155015511552155315541555155615571558155915601561156215631564156515661567156815691570157115721573157415751576157715781579158015811582158315841585158615871588158915901591159215931594159515961597159815991600160116021603160416051606160716081609161016111612161316141615161616171618161916201621162216231624162516261627162816291630163116321633163416351636163716381639164016411642164316441645164616471648164916501651165216531654165516561657165816591660166116621663166416651666166716681669167016711672167316741675167616771678167916801681168216831684168516861687168816891690169116921693169416951696169716981699170017011702170317041705170617071708170917101711171217131714171517161717171817191720172117221723172417251726172717281729173017311732173317341735173617371738173917401741174217431744174517461747174817491750175117521753175417551756175717581759176017611762176317641765176617671768176917701771177217731774177517761777177817791780178117821783178417851786178717881789179017911792179317941795179617971798179918001801180218031804180518061807180818091810181118121813181418151816181718181819182018211822182318241825182618271828182918301831183218331834183518361837183818391840184118421843184418451846184718481849185018511852185318541855185618571858185918601861186218631864186518661867186818691870187118721873187418751876187718781879188018811882188318841885188618871888188918901891189218931894189518961897189818991900190119021903190419051906190719081909191019111912191319141915191619171918191919201921192219231924192519261927192819291930193119321933193419351936193719381939194019411942194319441945194619471948194919501951195219531954195519561957195819591960196119621963196419651966196719681969197019711972197319741975197619771978197919801981198219831984198519861987198819891990199119921993199419951996199719981999200020012002200320042005200620072008200920102011201220132014201520162017201820192020202120222023202420252026202720282029203020312032203320342035203620372038203920402041204220432044204520462047204820492050205120522053205420552056205720582059206020612062206320642065206620672068206920702071207220732074207520762077207820792080208120822083208420852086208720882089209020912092209320942095209620972098209921002101210221032104210521062107210821092110211121122113211421152116211721182119212021212122212321242125212621272128212921302131213221332134213521362137213821392140214121422143214421452146214721482149215021512152215321542155215621572158215921602161216221632164216521662167216821692170217121722173217421752176217721782179218021812182218321842185218621872188218921902191219221932194219521962197219821992200220122022203220422052206220722082209221022112212221322142215221622172218221922202221222222232224222522262227222822292230223122322233223422352236223722382239224022412242224322442245224622472248224922502251225222532254225522562257225822592260226122622263226422652266226722682269227022712272227322742275227622772278227922802281228222832284228522862287228822892290229122922293229422952296229722982299230023012302230323042305230623072308230923102311231223132314231523162317231823192320232123222323232423252326
  1. import numpy as np
  2. import matplotlib.pyplot as plt
  3. from scipy.optimize import minimize
  4. # ============================================================
  5. # Stable sigmoid
  6. # ============================================================
  7. def _sigmoid_stable(z):
  8. z = np.asarray(z, float)
  9. z = np.clip(z, -50.0, 50.0)
  10. return 1.0 / (1.0 + np.exp(-z))
  11. # ============================================================
  12. # 1) Model
  13. # ============================================================
  14. def model_p(x, b):
  15. """p(x|b) = sigmoid(b0 + b1*x)."""
  16. x = np.asarray(x, float).reshape(-1)
  17. b0, b1 = np.asarray(b, float).reshape(2)
  18. return _sigmoid_stable(b0 + b1 * x)
  19. def design_matrix(x):
  20. """Design matrix X = [1, x]."""
  21. x = np.asarray(x, float).reshape(-1)
  22. return np.column_stack([np.ones_like(x), x])
  23. # ============================================================
  24. # 2) Likelihood
  25. # ============================================================
  26. def nll(x, y, b, l2=0.0):
  27. """
  28. Penalized negative log-likelihood:
  29. NLL(b) = -sum[y log p + (1-y) log(1-p)] + 0.5*l2*||b||^2
  30. """
  31. x = np.asarray(x, float).reshape(-1)
  32. y = np.asarray(y, float).reshape(-1)
  33. b = np.asarray(b, float).reshape(2)
  34. p = model_p(x, b)
  35. eps = 1e-12
  36. p = np.clip(p, eps, 1 - eps)
  37. base = -np.sum(y * np.log(p) + (1 - y) * np.log(1 - p))
  38. pen = 0.5 * l2 * float(np.dot(b, b))
  39. return base + pen
  40. def llf(x, y, b):
  41. """
  42. Ordinary (unpenalized) log-likelihood at fitted parameters.
  43. """
  44. x = np.asarray(x, float).reshape(-1)
  45. y = np.asarray(y, float).reshape(-1)
  46. b = np.asarray(b, float).reshape(2)
  47. p = model_p(x, b)
  48. eps = 1e-12
  49. p = np.clip(p, eps, 1 - eps)
  50. return float(np.sum(y * np.log(p) + (1 - y) * np.log(1 - p)))
  51. # ============================================================
  52. # 3) Gradient / Hessian / Covariance
  53. # ============================================================
  54. def grad_nll(x, y, b, l2=0.0):
  55. """
  56. Gradient of penalized NLL:
  57. g(b) = X^T (p - y) + l2*b
  58. """
  59. X = design_matrix(x)
  60. y = np.asarray(y, float).reshape(-1)
  61. b = np.asarray(b, float).reshape(2)
  62. p = model_p(x, b)
  63. return X.T @ (p - y) + l2 * b
  64. def hess_nll(x, b, l2=0.0):
  65. """
  66. Hessian of penalized NLL:
  67. H(b) = X^T W X + l2*I
  68. W = diag(p*(1-p))
  69. """
  70. X = design_matrix(x)
  71. b = np.asarray(b, float).reshape(2)
  72. p = model_p(x, b)
  73. w = p * (1 - p)
  74. return X.T @ (w[:, None] * X) + l2 * np.eye(2)
  75. def covariance(x, b, l2=0.0):
  76. """
  77. Cov(b) ≈ H(b)^(-1), where H is the penalized Hessian if l2 > 0.
  78. Robust to near-singular Hessians.
  79. """
  80. H = hess_nll(x, b, l2=l2)
  81. try:
  82. return np.linalg.inv(H)
  83. except np.linalg.LinAlgError:
  84. return np.linalg.pinv(H)
  85. def standard_errors(x, b, l2=0.0):
  86. """
  87. SE = sqrt(diag(Cov)).
  88. """
  89. C = covariance(x, b, l2=l2)
  90. return np.sqrt(np.maximum(np.diag(C), 0.0))
  91. # Compatibility alias
  92. def logit_poly_cov(x, b, l2=0.0):
  93. return covariance(x, b, l2=l2)
  94. # ============================================================
  95. # 4) Fit
  96. # ============================================================
  97. def fit_newton(x, y, b_start=None, max_iter=50, tol=1e-8, l2=0.0):
  98. """
  99. Newton updates for penalized NLL with backtracking line-search.
  100. Update:
  101. b_new = b - alpha * H^{-1} g
  102. alpha shrinks until NLL decreases.
  103. """
  104. x = np.asarray(x, float).reshape(-1)
  105. y = np.asarray(y, int).reshape(-1)
  106. if b_start is None:
  107. b = np.array([0.0, 0.0], float)
  108. else:
  109. b = np.asarray(b_start, float).reshape(2)
  110. f = nll(x, y, b, l2=l2)
  111. for _ in range(max_iter):
  112. g = grad_nll(x, y, b, l2=l2)
  113. H = hess_nll(x, b, l2=l2)
  114. try:
  115. step = np.linalg.solve(H, g)
  116. except np.linalg.LinAlgError:
  117. step = np.linalg.pinv(H) @ g
  118. alpha = 1.0
  119. while alpha > 1e-6:
  120. b_new = b - alpha * step
  121. f_new = nll(x, y, b_new, l2=l2)
  122. if np.isfinite(f_new) and f_new <= f:
  123. break
  124. alpha *= 0.5
  125. if alpha <= 1e-6:
  126. break
  127. if np.max(np.abs(b_new - b)) < tol:
  128. b = b_new
  129. break
  130. b, f = b_new, f_new
  131. return b
  132. # ============================================================
  133. # 14) Overlay plot (LOG left, RAW right)
  134. # ============================================================
  135. def plot_overlay_two_panels_final(
  136. r_log_full, r_log_trim, r_raw_full, r_raw_trim,
  137. dy_full=-0.010, dy_trim=0.010
  138. ):
  139. import numpy as np
  140. import matplotlib.pyplot as plt
  141. COL_NC = "#4C78A8"
  142. COL_AE = "#F58518"
  143. COL_FULL = "#1f77b4"
  144. COL_TRIM = "#ff7f0e"
  145. fig, axes = plt.subplots(1, 2, figsize=(9.5, 3.5), sharey=True)
  146. ax1, ax2 = axes
  147. def draw_panel(ax, r_full, r_trim, xlabel, panel_label,
  148. show_legend=False):
  149. xF = np.asarray(r_full["x"], float)
  150. yF = np.asarray(r_full["y"], int)
  151. xT = np.asarray(r_trim["x"], float)
  152. yT = np.asarray(r_trim["y"], int)
  153. xx = np.linspace(
  154. min(xF.min(), xT.min()),
  155. max(xF.max(), xT.max()),
  156. 500
  157. )
  158. # keep x-values unchanged
  159. xF_plot = xF
  160. xT_plot = xT
  161. # vertical offsets only
  162. yF_plot = yF + np.where(yF == 0, dy_full, -dy_full)
  163. yT_plot = yT + np.where(yT == 0, dy_trim, -dy_trim)
  164. # FULL = filled markers
  165. ax.scatter(
  166. xF_plot[yF == 0], yF_plot[yF == 0],
  167. s=16,
  168. color=COL_NC,
  169. alpha=0.70,
  170. edgecolors="none",
  171. label="data: FULL NC",
  172. zorder=3
  173. )
  174. ax.scatter(
  175. xF_plot[yF == 1], yF_plot[yF == 1],
  176. s=16,
  177. color=COL_AE,
  178. alpha=0.80,
  179. edgecolors="none",
  180. label="data: FULL AE",
  181. zorder=3
  182. )
  183. # TRIM = outlined markers
  184. ax.scatter(
  185. xT_plot[yT == 0], yT_plot[yT == 0],
  186. s=24,
  187. facecolors=COL_NC,
  188. edgecolors="black",
  189. linewidths=0.45,
  190. alpha=0.95,
  191. label="data: TRIM NC",
  192. zorder=4
  193. )
  194. ax.scatter(
  195. xT_plot[yT == 1], yT_plot[yT == 1],
  196. s=24,
  197. facecolors=COL_AE,
  198. edgecolors="black",
  199. linewidths=0.45,
  200. alpha=0.95,
  201. label="data: TRIM AE",
  202. zorder=4
  203. )
  204. # logistic fits
  205. ax.plot(
  206. xx,
  207. model_p(xx, r_full["b"]),
  208. lw=1.8,
  209. color=COL_FULL,
  210. label="fit FULL",
  211. zorder=2
  212. )
  213. ax.plot(
  214. xx,
  215. model_p(xx, r_trim["b"]),
  216. lw=1.8,
  217. color=COL_TRIM,
  218. label="fit TRIM",
  219. zorder=2
  220. )
  221. # legend only in panel B
  222. if show_legend:
  223. ax.legend(
  224. loc="lower right",
  225. fontsize=7,
  226. markerscale=0.9,
  227. frameon=True,
  228. framealpha=1.0,
  229. edgecolor="0.7",
  230. handlelength=1.8,
  231. borderpad=0.4,
  232. labelspacing=0.4,
  233. handletextpad=0.5
  234. )
  235. ax.set_xlabel(xlabel)
  236. ax.set_ylim(-0.08, 1.08)
  237. ax.text(
  238. 0.05, 0.90,
  239. panel_label,
  240. transform=ax.transAxes,
  241. fontsize=11
  242. )
  243. draw_panel(
  244. ax1,
  245. r_log_full,
  246. r_log_trim,
  247. "log(X)",
  248. "A",
  249. show_legend=False
  250. )
  251. draw_panel(
  252. ax2,
  253. r_raw_full,
  254. r_raw_trim,
  255. "X",
  256. "B",
  257. show_legend=True
  258. )
  259. ax1.set_ylabel("P(AE | X = x)")
  260. for ax in axes:
  261. ax.grid(False)
  262. ax.tick_params(labelsize=8)
  263. plt.tight_layout()
  264. plt.show()
  265. return fig, axes
  266. return fig, axes
  267. # ============================================================
  268. # 5) Goodness of fit
  269. # ============================================================
  270. def goodness_of_fit(x, y, b, thresh=0.5, l2=0.0):
  271. """
  272. Returns:
  273. LLF, NLL, AIC, BIC, Accuracy, n, k
  274. Notes
  275. -----
  276. Fit may use l2 > 0, but GOF metrics below are computed from the
  277. ordinary (unpenalized) likelihood evaluated at the fitted parameters.
  278. The argument l2 is kept only for interface consistency.
  279. """
  280. x = np.asarray(x, float).reshape(-1)
  281. y = np.asarray(y, int).reshape(-1)
  282. b = np.asarray(b, float).reshape(2)
  283. p = model_p(x, b)
  284. eps = 1e-12
  285. p = np.clip(p, eps, 1 - eps)
  286. LLF = np.sum(y * np.log(p) + (1 - y) * np.log(1 - p))
  287. NLL = -LLF
  288. n = len(x)
  289. k = len(b)
  290. AIC = 2 * k - 2 * LLF
  291. BIC = k * np.log(n) - 2 * LLF
  292. yhat = (p >= thresh).astype(int)
  293. acc = np.mean(yhat == y)
  294. return {
  295. "LLF": float(LLF),
  296. "NLL": float(NLL),
  297. "AIC": float(AIC),
  298. "BIC": float(BIC),
  299. "A": float(acc),
  300. "n": int(n),
  301. "k": int(k),
  302. }
  303. # ============================================================
  304. # 6) x50 / Wald helpers / compact fit
  305. # ============================================================
  306. def x50(b):
  307. """
  308. Model-scale midpoint:
  309. x50 = -b0 / b1
  310. For LOG panels, this is on the log(x) scale.
  311. Raw-scale SUV50 is exp(x50).
  312. """
  313. b0, b1 = np.asarray(b, float).reshape(2)
  314. return np.nan if np.abs(b1) < 1e-12 else (-b0 / b1)
  315. def check_x50_consistency(P):
  316. """
  317. Diagnostic check for x50 consistency.
  318. For LOG models:
  319. x50_model is on log(X) scale
  320. SUV50 is on raw X scale = exp(x50_model)
  321. For RAW models:
  322. x50_model = SUV50
  323. Correct result:
  324. P(x50_model) should be approximately 0.5
  325. """
  326. for key, pk in P.items():
  327. b = np.asarray(pk["b"], float).reshape(2)
  328. trans = pk.get("transform", "")
  329. x50_model = x50(b)
  330. suv50 = np.exp(x50_model) if trans == "log" else x50_model
  331. p_at_x50 = model_p(np.array([x50_model]), b)[0]
  332. print(
  333. key,
  334. "| transform =", trans,
  335. "| x50_model =", x50_model,
  336. "| SUV50 =", suv50,
  337. "| P(x50) =", p_at_x50
  338. )
  339. def x50_wald_ci(b, cov, z=1.959963984540054):
  340. """
  341. Wald CI for x50 = -b0/b1 via delta method.
  342. Returned on MODEL scale.
  343. """
  344. b = np.asarray(b, float).reshape(2)
  345. cov = np.asarray(cov, float).reshape(2, 2)
  346. b0, b1 = b
  347. if np.abs(b1) < 1e-12:
  348. return np.nan, np.nan
  349. xhat = -b0 / b1
  350. grad = np.array([-1.0 / b1, b0 / (b1 ** 2)], float)
  351. var = float(grad.T @ cov @ grad)
  352. se = np.sqrt(max(var, 0.0))
  353. return float(xhat - z * se), float(xhat + z * se)
  354. def wald_ci(b, cov, z=1.959963984540054):
  355. """
  356. Wald CI for parameters: b_i ± z*SE_i.
  357. """
  358. b = np.asarray(b, float).reshape(2)
  359. cov = np.asarray(cov, float).reshape(2, 2)
  360. se = np.sqrt(np.maximum(np.diag(cov), 0.0))
  361. return b - z * se, b + z * se
  362. def fit_pack(x, y, name="", thresh=0.5, l2=0.0, z=1.959963984540054):
  363. """
  364. Fit + covariance + GOF + parameter Wald CI.
  365. x should already be on the MODEL scale.
  366. """
  367. x = np.asarray(x, float).reshape(-1)
  368. y = np.asarray(y, int).reshape(-1)
  369. b = fit_newton(x, y, l2=l2)
  370. cov = covariance(x, b, l2=l2)
  371. gof = goodness_of_fit(x, y, b, thresh=thresh, l2=l2)
  372. lcl, ucl = wald_ci(b, cov, z=z)
  373. return {
  374. "name": name,
  375. "x": x,
  376. "y": y,
  377. "b": b,
  378. "cov": cov,
  379. "gof": gof,
  380. "LCL": lcl,
  381. "UCL": ucl,
  382. "l2": float(l2),
  383. }
  384. def trim_nc_by_value(x_raw, y, target=2.48, tol=0.05):
  385. """
  386. Remove ONE NC sample (y==0) with x_raw closest to target.
  387. """
  388. x_raw = np.asarray(x_raw, float).reshape(-1)
  389. y = np.asarray(y, int).reshape(-1)
  390. nc_idx = np.where(y == 0)[0]
  391. if len(nc_idx) == 0:
  392. raise ValueError("No NC samples found (y==0).")
  393. j = nc_idx[np.argmin(np.abs(x_raw[nc_idx] - target))]
  394. diff = float(np.abs(x_raw[j] - target))
  395. if diff > tol:
  396. print(f"[trim warning] closest NC to {target} is {x_raw[j]:.6f} (diff={diff:.6f}) > tol={tol}")
  397. mask = np.ones_like(y, dtype=bool)
  398. mask[j] = False
  399. print(f"[trim] removed index={j}, x_raw={x_raw[j]:.6f}, y={y[j]}")
  400. return x_raw[mask], y[mask]
  401. def s50_from_b(b, transform="raw"):
  402. """
  403. Local slope s50 = dp/dx at x50 on RAW x scale.
  404. """
  405. b = np.asarray(b, float).reshape(2)
  406. b0, b1 = b
  407. if np.abs(b1) < 1e-12:
  408. return np.nan
  409. x50_model = x50(b)
  410. if transform == "raw":
  411. return float(b1 / 4.0)
  412. elif transform == "log":
  413. x50_raw = np.exp(x50_model)
  414. return float(b1 / (4.0 * x50_raw))
  415. else:
  416. raise ValueError("transform must be 'raw' or 'log'")
  417. def s50_normal_ci_from_mvnorm(
  418. b, cov, transform="raw",
  419. M=200000, seed=123, alpha=0.05,
  420. enforce_positive_slope=True, slope_eps=1e-10
  421. ):
  422. """
  423. Normal-on-MLE CI for s50.
  424. """
  425. rng = np.random.default_rng(seed)
  426. b = np.asarray(b, float).reshape(2)
  427. cov = np.asarray(cov, float).reshape(2, 2)
  428. vals = []
  429. tries = 0
  430. max_tries = 20 * M
  431. while len(vals) < M and tries < max_tries:
  432. tries += 1
  433. bb = rng.multivariate_normal(mean=b, cov=cov)
  434. if not np.all(np.isfinite(bb)):
  435. continue
  436. if enforce_positive_slope and bb[1] <= slope_eps:
  437. continue
  438. val = s50_from_b(bb, transform=transform)
  439. if np.isfinite(val):
  440. vals.append(val)
  441. if len(vals) == 0:
  442. return np.nan, np.nan, np.nan, 0
  443. vals = np.asarray(vals, float)
  444. q = np.quantile(vals, [alpha / 2, 0.5, 1.0 - alpha / 2])
  445. return float(q[0]), float(q[1]), float(q[2]), int(len(vals))
  446. def s50_wald_ci_numeric(
  447. b, cov, transform="raw",
  448. z=1.959963984540054,
  449. eps=1e-5
  450. ):
  451. """
  452. Delta-method CI for s50 using numerical derivatives.
  453. """
  454. b = np.asarray(b, float).reshape(2)
  455. cov = np.asarray(cov, float).reshape(2, 2)
  456. s_hat = s50_from_b(b, transform=transform)
  457. if not np.isfinite(s_hat):
  458. return np.nan, np.nan
  459. grad = np.zeros(2, float)
  460. for j in range(2):
  461. step = eps * max(1.0, abs(b[j]))
  462. bp = b.copy()
  463. bm = b.copy()
  464. bp[j] += step
  465. bm[j] -= step
  466. sp = s50_from_b(bp, transform=transform)
  467. sm = s50_from_b(bm, transform=transform)
  468. grad[j] = (sp - sm) / (2.0 * step)
  469. var = float(grad.T @ cov @ grad)
  470. se = np.sqrt(max(var, 0.0))
  471. return float(s_hat - z * se), float(s_hat + z * se)
  472. # ============================================================
  473. # 7) Alternative confidence-interval estimation
  474. # ============================================================
  475. # Final terminology:
  476. # Wald : analytical approximation using the fitted covariance;
  477. # MC : Monte Carlo propagation from the local Gaussian approximation;
  478. # Nonparametric : ordinary patient-level nonparametric bootstrap;
  479. # Stratified : class-stratified nonparametric bootstrap, retained for comparison;
  480. # Parametric : model-based Bernoulli bootstrap.
  481. #
  482. # The delta method is used internally for analytical propagation under
  483. # the Wald approximation; it is not treated as a separate method.
  484. from collections import OrderedDict
  485. CI_METHODS = [
  486. "Wald",
  487. "MC",
  488. "Nonparametric",
  489. "Stratified",
  490. "Parametric",
  491. ]
  492. MC_DRAWS_BANDS = 20_000
  493. MC_DRAWS_TABLE = 200_000
  494. def eta_se_grid(x_grid, cov):
  495. """Standard error of eta(x) = b0 + b1*x on a model-scale grid."""
  496. x_grid = np.asarray(x_grid, float).reshape(-1)
  497. cov = np.asarray(cov, float).reshape(2, 2)
  498. Xg = design_matrix(x_grid)
  499. var_eta = np.einsum("ij,jk,ik->i", Xg, cov, Xg)
  500. return np.sqrt(np.maximum(var_eta, 0.0))
  501. def ci_band_wald(x_grid, b, cov, z=1.959963984540054):
  502. """
  503. Pointwise Wald confidence band for p(x).
  504. The fitted-parameter covariance is propagated to the probability scale
  505. using the first-order delta method.
  506. """
  507. x_grid = np.asarray(x_grid, float).reshape(-1)
  508. b = np.asarray(b, float).reshape(2)
  509. p = model_p(x_grid, b)
  510. se_eta = eta_se_grid(x_grid, cov)
  511. se_p = p * (1.0 - p) * se_eta
  512. lo = np.clip(p - z * se_p, 0.0, 1.0)
  513. hi = np.clip(p + z * se_p, 0.0, 1.0)
  514. return lo, p, hi
  515. # Historical alias retained for notebook compatibility.
  516. def ci_band_delta(x_grid, b, cov, z=1.959963984540054):
  517. return ci_band_wald(x_grid, b, cov, z=z)
  518. def gaussian_parameter_draws(
  519. b,
  520. cov,
  521. M=MC_DRAWS_BANDS,
  522. seed=123,
  523. enforce_positive_slope=True,
  524. slope_eps=1e-10,
  525. x50_bounds=None,
  526. ):
  527. """
  528. Draw beta* ~ N(beta_hat, Cov_hat) for Monte Carlo propagation.
  529. Parameters
  530. ----------
  531. x50_bounds : tuple(float, float) or None
  532. Optional admissible interval for model-scale x50. When supplied,
  533. draws with x50 outside [lower, upper] are rejected.
  534. """
  535. rng = np.random.default_rng(seed)
  536. b = np.asarray(b, float).reshape(2)
  537. cov = np.asarray(cov, float).reshape(2, 2)
  538. draws = []
  539. attempts = 0
  540. rejected_nonfinite = 0
  541. rejected_slope = 0
  542. rejected_x50 = 0
  543. max_attempts = max(50 * int(M), 1000)
  544. while len(draws) < int(M) and attempts < max_attempts:
  545. attempts += 1
  546. try:
  547. bb = rng.multivariate_normal(mean=b, cov=cov)
  548. except Exception:
  549. break
  550. if not np.all(np.isfinite(bb)):
  551. rejected_nonfinite += 1
  552. continue
  553. if enforce_positive_slope and bb[1] <= slope_eps:
  554. rejected_slope += 1
  555. continue
  556. if x50_bounds is not None:
  557. x50_draw = x50(bb)
  558. if not np.isfinite(x50_draw):
  559. rejected_x50 += 1
  560. continue
  561. x50_lower, x50_upper = x50_bounds
  562. if not (x50_lower <= x50_draw <= x50_upper):
  563. rejected_x50 += 1
  564. continue
  565. draws.append(bb)
  566. arr = np.asarray(draws, float) if draws else np.empty((0, 2), float)
  567. diagnostics = {
  568. "attempted": int(attempts),
  569. "successful": int(len(arr)),
  570. "rejected": int(attempts - len(arr)),
  571. "rejected_nonfinite": int(rejected_nonfinite),
  572. "rejected_slope": int(rejected_slope),
  573. "rejected_x50": int(rejected_x50),
  574. "success_rate": (
  575. float(len(arr) / attempts)
  576. if attempts > 0 else np.nan
  577. ),
  578. }
  579. return arr, diagnostics
  580. def bootstrap_band_from_params(x_grid, pars, alpha=0.05):
  581. """Convert parameter draws to pointwise confidence bands."""
  582. x_grid = np.asarray(x_grid, float).reshape(-1)
  583. pars = np.asarray(pars, float)
  584. if pars.ndim != 2 or pars.shape[0] == 0:
  585. nan = np.full_like(x_grid, np.nan, dtype=float)
  586. return nan, nan, nan
  587. curves = np.asarray([model_p(x_grid, bb) for bb in pars], float)
  588. q = np.quantile(curves, [alpha / 2, 0.5, 1.0 - alpha / 2], axis=0)
  589. return q[0], q[1], q[2]
  590. def ci_band_normal_mle_sim(
  591. x_grid,
  592. b,
  593. cov,
  594. M=MC_DRAWS_BANDS,
  595. seed=123,
  596. alpha=0.05,
  597. enforce_positive_slope=True,
  598. enforce_x50_in_grid=False,
  599. slope_eps=1e-10,
  600. ):
  601. """
  602. Monte Carlo confidence band from the local Gaussian approximation.
  603. The historical function name and signature are retained. The old
  604. x50-in-grid filter is intentionally ignored.
  605. """
  606. draws, _ = gaussian_parameter_draws(
  607. b,
  608. cov,
  609. M=M,
  610. seed=seed,
  611. enforce_positive_slope=enforce_positive_slope,
  612. slope_eps=slope_eps,
  613. )
  614. return bootstrap_band_from_params(x_grid, draws, alpha=alpha)
  615. def x50_normal_ci_from_mvnorm(
  616. b,
  617. cov,
  618. M=MC_DRAWS_TABLE,
  619. seed=123,
  620. alpha=0.05,
  621. enforce_positive_slope=True,
  622. slope_eps=1e-10,
  623. ):
  624. """Monte Carlo interval for model-scale x50 = -b0/b1."""
  625. draws, _ = gaussian_parameter_draws(
  626. b,
  627. cov,
  628. M=M,
  629. seed=seed,
  630. enforce_positive_slope=enforce_positive_slope,
  631. slope_eps=slope_eps,
  632. )
  633. if len(draws) == 0:
  634. return np.nan, np.nan, np.nan, 0
  635. vals = np.asarray([x50(bb) for bb in draws], float)
  636. vals = vals[np.isfinite(vals)]
  637. if len(vals) == 0:
  638. return np.nan, np.nan, np.nan, 0
  639. q = np.quantile(vals, [alpha / 2, 0.5, 1.0 - alpha / 2])
  640. return float(q[0]), float(q[1]), float(q[2]), int(len(vals))
  641. def s50_mc_ci_from_mvnorm(
  642. b,
  643. cov,
  644. transform="raw",
  645. M=MC_DRAWS_TABLE,
  646. seed=123,
  647. alpha=0.05,
  648. enforce_positive_slope=True,
  649. slope_eps=1e-10,
  650. ):
  651. """Monte Carlo interval for raw-scale midpoint slope s50."""
  652. draws, _ = gaussian_parameter_draws(
  653. b,
  654. cov,
  655. M=M,
  656. seed=seed,
  657. enforce_positive_slope=enforce_positive_slope,
  658. slope_eps=slope_eps,
  659. )
  660. if len(draws) == 0:
  661. return np.nan, np.nan, np.nan, 0
  662. vals = np.asarray(
  663. [s50_from_b(bb, transform=transform) for bb in draws],
  664. float,
  665. )
  666. vals = vals[np.isfinite(vals)]
  667. if len(vals) == 0:
  668. return np.nan, np.nan, np.nan, 0
  669. q = np.quantile(vals, [alpha / 2, 0.5, 1.0 - alpha / 2])
  670. return float(q[0]), float(q[1]), float(q[2]), int(len(vals))
  671. # Historical alias retained.
  672. def s50_normal_ci_from_mvnorm(
  673. b,
  674. cov,
  675. transform="raw",
  676. M=MC_DRAWS_TABLE,
  677. seed=123,
  678. alpha=0.05,
  679. enforce_positive_slope=True,
  680. slope_eps=1e-10,
  681. ):
  682. return s50_mc_ci_from_mvnorm(
  683. b,
  684. cov,
  685. transform=transform,
  686. M=M,
  687. seed=seed,
  688. alpha=alpha,
  689. enforce_positive_slope=enforce_positive_slope,
  690. slope_eps=slope_eps,
  691. )
  692. # ============================================================
  693. # 8) Bootstrap parameter generators
  694. # ============================================================
  695. def bootstrap_params_nonparametric(
  696. x,
  697. y,
  698. B=2000,
  699. seed=123,
  700. l2=0.0,
  701. b_start=None,
  702. ):
  703. """Ordinary patient-level nonparametric bootstrap."""
  704. rng = np.random.default_rng(seed)
  705. x = np.asarray(x, float).reshape(-1)
  706. y = np.asarray(y, int).reshape(-1)
  707. n = len(y)
  708. out = []
  709. failed = 0
  710. for _ in range(int(B)):
  711. idx = rng.choice(n, size=n, replace=True)
  712. xb = x[idx]
  713. yb = y[idx]
  714. if np.unique(yb).size < 2:
  715. failed += 1
  716. continue
  717. try:
  718. bb = fit_newton(xb, yb, b_start=b_start, l2=l2)
  719. if np.all(np.isfinite(bb)):
  720. out.append(bb)
  721. else:
  722. failed += 1
  723. except Exception:
  724. failed += 1
  725. arr = np.asarray(out, float) if out else np.empty((0, 2), float)
  726. diagnostics = {
  727. "attempted": int(B),
  728. "successful": int(len(arr)),
  729. "failed": int(failed),
  730. }
  731. return arr, diagnostics
  732. # Historical alias retained for older notebook cells.
  733. def bootstrap_params_ordinary(
  734. x,
  735. y,
  736. B=2000,
  737. seed=123,
  738. l2=0.0,
  739. b_start=None,
  740. ):
  741. return bootstrap_params_nonparametric(
  742. x,
  743. y,
  744. B=B,
  745. seed=seed,
  746. l2=l2,
  747. b_start=b_start,
  748. )
  749. def bootstrap_params_stratified(
  750. x,
  751. y,
  752. B=2000,
  753. seed=123,
  754. l2=0.0,
  755. b_start=None,
  756. ):
  757. """Class-stratified nonparametric bootstrap preserving class counts."""
  758. rng = np.random.default_rng(seed)
  759. x = np.asarray(x, float).reshape(-1)
  760. y = np.asarray(y, int).reshape(-1)
  761. x0 = x[y == 0]
  762. x1 = x[y == 1]
  763. n0 = len(x0)
  764. n1 = len(x1)
  765. if n0 == 0 or n1 == 0:
  766. return np.empty((0, 2), float), {
  767. "attempted": int(B),
  768. "successful": 0,
  769. "failed": int(B),
  770. }
  771. out = []
  772. failed = 0
  773. for _ in range(int(B)):
  774. xb0 = rng.choice(x0, size=n0, replace=True)
  775. xb1 = rng.choice(x1, size=n1, replace=True)
  776. xb = np.concatenate([xb0, xb1])
  777. yb = np.concatenate([
  778. np.zeros(n0, dtype=int),
  779. np.ones(n1, dtype=int),
  780. ])
  781. try:
  782. bb = fit_newton(xb, yb, b_start=b_start, l2=l2)
  783. if np.all(np.isfinite(bb)):
  784. out.append(bb)
  785. else:
  786. failed += 1
  787. except Exception:
  788. failed += 1
  789. arr = np.asarray(out, float) if out else np.empty((0, 2), float)
  790. diagnostics = {
  791. "attempted": int(B),
  792. "successful": int(len(arr)),
  793. "failed": int(failed),
  794. "success_rate": float(len(arr) / B) if B > 0 else np.nan,
  795. }
  796. return arr, diagnostics
  797. def bootstrap_params_parametric(
  798. x,
  799. b,
  800. B=2000,
  801. seed=123,
  802. l2=0.0,
  803. min_ae=2,
  804. ):
  805. """
  806. Parametric bootstrap with y* ~ Bernoulli[p_hat(x)].
  807. ``min_ae`` is retained for compatibility.
  808. """
  809. rng = np.random.default_rng(seed)
  810. x = np.asarray(x, float).reshape(-1)
  811. b = np.asarray(b, float).reshape(2)
  812. p = model_p(x, b)
  813. n = len(x)
  814. out = []
  815. tries = 0
  816. max_tries = max(10 * int(B), 1000)
  817. while len(out) < int(B) and tries < max_tries:
  818. tries += 1
  819. yb = rng.binomial(1, p, size=n).astype(int)
  820. n1 = int(np.sum(yb))
  821. n0 = n - n1
  822. if n1 < int(min_ae) or n0 < 1:
  823. continue
  824. try:
  825. bb = fit_newton(x, yb, b_start=b, l2=l2)
  826. if np.all(np.isfinite(bb)):
  827. out.append(bb)
  828. except Exception:
  829. pass
  830. arr = np.asarray(out, float) if out else np.empty((0, 2), float)
  831. diagnostics = {
  832. "attempted": int(tries),
  833. "successful": int(len(arr)),
  834. "failed_or_rejected": int(tries - len(arr)),
  835. "success_rate": float(len(arr) / tries) if tries > 0 else np.nan,
  836. }
  837. return arr, diagnostics
  838. # ============================================================
  839. # 9) High-level wrapper for one panel
  840. # ============================================================
  841. def fit_ci_pack_rawgrid(
  842. x_raw,
  843. y,
  844. transform="raw",
  845. xmax_raw=None,
  846. grid_n=500,
  847. name="",
  848. l2=0.0,
  849. B=2000,
  850. seed=123,
  851. min_ae=2,
  852. z=1.959963984540054,
  853. ):
  854. """
  855. Fit one validated logistic model and construct five uncertainty summaries.
  856. """
  857. x_raw = np.asarray(x_raw, float).reshape(-1)
  858. y = np.asarray(y, int).reshape(-1)
  859. if transform not in ("raw", "log"):
  860. raise ValueError("transform must be 'raw' or 'log'")
  861. x_raw = np.clip(x_raw, 1e-12, None)
  862. x_model = x_raw if transform == "raw" else np.log(x_raw)
  863. # Single validated fitting path.
  864. validated = fit_pack(
  865. x_model,
  866. y,
  867. name=name,
  868. l2=l2,
  869. z=z,
  870. )
  871. b = validated["b"]
  872. cov = validated["cov"]
  873. gof = validated["gof"]
  874. xmin_raw = float(np.min(x_raw))
  875. xmax0 = float(np.max(x_raw))
  876. xmax_use = xmax0 if xmax_raw is None else max(float(xmax_raw), xmax0)
  877. x_grid_raw = np.linspace(xmin_raw, xmax_use, int(grid_n))
  878. x_grid_raw = np.clip(x_grid_raw, 1e-12, None)
  879. x_grid_model = x_grid_raw if transform == "raw" else np.log(x_grid_raw)
  880. mc_x50_bounds = (
  881. float(np.min(x_grid_model)),
  882. float(np.max(x_grid_model)),
  883. )
  884. mc_draws, diag_mc = gaussian_parameter_draws(
  885. b,
  886. cov,
  887. M=MC_DRAWS_BANDS,
  888. seed=seed + 10,
  889. enforce_positive_slope=True,
  890. x50_bounds=mc_x50_bounds,
  891. )
  892. pars_np, diag_np = bootstrap_params_nonparametric(
  893. x_model,
  894. y,
  895. B=B,
  896. seed=seed + 1,
  897. l2=l2,
  898. b_start=b,
  899. )
  900. pars_str, diag_str = bootstrap_params_stratified(
  901. x_model,
  902. y,
  903. B=B,
  904. seed=seed + 2,
  905. l2=l2,
  906. b_start=b,
  907. )
  908. pars_pm, diag_pm = bootstrap_params_parametric(
  909. x_model,
  910. b,
  911. B=B,
  912. seed=seed + 3,
  913. l2=l2,
  914. min_ae=min_ae,
  915. )
  916. return {
  917. "name": name,
  918. "transform": transform,
  919. "l2": float(l2),
  920. "x_raw": x_raw,
  921. "x_model": x_model,
  922. "y": y,
  923. "x_grid_raw": x_grid_raw,
  924. "x_grid_model": x_grid_model,
  925. "b": b,
  926. "cov": cov,
  927. "gof": gof,
  928. "LCL": validated["LCL"],
  929. "UCL": validated["UCL"],
  930. "bands": OrderedDict([
  931. ("Wald", ci_band_wald(x_grid_model, b, cov, z=z)),
  932. ("MC", bootstrap_band_from_params(x_grid_model, mc_draws)),
  933. ("Nonparametric", bootstrap_band_from_params(x_grid_model, pars_np)),
  934. ("Stratified", bootstrap_band_from_params(x_grid_model, pars_str)),
  935. ("Parametric", bootstrap_band_from_params(x_grid_model, pars_pm)),
  936. ]),
  937. "pars_mc": mc_draws,
  938. "pars_nonparametric": pars_np,
  939. "pars_stratified": pars_str,
  940. "pars_parametric": pars_pm,
  941. # Historical aliases
  942. "pars_normal": mc_draws,
  943. "pars_nonparam": pars_np,
  944. "pars_nonparam_ordinary": pars_np,
  945. "pars_nonparam_stratified": pars_str,
  946. "bootstrap_diagnostics": OrderedDict([
  947. ("MC", diag_mc),
  948. ("Nonparametric", diag_np),
  949. ("Stratified", diag_str),
  950. ("Parametric", diag_pm),
  951. ]),
  952. }
  953. # ============================================================
  954. # 10) Model-band table with LL / UL
  955. # ============================================================
  956. def model_ci_table_4methods(
  957. P,
  958. keys=("FULL-RAW", "TRIM-RAW", "FULL-LOG", "TRIM-LOG"),
  959. ):
  960. """Long pointwise model-band table."""
  961. import pandas as pd
  962. rows = []
  963. for key in keys:
  964. pk = P[key]
  965. xg_raw = np.asarray(pk["x_grid_raw"], float)
  966. xg_mod = np.asarray(pk["x_grid_model"], float)
  967. trans = pk.get("transform", "")
  968. for method in CI_METHODS:
  969. if method not in pk["bands"]:
  970. continue
  971. lo, md, hi = pk["bands"][method]
  972. lo = np.asarray(lo, float)
  973. md = np.asarray(md, float)
  974. hi = np.asarray(hi, float)
  975. for i in range(len(xg_raw)):
  976. rows.append({
  977. "Panel": key,
  978. "Method": method,
  979. "transform": trans,
  980. "x_grid_raw": float(xg_raw[i]),
  981. "x_grid_model": float(xg_mod[i]),
  982. "fit": float(md[i]),
  983. "LL": float(lo[i]),
  984. "UL": float(hi[i]),
  985. "width": float(hi[i] - lo[i]),
  986. })
  987. return pd.DataFrame(rows)
  988. # ============================================================
  989. # 11) Parameter, x50, and s50 CI summary table
  990. # ============================================================
  991. def param_ci_table_4methods(
  992. P,
  993. keys=("FULL-RAW", "TRIM-RAW", "FULL-LOG", "TRIM-LOG"),
  994. z=1.959963984540054,
  995. alpha=0.05,
  996. include_point_est=True,
  997. M_normal=MC_DRAWS_TABLE,
  998. seed_normal=123,
  999. ):
  1000. """
  1001. Notebook-compatible CI summary using:
  1002. Wald, MC, Nonparametric, Stratified, Parametric.
  1003. """
  1004. import pandas as pd
  1005. def _quantile_ci(values):
  1006. values = np.asarray(values, float)
  1007. values = values[np.isfinite(values)]
  1008. if len(values) == 0:
  1009. return np.nan, np.nan, np.nan
  1010. q = np.quantile(values, [alpha / 2, 0.5, 1.0 - alpha / 2])
  1011. return float(q[0]), float(q[1]), float(q[2])
  1012. def _to_raw_x50_scalar(value, trans):
  1013. if not np.isfinite(value):
  1014. return np.nan
  1015. return float(np.exp(value)) if trans == "log" else float(value)
  1016. def _width(lo, hi):
  1017. if np.isfinite(lo) and np.isfinite(hi):
  1018. return float(hi - lo)
  1019. return np.nan
  1020. rows = []
  1021. for ik, key in enumerate(keys):
  1022. pk = P[key]
  1023. b = np.asarray(pk["b"], float).reshape(2)
  1024. cov = np.asarray(pk["cov"], float).reshape(2, 2)
  1025. trans = pk.get("transform", "")
  1026. l2 = float(pk.get("l2", 0.0))
  1027. b0_hat = float(b[0])
  1028. b1_hat = float(b[1])
  1029. x50_hat = float(x50(b))
  1030. suv50_hat = _to_raw_x50_scalar(x50_hat, trans)
  1031. s50_hat = float(s50_from_b(b, transform=trans))
  1032. def add_row(method, b0_ci, b1_ci, x_ci, suv_ci, s_ci, n_used):
  1033. row = {
  1034. "Panel": key,
  1035. "Method": method,
  1036. "b0_hat": b0_hat,
  1037. "b0_LCL": b0_ci[0],
  1038. "b0_UCL": b0_ci[2],
  1039. "b0_width": _width(b0_ci[0], b0_ci[2]),
  1040. "b1_hat": b1_hat,
  1041. "b1_LCL": b1_ci[0],
  1042. "b1_UCL": b1_ci[2],
  1043. "b1_width": _width(b1_ci[0], b1_ci[2]),
  1044. "x50_hat": x50_hat,
  1045. "x50_med": x_ci[1],
  1046. "x50_LCL": x_ci[0],
  1047. "x50_UCL": x_ci[2],
  1048. "x50_width": _width(x_ci[0], x_ci[2]),
  1049. "SUV50_hat": suv50_hat,
  1050. "SUV50_med": suv_ci[1],
  1051. "SUV50_LCL": suv_ci[0],
  1052. "SUV50_UCL": suv_ci[2],
  1053. "SUV50_width": _width(suv_ci[0], suv_ci[2]),
  1054. "s50_hat": s50_hat,
  1055. "s50_med": s_ci[1],
  1056. "s50_LCL": s_ci[0],
  1057. "s50_UCL": s_ci[2],
  1058. "s50_width": _width(s_ci[0], s_ci[2]),
  1059. "transform": trans,
  1060. "l2": l2,
  1061. "B_used": n_used,
  1062. }
  1063. if not include_point_est:
  1064. for col in (
  1065. "b0_hat",
  1066. "b1_hat",
  1067. "x50_hat",
  1068. "x50_med",
  1069. "SUV50_hat",
  1070. "SUV50_med",
  1071. "s50_hat",
  1072. "s50_med",
  1073. "transform",
  1074. "l2",
  1075. ):
  1076. row.pop(col, None)
  1077. rows.append(row)
  1078. # Wald
  1079. lcl, ucl = wald_ci(b, cov, z=z)
  1080. x_l, x_u = x50_wald_ci(b, cov, z=z)
  1081. suv_l = _to_raw_x50_scalar(x_l, trans)
  1082. suv_u = _to_raw_x50_scalar(x_u, trans)
  1083. s_l, s_u = s50_wald_ci_numeric(b, cov, transform=trans, z=z)
  1084. add_row(
  1085. "Wald",
  1086. (float(lcl[0]), b0_hat, float(ucl[0])),
  1087. (float(lcl[1]), b1_hat, float(ucl[1])),
  1088. (float(x_l), x50_hat, float(x_u)),
  1089. (float(suv_l), suv50_hat, float(suv_u)),
  1090. (float(s_l), s50_hat, float(s_u)),
  1091. np.nan,
  1092. )
  1093. # Draw-based methods
  1094. mc_draws = pk.get("pars_mc", pk.get("pars_normal"))
  1095. if mc_draws is None or len(mc_draws) < int(M_normal):
  1096. mc_draws, _ = gaussian_parameter_draws(
  1097. b,
  1098. cov,
  1099. M=M_normal,
  1100. seed=seed_normal + 1000 * ik,
  1101. enforce_positive_slope=True,
  1102. )
  1103. draw_sets = {
  1104. "MC": mc_draws,
  1105. "Nonparametric": pk.get(
  1106. "pars_nonparametric",
  1107. pk.get("pars_nonparam", np.empty((0, 2))),
  1108. ),
  1109. "Stratified": pk.get(
  1110. "pars_stratified",
  1111. pk.get("pars_nonparam_stratified", np.empty((0, 2))),
  1112. ),
  1113. "Parametric": pk.get(
  1114. "pars_parametric",
  1115. np.empty((0, 2)),
  1116. ),
  1117. }
  1118. for method in ("MC", "Nonparametric", "Stratified", "Parametric"):
  1119. pars = np.asarray(draw_sets[method], float)
  1120. if pars.ndim != 2 or len(pars) == 0:
  1121. nan3 = (np.nan, np.nan, np.nan)
  1122. add_row(method, nan3, nan3, nan3, nan3, nan3, 0)
  1123. continue
  1124. b0_ci = _quantile_ci(pars[:, 0])
  1125. b1_ci = _quantile_ci(pars[:, 1])
  1126. xvals = np.asarray([x50(bb) for bb in pars], float)
  1127. suvvals = np.asarray(
  1128. [_to_raw_x50_scalar(v, trans) for v in xvals],
  1129. float,
  1130. )
  1131. svals = np.asarray(
  1132. [s50_from_b(bb, transform=trans) for bb in pars],
  1133. float,
  1134. )
  1135. add_row(
  1136. method,
  1137. b0_ci,
  1138. b1_ci,
  1139. _quantile_ci(xvals),
  1140. _quantile_ci(suvvals),
  1141. _quantile_ci(svals),
  1142. int(len(pars)),
  1143. )
  1144. return pd.DataFrame(rows)
  1145. def combined_x50_model_bounds_table(
  1146. P,
  1147. keys=("FULL-RAW", "TRIM-RAW", "FULL-LOG", "TRIM-LOG"),
  1148. z=1.959963984540054,
  1149. alpha=0.05,
  1150. M_normal=MC_DRAWS_TABLE,
  1151. seed_normal=123,
  1152. ):
  1153. """Combine characteristic intervals with global curve-band summaries."""
  1154. import pandas as pd
  1155. param_df = param_ci_table_4methods(
  1156. P,
  1157. keys=keys,
  1158. z=z,
  1159. alpha=alpha,
  1160. include_point_est=True,
  1161. M_normal=M_normal,
  1162. seed_normal=seed_normal,
  1163. ).copy()
  1164. model_df = model_ci_table_4methods(P, keys=keys).copy()
  1165. global_df = (
  1166. model_df
  1167. .groupby(["Panel", "Method"], as_index=False)
  1168. .agg(
  1169. global_LL=("LL", "min"),
  1170. global_UL=("UL", "max"),
  1171. fit_min=("fit", "min"),
  1172. fit_max=("fit", "max"),
  1173. mean_width=("width", "mean"),
  1174. max_width=("width", "max"),
  1175. )
  1176. )
  1177. global_df["global_width"] = global_df["global_UL"] - global_df["global_LL"]
  1178. out = pd.merge(
  1179. param_df,
  1180. global_df,
  1181. on=["Panel", "Method"],
  1182. how="left",
  1183. )
  1184. preferred = [
  1185. "Panel", "Method", "transform",
  1186. "x50_hat", "x50_LCL", "x50_UCL", "x50_width",
  1187. "SUV50_hat", "SUV50_LCL", "SUV50_UCL", "SUV50_width",
  1188. "s50_hat", "s50_LCL", "s50_UCL", "s50_width",
  1189. "global_LL", "global_UL", "global_width",
  1190. "fit_min", "fit_max", "mean_width", "max_width", "B_used",
  1191. ]
  1192. cols = [c for c in preferred if c in out.columns] + [
  1193. c for c in out.columns if c not in preferred
  1194. ]
  1195. return out[cols]
  1196. def bootstrap_diagnostics_table(P):
  1197. """Return success information for MC and bootstrap methods."""
  1198. import pandas as pd
  1199. rows = []
  1200. for panel, pk in P.items():
  1201. for method, diag in pk.get("bootstrap_diagnostics", {}).items():
  1202. rows.append({"Panel": panel, "Method": method, **diag})
  1203. return pd.DataFrame(rows)
  1204. # ============================================================
  1205. # 12) CI figure
  1206. # ============================================================
  1207. def plot_ci_four_panels(P):
  1208. import numpy as np
  1209. import matplotlib.pyplot as plt
  1210. import matplotlib.lines as mlines
  1211. plt.style.use("default")
  1212. COL_NC = "#4c9ed9"
  1213. COL_AE = "#f28e2b"
  1214. COL_FIT = "#000000"
  1215. # Methods displayed in the main figure
  1216. FIGURE_METHODS = [
  1217. "Wald",
  1218. "MC",
  1219. "Nonparametric",
  1220. "Parametric",
  1221. ]
  1222. styles = OrderedDict([
  1223. ("Wald", ("#2ca02c", "-.", 0.12)),
  1224. ("MC", ("#d62728", "--", 0.14)),
  1225. ("Nonparametric", ("#1f77b4", ":", 0.16)),
  1226. ("Parametric", ("#17becf", (0, (6, 2)), 0.14)),
  1227. ])
  1228. panel_order = [
  1229. "FULL-LOG",
  1230. "FULL-RAW",
  1231. "TRIM-LOG",
  1232. "TRIM-RAW",
  1233. ]
  1234. panel_letters = ["A", "B", "C", "D"]
  1235. fig, axs = plt.subplots(
  1236. 2,
  1237. 2,
  1238. figsize=(15, 10),
  1239. dpi=180,
  1240. sharex="col",
  1241. sharey=True,
  1242. )
  1243. for ax, key, letter in zip(
  1244. axs.flat,
  1245. panel_order,
  1246. panel_letters,
  1247. ):
  1248. pk = P[key]
  1249. x_raw = np.asarray(pk["x_raw"], float)
  1250. y = np.asarray(pk["y"], int)
  1251. xg_raw = np.asarray(pk["x_grid_raw"], float)
  1252. xg_model = np.asarray(pk["x_grid_model"], float)
  1253. transform = pk["transform"]
  1254. if transform == "raw":
  1255. xs = x_raw
  1256. xg = xg_raw
  1257. else:
  1258. xs = np.log(x_raw)
  1259. xg = xg_model
  1260. rng = np.random.default_rng(123 + ord(letter))
  1261. jit = (rng.random(len(y)) - 0.5) * 0.04
  1262. ax.scatter(
  1263. xs[y == 0],
  1264. (y + jit)[y == 0],
  1265. s=22,
  1266. alpha=0.45,
  1267. color=COL_NC,
  1268. edgecolors="none",
  1269. zorder=5,
  1270. )
  1271. ax.scatter(
  1272. xs[y == 1],
  1273. (y + jit)[y == 1],
  1274. s=24,
  1275. alpha=0.85,
  1276. color=COL_AE,
  1277. edgecolors="none",
  1278. zorder=5,
  1279. )
  1280. for method in FIGURE_METHODS:
  1281. if method not in pk["bands"]:
  1282. continue
  1283. lo, _, hi = pk["bands"][method]
  1284. color, linestyle, fill_alpha = styles[method]
  1285. ax.fill_between(
  1286. xg,
  1287. lo,
  1288. hi,
  1289. color=color,
  1290. alpha=fill_alpha,
  1291. zorder=1,
  1292. )
  1293. ax.plot(
  1294. xg,
  1295. lo,
  1296. color=color,
  1297. linestyle=linestyle,
  1298. lw=1.6,
  1299. zorder=2,
  1300. )
  1301. ax.plot(
  1302. xg,
  1303. hi,
  1304. color=color,
  1305. linestyle=linestyle,
  1306. lw=1.6,
  1307. zorder=2,
  1308. )
  1309. fit_curve = model_p(
  1310. xg_model,
  1311. pk["b"],
  1312. )
  1313. ax.plot(
  1314. xg,
  1315. fit_curve,
  1316. color=COL_FIT,
  1317. lw=2.5,
  1318. zorder=6,
  1319. )
  1320. ax.text(
  1321. 0.03,
  1322. 0.95,
  1323. letter,
  1324. transform=ax.transAxes,
  1325. fontsize=15,
  1326. ha="left",
  1327. va="top",
  1328. )
  1329. ax.set_ylim(-0.05, 1.05)
  1330. ax.grid(False)
  1331. ax.tick_params(
  1332. axis="both",
  1333. which="major",
  1334. labelsize=11,
  1335. length=4,
  1336. width=0.8,
  1337. direction="out",
  1338. )
  1339. axs[0, 0].set_ylabel(
  1340. r"$\mathrm{P(AE \mid X = x)}$",
  1341. fontsize=13,
  1342. )
  1343. axs[1, 0].set_ylabel(
  1344. r"$\mathrm{P(AE \mid X = x)}$",
  1345. fontsize=13,
  1346. )
  1347. axs[1, 0].set_xlabel(
  1348. r"$\log(\mathrm{X})$",
  1349. fontsize=13,
  1350. )
  1351. axs[1, 1].set_xlabel(
  1352. r"$\mathrm{X}$",
  1353. fontsize=13,
  1354. )
  1355. for ax in axs[0, :]:
  1356. ax.tick_params(
  1357. axis="x",
  1358. which="both",
  1359. labelbottom=False,
  1360. )
  1361. for ax in axs[:, 1]:
  1362. ax.tick_params(
  1363. axis="y",
  1364. which="both",
  1365. labelleft=False,
  1366. )
  1367. labels = {
  1368. "Wald": "CI: Wald 95%",
  1369. "MC": "CI: MCA propagation 95%",
  1370. "Nonparametric": "CI: nonparametric bootstrap 95%",
  1371. "Parametric": "CI: parametric bootstrap 95%",
  1372. }
  1373. handles = [
  1374. mlines.Line2D(
  1375. [],
  1376. [],
  1377. marker="o",
  1378. color=COL_NC,
  1379. linestyle="None",
  1380. markersize=8,
  1381. label="data: NC",
  1382. ),
  1383. mlines.Line2D(
  1384. [],
  1385. [],
  1386. marker="o",
  1387. color=COL_AE,
  1388. linestyle="None",
  1389. markersize=8,
  1390. label="data: AE",
  1391. ),
  1392. mlines.Line2D(
  1393. [],
  1394. [],
  1395. color=COL_FIT,
  1396. lw=2.5,
  1397. label="fit",
  1398. ),
  1399. ]
  1400. for method in FIGURE_METHODS:
  1401. color, linestyle, _ = styles[method]
  1402. handles.append(
  1403. mlines.Line2D(
  1404. [],
  1405. [],
  1406. color=color,
  1407. lw=2,
  1408. linestyle=linestyle,
  1409. label=labels[method],
  1410. )
  1411. )
  1412. leg = axs[1, 1].legend(
  1413. handles=handles,
  1414. loc="lower right",
  1415. bbox_to_anchor=(0.98, 0.04),
  1416. fontsize=8.5,
  1417. frameon=True,
  1418. )
  1419. leg.get_frame().set_facecolor("white")
  1420. leg.get_frame().set_edgecolor("#bdbdbd")
  1421. leg.get_frame().set_linewidth(0.8)
  1422. fig.subplots_adjust(
  1423. left=0.08,
  1424. right=0.98,
  1425. bottom=0.08,
  1426. top=0.98,
  1427. wspace=0.06,
  1428. hspace=0.06,
  1429. )
  1430. return fig, axs
  1431. # ============================================================
  1432. # ELASTICITY ANALYSIS (x50 and s50)
  1433. # ============================================================
  1434. import numpy as np
  1435. import matplotlib.pyplot as plt
  1436. # ------------------------------------------------------------
  1437. # Core elasticity computation
  1438. # ------------------------------------------------------------
  1439. def elasticity_x50_s50(theta, mode="raw"):
  1440. """
  1441. Elasticity for x50 and s50 with respect to theta0 and theta1.
  1442. mode
  1443. ----
  1444. 'raw' : eta = theta0 + theta1*x
  1445. 'log' : eta = theta0 + theta1*log(x)
  1446. Returns
  1447. -------
  1448. dict with:
  1449. theta0, theta1,
  1450. x50, s50,
  1451. E_x50_theta0, E_x50_theta1,
  1452. E_s50_theta0, E_s50_theta1
  1453. """
  1454. theta0, theta1 = map(float, np.asarray(theta).reshape(2))
  1455. if np.abs(theta1) < 1e-12:
  1456. return dict(
  1457. theta0=theta0,
  1458. theta1=theta1,
  1459. x50=np.nan,
  1460. s50=np.nan,
  1461. E_x50_theta0=np.nan,
  1462. E_x50_theta1=np.nan,
  1463. E_s50_theta0=np.nan,
  1464. E_s50_theta1=np.nan,
  1465. )
  1466. # =========================
  1467. # RAW MODEL
  1468. # =========================
  1469. if mode == "raw":
  1470. x50 = -theta0 / theta1
  1471. s50 = theta1 / 4.0
  1472. E_x50_theta0 = 1.0
  1473. E_x50_theta1 = -1.0
  1474. E_s50_theta0 = 0.0
  1475. E_s50_theta1 = 1.0
  1476. # =========================
  1477. # LOG MODEL
  1478. # =========================
  1479. elif mode == "log":
  1480. x50 = float(np.exp(-theta0 / theta1))
  1481. s50 = theta1 / (4.0 * x50)
  1482. E_x50_theta0 = -theta0 / theta1
  1483. E_x50_theta1 = theta0 / theta1
  1484. E_s50_theta0 = -E_x50_theta0
  1485. E_s50_theta1 = 1.0 - E_x50_theta1
  1486. else:
  1487. raise ValueError("mode must be 'raw' or 'log'")
  1488. return dict(
  1489. theta0=theta0,
  1490. theta1=theta1,
  1491. x50=x50,
  1492. s50=s50,
  1493. E_x50_theta0=E_x50_theta0,
  1494. E_x50_theta1=E_x50_theta1,
  1495. E_s50_theta0=E_s50_theta0,
  1496. E_s50_theta1=E_s50_theta1,
  1497. )
  1498. # ------------------------------------------------------------
  1499. # Table for 4 panels
  1500. # ------------------------------------------------------------
  1501. def elasticity_table_4panels(P, keys=None, make_plots=True):
  1502. import pandas as pd
  1503. import numpy as np
  1504. if keys is None:
  1505. keys = list(P.keys())
  1506. rows = []
  1507. for key in keys:
  1508. pk = P[key]
  1509. if "b" not in pk:
  1510. print(f"[skip] {key}: no fitted parameter key 'b'")
  1511. continue
  1512. theta = np.asarray(pk["b"], float).reshape(2)
  1513. transform = pk.get("transform", "raw")
  1514. res = elasticity_x50_s50(theta, mode=transform)
  1515. rows.append({
  1516. "Panel": key,
  1517. "transform": transform,
  1518. **res
  1519. })
  1520. df = pd.DataFrame(rows)
  1521. if make_plots and len(df) > 0:
  1522. plot_x50_values(df)
  1523. plot_s50_values(df)
  1524. plot_x50_theta1_elasticity(df)
  1525. plot_s50_theta1_elasticity(df)
  1526. return df
  1527. # ------------------------------------------------------------
  1528. # Plots
  1529. # ------------------------------------------------------------
  1530. def plot_x50_values(df):
  1531. fig, ax = plt.subplots(figsize=(7, 4))
  1532. ax.bar(df["Panel"], df["x50"])
  1533. ax.set_ylabel("x50")
  1534. ax.set_title("x50 across panels")
  1535. plt.xticks(rotation=30)
  1536. plt.tight_layout()
  1537. plt.show()
  1538. def plot_s50_values(df):
  1539. fig, ax = plt.subplots(figsize=(7, 4))
  1540. ax.bar(df["Panel"], df["s50"])
  1541. ax.set_ylabel("s50")
  1542. ax.set_title("s50 across panels")
  1543. plt.xticks(rotation=30)
  1544. plt.tight_layout()
  1545. plt.show()
  1546. def plot_x50_theta1_elasticity(df):
  1547. fig, ax = plt.subplots(figsize=(7, 4))
  1548. ax.bar(df["Panel"], df["E_x50_theta1"])
  1549. ax.set_ylabel("Elasticity")
  1550. ax.set_title("Elasticity of x50 w.r.t. theta1")
  1551. plt.xticks(rotation=30)
  1552. plt.tight_layout()
  1553. plt.show()
  1554. def plot_s50_theta1_elasticity(df):
  1555. fig, ax = plt.subplots(figsize=(7, 4))
  1556. ax.bar(df["Panel"], df["E_s50_theta1"])
  1557. ax.set_ylabel("Elasticity")
  1558. ax.set_title("Elasticity of s50 w.r.t. theta1")
  1559. plt.xticks(rotation=30)
  1560. plt.tight_layout()
  1561. plt.show()
  1562. # ------------------------------------------------------------
  1563. # logistic helpers for noise analysis
  1564. # ------------------------------------------------------------
  1565. # ============================================================
  1566. # NOISE ANALYSIS FOR LOGISTIC MODEL
  1567. # Correct x50 for RAW and LOG models
  1568. # ============================================================
  1569. import os
  1570. import numpy as np
  1571. import matplotlib.pyplot as plt
  1572. import matplotlib.lines as mlines
  1573. # ------------------------------------------------------------
  1574. # logistic fit / prediction / x50
  1575. # CONSISTENT WITH MAIN LOGISTIC ANALYSIS
  1576. # ------------------------------------------------------------
  1577. def fit_logistic_x(x_raw, y, transform="raw", l2=1e-8):
  1578. x_raw = np.clip(np.asarray(x_raw, float).ravel(), 1e-12, None)
  1579. y = np.asarray(y, int).ravel()
  1580. if transform == "raw":
  1581. x_model = x_raw
  1582. elif transform == "log":
  1583. x_model = np.log(x_raw)
  1584. else:
  1585. raise ValueError("transform must be 'raw' or 'log'")
  1586. return fit_newton(x_model, y, l2=l2)
  1587. def predict_curve_x(b, x_grid_raw, transform="raw"):
  1588. x_grid_raw = np.clip(np.asarray(x_grid_raw, float), 1e-12, None)
  1589. if transform == "raw":
  1590. x_model = x_grid_raw
  1591. elif transform == "log":
  1592. x_model = np.log(x_grid_raw)
  1593. else:
  1594. raise ValueError("transform must be 'raw' or 'log'")
  1595. return model_p(x_model, b)
  1596. def x50_from_b(b, transform="raw"):
  1597. b = np.asarray(b, float).reshape(2)
  1598. x50_model = x50(b)
  1599. if not np.isfinite(x50_model):
  1600. return np.nan
  1601. if transform == "raw":
  1602. return float(x50_model)
  1603. elif transform == "log":
  1604. return float(np.exp(x50_model))
  1605. else:
  1606. raise ValueError("transform must be 'raw' or 'log'")
  1607. def check_noise_x50(pack):
  1608. b = np.asarray(pack["b_clean"], float).reshape(2)
  1609. transform = pack["transform"]
  1610. x50_raw = pack["x50"]
  1611. x50_model = x50_raw if transform == "raw" else np.log(x50_raw)
  1612. p50 = model_p(np.array([x50_model]), b)[0]
  1613. print(
  1614. "transform =", transform,
  1615. "| x50_raw =", x50_raw,
  1616. "| P(x50) =", p50
  1617. )
  1618. # ------------------------------------------------------------
  1619. # noise helpers
  1620. # ------------------------------------------------------------
  1621. def add_noise_mult(x, sigma, rng):
  1622. x = np.asarray(x, float)
  1623. return np.clip(x * np.exp(rng.normal(0, sigma, size=x.shape)), 1e-12, None)
  1624. def add_noise_add(x, sigma, rng):
  1625. x = np.asarray(x, float)
  1626. return np.clip(x + rng.normal(0, sigma, size=x.shape), 1e-12, None)
  1627. def band_quantiles(curves):
  1628. C = np.vstack(curves)
  1629. return np.quantile(C, [0.025, 0.5, 0.975], axis=0)
  1630. # ------------------------------------------------------------
  1631. # build noise bands
  1632. # ------------------------------------------------------------
  1633. def noise_logistic_bands(
  1634. x_raw,
  1635. y,
  1636. transform="raw",
  1637. sigma_mult=0.129,
  1638. sigma_add=0.144,
  1639. x_max=5,
  1640. grid_n=1000,
  1641. n_refit=200,
  1642. n_tta=3000,
  1643. seed=1234,
  1644. l2=1e-8,
  1645. ):
  1646. rng = np.random.default_rng(seed)
  1647. x_raw = np.clip(np.asarray(x_raw, float).ravel(), 1e-12, None)
  1648. y = np.asarray(y).astype(int).ravel()
  1649. xc = np.linspace(1e-12, x_max, grid_n)
  1650. b_clean = fit_logistic_x(x_raw, y, transform=transform, l2=l2)
  1651. clean = predict_curve_x(b_clean, xc, transform=transform)
  1652. x50_val = x50_from_b(b_clean, transform=transform)
  1653. curves = []
  1654. for _ in range(n_refit):
  1655. xn = add_noise_mult(x_raw, sigma_mult, rng)
  1656. bn = fit_logistic_x(xn, y, transform=transform, l2=l2)
  1657. curves.append(predict_curve_x(bn, xc, transform=transform))
  1658. mult_refit = band_quantiles(curves)
  1659. curves = []
  1660. for _ in range(n_tta):
  1661. xn = add_noise_mult(xc, sigma_mult, rng)
  1662. curves.append(predict_curve_x(b_clean, xn, transform=transform))
  1663. mult_tta = band_quantiles(curves)
  1664. curves = []
  1665. for _ in range(n_refit):
  1666. xn = add_noise_add(x_raw, sigma_add, rng)
  1667. bn = fit_logistic_x(xn, y, transform=transform, l2=l2)
  1668. curves.append(predict_curve_x(bn, xc, transform=transform))
  1669. add_refit = band_quantiles(curves)
  1670. curves = []
  1671. for _ in range(n_tta):
  1672. xn = add_noise_add(xc, sigma_add, rng)
  1673. curves.append(predict_curve_x(b_clean, xn, transform=transform))
  1674. add_tta = band_quantiles(curves)
  1675. return {
  1676. "xc": xc,
  1677. "clean": clean,
  1678. "x50": x50_val,
  1679. "b_clean": b_clean,
  1680. "transform": transform,
  1681. "l2": float(l2),
  1682. "mult_refit": mult_refit,
  1683. "mult_tta": mult_tta,
  1684. "add_refit": add_refit,
  1685. "add_tta": add_tta,
  1686. }
  1687. from scipy.ndimage import gaussian_filter1d
  1688. lo = np.quantile(curves, 0.025, axis=0)
  1689. md = np.quantile(curves, 0.500, axis=0)
  1690. hi = np.quantile(curves, 0.975, axis=0)
  1691. # smooth boundaries
  1692. lo = gaussian_filter1d(lo, sigma=8)
  1693. md = gaussian_filter1d(md, sigma=8)
  1694. hi = gaussian_filter1d(hi, sigma=8)
  1695. return lo, md, hi
  1696. # ------------------------------------------------------------
  1697. # legend
  1698. # ------------------------------------------------------------
  1699. def noise_legend_handles():
  1700. return [
  1701. mlines.Line2D([], [], marker="o", color="#2b8cbe",
  1702. linestyle="None", markersize=7, label="NC data"),
  1703. mlines.Line2D([], [], marker="o", color="#d7301f",
  1704. linestyle="None", markersize=7, label="AE data"),
  1705. mlines.Line2D([], [], color="black", lw=2.2, label="Initial fit"),
  1706. mlines.Line2D([], [], color="#1f78b4", lw=6, alpha=0.24,
  1707. label="refit band, multiplicative noise"),
  1708. mlines.Line2D([], [], color="#1f78b4", lw=6, alpha=0.10,
  1709. label="fixed-model band, multiplicative noise"),
  1710. mlines.Line2D([], [], color="#e66101", lw=6, alpha=0.24,
  1711. label="refit band, additive noise"),
  1712. mlines.Line2D([], [], color="#e66101", lw=6, alpha=0.10,
  1713. label="fixed-model band, additive noise"),
  1714. mlines.Line2D([], [], color="#666666", ls="--", lw=1.2,
  1715. label=r"$x_{50}$"),
  1716. ]
  1717. # ------------------------------------------------------------
  1718. # plot one panel
  1719. # ------------------------------------------------------------
  1720. def plot_noise_panel(ax, pack, kind="mult", label="A", X=None, y=None):
  1721. COL_MULT = "#1f78b4"
  1722. COL_ADD = "#e66101"
  1723. COL_NC = "#2b8cbe"
  1724. COL_AE = "#d7301f"
  1725. xc = pack["xc"]
  1726. clean = pack["clean"]
  1727. x50_val = pack["x50"]
  1728. if kind == "mult":
  1729. refit = pack["mult_refit"]
  1730. tta = pack["mult_tta"]
  1731. color = COL_MULT
  1732. elif kind == "add":
  1733. refit = pack["add_refit"]
  1734. tta = pack["add_tta"]
  1735. color = COL_ADD
  1736. else:
  1737. raise ValueError("kind must be 'mult' or 'add'")
  1738. lo_r, _, hi_r = refit
  1739. lo_t, _, hi_t = tta
  1740. ax.fill_between(xc, lo_t, hi_t, color=color, alpha=0.10, zorder=1)
  1741. ax.fill_between(xc, lo_r, hi_r, color=color, alpha=0.24, zorder=2)
  1742. ax.plot(xc, lo_r, color=color, lw=1.0, alpha=0.65, zorder=3)
  1743. ax.plot(xc, hi_r, color=color, lw=1.0, alpha=0.65, zorder=3)
  1744. ax.plot(xc, clean, color="black", lw=2.2, zorder=5)
  1745. ax.axvline(x50_val, color="#666666", ls="--", lw=1.2, alpha=0.9, zorder=4)
  1746. lo_r_x = np.interp(x50_val, xc, lo_r)
  1747. hi_r_x = np.interp(x50_val, xc, hi_r)
  1748. lo_t_x = np.interp(x50_val, xc, lo_t)
  1749. hi_t_x = np.interp(x50_val, xc, hi_t)
  1750. if X is not None and y is not None:
  1751. X = np.asarray(X).ravel()
  1752. y = np.asarray(y).astype(int).ravel()
  1753. ax.scatter(
  1754. X[y == 0], np.zeros(np.sum(y == 0)),
  1755. color=COL_NC, s=24, alpha=0.75,
  1756. edgecolors="none", zorder=7
  1757. )
  1758. ax.scatter(
  1759. X[y == 1], np.ones(np.sum(y == 1)),
  1760. color=COL_AE, s=24, alpha=0.75,
  1761. edgecolors="none", zorder=7
  1762. )
  1763. ax.text(0.03, 0.97, label, transform=ax.transAxes,
  1764. ha="left", va="top", fontsize=15)
  1765. variant_txt = "FULL" if label in ["A", "B"] else "TRIM"
  1766. d_ref = hi_r_x - lo_r_x
  1767. d_tta = hi_t_x - lo_t_x
  1768. info_txt = (
  1769. f"{variant_txt}\n"
  1770. f"$x_{{50}}$={x50_val:.2f}\n"
  1771. f"$\\Delta r$={d_ref:.2f} $\\Delta t$={d_tta:.2f}"
  1772. )
  1773. ax.text(
  1774. 0.02, 0.14,
  1775. info_txt,
  1776. transform=ax.transAxes,
  1777. fontsize=10,
  1778. color="#222",
  1779. ha="left", va="bottom",
  1780. bbox=dict(facecolor="white", edgecolor=color,
  1781. boxstyle="square,pad=0.25", alpha=0.9)
  1782. )
  1783. ax.set_xlim(0, xc.max())
  1784. ax.set_ylim(-0.05, 1.05)
  1785. ax.grid(alpha=0.25)
  1786. ax.tick_params(axis="both", labelsize=10)
  1787. # ------------------------------------------------------------
  1788. # full noise figure
  1789. # ------------------------------------------------------------
  1790. def plot_noise_figure(
  1791. pack_full,
  1792. pack_trim,
  1793. X_full,
  1794. y_full,
  1795. X_trim,
  1796. y_trim,
  1797. figsize=(12, 9),
  1798. dpi=300,
  1799. ):
  1800. fig, axes = plt.subplots(
  1801. 2, 2,
  1802. figsize=figsize,
  1803. dpi=dpi,
  1804. sharex=True,
  1805. sharey=True
  1806. )
  1807. axes = axes.ravel()
  1808. plot_noise_panel(axes[0], pack_full, kind="mult", label="A", X=X_full, y=y_full)
  1809. plot_noise_panel(axes[1], pack_full, kind="add", label="B", X=X_full, y=y_full)
  1810. plot_noise_panel(axes[2], pack_trim, kind="mult", label="C", X=X_trim, y=y_trim)
  1811. plot_noise_panel(axes[3], pack_trim, kind="add", label="D", X=X_trim, y=y_trim)
  1812. axes[0].set_ylabel(r"$P(\mathrm{AE}\mid X=x)$", fontsize=12)
  1813. axes[2].set_ylabel(r"$P(\mathrm{AE}\mid X=x)$", fontsize=12)
  1814. axes[2].set_xlabel(r"$X$", fontsize=12)
  1815. axes[3].set_xlabel(r"$X$", fontsize=12)
  1816. handles = noise_legend_handles()
  1817. leg = axes[3].legend(
  1818. handles=handles,
  1819. loc="lower right",
  1820. bbox_to_anchor=(0.97, 0.05),
  1821. fontsize=9,
  1822. frameon=True
  1823. )
  1824. frame = leg.get_frame()
  1825. frame.set_facecolor("white")
  1826. frame.set_edgecolor("#bdbdbd")
  1827. frame.set_linewidth(0.8)
  1828. fig.tight_layout()
  1829. return fig, axes
  1830. # ------------------------------------------------------------
  1831. # wrapper
  1832. # ------------------------------------------------------------
  1833. def make_noise_figure(
  1834. X_full,
  1835. y_full,
  1836. X_trim,
  1837. y_trim,
  1838. transform="raw",
  1839. sigma_mult=0.129,
  1840. sigma_add=0.144,
  1841. x_max=5,
  1842. grid_n=1000,
  1843. n_refit=1100,
  1844. n_tta=10000,
  1845. seed=1234,
  1846. l2=1e-8,
  1847. save_path=None,
  1848. ):
  1849. pack_full = noise_logistic_bands(
  1850. X_full, y_full,
  1851. transform=transform,
  1852. sigma_mult=sigma_mult,
  1853. sigma_add=sigma_add,
  1854. x_max=x_max,
  1855. grid_n=grid_n,
  1856. n_refit=n_refit,
  1857. n_tta=n_tta,
  1858. seed=seed,
  1859. l2=l2,
  1860. )
  1861. pack_trim = noise_logistic_bands(
  1862. X_trim, y_trim,
  1863. transform=transform,
  1864. sigma_mult=sigma_mult,
  1865. sigma_add=sigma_add,
  1866. x_max=x_max,
  1867. grid_n=grid_n,
  1868. n_refit=n_refit,
  1869. n_tta=n_tta,
  1870. seed=seed + 100,
  1871. l2=l2,
  1872. )
  1873. print("FULL n:", len(X_full), "x50:", pack_full["x50"])
  1874. print("TRIM n:", len(X_trim), "x50:", pack_trim["x50"])
  1875. check_noise_x50(pack_full)
  1876. check_noise_x50(pack_trim)
  1877. fig, axes = plot_noise_figure(
  1878. pack_full, pack_trim,
  1879. X_full, y_full,
  1880. X_trim, y_trim,
  1881. figsize=(12, 9),
  1882. dpi=300,
  1883. )
  1884. if save_path is not None:
  1885. folder = os.path.dirname(save_path)
  1886. if folder:
  1887. os.makedirs(folder, exist_ok=True)
  1888. fig.savefig(f"{save_path}.png", dpi=300, bbox_inches="tight")
  1889. fig.savefig(f"{save_path}.pdf", bbox_inches="tight")
  1890. return fig, axes, pack_full, pack_trim