Add Guard dataset curation pipeline

This commit is contained in:
Daniel Maddern 2026-08-17 02:04:43 +07:00
parent 557764388b
commit e27df771d5
9 changed files with 366 additions and 3 deletions

View file

@ -0,0 +1,27 @@
import json
from django.core.management.base import BaseCommand, CommandError
from control_plane.model_studio.models import DatasetVersion, TrainingProject
from control_plane.model_studio.services import ModelStudioService
from model_router.providers import providers_from_resources
from model_router.router import ModelRouter
class Command(BaseCommand):
help = "Ask the Qwen Dataset Curator for a structured, non-mutating Guard dataset improvement proposal."
def add_arguments(self, parser):
parser.add_argument("--project", required=True, help="TrainingProject slug")
parser.add_argument("--status", default="WARNING", help="Dataset validation status to review")
def handle(self, *args, **options):
project = TrainingProject.objects.filter(slug=options["project"]).first()
if project is None:
raise CommandError("TrainingProject not found.")
versions = list(DatasetVersion.objects.filter(dataset__training_project=project, validation_status=options["status"]).order_by("-record_count"))
if not versions:
raise CommandError("No matching DatasetVersions to curate.")
router = ModelRouter(providers_from_resources(), persist_requests=True)
proposal = ModelStudioService(router=router).propose_dataset_curation(project, versions)
self.stdout.write(json.dumps({"proposal_id": str(proposal.id), "title": proposal.title, "operations": proposal.proposed_operations, "validation_plan": proposal.validation_plan}, indent=2, default=str))

View file

@ -0,0 +1,25 @@
import json
from django.core.management.base import BaseCommand, CommandError
from control_plane.model_studio.models import DatasetVersion
from control_plane.model_studio.services import ModelStudioService
from model_router.providers import providers_from_resources
from model_router.router import ModelRouter
class Command(BaseCommand):
help = "Run a bounded Qwen quality audit over stratified samples from a curated Spark dataset version."
def add_arguments(self, parser):
parser.add_argument("--dataset-version", required=True)
parser.add_argument("--samples", type=int, default=8)
parser.add_argument("--ssh-alias", default="spark")
def handle(self, *args, **options):
version = DatasetVersion.objects.filter(id=options["dataset_version"]).select_related("dataset__training_project").first()
if version is None:
raise CommandError("DatasetVersion not found.")
service = ModelStudioService(router=ModelRouter(providers_from_resources(), persist_requests=True))
report = service.audit_curated_dataset_with_qwen(version.dataset.training_project, version, ssh_alias=options["ssh_alias"], sample_count=options["samples"])
self.stdout.write(json.dumps(report, indent=2, default=str))

View file

@ -0,0 +1,22 @@
import json
from django.core.management.base import BaseCommand, CommandError
from control_plane.model_studio.models import TrainingProject
from control_plane.model_studio.services import ModelStudioService
class Command(BaseCommand):
help = "Read-only import of Guard dataset manifest metadata from Spark."
def add_arguments(self, parser):
parser.add_argument("--project", required=True, help="TrainingProject slug")
parser.add_argument("--manifest", action="append", required=True, help="Exact Spark JSON manifest path. Repeat for each manifest; recursive directory scans are deliberately unsupported.")
parser.add_argument("--ssh-alias", default="spark")
def handle(self, *args, **options):
project = TrainingProject.objects.filter(slug=options["project"]).first()
if project is None:
raise CommandError("TrainingProject not found.")
report = ModelStudioService().import_spark_guard_datasets(project, options["manifest"], ssh_alias=options["ssh_alias"])
self.stdout.write(json.dumps({key: report[key] for key in ["references", "manifest_count", "record_count", "blocked", "warning"]}, indent=2))

View file

