import json, math, time import numpy as np from sklearn.neural_network import MLPRegressor # Adversarially calibrated residualization MVP. # The implementation follows the displayed constraint literally. def edit_weights(t, F, w0=None, tau=0.0, kappa=0.0): """Solve the stated local edit problem for the important w0=1 case. Since w=1 has objective zero and exactly zero (w-1) moments, it is the unique minimizer whenever w0=1 (the objective is strictly convex). """ n = len(t) if w0 is None: w0 = np.ones(n) w0 = np.asarray(w0) if np.allclose(w0, 1.0): w = np.ones(n) else: # Small projected-gradient fallback for non-unit initial weights. w = np.maximum(w0.copy(), 0.0) A = t[:, None] * F lr = 0.2 / (np.linalg.norm(A, 2)**2 / n + 1e-8) for _ in range(2000): mom = ((w-1)[:, None] * A).mean(0) viol = np.maximum(np.abs(mom)-tau, 0) grad = (w-w0)/n if np.any(viol): grad += ((A @ (np.sign(mom)*viol)) / n) w = np.maximum(w-lr*grad, 0) w *= len(w)/w.sum() energy = np.mean(w*t*t) # The stated rejection rule: retain a feasible-energy solution if possible. if energy < kappa: return np.ones(n) if np.mean(t*t) >= kappa else w return w def moments(w,t,F): return np.mean((w-1)[:,None]*t[:,None]*F, axis=0) def toy_verification(seed=123): rng=np.random.default_rng(seed); n=600 x=rng.normal(size=(n,8)); t=0.5*x[:,0]+rng.normal(size=n) F=np.tanh(x[:,:4]) out={"tau_sweep":[], "critic_scale_sweep":[], "dimension_sweep":[]} for tau in [0,1e-4,1e-2,0.1,1.0]: w=edit_weights(t,F,tau=tau,kappa=.01) out["tau_sweep"].append({"tau":tau,"edit_l2":float(np.linalg.norm(w-1)),"max_moment":float(np.max(np.abs(moments(w,t,F))))}) for scale in [0.1,1,10]: Fs=scale*F w=edit_weights(t,Fs,tau=.01,kappa=.01) out["critic_scale_sweep"].append({"scale":scale,"edit_l2":float(np.linalg.norm(w-1)),"max_moment":float(np.max(np.abs(moments(w,t,Fs))))}) for d in [1,2,4,8]: Fs=np.tanh(x[:,:d]) w=edit_weights(t,Fs,tau=.01,kappa=.01) out["dimension_sweep"].append({"dimension":d,"edit_l2":float(np.linalg.norm(w-1)),"max_moment":float(np.max(np.abs(moments(w,t,Fs))))}) return out def dml_once(seed, n=500, reverse=False): rng=np.random.default_rng(seed); p=20; beta=1.0 X=rng.normal(size=(n,p)); # Imbalanced nuisance difficulty: high-frequency outcome versus smooth treatment, # then the reverse. A small MLP intentionally underfits the high-frequency part. smooth_pi=0.8*np.sin(X[:,0])+0.3*X[:,1] hard_mu=1.5*np.sin(5*X[:,0])+0.7*np.cos(4*X[:,1])+0.3*X[:,2] if not reverse: mu, pi=hard_mu, smooth_pi else: mu, pi=smooth_pi, hard_mu/1.5 T=pi+rng.normal(size=n) Y=mu+beta*T+0.8*rng.normal(size=n) idx=np.arange(n); rng.shuffle(idx); folds=np.array_split(idx,2) ry=np.zeros(n); rt=np.zeros(n) for te in folds: tr=np.setdiff1d(idx,te,assume_unique=False) my=MLPRegressor(hidden_layer_sizes=(32,16),early_stopping=False,max_iter=100, random_state=seed, solver='adam', learning_rate_init=.003) mt=MLPRegressor(hidden_layer_sizes=(32,16),early_stopping=False,max_iter=100, random_state=seed+17, solver='adam', learning_rate_init=.003) my.fit(X[tr],Y[tr]); mt.fit(X[tr],T[tr]) ry[te]=Y[te]-my.predict(X[te]); rt[te]=T[te]-mt.predict(X[te]) beta_hat=float(np.sum(rt*ry)/np.sum(rt*rt)) F=np.tanh(X[:,:4]); w=edit_weights(rt,F,tau=.02,kappa=.1*np.mean(rt*rt)) cal=float(np.sum(w*rt*ry)/np.sum(w*rt*rt)) ess=float(w.sum()**2/np.sum(w*w)); viol=float(np.max(np.abs(moments(w,rt,F)))) return beta_hat,cal,ess,viol def mini_experiment(): rows=[] for reverse in [False,True]: vals=[dml_once(100+i,reverse=reverse) for i in range(8)] a=np.array(vals) rows.append({"regime":"hard_mu" if not reverse else "hard_pi", "baseline_abs_bias":float(np.mean(np.abs(a[:,0]-1))), "idea_abs_bias":float(np.mean(np.abs(a[:,1]-1))), "baseline_rmse":float(np.sqrt(np.mean((a[:,0]-1)**2))), "idea_rmse":float(np.sqrt(np.mean((a[:,1]-1)**2))), "mean_ess":float(np.mean(a[:,2])),"max_violation":float(np.max(a[:,3])), "paired_max_difference":float(np.max(np.abs(a[:,0]-a[:,1])))}) return rows if __name__=='__main__': t=time.time(); result={"toy":toy_verification(),"mini":mini_experiment(),"runtime_sec":time.time()-t} print(json.dumps(result,indent=2))