#!/usr/bin/env python3 '''Retry only held-out fixed-lever windows that ended at max_nfev.''' from __future__ import annotations import argparse,json,sys from concurrent.futures import ProcessPoolExecutor from pathlib import Path import numpy as np from scipy.spatial.transform import Rotation ROOT=Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0,str(ROOT)) from imu_lidar.rtk_imu_engineering import _height_reference from imu_lidar.rtk_imu_multisource import load_unified_sessions from imu_lidar.rtk_imu_node_graph import ( build_problem,fit_states_at_fixed_lever,summarize_fixed_state_values) from tools.audit_rtk_imu_factor_consistency import _jsonable from tools.run_rtk_imu_node_graph_free_selected import _restore_segments def _retry(task): index,problem,lever,state,max_nfev=task value,meta=fit_states_at_fixed_lever( problem,lever,max_nfev,initial_state_values=state) return index,value,meta def main(): p=argparse.ArgumentParser(description=__doc__) p.add_argument('--manifest',type=Path,required=True) p.add_argument('--calibration-selection',type=Path,required=True) p.add_argument('--all-selection',type=Path,required=True) p.add_argument('--engineering-result',type=Path,required=True) p.add_argument('--checkpoint-dir',type=Path,required=True) p.add_argument('--output',type=Path,required=True) p.add_argument('--retry-checkpoint-dir',type=Path,required=True) p.add_argument('--max-nfev',type=int,default=120) p.add_argument('--workers',type=int,default=4) p.add_argument('--rotation-rpy-deg',nargs=3,type=float, default=[.4543066225,-.0026392019,.0122384129]) args=p.parse_args() engineering=json.loads(args.engineering_result.read_text(encoding='utf-8')) calibration=json.loads(args.calibration_selection.read_text(encoding='utf-8')) selected=json.loads(args.all_selection.read_text(encoding='utf-8')) ids={x['candidate_id'] for x in calibration['selected_windows']} heldout=[x for x in selected['selected_windows'] if x['candidate_id'] not in ids] lever=np.asarray(engineering['engineering_l_I_m'],dtype=float) bad=[] for index,item in enumerate(heldout): data=np.load(args.checkpoint_dir/f'{index:04d}.npz',allow_pickle=False) meta=json.loads(str(data['optimizer'])) if not meta['success']: bad.append((index,item,data['state'].copy(),meta)) if len(bad)!=10: raise RuntimeError(f'expected 10 non-converged windows, got {len(bad)}') sessions=load_unified_sessions( args.manifest,selected_session_ids={item['session_id'] for _,item,_,_ in bad}) reference=_height_reference(sessions) selections=[item for _,item,_,_ in bad] segments=_restore_segments(sessions,reference,selections,1.) R=Rotation.from_euler('xyz',args.rotation_rpy_deg,degrees=True).as_matrix() problems=[build_problem(segment,R,lever,.006) for segment in segments] tasks=[(local,problem,lever,bad[local][2],args.max_nfev) for local,problem in enumerate(problems)] with ProcessPoolExecutor(max_workers=args.workers) as executor: retried=list(executor.map(_retry,tasks)) retried.sort(key=lambda x:x[0]) args.retry_checkpoint_dir.mkdir(parents=True,exist_ok=True) mapping={'0808_20260808_092827':'circle', '0808_20260808_082148':'left_right', '0815_20260812_123424':'slope'} results=[] for local,state,meta in retried: original_index,item,_,previous=bad[local] summary=summarize_fixed_state_values([problems[local]],[state],lever) np.savez_compressed(args.retry_checkpoint_dir/f'{original_index:04d}.npz', state=state,optimizer=json.dumps(meta)) results.append({'heldout_index':original_index, 'candidate_id':item['candidate_id'],'session_id':item['session_id'], 'time_range_s':[item['start_s'],item['end_s']], 'motion_class':mapping.get(item['session_id'],'other_recovered_dynamic'), 'original_termination_reason':previous['message'], 'original_nfev':previous['nfev'],'original_final_cost':previous['cost'], 'retry_started_from_original_final_state':True, 'retry_termination_reason':meta['message'],'retry_nfev':meta['nfev'], 'retry_initial_cost':meta['initial_cost'],'retry_final_cost':meta['cost'], 'retry_hit_max_nfev':bool(not meta['success'] and meta['nfev']>=args.max_nfev), 'retry_success':meta['success'],'residual':summary}) converged=sum(x['retry_success'] for x in results) payload={'scope':'10 held-out max_nfev windows; exact-model continuation retry', 'lever_reoptimized':False,'covariance_parameters_modified':False, 'parser_R0_modified':False,'windows_removed':False, 'fixed_l_I_m':lever,'retry_max_nfev':args.max_nfev, 'initial_nonconverged_count':len(results), 'converged_after_retry_count':converged, 'heldout_convergence_after_retry':(257+converged)/267., 'remaining_nonconverged_count':len(results)-converged, 'windows':results} args.output.write_text(json.dumps(_jsonable(payload),ensure_ascii=False,indent=2, allow_nan=False)+'\n',encoding='utf-8') print(json.dumps(_jsonable({'converged_after_retry':converged, 'remaining':len(results)-converged, 'heldout_convergence_after_retry':payload['heldout_convergence_after_retry'], 'windows':[{'index':x['heldout_index'],'session':x['session_id'], 'motion':x['motion_class'],'success':x['retry_success'], 'nfev':x['retry_nfev'],'cost':x['retry_final_cost']} for x in results]}),ensure_ascii=False,indent=2)) return 0 if __name__=='__main__': raise SystemExit(main())