Files
calibration/tools/retry_rtk_imu_heldout_nonconverged.py
T

105 lines
5.7 KiB
Python

#!/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())