Ligant/castle-code / castle-batch-v1 /batch_pipeline.py
Ligant's picture
download
raw
10.9 kB
#!/usr/bin/env python3
"""Short-lived Vertex Batch workers: prepare/submit, or check/collect/advance once."""
import argparse
import hashlib
import itertools
import json
import os
from pathlib import Path
import sys
import tempfile
from PIL import Image
from castle_pipeline.batch_cloud import BatchCloud, parse_gs_uri
from castle_pipeline.batch_engine import BatchEngine
from castle_pipeline.batch_tasks import TaskPlanner
from castle_pipeline.events import EventLog
from castle_pipeline.inputs import REPO, REVISION, source_metadata, download_source, remove_downloaded_source
from castle_pipeline.media import MediaExtractor, probe
from castle_pipeline.runner import atomic_json, cleanup_media, clip_windows, fingerprint
def code_hash():
root=Path(__file__).parent
paths=[root/'batch_pipeline.py',*(root/'castle_pipeline').glob('*.py')]
return fingerprint({p.name:hashlib.sha256(p.read_bytes()).hexdigest() for p in paths})
def prepare_media(source, metadata, config, planner):
"""Persist both model-sized stamped frames and native samples one clip at a time."""
info = probe(source)
extractor = MediaExtractor(threads=1, decode_slots=1)
rows = []
windows = itertools.islice(clip_windows(info['duration'],config['clip_seconds']),config['start_clip'],
config['start_clip']+config['max_clips'])
for index,start,duration in windows:
cid=hashlib.sha256((metadata['source_id']+':'+str(index)).encode()).hexdigest()[:24]
prefix=config['gcs_prefix'].rstrip('/')+'/media/'+cid
work=Path(tempfile.mkdtemp(prefix='batch-prepare-',dir=planner.scratch));owned=[]
try:
media=extractor.prepare(source,start,duration,work,fps=config['fps'],max_dim=max(info['width'],info['height']),stamp=False)
owned.extend(media.frame_paths)
if media.audio_path:owned.append(media.audio_path)
native,frames=[],[]
for i,(path,t) in enumerate(zip(media.frame_paths,media.frame_times)):
if config.get('review',True):
uri=f'{prefix}/native-{i:03d}.jpg';planner.cloud.upload(path,uri);native.append(uri)
stamped=work/f'model-{i:03d}.jpg';owned.append(stamped)
with Image.open(path) as image:
image.thumbnail((config['max_dim'],config['max_dim']-24))
image.save(stamped,'JPEG',quality=95)
MediaExtractor._footer(stamped,t,24)
uri=f'{prefix}/frame-{i:03d}.jpg';planner.cloud.upload(stamped,uri);frames.append(uri)
audio_uri=None
if media.audio_path:
audio_uri=f'{prefix}/audio.wav';planner.cloud.upload(media.audio_path,audio_uri)
meta={'clip_id':cid,'source':metadata,'clip_index':index,'start_offset_sec':start,'duration_sec':duration,
'frame_times_sec':media.frame_times,'frames':frames,'native_frames':native,'audio_uri':audio_uri}
uri=planner.put_json(f'{prefix}/metadata.json',meta)
row={'clip_id':cid,'metadata_uri':uri}
if audio_uri is None:
empty={'data':{'summary':'No audio track supplied.','utterances':[],'sound_events':[],
'uncertainties':['No audio evidence available.']},'usage':{},'skipped':True}
row['audio_uri']=planner.put_json(f'{config["gcs_prefix"].rstrip("/")}/results/audio/{cid}.json',empty)
rows.append(row)
planner.emit('batch_clip_prepared',clip_id=cid,clip_index=index,frame_count=len(frames))
finally:
cleanup_media(work,owned)
return rows
def summary(state):
return {'status':state['status'],'stage':state['current_stage'],
'selected':len(state['config'].get('rows',[])),'completed':len(state['completed']),
'failed_records':len(state['failures']),'eligible_next':len(state['eligible_rows']),
'batch_jobs':{k:v.get('job',{}).get('name') for k,v in state['batches'].items()},
'batch_job_states':{k:v.get('job',{}).get('state') for k,v in state['batches'].items()},
'batch_job_errors':{k:v['job']['error'] for k,v in state['batches'].items() if v.get('job',{}).get('error')},
'last_error':state.get('last_error'),'last_error_code':state.get('last_error_code'),
'attention_reason':state.get('attention_reason'),
'stop_schedule':state['status'] in {'complete','complete_with_errors','needs_attention','paused'}}
def build_parser():
parser=argparse.ArgumentParser(description=__doc__)
sub=parser.add_subparsers(dest='command',required=True)
for name in ['prepare','start','tick','status','reconcile']:
p=sub.add_parser(name)
p.add_argument('--project',default=os.environ.get('GOOGLE_CLOUD_PROJECT',''))
p.add_argument('--location',default=os.environ.get('GOOGLE_CLOUD_LOCATION','global'))
p.add_argument('--state-uri',required=True,help='gs://bucket/run-prefix/state.json')
p.add_argument('--output-dir',type=Path,required=True)
p.add_argument('--scratch-dir',type=Path,required=True)
if name=='prepare':
source=p.add_mutually_exclusive_group(required=True)
source.add_argument('--source',action='append')
source.add_argument('--local-video',type=Path)
p.add_argument('--revision',default=REVISION)
p.add_argument('--model',default='gemini-3.8-flash')
p.add_argument('--gcs-prefix',required=True,help='A dedicated, unused gs://bucket/run-prefix')
p.add_argument('--start-clip',type=int,default=0)
p.add_argument('--max-clips',type=int,default=3,help='Per source; explicit scope, no full-day default')
p.add_argument('--max-total-clips',type=int,default=2000)
p.add_argument('--fps',type=float,default=1)
p.add_argument('--clip-seconds',type=float,default=30)
p.add_argument('--max-dim',type=int,default=1440)
p.add_argument('--no-review',action='store_true')
p.add_argument('--max-review-regions',type=int,default=4)
p.add_argument('--submit',action='store_true',help='Submit initial audio batch after preparing; otherwise uploads/config only')
return parser
def main(argv=None):
args=build_parser().parse_args(argv)
parse_gs_uri(args.state_uri)
args.output_dir.mkdir(parents=True,exist_ok=True);args.scratch_dir.mkdir(parents=True,exist_ok=True)
log=EventLog();log.attach(args.output_dir/'batch-events.jsonl')
cloud=BatchCloud(args.project,args.location)
planner=TaskPlanner(cloud,args.scratch_dir,args.output_dir,log=log)
engine=BatchEngine(cloud,args.state_uri,planner)
if args.command=='prepare':
parse_gs_uri(args.gcs_prefix)
if args.state_uri!=args.gcs_prefix.rstrip('/')+'/state.json':
raise ValueError('state-uri must be gcs-prefix/state.json')
if args.start_clip<0 or not 1<=args.max_clips<=2000 or not 1<=args.max_total_clips<=2000:
raise ValueError('Invalid bounded batch scope')
if not 0<args.fps<=4 or not 0<args.clip_seconds<=30 or not 128<=args.max_dim<=2880 or not 0<=args.max_review_regions<=8:
raise ValueError('Invalid media settings')
previous,_=cloud.read_state(args.state_uri)
if previous is not None:
raise ValueError('Run already initialized; use start/tick, never replace its scope')
paths=args.source or [args.local_video]
if len(paths)!=len(set(paths)):
raise ValueError('Duplicate source paths would duplicate paid requests')
if len(paths)*args.max_clips>args.max_total_clips:
raise ValueError('Requested maximum exceeds max-total-clips')
prompts={p.stem:p.read_text(encoding='utf-8') for p in (Path(__file__).parent/'prompts').glob('*.md')}
cfg={'model':args.model,'project':args.project,'location':args.location,'gcs_prefix':args.gcs_prefix.rstrip('/'),
'fps':args.fps,'clip_seconds':args.clip_seconds,'max_dim':args.max_dim,
'start_clip':args.start_clip,'max_clips':args.max_clips,'review':not args.no_review,
'max_review_regions':args.max_review_regions,'prompts':prompts,'rows':[]}
cfg['code_hash']=code_hash()
if args.source:
from huggingface_hub import HfApi
revision=HfApi().dataset_info(REPO,revision=args.revision).sha
for path in paths:
if args.source:
meta=source_metadata(path,revision);source=download_source(path,revision,args.scratch_dir)
else:
source=Path(path).resolve(strict=True);s=source.stat()
meta={'source_id':str(source),'file_size':s.st_size,'mtime_ns':s.st_mtime_ns,'viewpoint':'egocentric'}
log.emit('batch_media_start',source=meta['source_id'])
try:cfg['rows'].extend(prepare_media(source,meta,cfg,planner))
finally:
if args.source:remove_downloaded_source(source,args.scratch_dir)
log.emit('batch_media_success',source=meta['source_id'],prepared_clips=len(cfg['rows']))
if not cfg['rows']:raise ValueError('No clips selected')
state=engine.initialize(cfg)
if state['config']!=cfg:
raise ValueError('Another worker initialized this prefix with a different scope; no submission performed')
if args.submit:state=engine.start()
else:
stored,_=cloud.read_state(args.state_uri)
if stored is None:raise ValueError('No batch state exists')
if stored['config']['project']!=args.project or stored['config']['location']!=args.location:
raise ValueError('Project/location must match the saved run')
if args.command=='status':state=stored
else:
expected=code_hash()
if stored['config'].get('code_hash')!=expected:
raise ValueError('Use the pinned code version that initialized this run')
state={'start':engine.start,'tick':engine.tick,'reconcile':engine.reconcile}[args.command]()
atomic_json(args.output_dir/'batch-state.snapshot.json',state)
result=summary(state);atomic_json(args.output_dir/'batch-summary.json',result)
log.emit('batch_tick',**result)
return 2 if state['status']=='needs_attention' else 0
if __name__=='__main__':
try:raise SystemExit(main())
except KeyboardInterrupt:raise SystemExit(130)
except Exception as error:
code=getattr(error,'code',None)
print(json.dumps({'ok':False,'error_type':type(error).__name__,
'code':code if isinstance(code,int) else None,
'message':'Batch worker failed; inspect configuration and persisted state without exposing credentials.'}),file=sys.stderr)
raise SystemExit(1)

Xet Storage Details

Size:
10.9 kB
·
Xet hash:
28e745413f58dd2598af8a1cf353b455ed2a0da533bee1676921d4958f8447f1

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.