@ -0,0 +1,26 @@
import json
from django.core.management.base import BaseCommand, CommandError
from control_plane.model_studio.models import DatasetCurationProposal
from control_plane.model_studio.services import ModelStudioService
class Command(BaseCommand):
help = "Materialize a Qwen curation proposal into a new immutable Spark dataset version and source-hash splits."
def add_arguments(self, parser):
parser.add_argument("--proposal", required=True)
parser.add_argument("--output-directory", required=True)
parser.add_argument("--ssh-alias", default="spark")
parser.add_argument("--strict-schema-repair", action="store_true")
def handle(self, *args, **options):
proposal = DatasetCurationProposal.objects.filter(id=options["proposal"]).first()
if proposal is None:
raise CommandError("DatasetCurationProposal not found.")
training_project = proposal.source_versions.first().dataset.training_project if proposal.source_versions.exists() else None
if training_project is None:
raise CommandError("Proposal has no source versions.")
version = ModelStudioService().materialize_spark_curation(training_project, proposal, output_directory=options["output_directory"], ssh_alias=options["ssh_alias"], strict=options["strict_schema_repair"])
self.stdout.write(json.dumps({"dataset_version": str(version.id), "reference": version.manifest_reference, "records": version.record_count, "splits": version.split_metadata}, indent=2))

View file

@ -0,0 +1,19 @@
from django.core.management.base import BaseCommand, CommandError
from control_plane.model_studio.models import TrainingProject
from control_plane.model_studio.services import ModelStudioService
class Command(BaseCommand):
help = "Remove malformed null-record-count rows from a previous Spark inventory import; never touches Spark files."
def add_arguments(self, parser):
parser.add_argument("--project", required=True)
parser.add_argument("--reference-prefix", required=True)
def handle(self, *args, **options):
project = TrainingProject.objects.filter(slug=options["project"]).first()
if project is None:
raise CommandError("TrainingProject not found.")
count = ModelStudioService().purge_malformed_spark_inventory(project, reference_prefix=options["reference_prefix"])
self.stdout.write(self.style.SUCCESS(f"Purged {count} malformed inventory rows."))

View file

@ -0,0 +1,39 @@
# Generated by Django 5.2.16 on 2026-08-16 18:16
import django.db.models.deletion
import uuid
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('model_studio', '0001_initial'),
]
operations = [
migrations.CreateModel(
name='DatasetCurationProposal',
fields=[
('id', models.UUIDField(default=uuid.uuid4, editable=False, primary_key=True, serialize=False)),
('created_at', models.DateTimeField(auto_now_add=True)),
('updated_at', models.DateTimeField(auto_now=True)),
('title', models.CharField(max_length=255)),
('hypothesis', models.TextField()),
('evidence', models.JSONField(default=dict)),
('proposed_operations', models.JSONField(default=list)),
('expected_capability_effect', models.TextField(blank=True)),
('expected_risks', models.JSONField(blank=True, default=list)),
('validation_plan', models.JSONField(default=dict)),
('contamination_plan', models.JSONField(default=dict)),
('status', models.CharField(choices=[('PROPOSED', 'Proposed'), ('VALIDATED', 'Validated'), ('REJECTED', 'Rejected'), ('MATERIALIZED', 'Materialized')], default='PROPOSED', max_length=32)),
('created_by_agent', models.CharField(default='DATASET_CURATOR', max_length=120)),
('model_evidence', models.JSONField(blank=True, default=dict)),
('materialized_version', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='materialized_by_proposals', to='model_studio.datasetversion')),
('source_versions', models.ManyToManyField(related_name='curation_proposals', to='model_studio.datasetversion')),
],
options={
'abstract': False,
},
),
]

View file

