Release best Gemini-tuned Whisper Base and Small with code, normalization and evaluation
cd9b2d8 verified Download training/embedding_probe_study/watch_study.py from laion/humaneness-ears-base-medium: direct link, hf CLI and curl.
- Browser
- Download file 11.3 kB
-
https://huggingface.co/laion/humaneness-ears-base-medium/resolve/main/training/embedding_probe_study/watch_study.py
- Command line
-
hf download hf://laion/humaneness-ears-base-medium/training/embedding_probe_study/watch_study.py
-
curl -L -o watch_study.py https://huggingface.co/laion/humaneness-ears-base-medium/resolve/main/training/embedding_probe_study/watch_study.py
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() | |