Skip to content

Commit bd6c0bb

Browse files
committed
feat(datasets): add dataset materialisation utility
1 parent b20434d commit bd6c0bb

9 files changed

Lines changed: 1420 additions & 13 deletions

File tree

‎cspell/library-words.txt‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
huggingface

‎pyproject.toml‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ requires-python = ">=3.12"
77
dependencies = [
88
"datasets>=4.5.0",
99
"dotenv>=0.9.9",
10+
"huggingface-hub>=1.4.1",
1011
]
1112

1213
[build-system]
@@ -23,6 +24,7 @@ dev = [
2324
"flake8-cognitive-complexity>=0.1.0",
2425
"isort>=7.0.0",
2526
"mypy>=1.19.1",
27+
"notebook>=7.5.3",
2628
"pre-commit>=4.5.1",
2729
"pydoclint>=0.8.3",
2830
"pytest>=9.0.2",
@@ -83,7 +85,7 @@ module = ["tests.*"]
8385
ignore_errors = true
8486

8587
[[tool.mypy.overrides]]
86-
module = ["datasets", "datasets.*"]
88+
module = ["datasets", "datasets.*", "huggingface_hub", "huggingface_hub.*"]
8789
ignore_missing_imports = true
8890

8991
[tool.isort]

‎src/voice/__init__.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
LLMs for stylistic fidelity.
66
"""
77

8-
from voice.datasets import DatasetSpec
8+
from voice.datasets import DatasetSpec, get_dataset
99
from voice.stylometry import get_metrics
1010

11-
__all__: list[str] = ["DatasetSpec", "get_metrics"]
11+
__all__: list[str] = ["DatasetSpec", "get_metrics", "get_dataset"]

‎src/voice/datasets/__init__.py‎

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,15 @@
1-
"""Placeholder docstring."""
1+
"""
2+
Datasets package for VOICE.
23
3-
from voice.datasets.dataset import DatasetSpec
4+
This package contains user-facing abstractions for dataset management,
5+
in particular we expose:
46
5-
__all__: list[str] = [
6-
"DatasetSpec",
7-
]
7+
- `DatasetSpec` - Lightweight, declarative specification of a dataset.
8+
- `VoiceDataset` - Split-aware wrapper around a Hugging Face dataset.
9+
- `get_dataset` - Entrypoint for materialising a dataset from a specification.
10+
"""
11+
12+
from voice.datasets.dataset import DatasetSpec, VoiceDataset
13+
from voice.datasets.get_dataset import get_dataset
14+
15+
__all__: list[str] = ["DatasetSpec", "VoiceDataset", "get_dataset"]

‎src/voice/datasets/_schema.py‎

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,11 @@
1-
"""Placeholder docstring."""
1+
"""
2+
Dataset schema primitives.
3+
4+
This module defines lightweight schema-level abstractions shared across
5+
the dataset loading pipeline.
6+
7+
This module is private by convention and not part of the public API.
8+
"""
29

310
from __future__ import annotations
411

‎src/voice/datasets/_specs.py‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,16 @@
1+
"""
2+
Pre-defined dataset specifications for testing.
3+
4+
All specifications are of type `voice.datasets.DatasetSpec`.
5+
6+
This module is private by convention and not part of the public API.
7+
"""
8+
9+
from voice.datasets import DatasetSpec
10+
from voice.datasets._schema import Split
11+
12+
BUSH_LATEST = DatasetSpec(
13+
repo_id="AccelerateScience/bush-dataset",
14+
revision="main",
15+
splits=(Split.TRAIN, Split.VALIDATION, Split.TEST),
16+
)

‎src/voice/datasets/dataset.py‎

Lines changed: 14 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,17 @@
1-
"""Placeholder docstring."""
1+
"""
2+
Core dataset abstractions for the `voice.datasets` package.
3+
4+
This module defines:
5+
6+
- `DatasetSpec` - A lightweight declarative specification of a dataset
7+
hosted on Hugging Face.
8+
- `_PinnedDatasetSpec` - An internal revision-pinned specification
9+
resolved to a specific git commit SHA.
10+
- `Example` - A representation of a single question-answer example
11+
with provenance metadata.
12+
- `VoiceDataset` — A split-aware wrapper providing structured access
13+
to examples across dataset splits.
14+
"""
215

316
from __future__ import annotations
417

@@ -24,9 +37,6 @@ class DatasetSpec:
2437
"""
2538
Lightweight specification of a Hugging Face dataset.
2639
27-
`_pin()` is used to resolve symbolic revisions to a commit hash.
28-
It is private by convention and not part of the public API
29-
3040
.. attribute :: repo_id
3141
3242
Hugging Face dataset repository id, i.e. namespace/name

‎src/voice/datasets/get_dataset.py‎

Lines changed: 181 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,181 @@
1+
"""
2+
Dataset loading and canonicalisation utilities.
3+
4+
This module implements the dataset materialisation pipeline
5+
used by the public entrypoint `get_dataset`.
6+
7+
Responsibilities include:
8+
9+
- Resolving symbolic revisions (branches/tags) to commit SHAs.
10+
- Loading splits from Hugging Face via `datasets.load_dataset`.
11+
- Canonicalising dataset schemas into a uniform
12+
(system, question, answer) format.
13+
- Validating dataset structure.
14+
- Constructing a split-aware `VoiceDataset` instance.
15+
"""
16+
17+
from __future__ import annotations
18+
19+
from datasets import Dataset, load_dataset
20+
from huggingface_hub import HfApi
21+
from huggingface_hub.utils import HfHubHTTPError
22+
23+
from voice.datasets._schema import Split
24+
from voice.datasets.dataset import (
25+
DatasetSpec,
26+
VoiceDataset,
27+
_PinnedDatasetSpec,
28+
)
29+
30+
# -----------------------------------------------------------------------------
31+
# Utilities (revision resolution, schema canonicalisation .etc)
32+
# -----------------------------------------------------------------------------
33+
34+
35+
def _resolve_revision(repo_id: str, revision: str) -> str:
36+
"""
37+
Resolve a symbolic revision (branch or tag) to a commit SHA.
38+
39+
:param repo_id: Hugging Face dataset repository id
40+
:param revision: Branch name, tag, or commit hash
41+
:return: Commit SHA corresponding to the revision
42+
:raises RuntimeError: If dataset or revision does not exist
43+
"""
44+
api = HfApi()
45+
46+
try:
47+
info = api.dataset_info(repo_id=repo_id, revision=revision)
48+
except HfHubHTTPError as e:
49+
raise RuntimeError(
50+
f"Failed to resolve revision '{revision}' for dataset '{repo_id}'."
51+
) from e
52+
53+
return str(info.sha)
54+
55+
56+
def _extract_chat_style(
57+
column: dict[str, list[dict[str, str]]],
58+
) -> dict[str, str]:
59+
"""
60+
Extract (system, question, answer) dict from chat style example.
61+
62+
:param column: column containing chat-style messages
63+
:return: canonical dict
64+
"""
65+
messages = column["messages"]
66+
67+
system = ""
68+
question = ""
69+
answer = ""
70+
71+
for msg in messages:
72+
role = msg.get("role")
73+
content = msg.get("content", "")
74+
75+
if role == "system":
76+
system = content
77+
elif role == "user":
78+
question = content
79+
elif role == "assistant":
80+
answer = content
81+
82+
return {
83+
"system": system,
84+
"question": question,
85+
"answer": answer,
86+
}
87+
88+
89+
def _canonicalise_dataset(ds: Dataset) -> Dataset:
90+
"""
91+
Convert a HF dataset into canonical (system, question, answer) format.
92+
93+
Currently supported schemas:
94+
- Already canonical: columns {system, question, answer}
95+
- Chat style: column 'messages' with list[{role, content}]
96+
97+
:param ds: Raw Hugging Face dataset
98+
:return: Canonicalised QA dataset
99+
:raises NotImplementedError: If schema not recognised
100+
"""
101+
cols = set(ds.column_names)
102+
103+
if {"system", "question", "answer"}.issubset(cols):
104+
return ds
105+
106+
if "messages" in cols and len(cols) == 1:
107+
return ds.map(
108+
_extract_chat_style,
109+
remove_columns=ds.column_names,
110+
)
111+
112+
raise NotImplementedError(
113+
f"Unsupported dataset schema. Columns found: {sorted(cols)}"
114+
)
115+
116+
117+
def _validate_canonical(ds: Dataset) -> None:
118+
"""
119+
Ensure dataset contains required canonical columns.
120+
121+
:param ds: Canonicalised dataset
122+
:return: None
123+
:raises ValueError: If required columns missing
124+
"""
125+
required = {"system", "question", "answer"}
126+
missing = required - set(ds.column_names)
127+
if not required.issubset(ds.column_names):
128+
raise ValueError(f"Missing required columns: {sorted(missing)}")
129+
130+
131+
# -----------------------------------------------------------------------------
132+
# Public API
133+
# -----------------------------------------------------------------------------
134+
135+
136+
def get_dataset(spec: DatasetSpec) -> VoiceDataset:
137+
"""
138+
Materialise a Hugging Face dataset as a VoiceDataset.
139+
140+
Steps:
141+
1. Resolve DatasetSpec to _PinnedDatasetSpec.
142+
2. Load requested splits from Hugging Face.
143+
3. Canonicalise each split into (system, question, answer).
144+
4. Validate canonical structure.
145+
5. Construct VoiceDataset.
146+
147+
:param spec: Dataset specification
148+
:return: Split-aware VoiceDataset
149+
:raises RuntimeError: If a split cannot be loaded
150+
"""
151+
# Resolve dataset revision
152+
pinned: _PinnedDatasetSpec = spec._pin(_resolve_revision)
153+
154+
# Load dataset splits
155+
datasets: dict[Split, Dataset] = {}
156+
157+
for split in pinned.splits:
158+
try:
159+
ds = load_dataset(
160+
pinned.repo_id,
161+
split=split.value,
162+
revision=pinned.revision,
163+
)
164+
except Exception as e:
165+
raise RuntimeError(
166+
f"Failed to load split '{split.value}' "
167+
f"from {pinned.repo_id}@{pinned.revision}"
168+
) from e
169+
170+
# Canonicalise and validate dataset
171+
ds = _canonicalise_dataset(ds)
172+
173+
_validate_canonical(ds)
174+
175+
datasets[split] = ds
176+
177+
# Build VoiceDataset object
178+
return VoiceDataset(
179+
datasets=datasets,
180+
spec=pinned,
181+
)

0 commit comments

Comments
 (0)