@ -175,6 +175,29 @@ class DatasetVersion(TimestampedModel):
super().save(*args, **kwargs) super().save(*args, **kwargs)
class DatasetCurationProposalStatus(models.TextChoices):
PROPOSED = "PROPOSED"
VALIDATED = "VALIDATED"
REJECTED = "REJECTED"
MATERIALIZED = "MATERIALIZED"
class DatasetCurationProposal(TimestampedModel):
source_versions = models.ManyToManyField(DatasetVersion, related_name="curation_proposals")
title = models.CharField(max_length=255)
hypothesis = models.TextField()
evidence = models.JSONField(default=dict)
proposed_operations = models.JSONField(default=list)
expected_capability_effect = models.TextField(blank=True)
expected_risks = models.JSONField(default=list, blank=True)
validation_plan = models.JSONField(default=dict)
contamination_plan = models.JSONField(default=dict)
status = models.CharField(max_length=32, choices=DatasetCurationProposalStatus.choices, default=DatasetCurationProposalStatus.PROPOSED)
created_by_agent = models.CharField(max_length=120, default="DATASET_CURATOR")
model_evidence = models.JSONField(default=dict, blank=True)
materialized_version = models.ForeignKey(DatasetVersion, on_delete=models.SET_NULL, null=True, blank=True, related_name="materialized_by_proposals")
class TrainingRecipe(TimestampedModel): class TrainingRecipe(TimestampedModel):
training_project = models.ForeignKey(TrainingProject, on_delete=models.CASCADE, related_name="recipes") training_project = models.ForeignKey(TrainingProject, on_delete=models.CASCADE, related_name="recipes")
name = models.CharField(max_length=200) name = models.CharField(max_length=200)

View file

