"""Build isolated context-fallback diagnostic workers; no implementation edits.""" from pathlib import Path from concurrent.futures import ThreadPoolExecutor import argparse,hashlib,json,os,re,shutil,subprocess,time import local_probe_experiment as ex ROOT=ex.ROOT;SOURCE=ROOT/'test/local-probe-20260917/worker';OUT=ROOT/'test/context-fallback-20260917';HERE=Path(__file__).parent KERNELS=[('properties','property_pt'),('properties','property_density'),('properties','property_viscosity'), ('properties','native_temperature_ph_context'),('properties','local_isentropic'),('properties','state_valve'), ('properties','native_jacobian_scalar_get'),('properties','native_temperature_ph'),('properties','native_density'),('properties','native_viscosity'), ('pipe','native_pipe_flow_cached_context'),('pipe','native_pipe_flow_context'),('pipe','native_pipe_resistance'), ('orifice','native_medium_orifice_context')] replace=ex.replace def function(s,name):a,b,e=ex.function_span(s,name);return s[a:e] def vec(refs):return '(double[]){'+(','.join(refs) or '0')+'}' def generate_model(source,ops,meta): original=function(source,'model_eval_local_internal');start=original.index('if(lp_capture){');end=original.index('double node_energy[') checkpoints=meta['contextCheckpointSlots'];capture=[] for pos,op in enumerate(ops): if str(pos) in checkpoints:capture.append(f'lp_snapshot({checkpoints[str(pos)]},properties,pipe_cache);') capture += [f'dx_op_begin({pos},{vec(sorted(op.inputs))});',*op.code,f'dx_op_end({pos},{vec(op.outputs)},properties,pipe_cache);'] capture += [f'lp_snapshot({checkpoints[str(len(ops))]},properties,pipe_cache);','lp_save(p,h,q,w,fb);'] # Use the original capture branch in time-only runs; it is not profiled. orig_capture=original[start+len('if(lp_capture){'):original.index('}else{',start)] schedule=['dx_schedule(properties,pipe_cache);','if(lp_capture){','if(dx_mode==1 || dx_mode==5){',*capture,'}else{',orig_capture,'}','}else{','for(int pos=0;pos=0 && dx_reuse(region,properties,pipe_cache,p,h,q,w,fb)){pos=lp_end[region];continue;}', 'switch(pos){'] for pos,op in enumerate(ops): schedule += [f'case {pos}:{{',f'dx_op_begin({pos},(dx_mode==1 || dx_mode==5)?{vec(sorted(op.inputs))}:NULL);',*op.code, f'dx_op_end({pos},(dx_mode==1 || dx_mode==5)?{vec(op.outputs)}:NULL,properties,pipe_cache);','break;}'] schedule += ['default:return 0;}','pos++;','if(dx_region>=0 && pos==lp_end[dx_region])dx_region_end(p,h,q,w,fb,properties,pipe_cache);','}}'] clone=(original[:start]+'\n'.join(schedule)+'\ndx_position=-1;\n'+original[end:]).replace('model_eval_local_internal(', 'dx_model_eval_local_internal(',1) wrapper=function(source,'lp_eval').replace('int lp_eval(', 'int dx_eval(',1).replace('model_eval_local_internal(', 'dx_model_eval_local_internal(') a=wrapper.index('{')+1;wrapper=wrapper[:a]+'\nif(!dx_selected)return lp_eval(t,y,dy,w,workspace);\ndx_eval_begin(t,y);\n'+wrapper[a:] wrapper=replace(wrapper,'lp_active=-1;return result;','lp_active=-1;dx_eval_end();return result;') return source+'\n'+clone+'\n'+wrapper def kernel_wrapper(source,name,index): a,b,e=ex.function_span(source,name);sig=source[a:b].strip();body=source[a:e] args=sig[sig.index('(')+1:sig.rindex(')')] params=[re.search(r'([A-Za-z_]\w*)\s*(?:\[[^]]*\])?$',x.strip())[1] for x in args.split(',')] prefix=sig[:sig.index(name)].strip();typ=re.sub(r'^(?:static|NATIVE_COMPONENT_INTERNAL)\s+','',prefix).strip() impl=body.replace(name+'(', 'dx_impl_'+name+'(',1) call='dx_impl_'+name+'('+','.join(params)+')' action=(call+';dx_kernel_end('+str(index)+',ticket);') if typ=='void' else (typ+' result='+call+';dx_kernel_end('+str(index)+',ticket);return result;') wrapper=sig+'{uint64_t ticket=dx_kernel_begin('+str(index)+');'+action+'}' # Forward declaration preserves recursive calls and cross-calls. return source[:a]+sig+';\n'+impl+'\n'+wrapper+source[e:] def prepare(kernels=False,trace=False): OUT.mkdir(exist_ok=True);work=OUT/('trace-worker' if trace else 'kernels' if kernels else 'worker');work.mkdir(exist_ok=True) program,saved=ex.capture(ROOT/'tests/data/test-mql-8-corrected.json') assert program.source==(SOURCE.parent/'original-model.c').read_text(encoding='utf-8') schedule=saved['schedule'];ops=[schedule.computations[b.members[0]] for b in schedule.blocks] meta=json.loads((SOURCE.parent/'plan.json').read_text(encoding='utf-8')) versions=[0];pure={'if','for','sizeof','fmax','fmin','fabs','sqrt','copysign','pow'} for op in ops:versions.append(versions[-1]+int(bool(set(re.findall(r'\b([A-Za-z_]\w*)\s*\(', '\n'.join(op.code)))-pure))) assert all(versions[int(pos)]==slot for pos,slot in meta['contextCheckpointSlots'].items()) ins=[0];outs=[0] for op in ops:ins.append(ins[-1]+len(op.inputs));outs.append(outs[-1]+len(op.outputs)) arrays={'dx_start_pos':[a for a,b in meta['regions']],'dx_contextual':[int(versions[a]!=versions[b]) for a,b in meta['regions']], 'dx_version':versions,'dx_mutates':[versions[i+1]!=versions[i] for i in range(len(ops))],'dx_in_offset':ins,'dx_out_offset':outs} tables='\n'.join('const int '+name+'[]={'+','.join(str(int(v)) for v in values)+'};' for name,values in arrays.items()) tables+=f'\n#define DX_NIN {ins[-1]}\n#define DX_NOUT {outs[-1]}\n' sources={p.name:p.read_text(encoding='utf-8') for p in SOURCE.glob('*.c')};before={n:hashlib.sha256(s.encode()).hexdigest() for n,s in sources.items()} sources['model.c']=generate_model(sources['model.c'],ops,meta) sources['local_probe_support.c']+='\n'+(HERE/'context_fallback_diag.c').read_text(encoding='utf-8').replace('/* DIAG_TABLES */',tables) sources['common.c']=replace(sources['common.c'],'lp_start();','lp_start();dx_initialize();') sources['common.c']=replace(sources['common.c'],'lp_finish();','lp_finish();dx_finish();') sources['cvode_solver.c']=replace(sources['cvode_solver.c'],'lp_eval(t,N_VGetArrayPointer(y),N_VGetArrayPointer(f),outputs,workspace)','dx_eval(t,N_VGetArrayPointer(y),N_VGetArrayPointer(f),outputs,workspace)') sources['cvode_solver.c']=replace(sources['cvode_solver.c'],'{lp_color=-1;uint64_t start=lp_tick();','{lp_color=-1;dx_jacobian();uint64_t start=lp_tick();') sources['cvode_solver.c']=replace(sources['cvode_solver.c'],'if(!result)lp_matrix(t,N_VGetArrayPointer(y),SUNDenseMatrix_Data(matrix));', 'if(!result){lp_matrix(t,N_VGetArrayPointer(y),SUNDenseMatrix_Data(matrix));dx_validate_matrix(t,N_VGetArrayPointer(y),SUNDenseMatrix_Data(matrix));}') if kernels: for i,(module,name) in enumerate(KERNELS):sources[module+'.c']=kernel_wrapper(sources[module+'.c'],name,i) if trace: a,b,e=ex.function_span(sources['properties.c'],'property_new') fn=sources['properties.c'][a:e];fn=replace(fn,'return s;','dx_property_created(cache,s);return s;') sources['properties.c']=sources['properties.c'][:a]+fn+sources['properties.c'][e:] for name in ('model.h','local_probe.h'):shutil.copyfile(SOURCE/name,work/name) (work/'context_fallback_diag.h').write_text((HERE/'context_fallback_diag.h').read_text(encoding='utf-8').replace('#include "local_probe.h"','#include "local_probe.h"\n#define DX_NK '+str(len(KERNELS))),encoding='utf-8') cc,sun,_=ex.builder.toolchain();flags,libs,dlls,exe=ex.builder.platform_build_inputs(sun);flags+=['-DLP_OBSERVE=0'];started=time.perf_counter() def compile_one(item): i,(name,s)=item;path=work/name;path.write_text('#include "context_fallback_diag.h"\n'+s,encoding='utf-8',newline='\n');obj=work/f'diag-{i}.o';log=[] ex.builder._command([cc,*flags,'-I',str(work),'-I',str(ex.builder.NATIVE/'include'),'-I',str(sun/'include'),'-c',str(path),'-o',str(obj)],log=log,timeout=240) return obj,log with ThreadPoolExecutor(max_workers=4) as pool:objs=list(pool.map(compile_one,enumerate(sources.items()))) log=[];ex.builder._command([cc,*flags,*[str(o) for o,_ in objs],*ex.builder.link_library_arguments(libs),'-lm','-o',str(work/exe)],log=log) for dll in dlls:shutil.copyfile(dll,work/dll.name) (work/'build.log').write_text('\n'.join(sum([v for _,v in objs],[])+log),encoding='utf-8') ex.write(work/'build.json',dict(sourceHashes=before,seconds=time.perf_counter()-started,kernels=kernels)) ex.write(OUT/'plan.json',dict(**meta,versions=versions,kernels=KERNELS,code=[list(o.code) for o in ops],inputs=[sorted(o.inputs) for o in ops])) print('BUILT',work.name,time.perf_counter()-started,flush=True) def run(label,mode=1,stride=16,seed=1,kernels=False,matrices=False,control=False,trace=False): work=OUT/label;work.mkdir(exist_ok=True);exe=(SOURCE if control else OUT/('trace-worker' if trace else 'kernels' if kernels else 'worker'))/'model.exe' env=os.environ.copy();env.update(LOCAL_PROBE_MASK='0x7ffffff',CONTEXT_DIAG_MODE=str(mode),CONTEXT_DIAG_STRIDE=str(stride),CONTEXT_DIAG_SEED=str(seed),CONTEXT_DIAG_MATRICES=str(int(matrices))) args=[str(exe),'--method','BDF','--start','0','--stop','10','--sample-step','.01','--max-step','1e30','--rtol','1e-8','--timeout','300', '--sample-file',str(work/'states.bin'),'--output-block-file',str(work/'outputs.bin'),'--output',str(work/'result.json')] start=time.perf_counter() with (work/'stderr.log').open('wb') as f:p=subprocess.run(args,cwd=work,env=env,stdout=subprocess.PIPE,stderr=f,timeout=330,creationflags=subprocess.CREATE_NO_WINDOW) elapsed=time.perf_counter()-start if p.returncode:raise RuntimeError((label,p.returncode,(work/'stderr.log').read_text()[-5000:])) result=json.loads((work/'result.json').read_text());diag=json.loads((work/'probe.json').read_text());ref=json.loads((SOURCE.parent/'all-run-0/measurement.json').read_text(encoding='utf-8')) keys=['success','finalState','final','propertyWarnings','acceptedSteps','rejectedSteps','stateTransitions','solverStarts','nfev','njev','nlu'] assert all(result[k]==ref[k] for k in keys),(label,'result differs') assert all(diag[k]==ref['diagnostic'][k] for k in ['newtonIterations','newtonConvergenceFailures','modelCalls','groups','contextCopiedBytes','contextComparedBytes']),(label,'counters differ') hashes={} for name in ('states','outputs','events','jacobians'): p=work/(name+'.bin') if p.exists(): with p.open('rb') as f:hashes[name]=hashlib.file_digest(f,'sha256').hexdigest() if name=='jacobians': with (SOURCE.parent/'all-audit/jacobians.bin').open('rb') as f:assert hashes[name]==hashlib.file_digest(f,'sha256').hexdigest() else:assert hashes[name]==ref[name+'Sha256'],(label,name) record=dict(label=label,mode=mode,stride=stride,seed=seed,kernels=kernels,control=control,processSeconds=elapsed, solveSeconds=result['solveSeconds'],solveCpuSeconds=result['solveCpuSeconds'],jacobianSeconds=diag['jacobianSeconds'],hashes=hashes) ex.write(work/'measurement.json',record);print('RUN',label,'Jac',diag['jacobianSeconds'],'solve',result['solveSeconds'],'exact OK',flush=True) return record if __name__=='__main__': p=argparse.ArgumentParser();p.add_argument('action',choices=['prepare','run']);p.add_argument('--kernels',action='store_true');p.add_argument('--matrices',action='store_true');p.add_argument('--control',action='store_true');p.add_argument('--trace',action='store_true');p.add_argument('--label',default='census');p.add_argument('--mode',type=int,default=1);p.add_argument('--stride',type=int,default=16);p.add_argument('--seed',type=int,default=1);a=p.parse_args() if a.action=='prepare':prepare(a.kernels,a.trace) else:run(a.label,a.mode,a.stride,a.seed,a.kernels,a.matrices,a.control,a.trace)