ark-shopify commited on
Commit
05fc91d
·
unverified ·
1 Parent(s): 48e5c2e

refactor: Orchestrator - Added the HF launcher and storage provider to the dedicated HF repo

Browse files
huggingface_overlay/cloud_pipelines_backend/launchers/huggingface_launchers.py ADDED
@@ -0,0 +1,481 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import copy
2
+ import dataclasses
3
+ import datetime
4
+ import logging
5
+ import pathlib
6
+ import typing
7
+ from typing import Any, Optional
8
+
9
+ import huggingface_hub
10
+
11
+ from cloud_pipelines.orchestration.launchers import naming_utils
12
+ from ..storage_providers import huggingface_repo_storage
13
+ from .. import component_structures as structures
14
+ from . import container_component_utils
15
+ from . import interfaces
16
+
17
+
18
+ _logger = logging.getLogger(__name__)
19
+
20
+ _MAX_INPUT_VALUE_SIZE = 10000
21
+
22
+ _CONTAINER_FILE_NAME = "data"
23
+
24
+
25
+ class HuggingFaceJobsContainerLauncher(
26
+ interfaces.ContainerTaskLauncher["LaunchedHuggingFaceJobContainer"]
27
+ ):
28
+ """Launcher that uses HuggingFace Jobs installed locally"""
29
+
30
+ def __init__(
31
+ self,
32
+ *,
33
+ client: Optional[huggingface_hub.HfApi] = None,
34
+ namespace: Optional[str] = None,
35
+ hf_token: Optional[str] = None,
36
+ hf_job_token: Optional[str] = None,
37
+ job_timeout: Optional[int | float | str] = None,
38
+ ):
39
+ # The HF Jobs that we launch need token to write the output artifacts and logs
40
+ hf_token = hf_token or huggingface_hub.get_token()
41
+ hf_job_token = hf_job_token or hf_token
42
+ self._api_client = client or huggingface_hub.HfApi(token=hf_token)
43
+ self._namespace: str = namespace or self._api_client.whoami()["name"]
44
+ self._storage_provider = (
45
+ huggingface_repo_storage.HuggingFaceRepoStorageProvider(
46
+ client=self._api_client
47
+ )
48
+ )
49
+ self._job_timeout = job_timeout
50
+ self._hf_job_token = hf_job_token
51
+
52
+ def self_check(self):
53
+ _ = self._api_client.list_jobs(namespace=self._namespace)
54
+
55
+ def launch_container_task(
56
+ self,
57
+ *,
58
+ component_spec: structures.ComponentSpec,
59
+ # Input arguments may be updated with new downloaded values and new URIs of uploaded values.
60
+ input_arguments: dict[str, interfaces.InputArgument],
61
+ output_uris: dict[str, str],
62
+ log_uri: str,
63
+ annotations: dict[str, Any] | None = None,
64
+ ) -> "LaunchedHuggingFaceJobContainer":
65
+ if not isinstance(
66
+ component_spec.implementation, structures.ContainerImplementation
67
+ ):
68
+ raise TypeError(
69
+ f"Container launchers only support container implementations. Got {component_spec=}"
70
+ )
71
+ container_spec = component_spec.implementation.container
72
+
73
+ # TODO: Validate the input/output URIs.
74
+ container_inputs_root = pathlib.PurePosixPath("/tmp/component/inputs")
75
+ container_outputs_root = pathlib.PurePosixPath("/tmp/component/outputs")
76
+
77
+ # download_input_uris: dict[huggingface_repo_storage.HuggingFaceRepoUri, str] = {}
78
+ # upload_output_uris: dict[str, huggingface_repo_storage.HuggingFaceRepoUri] = {}
79
+ download_input_uris: dict[str, str] = {}
80
+ upload_output_uris: dict[str, str] = {}
81
+ # TODO: Derive common prefix for the upload_output_uris (also log_uri) and upload everything at once
82
+
83
+ # Callbacks for the command-line resolving
84
+ # Their main purpose is to return input/output path or value.
85
+ # They add volumes and volume mounts when needed.
86
+ # They also upload/download artifact data when needed.
87
+ def get_input_value(input_name: str) -> str:
88
+ input_argument = input_arguments[input_name]
89
+ if input_argument.is_dir:
90
+ raise interfaces.LauncherError(
91
+ f"Cannot consume directory as value. {input_name=}, {input_argument=}"
92
+ )
93
+ if input_argument.total_size > _MAX_INPUT_VALUE_SIZE:
94
+ raise interfaces.LauncherError(
95
+ f"Artifact is too big to consume as value. Consume it as file instead. {input_name=}, {input_argument=}"
96
+ )
97
+ value = input_argument.value
98
+ if value is None:
99
+ # Download artifact data
100
+ if not input_argument.uri:
101
+ raise interfaces.LauncherError(
102
+ f"Artifact data has no value and no uri. This cannot happen. {input_name=}, {input_argument=}"
103
+ )
104
+ uri_reader = self._storage_provider.make_uri(
105
+ input_argument.uri
106
+ ).get_reader()
107
+ try:
108
+ data = uri_reader.download_as_bytes()
109
+ except Exception as ex:
110
+ raise interfaces.LauncherError(
111
+ f"Error downloading artifact data. {input_name=}, {input_argument.uri=}"
112
+ ) from ex
113
+ try:
114
+ value = data.decode("utf-8")
115
+ except Exception as ex:
116
+ raise interfaces.LauncherError(
117
+ f"Error converting artifact data to text. {input_name=}, {input_argument.uri=}"
118
+ ) from ex
119
+ # Updating the input_arguments with the downloaded value
120
+ input_argument.value = value
121
+ return value
122
+
123
+ def get_input_path(input_name: str) -> str:
124
+ input_argument = input_arguments[input_name]
125
+ uri = input_argument.uri
126
+ if not uri:
127
+ if input_argument.value is None:
128
+ raise interfaces.LauncherError(
129
+ f"Artifact data has no value and no uri. This cannot happen. {input_name=}, {input_argument=}"
130
+ )
131
+ uri_writer = self._storage_provider.make_uri(
132
+ input_argument.staging_uri
133
+ ).get_writer()
134
+ try:
135
+ uri_writer.upload_from_text(input_argument.value)
136
+ except Exception as ex:
137
+ raise interfaces.LauncherError(
138
+ f"Error uploading argument value. {input_name=}, {input_argument=}"
139
+ ) from ex
140
+ uri = input_argument.staging_uri
141
+ # Updating the input_arguments with the URI of the uploaded value
142
+ input_argument.uri = uri
143
+
144
+ container_path = (
145
+ container_inputs_root
146
+ / naming_utils.sanitize_file_name(input_name)
147
+ / _CONTAINER_FILE_NAME
148
+ ).as_posix()
149
+ # hf_uri = huggingface_repo_storage.HuggingFaceRepoUri.parse(uri)
150
+ download_input_uris[uri] = container_path
151
+ return container_path
152
+
153
+ def get_output_path(output_name: str) -> str:
154
+ uri = output_uris[output_name]
155
+ # container_path = (
156
+ # container_outputs_root
157
+ # / naming_utils.sanitize_file_name(output_name)
158
+ # / _CONTAINER_FILE_NAME
159
+ # ).as_posix()
160
+ hf_uri = huggingface_repo_storage.HuggingFaceRepoUri.parse(uri)
161
+ uri_path_in_repo = hf_uri.path
162
+ container_path = str(container_outputs_root / uri_path_in_repo)
163
+ upload_output_uris[container_path] = uri
164
+ return container_path
165
+
166
+ def get_log_path() -> str:
167
+ # TODO: Use common URI here
168
+ hf_uri = huggingface_repo_storage.HuggingFaceRepoUri.parse(log_uri)
169
+ uri_path_in_repo = hf_uri.path
170
+ container_path = str(container_outputs_root / uri_path_in_repo)
171
+ return container_path
172
+
173
+ def get_exit_code_path() -> str:
174
+ # TODO: Use common URI here
175
+ hf_uri = huggingface_repo_storage.HuggingFaceRepoUri.parse(log_uri)
176
+ uri_path_in_repo = hf_uri.path
177
+ container_path = str(
178
+ (container_outputs_root / uri_path_in_repo).with_name("exit_code.txt")
179
+ )
180
+ return container_path
181
+
182
+ container_log_path = get_log_path()
183
+ exit_code_path = get_exit_code_path()
184
+
185
+ # Resolving the command line.
186
+ # Also indirectly populates volumes and volume_mounts.
187
+ resolved_cmd = container_component_utils.resolve_container_command_line(
188
+ component_spec=component_spec,
189
+ provided_input_names=set(input_arguments.keys()),
190
+ get_input_value=get_input_value,
191
+ get_input_path=get_input_path,
192
+ get_output_path=get_output_path,
193
+ )
194
+
195
+ # Preparing the artifact uploader wrapper
196
+ # TODO: Use common URI here
197
+ # TODO: Add: --commit-message '{path_in_repo}' once path_in_repo becomes non-empty
198
+ path_in_repo = ""
199
+ # commit_message = path_in_repo
200
+
201
+ hf_repo_uri = huggingface_repo_storage.HuggingFaceRepoUri.parse(log_uri)
202
+
203
+ # It's hard to download data from HuggingFace.
204
+ # First, there is no way to download a directory:
205
+ # 1. The CLI only downloads to cache. So we have to correctly find the data location in the cache sna copy the data out.
206
+ # 2. We cannot specify which directory to download, so we have to use the --include filter.
207
+ # But `hf download --include path` has an issue that it cannot download a directory unless path ends with slash (and for files, there should be no slash).
208
+ # Adding "*" works for both files and directories. It's imperfect, but fine.
209
+ # Another problem: The files in the snapshot are actually symlinks.
210
+ # cp options:
211
+ # -H follow command-line symbolic links in SOURCE
212
+ # -l, --link hard link files instead of copying
213
+ # -L, --dereference always follow symbolic links in SOURCE
214
+ input_download_lines = [
215
+ f'mkdir -p "$(dirname "{container_path}")"'
216
+ f' && snapshot_dir=`hf download --repo-type "{hf_uri.repo_type}" "{hf_uri.repo_id}" --include "{hf_uri.path}*"`'
217
+ f' && cp -r -L "$snapshot_dir/{hf_uri.path}" "{container_path}"'
218
+ for hf_uri, container_path in (
219
+ (huggingface_repo_storage.HuggingFaceRepoUri.parse(uri), container_path)
220
+ for uri, container_path in download_input_uris.items()
221
+ )
222
+ ]
223
+ input_download_code = "\n".join(input_download_lines)
224
+
225
+ artifact_uploader_script = f"""
226
+ set -e -x
227
+ # Workaround for Dash and other shells that do not support -o pipefail.
228
+ if (set -o pipefail 2>/dev/null); then
229
+ set -o pipefail
230
+ fi
231
+
232
+ # Installing uv
233
+ url="https://astral.sh/uv/install.sh"
234
+ if command -v curl 2>/dev/null; then
235
+ # -s: silent, -L: follow redirects
236
+ curl -s -L "$url" | sh
237
+ elif command -v wget 2>/dev/null; then
238
+ wget -q -O - "$url" | sh
239
+ else
240
+ echo "Error: Neither curl nor wget was found. Trying apt-get install" >&2
241
+ apt-get update --quiet && apt-get install -y --no-install-recommends --quiet curl
242
+ curl -s -L "$url" | sh
243
+ fi
244
+
245
+ export PATH="$HOME/.local/bin:$PATH"
246
+
247
+ uv tool install 'huggingface_hub>=1.0.0' --python '>=3.9'
248
+ hf version
249
+
250
+ # Downloading the input data
251
+ {input_download_code}
252
+
253
+ # Running the program
254
+ log_path='{container_log_path}'
255
+ exit_code_path='{exit_code_path}'
256
+ mkdir -p "$(dirname "$log_path")"
257
+ mkdir -p "$(dirname "$exit_code_path")"
258
+ # We need to capture the exit code while piping the stderr and stdout to a log file. Not all shells support `${{PIPEFAIL[0]}}`
259
+ set +e +x
260
+ {{ "$0" "$@"; echo $? >"$exit_code_path";}} 2>&1 | tee "$log_path"
261
+ set -e +x
262
+
263
+ exit_code=`cat "$exit_code_path"`
264
+
265
+ hf upload --repo-type '{hf_repo_uri.repo_type}' '{hf_repo_uri.repo_id}' '{container_outputs_root}' '{path_in_repo}'
266
+ exit "$exit_code"
267
+ """
268
+
269
+ container_env = container_spec.env or {}
270
+
271
+ # Passing HF token to the Job
272
+ secrets: dict[str, str] = {}
273
+ if self._hf_job_token:
274
+ secrets["HF_TOKEN"] = self._hf_job_token
275
+
276
+ command_line = list(resolved_cmd.command or []) + list(resolved_cmd.args or [])
277
+ command_line = ["sh", "-c", artifact_uploader_script] + command_line
278
+ job = self._api_client.run_job(
279
+ image=container_spec.image,
280
+ command=command_line,
281
+ env=dict(container_env),
282
+ timeout=self._job_timeout,
283
+ namespace=self._namespace,
284
+ secrets=secrets,
285
+ # flavor=...,
286
+ )
287
+
288
+ _logger.info(f"Launched HF Job {job.id=}, {job.url=}")
289
+ launched_container = LaunchedHuggingFaceJobContainer(
290
+ id=job.id,
291
+ namespace=self._namespace,
292
+ job=job,
293
+ output_uris=output_uris,
294
+ log_uri=log_uri,
295
+ )
296
+ return launched_container
297
+
298
+ def deserialize_launched_container_from_dict(
299
+ self, launched_container_dict: dict[str, Any]
300
+ ) -> "LaunchedHuggingFaceJobContainer":
301
+ launched_container = LaunchedHuggingFaceJobContainer.from_dict(
302
+ launched_container_dict, api_client=self._api_client
303
+ )
304
+ return launched_container
305
+
306
+ def get_refreshed_launched_container_from_dict(
307
+ self, launched_container_dict: dict[str, Any]
308
+ ) -> "LaunchedHuggingFaceJobContainer":
309
+ launched_container = LaunchedHuggingFaceJobContainer.from_dict(
310
+ launched_container_dict, api_client=self._api_client
311
+ )
312
+ job = self._api_client.inspect_job(
313
+ job_id=launched_container.id,
314
+ namespace=launched_container._namespace,
315
+ )
316
+ new_launched_container = copy.copy(launched_container)
317
+ new_launched_container._job = job
318
+ return new_launched_container
319
+
320
+
321
+ class LaunchedHuggingFaceJobContainer(interfaces.LaunchedContainer):
322
+ def __init__(
323
+ self,
324
+ id: str,
325
+ namespace: str,
326
+ job: huggingface_hub.JobInfo,
327
+ output_uris: dict[str, str],
328
+ log_uri: str,
329
+ api_client: huggingface_hub.HfApi | None = None,
330
+ ):
331
+ self._id: str = id
332
+ self._namespace: str = namespace
333
+ self._job = job
334
+ self._output_uris: dict[str, str] = output_uris
335
+ self._log_uri: str = log_uri
336
+ self._api_client: huggingface_hub.HfApi | None = api_client
337
+
338
+ def _get_api_client(self):
339
+ if not self._api_client:
340
+ raise interfaces.LauncherError(
341
+ "This action requires an API client, but this instance was constructed without one."
342
+ )
343
+ return self._api_client
344
+
345
+ @property
346
+ def id(self) -> str:
347
+ return self._id
348
+
349
+ @property
350
+ def status(self) -> interfaces.ContainerStatus:
351
+ status_str = self._job.status.stage
352
+ # status_message = self._job.status.message
353
+ if status_str == huggingface_hub.JobStage.RUNNING:
354
+ return interfaces.ContainerStatus.RUNNING
355
+ elif status_str == huggingface_hub.JobStage.COMPLETED:
356
+ return interfaces.ContainerStatus.SUCCEEDED
357
+ elif status_str == huggingface_hub.JobStage.ERROR:
358
+ return interfaces.ContainerStatus.FAILED
359
+ elif status_str == huggingface_hub.JobStage.CANCELED:
360
+ return interfaces.ContainerStatus.FAILED
361
+ else: # "DELETED"
362
+ return interfaces.ContainerStatus.ERROR
363
+
364
+ @property
365
+ def exit_code(self) -> Optional[int]:
366
+ # HF Jobs do not provide exit code
367
+ if not self.has_ended:
368
+ return None
369
+ return None
370
+
371
+ @property
372
+ def has_ended(self) -> bool:
373
+ return self.status in (
374
+ interfaces.ContainerStatus.SUCCEEDED,
375
+ interfaces.ContainerStatus.FAILED,
376
+ interfaces.ContainerStatus.ERROR,
377
+ )
378
+
379
+ @property
380
+ def has_succeeded(self) -> bool:
381
+ return self.status == interfaces.ContainerStatus.SUCCEEDED
382
+
383
+ @property
384
+ def has_failed(self) -> bool:
385
+ return self.status == interfaces.ContainerStatus.FAILED
386
+
387
+ @property
388
+ def started_at(self) -> datetime.datetime | None:
389
+ # HF Jobs do not provide started_at, so using created_at
390
+ return self._job.created_at
391
+
392
+ @property
393
+ def ended_at(self) -> datetime.datetime | None:
394
+ # HF Jobs do not provide ended_at
395
+ # Fudging the value by returning the current time.
396
+ if self.has_ended:
397
+ return datetime.datetime.now(datetime.timezone.utc)
398
+ return None
399
+
400
+ @property
401
+ def launcher_error_message(self) -> str | None:
402
+ if self._job.status.message:
403
+ # TODO: Check what kind of messages this returns and when.
404
+ _logger.info(
405
+ f"launcher_error_message: {self._id=}: {self._job.status.message=}"
406
+ )
407
+ return self._job.status.message
408
+ return None
409
+
410
+ def get_log(self) -> str:
411
+ if self.has_ended:
412
+ try:
413
+ return (
414
+ huggingface_repo_storage.HuggingFaceRepoStorageProvider(
415
+ client=self._get_api_client()
416
+ )
417
+ .make_uri(self._log_uri)
418
+ .get_reader()
419
+ .download_as_text()
420
+ )
421
+ except Exception as ex:
422
+ _logger.warning(
423
+ f"get_log: {self._id=}: Error getting log from URI: {self._log_uri}",
424
+ ex,
425
+ )
426
+ return "\n".join(
427
+ self._get_api_client().fetch_job_logs(
428
+ job_id=self._id,
429
+ )
430
+ )
431
+
432
+ def upload_log(self):
433
+ # Logs should be uploaded automatically by the modified command-line wrapper
434
+ pass
435
+
436
+ def stream_log_lines(self) -> typing.Iterator[str]:
437
+ return (
438
+ self._get_api_client()
439
+ .fetch_job_logs(
440
+ job_id=self._id,
441
+ namespace=self._namespace,
442
+ )
443
+ .__iter__()
444
+ )
445
+
446
+ def terminate(self):
447
+ self._get_api_client().cancel_job(job_id=self._id, namespace=self._namespace)
448
+
449
+ def to_dict(self) -> dict[str, Any]:
450
+ debug_job_info = dataclasses.asdict(self._job)
451
+ # Fix JSON serialization of datetime
452
+ del debug_job_info["created_at"]
453
+ return dict(
454
+ huggingface_job=dict(
455
+ id=self.id,
456
+ namespace=self._namespace,
457
+ output_uris=self._output_uris,
458
+ log_uri=self._log_uri,
459
+ # For debugging purposes, not needed otherwise
460
+ debug_job_info=debug_job_info,
461
+ )
462
+ )
463
+
464
+ @classmethod
465
+ def from_dict(
466
+ cls,
467
+ d: dict[str, Any],
468
+ api_client: huggingface_hub.HfApi | None = None,
469
+ ) -> "LaunchedHuggingFaceJobContainer":
470
+ container_dict = d["huggingface_job"]
471
+ job_info = huggingface_hub.JobInfo(
472
+ **container_dict["debug_job_info"],
473
+ )
474
+ return LaunchedHuggingFaceJobContainer(
475
+ id=container_dict["id"],
476
+ namespace=container_dict["namespace"],
477
+ job=job_info,
478
+ output_uris=container_dict["output_uris"],
479
+ log_uri=container_dict["log_uri"],
480
+ api_client=api_client,
481
+ )
huggingface_overlay/cloud_pipelines_backend/storage_providers/huggingface_repo_storage.py ADDED
@@ -0,0 +1,189 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import copy
2
+ import dataclasses
3
+ import logging
4
+ import pathlib
5
+ from typing import Optional
6
+
7
+ import huggingface_hub
8
+ from huggingface_hub import hf_api
9
+
10
+ from cloud_pipelines.orchestration.storage_providers import interfaces
11
+
12
+ _LOGGER = logging.getLogger(name=__name__)
13
+
14
+
15
+ # hf://(model|dataset|space)s/user/repo@branch/path
16
+
17
+
18
+ @dataclasses.dataclass
19
+ class HuggingFaceRepoUri(interfaces.DataUri):
20
+ # uri: str
21
+ repo_type: str # model | dataset | space
22
+ user: str
23
+ repo: str
24
+ path: str
25
+ branch: str | None = None
26
+
27
+ def join_path(self, relative_path: str) -> "HuggingFaceRepoUri":
28
+ new_uri = copy.copy(self)
29
+ new_uri.path = new_uri.path.rstrip("/") + "/" + relative_path
30
+ return new_uri
31
+
32
+ def __str__(self):
33
+ return f"hf://{self.repo_type}s/{self.user}/{self.repo}{'@' + self.branch if self.branch else ''}/{self.path}"
34
+
35
+ @classmethod
36
+ def parse(cls, uri_string: str) -> "HuggingFaceRepoUri":
37
+ # Validating the URI
38
+ if not uri_string.startswith("hf://"):
39
+ raise ValueError(
40
+ f"HuggingFace URI must start with hf://, but got {uri_string}"
41
+ )
42
+ parts = uri_string.split("/", 5)
43
+ repo_type = parts[2].rstrip("s") # Making type singular
44
+ user = parts[3]
45
+ repo, _, branch = parts[4].partition("@")
46
+ path = parts[5]
47
+ if repo_type not in ("model", "dataset", "space"):
48
+ raise ValueError(
49
+ f"HuggingFace URI repo_type must be (model | dataset | space), but got {uri_string}"
50
+ )
51
+ if not user:
52
+ raise ValueError(f"HuggingFace URI must have user, but got {uri_string}")
53
+ if not repo:
54
+ raise ValueError(f"HuggingFace URI must have repo, but got {uri_string}")
55
+ if not path:
56
+ raise ValueError(f"HuggingFace URI must have path, but got {uri_string}")
57
+
58
+ return HuggingFaceRepoUri(
59
+ repo_type=repo_type,
60
+ user=user,
61
+ repo=repo,
62
+ branch=branch,
63
+ path=path,
64
+ )
65
+
66
+ @property
67
+ def repo_id(self):
68
+ return f"{self.user}/{self.repo}"
69
+
70
+
71
+ HuggingFaceRepoUri._register_subclass("huggingface_repo_storage")
72
+
73
+
74
+ class HuggingFaceRepoStorageProvider(interfaces.StorageProvider):
75
+ def __init__(self, client: Optional[huggingface_hub.HfApi] = None) -> None:
76
+ self._client = client or huggingface_hub.HfApi()
77
+
78
+ def make_uri(self, uri: str) -> interfaces.UriAccessor:
79
+ return interfaces.UriAccessor(
80
+ uri=HuggingFaceRepoUri.parse(uri),
81
+ provider=self,
82
+ )
83
+
84
+ def parse_uri_get_accessor(self, uri_string: str) -> interfaces.UriAccessor:
85
+ return interfaces.UriAccessor(
86
+ uri=HuggingFaceRepoUri.parse(uri_string),
87
+ provider=self,
88
+ )
89
+
90
+ def upload(self, source_path: str, destination_uri: HuggingFaceRepoUri):
91
+ _LOGGER.debug(f"Uploading from {source_path} to {destination_uri}")
92
+ if pathlib.Path(source_path).is_dir:
93
+ self._client.upload_folder(
94
+ repo_type=destination_uri.repo_type,
95
+ repo_id=destination_uri.repo_id,
96
+ path_in_repo=destination_uri.path,
97
+ folder_path=source_path,
98
+ commit_message=destination_uri.path,
99
+ )
100
+ else:
101
+ self._client.upload_file(
102
+ repo_type=destination_uri.repo_type,
103
+ repo_id=destination_uri.repo_id,
104
+ path_in_repo=destination_uri.path,
105
+ path_or_fileobj=source_path,
106
+ commit_message=destination_uri.path,
107
+ )
108
+
109
+ def download(self, source_uri: HuggingFaceRepoUri, destination_path: str):
110
+ _LOGGER.debug(f"Downloading from {source_uri} to {destination_path}")
111
+ cache_path = self._client.snapshot_download(
112
+ repo_type=source_uri.repo_type,
113
+ repo_id=source_uri.repo_id,
114
+ allow_patterns=source_uri.path + "*",
115
+ )
116
+ cache_data_path = pathlib.Path(cache_path, source_uri.path).resolve()
117
+ _LOGGER.debug(
118
+ f"Downloaded data to from {source_uri} to cache ({cache_path}). The data should be in {cache_data_path}"
119
+ )
120
+ import shutil
121
+
122
+ if cache_data_path.is_dir:
123
+ shutil.copytree(cache_data_path, destination_path, dirs_exist_ok=True)
124
+ else:
125
+ pathlib.Path(destination_path).parent.mkdir(parents=True, exist_ok=True)
126
+ shutil.copy(cache_data_path, destination_path)
127
+
128
+ def download_bytes(self, source_uri: HuggingFaceRepoUri) -> bytes:
129
+ cache_path = self._client.hf_hub_download(
130
+ repo_type=source_uri.repo_type,
131
+ repo_id=source_uri.repo_id,
132
+ filename=source_uri.path,
133
+ )
134
+ return pathlib.Path(cache_path).read_bytes()
135
+
136
+ def exists(self, uri: HuggingFaceRepoUri) -> bool:
137
+ repo_objects = self._client.get_paths_info(
138
+ repo_type=uri.repo_type, repo_id=uri.repo_id, paths=[uri.path]
139
+ )
140
+ return len(repo_objects) > 0
141
+
142
+ def calculate_data_hash(self, *, data: bytes) -> dict[str, str]:
143
+ import hashlib
144
+
145
+ header = f"blob {len(data)}\0".encode("utf-8")
146
+ hasher = hashlib.sha1(header)
147
+ hasher.update(data)
148
+ return {"GitBlobHash": hasher.hexdigest()}
149
+
150
+ def get_info(self, uri: HuggingFaceRepoUri) -> interfaces.DataInfo:
151
+ repo_objects = self._client.get_paths_info(
152
+ repo_type=uri.repo_type, repo_id=uri.repo_id, paths=[uri.path]
153
+ )
154
+ if not repo_objects:
155
+ raise ValueError(f"Uri {uri} was not found.")
156
+ if len(repo_objects) > 1:
157
+ raise ValueError(
158
+ f"Uri {uri} was found more than once. This cannot happen. {repo_objects}"
159
+ )
160
+ repo_object = repo_objects[0]
161
+
162
+ if isinstance(repo_object, hf_api.RepoFile):
163
+ return interfaces.DataInfo(
164
+ is_dir=False,
165
+ total_size=repo_object.size,
166
+ hashes={"GitBlobHash": repo_object.blob_id},
167
+ )
168
+ elif isinstance(repo_object, hf_api.RepoFolder):
169
+ # Calculating the total size of files in the folder
170
+ child_repo_objects = self._client.list_repo_tree(
171
+ repo_id=uri.repo_id,
172
+ repo_type=uri.repo_type,
173
+ path_in_repo=uri.path,
174
+ recursive=True,
175
+ )
176
+ total_size = sum(
177
+ child_repo_object.size
178
+ for child_repo_object in child_repo_objects
179
+ if isinstance(child_repo_object, hf_api.RepoFile)
180
+ )
181
+ return interfaces.DataInfo(
182
+ is_dir=True,
183
+ total_size=total_size,
184
+ hashes={"GitTreeHash": repo_object.tree_id},
185
+ )
186
+ else:
187
+ raise ValueError(
188
+ f"Got repo object that is neither RepoFile nor RepoFolder: {repo_object}"
189
+ )