@ -2,6 +2,8 @@ from __future__ import annotations
import hashlib import hashlib
import json import json
import shlex
import subprocess
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import Any, Protocol from typing import Any, Protocol
@ -73,3 +75,91 @@ class GuardModelProfile:
for block in iter(lambda: handle.read(1024 * 1024), b""): for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block) digest.update(block)
return digest.hexdigest() return digest.hexdigest()
class SparkGuardDatasetInventory:
"""Read-only manifest metadata probe for a registered Spark Guard workspace."""
def __init__(self, ssh_alias: str = "spark") -> None:
self.ssh_alias = ssh_alias
def inspect(self, references: list[str]) -> list[dict[str, Any]]:
script = """import hashlib,json,os,sys
keys=['dataset_status','final_model_eligible','production_training_authorized','benchmark_source_included','synthetic_code_included','split','c4_invalid','teacher','vulnerability_type']
for path in sys.argv[1:]:
try:
raw=open(path,'rb').read(); value=json.loads(raw); first=value[0] if isinstance(value,list) and value else {}
print(json.dumps({'reference':path,'content_hash':hashlib.sha256(raw).hexdigest(),'record_count':len(value) if isinstance(value,list) else None,'first':{key:first.get(key) for key in keys if key in first}},ensure_ascii=True))
except Exception as exc: print(json.dumps({'reference':path,'error':str(exc)},ensure_ascii=True))
"""
if not references:
return []
output = self._remote("python3 -c " + shlex.quote(script) + " " + " ".join(self._quote(item) for item in references), timeout=120)
records = []
for line in output.splitlines():
try:
item = json.loads(line)
except json.JSONDecodeError:
continue
if "content_hash" in item and item.get("record_count") is not None:
records.append(item)
return records
def _remote(self, command: str, *, timeout: int = 120) -> str:
completed = subprocess.run(["ssh", self.ssh_alias, command], capture_output=True, text=True, encoding="utf-8", errors="replace", timeout=timeout, check=False)
if completed.returncode:
raise RuntimeError(completed.stderr.strip() or f"Spark inventory command failed: {command}")
return completed.stdout.strip()
@staticmethod
def _quote(value: str) -> str:
return "'" + value.replace("'", "'\\''") + "'"
class SparkGuardDatasetMaterializer(SparkGuardDatasetInventory):
"""Creates new versioned JSON manifests from authorized Guard source records on Spark."""
def materialize(self, source_references: list[str], output_directory: str, *, strict: bool = False) -> dict[str, Any]:
script = """import hashlib,json,os,sys
out=sys.argv[1]; strict=sys.argv[2]=='1'; sources=sys.argv[3:]; os.makedirs(out,exist_ok=True)
train=[]; validation=[]; regression=[]; seen=set(); stats={'sources':len(sources),'read':0,'accepted':0,'malformed':0,'duplicates':0}
for source in sources:
value=json.load(open(source,encoding='utf-8'))
if not isinstance(value,list): continue
for record in value:
stats['read']+=1
if not isinstance(record,dict) or not isinstance(record.get('input'),str) or not record['input'].strip() or not isinstance(record.get('output'),str) or not record['output'].strip(): stats['malformed']+=1; continue
if strict:
try: target=json.loads(record['output'])
except Exception: stats['malformed']+=1; continue
if not isinstance(target,dict) or not isinstance(target.get('findings'),list) or record.get('c4_invalid') is True or 'pragma solidity' not in record['input'].lower(): stats['malformed']+=1; continue
normalized=[]
for finding in target['findings']:
if not isinstance(finding,dict): continue
kind=finding.get('canonical_type') or finding.get('type')
if kind: normalized.append({key:finding[key] for key in ['canonical_type','type','severity','location','description','evidence','vulnerable_line_spans'] if key in finding})
record={**record,'output':json.dumps({'findings':normalized},ensure_ascii=False),'artifex_schema':'guard_finding_v1'}
source_key=str(record.get('source_sha256') or record.get('parent_source_sha256') or hashlib.sha256(record['input'].encode()).hexdigest())
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)]:
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}))
"""
if not source_references:
raise ValueError("Dataset materialization requires source manifests.")
output = self._remote("python3 -c " + shlex.quote(script) + " " + self._quote(output_directory) + " " + ("1" if strict else "0") + " " + " ".join(self._quote(item) for item in source_references), timeout=900)
return json.loads(output)
def sample(self, reference: str, *, count: int = 8, maximum_field_chars: int = 3000) -> list[dict[str, Any]]:
script = """import json,sys
rows=json.load(open(sys.argv[1],encoding='utf-8')); count=int(sys.argv[2]); limit=int(sys.argv[3]); step=max(1,len(rows)//max(1,count)); out=[]
for index in range(0,len(rows),step):
record=rows[index]; out.append({key:(value[:limit] if isinstance(value,str) else value) for key,value in record.items() if key in ['input','output','record_id','source_sha256','parent_source_sha256','vulnerability_type','artifex_source_manifest']})
if len(out)>=count: break
print(json.dumps(out,ensure_ascii=True))
"""
output = self._remote("python3 -c " + shlex.quote(script) + " " + self._quote(reference) + f" {count} {maximum_field_chars}", timeout=120)
return json.loads(output)

View file

