Prepare Guard curated dataset for training
This commit is contained in:
parent
bc0c5214ab
commit
c39625227c
2 changed files with 3 additions and 3 deletions
|
|
@ -143,8 +143,8 @@ for source in sources:
|
|||
key=hashlib.sha256((record['input']+'\\0'+record['output']).encode()).hexdigest()
|
||||
if key in seen: stats['duplicates']+=1; continue
|
||||
seen.add(key); record={**record,'artifex_source_key':source_key,'artifex_record_hash':key,'artifex_source_manifest':source}; bucket=int(hashlib.sha256(source_key.encode()).hexdigest(),16)%10
|
||||
(train if bucket<8 else validation if bucket<9 else regression).append(record); stats['accepted']+=1
|
||||
for name,rows in [('train',train),('validation',validation),('regression',regression)]:
|
||||
split='train' if bucket<8 else 'validation' if bucket<9 else 'regression'; record['split']=split; (train if split=='train' else validation if split=='validation' else regression).append(record); stats['accepted']+=1
|
||||
for name,rows in [('train',train),('validation',validation),('regression',regression),('training',train+validation)]:
|
||||
path=os.path.join(out,name+'.json'); json.dump(rows,open(path,'w',encoding='utf-8'),ensure_ascii=False); stats[name]=len(rows); stats[name+'_sha256']=hashlib.sha256(open(path,'rb').read()).hexdigest()
|
||||
json.dump(stats,open(os.path.join(out,'manifest.json'),'w',encoding='utf-8'),indent=2); print(json.dumps({'output_directory':out,**stats}))
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -199,7 +199,7 @@ class ModelStudioService:
|
|||
sources = [item.manifest_reference for item in proposal.source_versions.all() if item.manifest_reference.startswith("/")]
|
||||
report = SparkGuardDatasetMaterializer(ssh_alias).materialize(sources, output_directory, strict=strict)
|
||||
dataset, _ = Dataset.objects.get_or_create(training_project=training_project, name=f"curated-{proposal.id.hex[:12]}")
|
||||
version = DatasetVersion.objects.create(dataset=dataset, version=report["train_sha256"][:12], manifest_reference=output_directory + "/train.json", content_hash=report["train_sha256"], record_count=report["train"], split_metadata={"train": report["train"], "validation": report["validation"], "regression": report["regression"], "validation_reference": output_directory + "/validation.json", "regression_reference": output_directory + "/regression.json", "source_group_split": "sha256(source_sha256 or input hash) mod 10"}, source_metadata={"source_manifests": sources, "curation_proposal": str(proposal.id), "materialization_report": report}, generation_metadata={"operations": proposal.proposed_operations, "authorized_source_mandate": True, "old_holdouts_resplit": True, "strict_schema_repair": strict}, tags=["curated", "train", "spark", "requires_new_evaluation_suite", *( ["strict_schema_repaired"] if strict else [])], validation_status=DatasetValidationStatus.VALID, contamination_status=DatasetValidationStatus.WARNING)
|
||||
version = DatasetVersion.objects.create(dataset=dataset, version=report["training_sha256"][:12], manifest_reference=output_directory + "/training.json", content_hash=report["training_sha256"], record_count=report["training"], split_metadata={"train": report["train"], "validation": report["validation"], "regression": report["regression"], "validation_reference": output_directory + "/validation.json", "regression_reference": output_directory + "/regression.json", "source_group_split": "sha256(source_sha256 or input hash) mod 10"}, source_metadata={"source_manifests": sources, "curation_proposal": str(proposal.id), "materialization_report": report}, generation_metadata={"operations": proposal.proposed_operations, "authorized_source_mandate": True, "old_holdouts_resplit": True, "strict_schema_repair": strict}, tags=["curated", "train", "spark", "requires_new_evaluation_suite", *( ["strict_schema_repaired"] if strict else [])], validation_status=DatasetValidationStatus.VALID, contamination_status=DatasetValidationStatus.WARNING)
|
||||
proposal.materialized_version = version
|
||||
proposal.status = "MATERIALIZED"
|
||||
proposal.save(update_fields=["materialized_version", "status", "updated_at"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue