"""Parallel sentence encoding with compact provenance and resumable local parts.""" import argparse import contextlib import hashlib import json import math import os import shutil import struct import sys import time from array import array from collections import Counter from concurrent.futures import ProcessPoolExecutor, as_completed from pathlib import Path from vimeml.tokenizer.encode import SPLITS, dump_json, file_sha, write_numbers from vimeml.tokenizer.monitor import RunMonitor from vimeml.tokenizer.train import hash_inputs ROOT = Path(__file__).resolve().parents[3] FORMAT = 'vimeml_sentence_tokens_v1' @contextlib.contextmanager def output_lock(output): with (output/'.encode.lock').open('a+b') as stream: stream.seek(0,2) if not stream.tell(): stream.write(b'1') stream.flush() stream.seek(0) try: if os.name=='nt': import msvcrt msvcrt.locking(stream.fileno(),msvcrt.LK_NBLCK,1) else: import fcntl fcntl.flock(stream,fcntl.LOCK_EX|fcntl.LOCK_NB) except OSError: raise RuntimeError('Another encoder is using this output directory.') from None try: yield finally: stream.seek(0) if os.name=='nt': msvcrt.locking(stream.fileno(),msvcrt.LK_UNLCK,1) else: fcntl.flock(stream,fcntl.LOCK_UN) def part_paths(job): prefix = Path(job['parts'])/job['name'] return {suffix:Path(str(prefix)+'.'+suffix) for suffix in ('tokens.bin','offsets.bin','rows.bin','sources.bin','done.json')} def encode_part(job): import sentencepiece as spm files = part_paths(job) # An interrupted part has no commit marker and is safely overwritten. files['done.json'].unlink(missing_ok=True) processor = spm.SentencePieceProcessor(model_file=job['model']) bos,eos,unk,pad = (getattr(processor,name+'_id')() for name in ('bos','eos','unk','pad')) special = {bos,eos,unk,pad} count = total = content = characters = longest = 0 sources = Counter() text_digest = hashlib.sha256() with Path(job['input']).open('rb') as rows, \ files['tokens.bin'].open('wb') as tokens, files['offsets.bin'].open('wb') as offsets, \ files['rows.bin'].open('wb') as positions, files['sources.bin'].open('wb') as source_ids: if job['start']: rows.seek(job['start']-1) if rows.read(1)!=b'\n': rows.readline() write_numbers(offsets,'Q',[0]) while rows.tell()=job['end']: break raw = rows.readline() if not raw: break row = json.loads(raw) text = row['text'] if not isinstance(text,str) or not text or '\n' in text or '\r' in text: raise ValueError(f"Invalid text: {job['name']} at byte {position}") if hashlib.sha256(text.encode('utf-8')).hexdigest()!=row['text_hash']: raise ValueError(f"Text hash mismatch: {job['name']} at byte {position}") if row['source'] not in job['source_ids']: raise ValueError(f"Unknown source: {row['source']}") byte_positions.append(position) batch.append(row) if not batch: break texts = [row['text'] for row in batch] encoded = processor.encode(texts,out_type=int,num_threads=1) decoded = processor.decode(encoded,num_threads=1) batch_tokens,batch_offsets,byte_sources = [],[],bytearray() for row,ids,restored in zip(batch,encoded,decoded,strict=True): text = row['text'] if restored!=text or any(token in special for token in ids): raise ValueError(f"Roundtrip or special token error: {job['name']} row {count}") batch_tokens.extend((bos,*ids,eos)) total += len(ids)+2 batch_offsets.append(total) count += 1 content += len(ids) characters += len(text) longest = max(longest,len(ids)+2) sources[row['source']] += 1 byte_sources.append(job['source_ids'][row['source']]) text_digest.update((text+'\n').encode('utf-8')) write_numbers(tokens,'H',batch_tokens) write_numbers(offsets,'Q',batch_offsets) write_numbers(positions,'Q',byte_positions) source_ids.write(byte_sources) sizes = {suffix:path.stat().st_size for suffix,path in files.items() if suffix!='done.json'} if sizes!={'tokens.bin':total*2,'offsets.bin':(count+1)*8,'rows.bin':count*8,'sources.bin':count}: raise ValueError('Part size mismatch.') result = {'name':job['name'],'signature':job['signature'],'split':job['split'], 'start':job['start'],'end':job['end'],'sentences':count,'stored_tokens':total, 'content_tokens':content,'characters':characters,'max_sequence_tokens':longest, 'primary_sources':dict(sources),'canonical_txt_sha256':text_digest.hexdigest(), 'files':{suffix:{'bytes':path.stat().st_size,'sha256':file_sha(path)} for suffix,path in files.items() if suffix!='done.json'}} dump_json(files['done.json'],result) return result def cached_part(job): files = part_paths(job) if not files['done.json'].is_file(): return None result = json.loads(files['done.json'].read_text(encoding='utf-8')) if result.get('signature')!=job['signature'] or result.get('name')!=job['name']: raise ValueError('Cached part belongs to a different encoding run.') for suffix,expected in result['files'].items(): path = files[suffix] if not path.is_file() or path.stat().st_size!=expected['bytes'] or file_sha(path)!=expected['sha256']: raise ValueError(f'Cached part is damaged: {path}; preserve evidence and use a new output directory.') return result def make_jobs(corpus,model,parts,source_ids,batch_size,chunk_bytes,signature): jobs = [] for split in SPLITS: path = corpus/f'{split}.jsonl' size = path.stat().st_size partitions = max(1,math.ceil(size/chunk_bytes)) for index in range(partitions): jobs.append({'name':f'{split}-{index:05d}','split':split,'input':str(path), 'start':size*index//partitions,'end':size*(index+1)//partitions, 'model':str(model),'parts':str(parts),'source_ids':source_ids, 'batch_size':batch_size,'signature':signature}) return jobs def collect_parts(jobs,workers,monitor,total_rows): results,pending_jobs = {},[] for job in jobs: cached = cached_part(job) if cached: results[job['name']]=cached else: pending_jobs.append(job) done = sum(result['sentences'] for result in results.values()) print(f'[encode] cached_parts={len(results)}; pending_parts={len(pending_jobs)}; sentences={done:,}/{total_rows:,}',flush=True) monitor.update('encoding',completed_sentences=done,total_sentences=total_rows) started = time.monotonic() initial_done = done pool = ProcessPoolExecutor(max_workers=workers) try: futures = {pool.submit(encode_part,job):job for job in pending_jobs} for future in as_completed(futures): job = futures[future] result = future.result() results[job['name']]=result done += result['sentences'] elapsed = time.monotonic()-started speed = (done-initial_done)/elapsed if elapsed else 0 eta = (total_rows-done)/speed if speed else None print(f'[encode] {done:,}/{total_rows:,} ({100*done/total_rows:.1f}%); {speed:,.0f} sentences/s; ETA={eta:.0f}s' if eta is not None else '[encode] no new rows yet',flush=True) monitor.update('encoding',completed_sentences=done,total_sentences=total_rows, encoding_sentences_per_second=speed,encoding_eta_seconds=eta) except BaseException: # Python 3.13 lacks the public terminate_workers API. Stop only our pool. processes = list((getattr(pool,'_processes',None) or {}).values()) for process in processes: if process.is_alive(): process.terminate() pool.shutdown(wait=True,cancel_futures=True) raise else: pool.shutdown(wait=True) return results def merge_split(corpus,output,split,jobs,results,expected,monitor,done_before,total_rows): counts = Counter() sources = Counter() longest = 0 target = {suffix:output/f'{split}.{suffix}' for suffix in ('tokens.bin','offsets.bin','rows.bin','sources.bin')} with (corpus/f'{split}.txt').open(encoding='utf-8') as text, \ target['tokens.bin'].open('wb') as tokens,target['offsets.bin'].open('wb') as offsets, \ target['rows.bin'].open('wb') as positions,target['sources.bin'].open('wb') as source_ids: write_numbers(offsets,'Q',[0]) for job in jobs: result = results[job['name']] files = part_paths(job) # Match every source TXT row, including exact whitespace, in order. digest = hashlib.sha256() for _ in range(result['sentences']): line = text.readline() if not line: raise ValueError(f'{split}: TXT has fewer rows than JSONL.') digest.update((line.removesuffix('\n')+'\n').encode('utf-8')) if digest.hexdigest()!=result['canonical_txt_sha256']: raise ValueError(f"{split}: TXT/JSONL alignment mismatch in {job['name']}") for suffix,stream in [('tokens.bin',tokens),('rows.bin',positions),('sources.bin',source_ids)]: with files[suffix].open('rb') as source: shutil.copyfileobj(source,stream,1024*1024) with files['offsets.bin'].open('rb') as source: if source.read(8)!=b'\0'*8: raise ValueError('Invalid part initial offset.') previous = 0 while raw:=source.read(1024*1024): values = array('Q') values.frombytes(raw) if sys.byteorder!='little': values.byteswap() adjusted = [] for value in values: if not previous256: raise ValueError('uint8 source IDs support at most 256 sources.') source_ids = {name:i for i,name in enumerate(source_names)} estimated = sum(stats['splits'][split]['tokens_with_bos_eos_per_sentence']*2+ stats['splits'][split]['sentences']*17+8 for split in SPLITS) total_rows = sum(stats['splits'][split]['sentences'] for split in SPLITS) plan = {'corpus':str(corpus),'tokenizer':str(tokenizer),'output':str(output),'workers':workers, 'native_threads_per_worker':1,'batch_size':batch_size,'chunk_mib':chunk_mib, 'sentences':total_rows,'estimated_binary_bytes':estimated,'estimated_peak_binary_bytes':estimated*2, 'source_ids':source_ids,'resume':resume,'dry_run':dry_run,'encoding_started':False, 'tensorboard_log_root':str(ROOT/'runs/tokenizer'/f'{output.name}-encode')} print(json.dumps(plan,ensure_ascii=False,indent=2),flush=True) if dry_run: return plan if output.exists() and (not output.is_dir() or any(output.iterdir())) and not resume: raise FileExistsError('Nonempty output exists; use --resume for this run or choose a new output directory.') if not output.exists() and resume: raise FileNotFoundError('Cannot resume an absent output directory.') output.mkdir(parents=True,exist_ok=True) with output_lock(output): if (output/'manifest.json').is_file(): existing = json.loads((output/'manifest.json').read_text(encoding='utf-8')) if existing.get('status')=='complete': print('Token data already has a complete manifest; no re-encoding performed.',flush=True) return existing['splits'] with RunMonitor(output,Path(plan['tensorboard_log_root'])/str(time.time_ns()),True) as monitor: monitor.update('checking_inputs') paths = [corpus/f'{split}.{suffix}' for split in SPLITS for suffix in ('txt','jsonl')] paths += [corpus/'manifest.json',corpus/'stats.json',model,tokenizer/'manifest.json',tokenizer/'stats.json'] paths += [Path(__file__).with_name(name) for name in ('encode_parallel.py','encode.py','monitor.py','train.py')] print('[inputs] Hashing corpus and tokenizer files...',flush=True) hashes = hash_inputs(paths) for name in ('train.txt','validation.txt','test.txt','manifest.json','stats.json'): if hashes[str(corpus/name)]!=recorded_hash(tm,name): raise ValueError(f'Corpus changed since tokenizer training: {name}') signature = hashlib.sha256(json.dumps({'inputs':hashes,'batch_size':batch_size,'chunk_mib':chunk_mib, 'source_ids':source_ids,'code_sha256':file_sha(Path(__file__))},sort_keys=True).encode()).hexdigest() request = output/'run.json' if resume: if not request.is_file() or json.loads(request.read_text(encoding='utf-8'))['signature']!=signature: raise ValueError('Resume inputs or encoding parameters changed; use the original parameters or a new output directory.') else: dump_json(request,{'signature':signature,'plan':plan,'input_sha256':hashes}) parts = output/'.parts' parts.mkdir(exist_ok=True) jobs = make_jobs(corpus,model,parts,source_ids,batch_size,chunk_mib*1024*1024,signature) results = collect_parts(jobs,workers,monitor,total_rows) exported = {} done = 0 for split in SPLITS: selected = [job for job in jobs if job['split']==split] exported[split] = merge_split(corpus,output,split,selected,results,stats['splits'][split],monitor,done,total_rows) if exported[split]['primary_sources']!=cs['splits'][split]['primary_sources']: raise ValueError(f'{split}: primary source counts differ from corpus stats.') done += exported[split]['sentences'] monitor.update('verifying_inputs') if hash_inputs(paths)!=hashes: raise ValueError('Inputs changed during encoding.') dump_json(output/'stats.json',{'splits':exported}) artifacts = [output/f'{split}.{suffix}' for split in SPLITS for suffix in ('tokens.bin','offsets.bin','rows.bin','sources.bin')] artifacts += [output/'stats.json'] artifact_hashes = {Path(path).name:value for path,value in hash_inputs(artifacts).items()} # Remove only known temporary files created by this run. Keep final data. assert parts.resolve().parent==output.resolve() for job in jobs: for path in part_paths(job).values(): assert path.resolve().parent==parts.resolve() path.unlink() parts.rmdir() dump_json(output/'manifest.json',{'status':'complete','format':FORMAT,'vocab_size':vocab_size, 'token_dtype':'uint16_le','offset_dtype':'uint64_le','offset_unit':'tokens','special_ids':special, 'sentencepiece_version':spm.__version__,'splits':exported,'source_ids':source_ids, 'sequence_policy':'Each sequence is BOS + complete encoded sentence + EOS; no truncation or concatenation.', 'prediction_policy':'x=s[:-1], y=s[1:]; sentence boundaries remain separate.', 'provenance_policy':'Index equals zero-based corpus JSONL/TXT row. rows.bin stores uint64 little-endian byte positions in the source JSONL; sources.bin stores uint8 primary source IDs. Full original metadata stays in corpus JSONL.', 'corpus_dir':str(corpus),'tokenizer_dir':str(tokenizer),'input_sha256':hashes,'output_sha256':artifact_hashes, 'corpus_quality_mode':cm.get('quality_mode'),'corpus_ready_for_lm_training':tm.get('corpus_ready_for_lm_training',False), 'purpose':'Full token data for the accepted first corpus baseline; known extraction and near-duplicate limitations remain.', 'workers':workers,'native_threads_per_worker':1,'signature':signature, 'script_sha256':file_sha(Path(__file__))}) monitor.update('complete',completed_sentences=total_rows,total_sentences=total_rows) print(json.dumps({'splits':exported},ensure_ascii=False,indent=2)) print(f'Token data: {output}',flush=True) return exported def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument('--corpus',type=Path,default=ROOT/'outputs/corpus-fast-v1') parser.add_argument('--tokenizer',type=Path,default=ROOT/'artifacts/tokenizers/ja-unigram-16k-v1') parser.add_argument('--output',type=Path,default=ROOT/'artifacts/token-data/corpus-v1-16k') parser.add_argument('--workers',type=int,default=8) parser.add_argument('--batch-size',type=int,default=1024) parser.add_argument('--chunk-mib',type=int,default=128) parser.add_argument('--resume',action='store_true') parser.add_argument('--dry-run',action='store_true') args = parser.parse_args() run(args.corpus,args.tokenizer,args.output,args.workers,args.batch_size,args.chunk_mib,args.resume,args.dry_run) if __name__=='__main__': main()