105 lines
5.7 KiB
Python
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())
|