@ -13,20 +13,24 @@ from control_plane.events.bus import EventBus
from control_plane.model_studio.backends import BackendResult, FakeTrainingBackend, TrainingBackend from control_plane.model_studio.backends import BackendResult, FakeTrainingBackend, TrainingBackend
from control_plane.model_studio.models import ( from control_plane.model_studio.models import (
BenchmarkResult, CheckpointType, CheckpointValidityStatus, Conclusion, Dataset, DatasetValidationStatus, BenchmarkResult, CheckpointType, CheckpointValidityStatus, Conclusion, Dataset, DatasetValidationStatus,
DatasetVersion, EvaluationRun, EvaluationRunStatus, EvaluationSuite, EvaluationSuiteVersion, ExperimentStatus, DatasetVersion, DatasetCurationProposal, EvaluationRun, EvaluationRunStatus, EvaluationSuite, EvaluationSuiteVersion, ExperimentStatus,
FailureCluster, ModelCheckpoint, ModelPromotionDecision, ModelPromotionPolicy, ModelStudioArtifact, FailureCluster, ModelCheckpoint, ModelPromotionDecision, ModelPromotionPolicy, ModelStudioArtifact,
OvernightResearchReport, OvernightTrainingProgram, ProgramStatus, PromotionDecision, TrainingExperiment, OvernightResearchReport, OvernightTrainingProgram, ProgramStatus, PromotionDecision, TrainingExperiment,
TrainingProject, TrainingProjectStatus, TrainingRecipe, TrainingRun, TrainingRunStatus, TrainingProject, TrainingProjectStatus, TrainingRecipe, TrainingRun, TrainingRunStatus,
) )
from control_plane.model_studio.profiles import GuardModelProfile, ModelProjectProfile from control_plane.model_studio.profiles import GuardModelProfile, ModelProjectProfile, SparkGuardDatasetInventory, SparkGuardDatasetMaterializer
from control_plane.projects.models import Project, ProjectStatus from control_plane.projects.models import Project, ProjectStatus
from model_router.providers import extract_json_object
from model_router.router import ModelCapability, ModelRequestContract, ModelRouter
class ModelStudioService: class ModelStudioService:
def __init__(self, *, profile: ModelProjectProfile | None = None, backend: TrainingBackend | None = None, bus: EventBus | None = None) -> None: def __init__(self, *, profile: ModelProjectProfile | None = None, backend: TrainingBackend | None = None, bus: EventBus | None = None, router: ModelRouter | None = None, dataset_curator_model_hint: str = "qwen") -> None:
self.profile = profile or GuardModelProfile() self.profile = profile or GuardModelProfile()
self.backend = backend or FakeTrainingBackend() self.backend = backend or FakeTrainingBackend()
self.bus = bus or EventBus() self.bus = bus or EventBus()
self.router = router
self.dataset_curator_model_hint = dataset_curator_model_hint
def import_guard(self, *, project: Project | None = None, repository_path: str, spark_working_directory: str = "", slug: str = "guard-3b") -> TrainingProject: def import_guard(self, *, project: Project | None = None, repository_path: str, spark_working_directory: str = "", slug: str = "guard-3b") -> TrainingProject:
project = project or Project.objects.create(name="Guard 3B Model Studio", project_type="MODEL", goal="Reconstruct and safely improve the Guard 3B model.", repository_path=repository_path, status=ProjectStatus.ARCHAEOLOGY) project = project or Project.objects.create(name="Guard 3B Model Studio", project_type="MODEL", goal="Reconstruct and safely improve the Guard 3B model.", repository_path=repository_path, status=ProjectStatus.ARCHAEOLOGY)
@ -129,6 +133,94 @@ class ModelStudioService:
self._artifact(training_project, "DATASET_CURATION_REPORT", "Guard dataset curation", summary) self._artifact(training_project, "DATASET_CURATION_REPORT", "Guard dataset curation", summary)
return summary return summary
def import_spark_guard_datasets(self, training_project: TrainingProject, references: list[str], *, ssh_alias: str = "spark") -> dict[str, Any]:
inventory = SparkGuardDatasetInventory(ssh_alias).inspect(references)
imported = []
for item in inventory:
name = Path(item["reference"]).stem
dataset, _ = Dataset.objects.get_or_create(training_project=training_project, name=f"spark-{name}")
first = item["first"]
tags = ["spark", "imported"]
status = DatasetValidationStatus.WARNING
reason = "Remote manifest requires record-level provenance and contamination review."
lowered = item["reference"].lower()
if "/research_only/" in lowered or "holdout" in lowered:
tags.extend(["curation_source", "requires_evaluation_resplit"])
status, reason = DatasetValidationStatus.WARNING, "Authorized holdout source requires a newly versioned evaluation split before it may enter training."
elif first.get("final_model_eligible") is False:
tags.extend(["curation_source", "not_direct_training"])
status, reason = DatasetValidationStatus.WARNING, "Authorized source material requires curation into a new validated DatasetVersion before training."
elif first.get("benchmark_source_included") is True:
tags.extend(["benchmark_exclusion", "not_training"])
status, reason = DatasetValidationStatus.BLOCKED, "Manifest declares benchmark source inclusion."
elif first.get("c4_invalid") is True:
tags.extend(["weak_or_invalid_label", "curation_source", "not_direct_training"])
status, reason = DatasetValidationStatus.WARNING, "Authorized source material has weak/invalid labels and requires repair or exclusion before training."
DatasetVersion.objects.update_or_create(dataset=dataset, version=item["content_hash"][:12], defaults={"manifest_reference": item["reference"], "content_hash": item["content_hash"], "record_count": item["record_count"], "source_metadata": {"spark_inventory": first, "curation_reason": reason}, "tags": tags, "validation_status": status, "contamination_status": DatasetValidationStatus.UNKNOWN})
imported.append({"reference": item["reference"], "record_count": item["record_count"], "validation_status": status, "reason": reason})
report = {"references": references, "manifest_count": len(imported), "record_count": sum(item["record_count"] or 0 for item in imported), "blocked": sum(item["validation_status"] == DatasetValidationStatus.BLOCKED for item in imported), "warning": sum(item["validation_status"] == DatasetValidationStatus.WARNING for item in imported), "manifests": imported}
self._artifact(training_project, "SPARK_DATASET_INVENTORY", "Spark Guard dataset inventory", report)
return report
def purge_malformed_spark_inventory(self, training_project: TrainingProject, *, reference_prefix: str) -> int:
stale = DatasetVersion.objects.filter(dataset__training_project=training_project, dataset__name__startswith="spark-", manifest_reference__startswith=reference_prefix, record_count__isnull=True)
count = stale.count()
while ids := list(stale.values_list("id", flat=True)[:500]):
DatasetVersion.objects.filter(id__in=ids).delete()
Dataset.objects.filter(training_project=training_project, name__startswith="spark-", versions__isnull=True).delete()
self._artifact(training_project, "SPARK_INVENTORY_PURGE", "Malformed Spark inventory purge", {"reference_prefix": reference_prefix, "deleted_dataset_versions": count})
return count
def propose_dataset_curation(self, training_project: TrainingProject, versions: list[DatasetVersion]) -> DatasetCurationProposal:
if not versions:
raise ValueError("Dataset curation requires at least one source DatasetVersion.")
inventory = [{"reference": item.manifest_reference, "records": item.record_count, "validation": item.validation_status, "contamination": item.contamination_status, "tags": item.tags, "evidence": item.source_metadata.get("curation", item.source_metadata.get("spark_inventory", {}))} for item in versions]
prompt = "DATASET_CURATOR V0.1. Analyze existing Guard dataset manifest metadata only. Do not invent source quality, labels, coverage, or benchmark results. Return JSON with title, hypothesis, evidence object, proposed_operations list, expected_capability_effect, expected_risks list, validation_plan object, contamination_plan object. Proposed operations must create a NEW immutable DatasetVersion and may filter/reweight/split/select existing records. All legacy data, including former holdouts, is authorized source material. If a former holdout enters training, explicitly require a new source-disjoint evaluation suite version and retire the old comparison split. Inventory: " + json.dumps(inventory, default=str)
review: dict[str, Any] = {}
if self.router is not None and self.dataset_curator_model_hint in self.router.providers:
try:
response = self.router.complete(ModelRequestContract(purpose=ModelCapability.REASONING, model_hint=self.dataset_curator_model_hint, prompt=prompt))
parsed = extract_json_object(response.content)
review = parsed if isinstance(parsed, dict) else {}
except Exception as exc:
review = {"error": str(exc)}
if not review:
review = {"title": "Manual provenance and coverage curation required", "hypothesis": "A source-disjoint, schema-valid subset may improve Guard without benchmark leakage.", "evidence": {"inventory": inventory}, "proposed_operations": [{"operation": "REVIEW_ONLY", "reason": "No model-backed curation response available."}], "expected_capability_effect": "Unknown until reviewed records are validated.", "expected_risks": ["provenance gaps", "benchmark contamination", "weak labels"], "validation_plan": {"required": ["schema", "source provenance", "exact and normalized benchmark overlap", "source-group split", "coverage matrix"]}, "contamination_plan": {"required": ["exact source hash", "normalized text", "repository", "benchmark identifier"]}}
required = ["title", "hypothesis", "evidence", "proposed_operations", "validation_plan", "contamination_plan"]
missing = [field for field in required if review.get(field) in (None, "", {}, [])]
if missing:
raise ValueError("Dataset curator response incomplete: " + ", ".join(missing))
proposal = DatasetCurationProposal.objects.create(title=str(review["title"])[:255], hypothesis=str(review["hypothesis"]), evidence=review["evidence"], proposed_operations=review["proposed_operations"], expected_capability_effect=str(review.get("expected_capability_effect", "")), expected_risks=review.get("expected_risks", []), validation_plan=review["validation_plan"], contamination_plan=review["contamination_plan"], model_evidence=review)
proposal.source_versions.set(versions)
self._artifact(training_project, "DATASET_CURATION_PROPOSAL", proposal.title, {"proposal_id": str(proposal.id), "source_versions": [str(item.id) for item in versions], **review})
return proposal
def materialize_spark_curation(self, training_project: TrainingProject, proposal: DatasetCurationProposal, *, output_directory: str, ssh_alias: str = "spark", strict: bool = False) -> DatasetVersion:
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)
proposal.materialized_version = version
proposal.status = "MATERIALIZED"
proposal.save(update_fields=["materialized_version", "status", "updated_at"])
self._artifact(training_project, "CURATED_DATASET_VERSION", dataset.name, {"dataset_version": str(version.id), **report})
return version
def audit_curated_dataset_with_qwen(self, training_project: TrainingProject, version: DatasetVersion, *, ssh_alias: str = "spark", sample_count: int = 8) -> dict[str, Any]:
samples = SparkGuardDatasetMaterializer(ssh_alias).sample(version.manifest_reference, count=sample_count)
prompt = "DATASET_CURATOR QUALITY AUDIT V0.1. Review these bounded Guard SFT samples. Return JSON only with overall_assessment, schema_issues list, label_risks list, provenance_risks list, leakage_risks list, recommended_operations list, and confidence. Do not invent evidence beyond samples. Do not modify data; recommendations must create a new DatasetVersion. Samples: " + json.dumps(samples, default=str)
review: dict[str, Any] = {"overall_assessment": "MODEL_UNAVAILABLE", "schema_issues": [], "label_risks": [], "provenance_risks": [], "leakage_risks": [], "recommended_operations": [], "confidence": "LOW"}
if self.router is not None and self.dataset_curator_model_hint in self.router.providers:
try:
response = self.router.complete(ModelRequestContract(purpose=ModelCapability.REASONING, model_hint=self.dataset_curator_model_hint, prompt=prompt))
parsed = extract_json_object(response.content)
if isinstance(parsed, dict):
review = parsed
except Exception as exc:
review["error"] = str(exc)
artifact = self._artifact(training_project, "DATASET_QUALITY_AUDIT", f"Qwen audit {version.version}", {"dataset_version": str(version.id), "sample_count": len(samples), "samples": samples, "review": review})
return {"artifact_id": str(artifact.id), "review": review}
def establish_champion(self, training_project: TrainingProject, checkpoint: ModelCheckpoint) -> ModelCheckpoint: def establish_champion(self, training_project: TrainingProject, checkpoint: ModelCheckpoint) -> ModelCheckpoint:
suite = self.validate_benchmark(training_project) suite = self.validate_benchmark(training_project)
if suite.integrity_status != DatasetValidationStatus.VALID: if suite.integrity_status != DatasetValidationStatus.VALID: