ChristophSchuhmann's picture
Release best Gemini-tuned Whisper Base and Small with code, normalization and evaluation
cd9b2d8 verified
Raw History Blame Contribute Delete
11.3 kB
#!/usr/bin/env python3
"""Advance per-model probes; give full Whisper FT an independent GPU queue."""
import fcntl
import os
import socket
import subprocess
import time
from datetime import datetime, timezone
from study_paths import ROOT, HERE, CODE, GEMINI, read, write
ACTIVE={'PENDING','RUNNING','CONFIGURING','COMPLETING','SUSPENDED','REQUEUED'}
RESUMABLE={'TIMEOUT','PREEMPTED','NODE_FAIL','BOOT_FAIL'}
def retry_allocations(attempts):
# A cancelled held script or a diagnosed/repaired startup/export failure
# does not consume the three-allocation limit for interruption recovery.
# Concrete failures still require a repair below before any retry.
return sum(not (a.get('state') in {'FAILED','CANCELLED'} and
(a.get('repair') or a.get('replacement_reason'))) for a in attempts)
def state(job):
r=subprocess.run(['sacct','-j',str(job),'-n','-X','--format=State','--parsable2'],
capture_output=True,text=True,check=True,timeout=15)
lines=r.stdout.strip().splitlines()
return lines[0].split('|')[0].split()[0].rstrip('+') if lines else 'UNKNOWN'
def submit(script,env,name,extra=()):
r=subprocess.run(['sbatch','--parsable','--job-name='+name,
'--export=ALL,'+','.join(k+'='+str(v) for k,v in env.items()),
*extra,str(script)],capture_output=True,text=True,check=True,timeout=20)
return int(r.stdout.strip().split(';')[0])
def domains_done(model,domains):
return all((ROOT/'features'/model/(d+'-rank'+str(r)+'-COMPLETE.json')).exists()
for d in domains for r in range(4))
def done_cache(model,cfg):
return domains_done(model,['ladder','p3_train','p3_validation','p3_test','gemini',*cfg['benchmarks']])
def probe_done(phase,model,evaluation=False):
out=ROOT/'probes'/phase/model
if evaluation:
return all((out/v/'public_metrics.json').exists() and
all((out/v/(k+'_matched_adapter.json')).exists() for k in ('crema','ravdess'))
for v in ('linear','mlp'))
return all((out/v/'COMPLETE.json').exists() for v in ('linear','mlp'))
def tick(w,cfg):
stamp=datetime.now(timezone.utc).isoformat()
w.update(updated_utc=stamp,host=socket.gethostname(),pid=os.getpid())
w['maximum_active_study_nodes']=5
known={}
for group in ('tasks','cache_jobs'):
for name,attempts in w.get(group,{}).items():
if attempts:
a=attempts[-1];a['state']=state(a['id']);known[a['id']]=a['state']
gate=(GEMINI/'prepared/TRAINING_READY.json').exists()
# Release the already submitted four-GPU FT jobs immediately on final data.
# Cache/probe capacity cannot delay this release; held jobs reserve no node.
for model in ('base','small'):
name='whisper-gemini-'+model;attempts=w['tasks'].get(name,[])
if not attempts:raise RuntimeError('Missing pre-submitted full-FT chain: '+name)
a=attempts[-1];out=GEMINI/'training'/('whisper_'+model)
if a.get('held_for_targets') and gate and a['state']=='PENDING':
subprocess.run(['scontrol','release',str(a['id'])],check=True,capture_output=True,text=True,timeout=20)
a.update(held_for_targets=False,released_utc=stamp);write(ROOT/'workflow.json',w)
elif a['state'] not in ACTIVE|{'UNKNOWN'} and not (out/'COMPLETE.json').exists():
if retry_allocations(attempts)>=3:
w.setdefault('errors',{})[name]='Three incomplete FT allocations; inspect saved logs.'
elif a['state'] in RESUMABLE:
env={'WHISPER_MODEL':model,'RESUME_CHECKPOINT':str(out/'latest.pt') if (out/'latest.pt').exists() else ''}
job=submit(CODE/'gemini_finetune/train_whisper.sbatch',env,'flash38-fullft-'+model)
attempts.append({'id':job,'state':'PENDING','submitted_utc':stamp,'resumed_from':a['id']})
write(ROOT/'workflow.json',w)
else:w.setdefault('errors',{})[name]='FT ended '+a['state']+'; repair the concrete error before another allocation.'
parents=[w['tasks']['whisper-gemini-'+m][-1]['id'] for m in ('base','small')]
eval_attempts=w['tasks'].setdefault('whisper-eval-all',[])
done_whisper=all((ROOT/'whisper/gemini'/m/'COMPLETE.json').exists() for m in ('base','small'))
# Reconnect evaluation to resumed FT jobs, rather than leaving a broken
# afterok dependency attached to an obsolete wall-time-limited allocation.
if not done_whisper:
a=eval_attempts[-1] if eval_attempts else None
if a and a['state']=='PENDING' and a.get('wait_for_ids')!=parents:
subprocess.run(['scontrol','update','JobId='+str(a['id']),
'Dependency=afterok:'+':'.join(map(str,parents))],
check=True,capture_output=True,text=True,timeout=20)
a['wait_for_ids']=parents;write(ROOT/'workflow.json',w)
elif not a or a['state'] not in ACTIVE|{'UNKNOWN'}:
if retry_allocations(eval_attempts)<3 and (not a or a['state'] in RESUMABLE|{'CANCELLED','COMPLETED'}):
job=submit(HERE/'whisper_eval.sbatch',{'WHISPER_MODEL':'all'},'flash38-fullft-evaluation',
['--dependency=afterok:'+':'.join(map(str,parents)),'--kill-on-invalid-dep=yes'])
eval_attempts.append({'id':job,'state':'PENDING','wait_for_ids':parents,'submitted_utc':stamp})
write(ROOT/'workflow.json',w)
else:w.setdefault('errors',{})['whisper-eval-all']='Evaluation incomplete; inspect its concrete error.'
# One probe/evaluation node runs alongside two cache nodes and two dedicated
# FT nodes. Evaluation with an unmet dependency consumes no node or slot.
busy_aux=False;active_aux=0
for name,attempts in w['tasks'].items():
if name.startswith('whisper-gemini-') or not attempts:continue
a=attempts[-1]
if a['state'] not in ACTIVE|{'UNKNOWN'}:continue
if a.get('wait_for_ids') and a['state']=='PENDING' and not all(known.get(j)=='COMPLETED' for j in a['wait_for_ids']):continue
busy_aux=True;active_aux+=1
def probe_task(phase,models,evaluation):
nonlocal busy_aux,active_aux
if busy_aux or not models:return
models=models[:2];name=('eval-' if evaluation else 'train-')+phase+'-'+'-'.join(models)
attempts=w['tasks'].setdefault(name,[])
if attempts and attempts[-1]['state'] in ACTIVE|{'UNKNOWN'}:return
if retry_allocations(attempts)>=3 or (attempts and attempts[-1]['state'] not in RESUMABLE|{'COMPLETED'}):
w.setdefault('errors',{})[name]='Incomplete allocation; repair its concrete error before retrying.';return
env={'PROBE_PHASE':phase,'PROBE_MODELS':':'.join(models),'PROBE_EVAL':0}
job=submit(HERE/('probe_eval.sbatch' if evaluation else 'train.sbatch'),env,name)
attempts.append({'id':job,'state':'PENDING','submitted_utc':stamp,'models':models,'phase':phase,
'operation':'evaluation' if evaluation else 'training'})
busy_aux=True;active_aux+=1;write(ROOT/'workflow.json',w)
for phase in ('gemini','legacy'):
ready=[m['id'] for m in cfg['models'] if probe_done(phase,m['id']) and not probe_done(phase,m['id'],True)
and domains_done(m['id'],cfg['benchmarks'])]
probe_task(phase,ready,True)
ready=[m['id'] for m in cfg['models'] if gate and domains_done(m['id'],['gemini'])
and probe_done('legacy',m['id']) and not probe_done('gemini',m['id'])]
probe_task('gemini',ready,False)
legacy_domains=['ladder','p3_train','p3_validation','p3_test']
ready=[m['id'] for m in cfg['models'] if domains_done(m['id'],legacy_domains) and not probe_done('legacy',m['id'])]
probe_task('legacy',ready,False)
active_cache=sum(bool(a) and a[-1]['state'] in ACTIVE|{'UNKNOWN'} for a in w['cache_jobs'].values())
active_ft=sum(w['tasks']['whisper-gemini-'+m][-1]['state'] in ACTIVE|{'UNKNOWN'} for m in ('base','small'))
ft_finished=all((GEMINI/'training'/('whisper_'+m)/'COMPLETE.json').exists() for m in ('base','small'))
# Keep FT capacity available until both full runs finish. Then reuse those
# slots for remaining frozen encoders without exceeding five study nodes.
cache_limit=min(4,max(0,5-active_ft-active_aux)) if ft_finished else 2
w['cache_node_limit_now']=cache_limit
w['scheduling_policy']['cache_after_full_ft']='Up to four; total study allocations stay at most five'
pilot=state(w['pilot_job']);w['pilot_state']=pilot
pilot_ok=pilot=='COMPLETED' and all((ROOT/'smoke/clapv2_xxs'/('CONTRACT-rank'+str(r)+'.json')).exists() for r in range(4))
if pilot_ok:
for m in cfg['models']:
mid=m['id']
if active_cache>=cache_limit:break
if done_cache(mid,cfg):continue
if m.get('repo') and not ((ROOT/'models'/mid/'MODEL_READY.json').exists() or (ROOT/'MODELS_READY.json').exists()):continue
attempts=w['cache_jobs'].setdefault(mid,[])
if attempts and attempts[-1]['state'] in ACTIVE|{'UNKNOWN'}:continue
if retry_allocations(attempts)>=3 or (attempts and attempts[-1]['state'] not in RESUMABLE):
w.setdefault('errors',{})[mid]='Encoder cache failed; inspect smoke/full-cache logs.';continue
job=submit(HERE/'cache.sbatch',{'PROBE_MODEL':mid,'PROBE_SMOKE':0},'probe-cache-'+mid)
attempts.append({'id':job,'state':'PENDING','submitted_utc':stamp});active_cache+=1
write(ROOT/'workflow.json',w)
else:w.setdefault('errors',{})['pilot']='Encoder/head smoke contract has not passed.'
w['cache_states']={}
for m in cfg['models']:
attempts=w['cache_jobs'].get(m['id'],[])
w['cache_states'][m['id']]='COMPLETE' if done_cache(m['id'],cfg) else attempts[-1]['state'] if attempts else 'NOT_SUBMITTED'
complete=done_whisper and all(probe_done(p,m['id'],True) for m in cfg['models'] for p in ('legacy','gemini'))
w['state']='COMPLETE' if complete else 'RUNNING_WITH_ERRORS' if w.get('errors') else 'RUNNING_OR_WAITING_FOR_DATA'
write(ROOT/'workflow.json',w)
return complete
def main():
with (ROOT/'watcher.lock').open('a') as lock:
try:fcntl.flock(lock,fcntl.LOCK_EX|fcntl.LOCK_NB)
except BlockingIOError:return
w=read(ROOT/'workflow.json');last_report=0
while True:
try:
complete=tick(w,read(ROOT/'study.json'));w.pop('watcher_error',None)
if complete or time.monotonic()-last_report>300:
subprocess.run(['python3',str(HERE/'write_report.py')],check=True,timeout=30)
r=subprocess.run(['bash',str(HERE/'publish.sh')],capture_output=True,text=True,timeout=120)
if r.returncode:w['publication_error']=r.stderr[-1200:]
else:w.pop('publication_error',None)
write(ROOT/'workflow.json',w);last_report=time.monotonic()
if complete:return
except Exception as error:
w['watcher_error']=str(error)[-1800:];write(ROOT/'workflow.json',w)
time.sleep(30)
if __name__=='__main__':
main()