dataflux-core 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- dataflux/__init__.py +4 -0
- dataflux/__version__.py +1 -0
- dataflux/cache.py +122 -0
- dataflux/config.py +18 -0
- dataflux/display/__init__.py +0 -0
- dataflux/display/console.py +4 -0
- dataflux/display/info.py +150 -0
- dataflux/display/search.py +122 -0
- dataflux/display/theme.py +32 -0
- dataflux/display/utils.py +138 -0
- dataflux/exceptions.py +150 -0
- dataflux/export.py +62 -0
- dataflux/flux.py +55 -0
- dataflux/models/__init__.py +0 -0
- dataflux/models/dataset.py +15 -0
- dataflux/models/search_result.py +10 -0
- dataflux/providers/__init__.py +0 -0
- dataflux/providers/base.py +23 -0
- dataflux/providers/huggingface.py +91 -0
- dataflux/providers/kaggle.py +85 -0
- dataflux/providers/seaborn.py +66 -0
- dataflux/providers/sklearn.py +67 -0
- dataflux/providers/statsmodels.py +66 -0
- dataflux/providers/torchvision.py +80 -0
- dataflux/providers/uci.py +78 -0
- dataflux/providers/vega_datasets.py +64 -0
- dataflux/providers/worldbank.py +73 -0
- dataflux/registry.py +38 -0
- dataflux/resolver.py +30 -0
- dataflux/resources/kaggle_index.json +365901 -0
- dataflux/resources/seaborn_dataset.py +24 -0
- dataflux/resources/sklearn_dataset.py +17 -0
- dataflux/resources/statsmodels_dataset.py +30 -0
- dataflux/resources/torch_vision_dataset.py +89 -0
- dataflux/tests/conftest.py +33 -0
- dataflux/tests/display/test_info.py +159 -0
- dataflux/tests/display/test_search.py +83 -0
- dataflux/tests/models/test_dataset_info.py +46 -0
- dataflux/tests/models/test_search_result.py +28 -0
- dataflux/tests/providers/test_hugging.py +93 -0
- dataflux/tests/providers/test_kaggle.py +94 -0
- dataflux/tests/providers/test_seaborn.py +88 -0
- dataflux/tests/providers/test_sklearn.py +92 -0
- dataflux/tests/providers/test_statsmodels.py +89 -0
- dataflux/tests/providers/test_torchvision.py +90 -0
- dataflux/tests/providers/test_uci.py +92 -0
- dataflux/tests/providers/test_vega.py +90 -0
- dataflux/tests/providers/test_worldbank.py +111 -0
- dataflux/tests/resources/test_resources.py +80 -0
- dataflux/tests/test_exceptions.py +20 -0
- dataflux/tests/test_export.py +146 -0
- dataflux/tests/test_flux.py +40 -0
- dataflux/tests/test_models.py +30 -0
- dataflux/tests/test_registry.py +149 -0
- dataflux/tests/test_resolver.py +77 -0
- dataflux/tests/utils/test_cache.py +173 -0
- dataflux/tests/utils/test_download.py +135 -0
- dataflux/tests/utils/test_filesystem.py +128 -0
- dataflux/tests/utils/test_fingerprint.py +114 -0
- dataflux/tests/utils/test_search_utils.py +152 -0
- dataflux/utils/__init__.py +0 -0
- dataflux/utils/download.py +110 -0
- dataflux/utils/filesystem.py +89 -0
- dataflux/utils/fingerprint.py +155 -0
- dataflux/utils/search.py +66 -0
- dataflux_core-0.1.0.dist-info/METADATA +400 -0
- dataflux_core-0.1.0.dist-info/RECORD +69 -0
- dataflux_core-0.1.0.dist-info/WHEEL +4 -0
- dataflux_core-0.1.0.dist-info/licenses/LICENSE +21 -0
dataflux/exceptions.py
ADDED
|
@@ -0,0 +1,150 @@
|
|
|
1
|
+
class DataFluxError(Exception):
|
|
2
|
+
"""Base exception for all DataFlux errors."""
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
# ============================================================================
|
|
6
|
+
# Provider Errors
|
|
7
|
+
# ============================================================================
|
|
8
|
+
|
|
9
|
+
class InvalidProviderError(DataFluxError):
|
|
10
|
+
"""Raised when an invalid provider is supplied."""
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class ProviderAlreadyRegisteredError(DataFluxError):
|
|
14
|
+
"""Raised when a provider with the same name is already registered."""
|
|
15
|
+
|
|
16
|
+
def __init__(self, provider_name: str):
|
|
17
|
+
super().__init__(
|
|
18
|
+
f"Provider '{provider_name}' is already registered."
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class ProviderNotFoundError(DataFluxError):
|
|
23
|
+
"""Raised when a requested provider does not exist."""
|
|
24
|
+
|
|
25
|
+
def __init__(self, provider_name: str):
|
|
26
|
+
super().__init__(
|
|
27
|
+
f"Provider '{provider_name}' was not found."
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
# ============================================================================
|
|
32
|
+
# Dataset Errors
|
|
33
|
+
# ============================================================================
|
|
34
|
+
|
|
35
|
+
class DatasetNotFoundError(DataFluxError):
|
|
36
|
+
"""Raised when no dataset matches the query."""
|
|
37
|
+
|
|
38
|
+
def __init__(self, dataset: str | int):
|
|
39
|
+
super().__init__(
|
|
40
|
+
f"Dataset '{dataset}' was not found."
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class InvalidDatasetError(DataFluxError):
|
|
45
|
+
"""Raised when a dataset identifier is invalid."""
|
|
46
|
+
|
|
47
|
+
def __init__(self, dataset: str | int):
|
|
48
|
+
super().__init__(
|
|
49
|
+
f"'{dataset}' is not a valid dataset identifier."
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class DatasetLoadError(DataFluxError):
|
|
54
|
+
"""Raised when a dataset cannot be loaded."""
|
|
55
|
+
|
|
56
|
+
def __init__(self, dataset: str | int):
|
|
57
|
+
super().__init__(
|
|
58
|
+
f"Failed to load dataset '{dataset}'."
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
# ============================================================================
|
|
63
|
+
# Search Errors
|
|
64
|
+
# ============================================================================
|
|
65
|
+
|
|
66
|
+
class SearchError(DataFluxError):
|
|
67
|
+
"""Raised when a search operation fails."""
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
class EmptySearchQueryError(SearchError):
|
|
71
|
+
"""Raised when an empty search query is provided."""
|
|
72
|
+
|
|
73
|
+
def __init__(self):
|
|
74
|
+
super().__init__(
|
|
75
|
+
"Search query cannot be empty."
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
# ============================================================================
|
|
80
|
+
# Pull Errors
|
|
81
|
+
# ============================================================================
|
|
82
|
+
|
|
83
|
+
class PullError(DataFluxError):
|
|
84
|
+
"""Raised when pulling a dataset fails."""
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
# ============================================================================
|
|
88
|
+
# Download Errors
|
|
89
|
+
# ============================================================================
|
|
90
|
+
|
|
91
|
+
class DownloadError(DataFluxError):
|
|
92
|
+
"""Raised when a download operation fails."""
|
|
93
|
+
|
|
94
|
+
def __init__(self, url: str):
|
|
95
|
+
super().__init__(
|
|
96
|
+
f"Failed to download resource from '{url}'."
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
# ============================================================================
|
|
101
|
+
# Cache Errors
|
|
102
|
+
# ============================================================================
|
|
103
|
+
|
|
104
|
+
class CacheError(DataFluxError):
|
|
105
|
+
"""Raised when a cache operation fails."""
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
# ============================================================================
|
|
109
|
+
# Filesystem Errors
|
|
110
|
+
# ============================================================================
|
|
111
|
+
|
|
112
|
+
class FileSystemError(DataFluxError):
|
|
113
|
+
"""Raised when a filesystem operation fails."""
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
# ============================================================================
|
|
117
|
+
# Export Errors
|
|
118
|
+
# ============================================================================
|
|
119
|
+
|
|
120
|
+
class ExportError(DataFluxError):
|
|
121
|
+
"""Raised when exporting a dataset fails."""
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
class UnsupportedExportFormatError(ExportError):
|
|
125
|
+
"""Raised when an unsupported export format is requested."""
|
|
126
|
+
|
|
127
|
+
def __init__(self, extension: str):
|
|
128
|
+
super().__init__(
|
|
129
|
+
f"Export format '{extension}' is not supported."
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
class ExportFileExistsError(ExportError):
|
|
134
|
+
"""Raised when the export file already exists."""
|
|
135
|
+
|
|
136
|
+
def __init__(self, path: str):
|
|
137
|
+
super().__init__(
|
|
138
|
+
f"'{path}' already exists. "
|
|
139
|
+
"Use overwrite=True to replace it."
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
class InvalidExportDataError(ExportError):
|
|
144
|
+
"""Raised when the object being exported is invalid."""
|
|
145
|
+
|
|
146
|
+
def __init__(
|
|
147
|
+
self,
|
|
148
|
+
message: str = "export() expects a Polars DataFrame.",
|
|
149
|
+
):
|
|
150
|
+
super().__init__(message)
|
dataflux/export.py
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
import polars as pl
|
|
4
|
+
from dataflux.exceptions import ExportFileExistsError,UnsupportedExportFormatError,InvalidExportDataError
|
|
5
|
+
SUPPORTED_EXPORTS = {
|
|
6
|
+
".csv",
|
|
7
|
+
".parquet",
|
|
8
|
+
".json",
|
|
9
|
+
".ndjson",
|
|
10
|
+
".ipc",
|
|
11
|
+
".feather",
|
|
12
|
+
".arrow",
|
|
13
|
+
}
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def export(df:pl.DataFrame,path:str | Path,*,overwrite:bool = False,mkdir:bool=True,)->Path:
|
|
17
|
+
if not isinstance(df,pl.DataFrame):
|
|
18
|
+
raise InvalidExportDataError(
|
|
19
|
+
"export() expect a Polars DataFrame"
|
|
20
|
+
)
|
|
21
|
+
path = Path(path)
|
|
22
|
+
|
|
23
|
+
if mkdir:
|
|
24
|
+
path.parent.mkdir(
|
|
25
|
+
parents=True,
|
|
26
|
+
exist_ok=True,
|
|
27
|
+
)
|
|
28
|
+
if path.exists() and not overwrite:
|
|
29
|
+
raise ExportFileExistsError(
|
|
30
|
+
f"'{path}' already exists."
|
|
31
|
+
)
|
|
32
|
+
|
|
33
|
+
suffix = path.suffix.lower()
|
|
34
|
+
|
|
35
|
+
if suffix not in SUPPORTED_EXPORTS:
|
|
36
|
+
raise UnsupportedExportFormatError(
|
|
37
|
+
f"Unsupported export format '{suffix}'. "
|
|
38
|
+
f"Supported formats:{', '.join(sorted(SUPPORTED_EXPORTS))} "
|
|
39
|
+
)
|
|
40
|
+
match suffix:
|
|
41
|
+
case ".csv":
|
|
42
|
+
df.write_csv(path)
|
|
43
|
+
|
|
44
|
+
case ".parquet":
|
|
45
|
+
df.write_parquet(path)
|
|
46
|
+
|
|
47
|
+
case ".json":
|
|
48
|
+
df.write_json(path)
|
|
49
|
+
|
|
50
|
+
case ".ndjson":
|
|
51
|
+
df.write_ndjson(path)
|
|
52
|
+
|
|
53
|
+
case ".ipc":
|
|
54
|
+
df.write_ipc(path)
|
|
55
|
+
|
|
56
|
+
case ".feather":
|
|
57
|
+
df.write_ipc(path)
|
|
58
|
+
|
|
59
|
+
case ".arrow":
|
|
60
|
+
df.write_ipc(path)
|
|
61
|
+
|
|
62
|
+
return path
|
dataflux/flux.py
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
from dataflux.models.search_result import SearchResult
|
|
2
|
+
from dataflux.registry import ProviderRegistry
|
|
3
|
+
from dataflux.resolver import Resolver
|
|
4
|
+
from dataflux.providers.uci import UCIProvider
|
|
5
|
+
from dataflux.models.dataset import DatasetInfo
|
|
6
|
+
from dataflux.providers.sklearn import SKLearnProvider
|
|
7
|
+
from dataflux.providers.kaggle import KaggleProvider
|
|
8
|
+
from dataflux.providers.huggingface import HuggingFaceProvider
|
|
9
|
+
from dataflux.providers.torchvision import TorchVisionProvider
|
|
10
|
+
from dataflux.providers.seaborn import SeabornProvider
|
|
11
|
+
from dataflux.providers.statsmodels import StatsModelsProvider
|
|
12
|
+
from dataflux.providers.vega_datasets import VegaDatasetsProvider
|
|
13
|
+
from dataflux.providers.worldbank import WorldBankProvider
|
|
14
|
+
from dataflux.display.search import display_search
|
|
15
|
+
from dataflux.display.info import display_info
|
|
16
|
+
from dataflux.export import export
|
|
17
|
+
import polars as pl
|
|
18
|
+
|
|
19
|
+
class Flux:
|
|
20
|
+
|
|
21
|
+
def __init__(self):
|
|
22
|
+
self._registry = ProviderRegistry()
|
|
23
|
+
self._resolver = Resolver(self._registry)
|
|
24
|
+
|
|
25
|
+
self._registry.register(UCIProvider())
|
|
26
|
+
self._registry.register(SKLearnProvider())
|
|
27
|
+
self._registry.register(KaggleProvider())
|
|
28
|
+
self._registry.register(HuggingFaceProvider())
|
|
29
|
+
self._registry.register(TorchVisionProvider())
|
|
30
|
+
self._registry.register(SeabornProvider())
|
|
31
|
+
self._registry.register(StatsModelsProvider())
|
|
32
|
+
self._registry.register(VegaDatasetsProvider())
|
|
33
|
+
self._registry.register(WorldBankProvider())
|
|
34
|
+
|
|
35
|
+
def search(self, query: str,*,raw:bool=False,display:bool=True,limit:int | None=10) -> list[SearchResult]:
|
|
36
|
+
results = self._resolver.resolve_dataset(query)
|
|
37
|
+
if display and not raw:
|
|
38
|
+
display_search(results,limit)
|
|
39
|
+
return results
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def info(self,result:SearchResult,*,raw:bool=False,display:bool=True)-> DatasetInfo:
|
|
43
|
+
provider = self._resolver.resolve_provider(result.provider)
|
|
44
|
+
if display and not raw:
|
|
45
|
+
display_info(provider.info(result.id))
|
|
46
|
+
return provider.info(result.id)
|
|
47
|
+
|
|
48
|
+
def pull(self,result:SearchResult)->pl.DataFrame:
|
|
49
|
+
provider = self._resolver.resolve_provider(result.provider)
|
|
50
|
+
return provider.pull(result.id)
|
|
51
|
+
|
|
52
|
+
def export(self,df:pl.DataFrame,path:str,**kwargs,):
|
|
53
|
+
return export(df,path,**kwargs)
|
|
54
|
+
|
|
55
|
+
flux = Flux()
|
|
File without changes
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from typing import Any
|
|
3
|
+
@dataclass(slots=True)
|
|
4
|
+
class DatasetInfo:
|
|
5
|
+
id: str | int
|
|
6
|
+
name: str
|
|
7
|
+
description: str | None
|
|
8
|
+
instances: int | None
|
|
9
|
+
features: int | None
|
|
10
|
+
tasks: list[str]
|
|
11
|
+
target: list[str]
|
|
12
|
+
has_missing_values: bool | None
|
|
13
|
+
provider: str
|
|
14
|
+
url: str | None
|
|
15
|
+
extra: dict[str, Any]
|
|
File without changes
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class BaseProvider(ABC):
|
|
7
|
+
|
|
8
|
+
@property
|
|
9
|
+
@abstractmethod
|
|
10
|
+
def name(self) -> str:
|
|
11
|
+
"""Unique provider name."""
|
|
12
|
+
|
|
13
|
+
@abstractmethod
|
|
14
|
+
def search(self, query: str) -> Any:
|
|
15
|
+
"""Search datasets."""
|
|
16
|
+
|
|
17
|
+
@abstractmethod
|
|
18
|
+
def pull(self, dataset: str) -> Path:
|
|
19
|
+
"""Download the requested dataset and return its local path."""
|
|
20
|
+
|
|
21
|
+
@abstractmethod
|
|
22
|
+
def info(self, dataset: str) -> Any:
|
|
23
|
+
"""Return dataset metadata."""
|
|
@@ -0,0 +1,91 @@
|
|
|
1
|
+
from dataflux.providers.base import BaseProvider
|
|
2
|
+
from dataflux.models.search_result import SearchResult
|
|
3
|
+
from dataflux.models.dataset import DatasetInfo
|
|
4
|
+
from huggingface_hub import HfApi
|
|
5
|
+
from datasets import load_dataset
|
|
6
|
+
from dataflux.utils.search import search_score
|
|
7
|
+
import polars as pl
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class HuggingFaceProvider(BaseProvider):
|
|
11
|
+
|
|
12
|
+
def __init__(self):
|
|
13
|
+
self.api = HfApi()
|
|
14
|
+
|
|
15
|
+
@property
|
|
16
|
+
def name(self) -> str:
|
|
17
|
+
return "huggingface"
|
|
18
|
+
|
|
19
|
+
def search(self, query: str) -> list[SearchResult]:
|
|
20
|
+
matches:list[tuple[int,SearchResult]] =[]
|
|
21
|
+
|
|
22
|
+
datasets = self.api.list_datasets(
|
|
23
|
+
search=query,
|
|
24
|
+
limit=20
|
|
25
|
+
)
|
|
26
|
+
for dataset in datasets:
|
|
27
|
+
score = search_score(
|
|
28
|
+
query,
|
|
29
|
+
dataset.id
|
|
30
|
+
)
|
|
31
|
+
if score == -1:
|
|
32
|
+
continue
|
|
33
|
+
matches.append(
|
|
34
|
+
( score,
|
|
35
|
+
SearchResult(
|
|
36
|
+
id=dataset.id,
|
|
37
|
+
name=dataset.id.split("/")[-1].replace("-"," ").title(),
|
|
38
|
+
provider=self.name,
|
|
39
|
+
description=None,
|
|
40
|
+
relevance=score,
|
|
41
|
+
)
|
|
42
|
+
)
|
|
43
|
+
)
|
|
44
|
+
matches.sort(key=lambda x:x[0],reverse=True)
|
|
45
|
+
return [result for _, result in matches]
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def info(self, dataset_id: str) -> DatasetInfo:
|
|
49
|
+
dataset = self.api.dataset_info(dataset_id)
|
|
50
|
+
return DatasetInfo(
|
|
51
|
+
id=dataset_id,
|
|
52
|
+
name=dataset.id.split("/")[-1].replace("-"," ").title(),
|
|
53
|
+
description=dataset.description,
|
|
54
|
+
instances=(
|
|
55
|
+
dataset.card_data["dataset_info"]["splits"][0]["num_examples"]
|
|
56
|
+
if dataset.card_data and "dataset_info" in dataset.card_data
|
|
57
|
+
else None
|
|
58
|
+
),
|
|
59
|
+
features=(
|
|
60
|
+
len(dataset.card_data["dataset_info"]["features"])
|
|
61
|
+
if dataset.card_data and "dataset_info" in dataset.card_data
|
|
62
|
+
else None
|
|
63
|
+
),
|
|
64
|
+
tasks=(
|
|
65
|
+
dataset.card_data.get("task_categories")
|
|
66
|
+
if dataset.card_data
|
|
67
|
+
else None
|
|
68
|
+
),
|
|
69
|
+
target=None,
|
|
70
|
+
has_missing_values=None,
|
|
71
|
+
provider=self.name,
|
|
72
|
+
url=f"https://huggingface.co/datasets/{dataset_id}",
|
|
73
|
+
extra={
|
|
74
|
+
"author":dataset.author,
|
|
75
|
+
"tags":dataset.tags,
|
|
76
|
+
"downloads":dataset.downloads,
|
|
77
|
+
"likes":dataset.likes,
|
|
78
|
+
"created_at":str(dataset.created_at),
|
|
79
|
+
"last_modified": str(dataset.last_modified),
|
|
80
|
+
}
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def pull(self, dataset_id: str):
|
|
86
|
+
dataset = load_dataset(dataset_id)
|
|
87
|
+
if hasattr(dataset,"keys"):
|
|
88
|
+
dataset = dataset[list(dataset.keys())[0]]
|
|
89
|
+
|
|
90
|
+
return dataset.to_polars()
|
|
91
|
+
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
import json
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
import kagglehub
|
|
5
|
+
import polars as pl
|
|
6
|
+
from kagglehub import KaggleDatasetAdapter
|
|
7
|
+
|
|
8
|
+
from dataflux.models.dataset import DatasetInfo
|
|
9
|
+
from dataflux.models.search_result import SearchResult
|
|
10
|
+
from dataflux.providers.base import BaseProvider
|
|
11
|
+
from dataflux.utils.search import search_score
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class KaggleProvider(BaseProvider):
|
|
15
|
+
|
|
16
|
+
@property
|
|
17
|
+
def name(self) -> str:
|
|
18
|
+
return "kaggle"
|
|
19
|
+
|
|
20
|
+
def _load_index(self) -> list[dict]:
|
|
21
|
+
with open(
|
|
22
|
+
Path(__file__).parent.parent / "resources" / "kaggle_index.json",
|
|
23
|
+
"r",
|
|
24
|
+
encoding="utf-8",
|
|
25
|
+
) as f:
|
|
26
|
+
return json.load(f)
|
|
27
|
+
|
|
28
|
+
def search(self, query: str) -> list[SearchResult]:
|
|
29
|
+
matches: list[tuple[int, SearchResult]] = []
|
|
30
|
+
|
|
31
|
+
for dataset in self._load_index():
|
|
32
|
+
score = search_score(
|
|
33
|
+
query,
|
|
34
|
+
dataset["title"],
|
|
35
|
+
dataset["subtitle"],
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
if score == -1:
|
|
39
|
+
continue
|
|
40
|
+
|
|
41
|
+
matches.append(
|
|
42
|
+
(
|
|
43
|
+
score,
|
|
44
|
+
SearchResult(
|
|
45
|
+
id=f"{dataset['owner_slug']}/{dataset['dataset_slug']}",
|
|
46
|
+
name=dataset["title"],
|
|
47
|
+
provider=self.name,
|
|
48
|
+
description=dataset["subtitle"],
|
|
49
|
+
relevance=score
|
|
50
|
+
),
|
|
51
|
+
)
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
matches.sort(key=lambda x: x[0], reverse=True)
|
|
55
|
+
|
|
56
|
+
return [result for _, result in matches]
|
|
57
|
+
|
|
58
|
+
def info(self, dataset_id: str) -> DatasetInfo:
|
|
59
|
+
for dataset in self._load_index():
|
|
60
|
+
if f"{dataset['owner_slug']}/{dataset['dataset_slug']}" == dataset_id:
|
|
61
|
+
return DatasetInfo(
|
|
62
|
+
id=dataset_id,
|
|
63
|
+
name=dataset["title"],
|
|
64
|
+
description=dataset["subtitle"],
|
|
65
|
+
instances=None,
|
|
66
|
+
features=None,
|
|
67
|
+
tasks=None,
|
|
68
|
+
target=None,
|
|
69
|
+
has_missing_values=None,
|
|
70
|
+
provider=self.name,
|
|
71
|
+
url=f"https://www.kaggle.com/datasets/{dataset_id}",
|
|
72
|
+
extra=dataset,
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
raise ValueError(f"Dataset '{dataset_id}' not found.")
|
|
76
|
+
|
|
77
|
+
def pull(self, dataset_id: str) -> pl.DataFrame:
|
|
78
|
+
dataset_path = Path(kagglehub.dataset_download(dataset_id))
|
|
79
|
+
csv_file = next(dataset_path.rglob("*.csv"))
|
|
80
|
+
|
|
81
|
+
return kagglehub.dataset_load(
|
|
82
|
+
KaggleDatasetAdapter.POLARS,
|
|
83
|
+
handle=dataset_id,
|
|
84
|
+
path=csv_file.name,
|
|
85
|
+
).collect()
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
from dataflux.providers.base import BaseProvider
|
|
2
|
+
from dataflux.models.search_result import SearchResult
|
|
3
|
+
from dataflux.models.dataset import DatasetInfo
|
|
4
|
+
from dataflux.resources.seaborn_dataset import SEABORN_DATASETS
|
|
5
|
+
from dataflux.utils.search import search_score
|
|
6
|
+
import seaborn as sns
|
|
7
|
+
import polars as pl
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class SeabornProvider(BaseProvider):
|
|
11
|
+
|
|
12
|
+
@property
|
|
13
|
+
def name(self) -> str:
|
|
14
|
+
return "seaborn"
|
|
15
|
+
|
|
16
|
+
def _load_df(self, dataset_id: str):
|
|
17
|
+
return sns.load_dataset(dataset_id)
|
|
18
|
+
|
|
19
|
+
def search(self, query: str) -> list[SearchResult]:
|
|
20
|
+
matches: list[tuple[int, SearchResult]] = []
|
|
21
|
+
|
|
22
|
+
for dataset in SEABORN_DATASETS:
|
|
23
|
+
score = search_score(query, dataset)
|
|
24
|
+
|
|
25
|
+
if score == -1:
|
|
26
|
+
continue
|
|
27
|
+
|
|
28
|
+
matches.append(
|
|
29
|
+
(
|
|
30
|
+
score,
|
|
31
|
+
SearchResult(
|
|
32
|
+
id=dataset,
|
|
33
|
+
name=dataset.replace("_", " ").title(),
|
|
34
|
+
provider=self.name,
|
|
35
|
+
description=None,
|
|
36
|
+
relevance=score
|
|
37
|
+
),
|
|
38
|
+
)
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
matches.sort(key=lambda x: x[0], reverse=True)
|
|
42
|
+
|
|
43
|
+
return [result for _, result in matches]
|
|
44
|
+
|
|
45
|
+
def info(self, dataset_id: str) -> DatasetInfo:
|
|
46
|
+
dataset = self._load_df(dataset_id)
|
|
47
|
+
|
|
48
|
+
return DatasetInfo(
|
|
49
|
+
id=dataset_id,
|
|
50
|
+
name=dataset_id.replace("_", " ").title(),
|
|
51
|
+
description=None,
|
|
52
|
+
instances=dataset.shape[0],
|
|
53
|
+
features=dataset.shape[1],
|
|
54
|
+
tasks=None,
|
|
55
|
+
target=None,
|
|
56
|
+
has_missing_values=dataset.isnull().values.any(),
|
|
57
|
+
provider=self.name,
|
|
58
|
+
url=None,
|
|
59
|
+
extra={
|
|
60
|
+
"columns": list(dataset.columns),
|
|
61
|
+
"dtypes": {k: str(v) for k, v in dataset.dtypes.items()},
|
|
62
|
+
},
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
def pull(self, dataset_id: str) -> pl.DataFrame:
|
|
66
|
+
return pl.from_pandas(self._load_df(dataset_id))
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
from dataflux.providers.base import BaseProvider
|
|
2
|
+
from dataflux.models.search_result import SearchResult
|
|
3
|
+
from dataflux.models.dataset import DatasetInfo
|
|
4
|
+
from dataflux.resources.sklearn_dataset import SKLEARN_DATASETS
|
|
5
|
+
from dataflux.utils.search import search_score
|
|
6
|
+
import polars as pl
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class SKLearnProvider(BaseProvider):
|
|
10
|
+
|
|
11
|
+
@property
|
|
12
|
+
def name(self) -> str:
|
|
13
|
+
return "sklearn"
|
|
14
|
+
|
|
15
|
+
def _load_dataset(self, dataset_id: str, *, as_frame: bool = False):
|
|
16
|
+
loader = SKLEARN_DATASETS[dataset_id]
|
|
17
|
+
return loader(as_frame=as_frame) if as_frame else loader()
|
|
18
|
+
|
|
19
|
+
def search(self, query: str) -> list[SearchResult]:
|
|
20
|
+
matches: list[tuple[int, SearchResult]] = []
|
|
21
|
+
|
|
22
|
+
for dataset in SKLEARN_DATASETS:
|
|
23
|
+
score = search_score(query, dataset)
|
|
24
|
+
|
|
25
|
+
if score == -1:
|
|
26
|
+
continue
|
|
27
|
+
|
|
28
|
+
matches.append(
|
|
29
|
+
(
|
|
30
|
+
score,
|
|
31
|
+
SearchResult(
|
|
32
|
+
id=dataset,
|
|
33
|
+
name=dataset.replace("_", " ").title(),
|
|
34
|
+
provider=self.name,
|
|
35
|
+
relevance=score
|
|
36
|
+
),
|
|
37
|
+
)
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
matches.sort(key=lambda x: x[0], reverse=True)
|
|
41
|
+
|
|
42
|
+
return [result for _, result in matches]
|
|
43
|
+
|
|
44
|
+
def info(self, dataset_id: str) -> DatasetInfo:
|
|
45
|
+
dataset = self._load_dataset(dataset_id)
|
|
46
|
+
|
|
47
|
+
return DatasetInfo(
|
|
48
|
+
id=dataset_id,
|
|
49
|
+
name=dataset_id.replace("_", " ").title(),
|
|
50
|
+
description=dataset.DESCR,
|
|
51
|
+
instances=dataset.data.shape[0],
|
|
52
|
+
features=dataset.data.shape[1],
|
|
53
|
+
tasks=None,
|
|
54
|
+
target=list(dataset.target_names),
|
|
55
|
+
has_missing_values=False,
|
|
56
|
+
provider=self.name,
|
|
57
|
+
url=None,
|
|
58
|
+
extra={
|
|
59
|
+
"feature_names": dataset.feature_names,
|
|
60
|
+
"filename": dataset.filename,
|
|
61
|
+
"data_module": dataset.data_module,
|
|
62
|
+
},
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
def pull(self, dataset_id: str) -> pl.DataFrame:
|
|
66
|
+
dataset = self._load_dataset(dataset_id, as_frame=True)
|
|
67
|
+
return pl.from_pandas(dataset.frame)
|