Aurelien-Morgan-Bot commited on
Commit
65de176
·
verified ·
1 Parent(s): fab62ce

source-code for model version v0.20_20250326_233239071_UTC- retrain-pipelines 0.1.1

Browse files
v0.20_20250326_233239071_UTC/requirements.txt ADDED
@@ -0,0 +1,628 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ absl-py==1.4.0
2
+ accelerate==1.5.2
3
+ aiohappyeyeballs==2.6.1
4
+ aiohttp==3.11.14
5
+ aiosignal==1.3.2
6
+ alabaster==1.0.0
7
+ albucore==0.0.23
8
+ albumentations==2.0.5
9
+ ale-py==0.10.2
10
+ altair==5.5.0
11
+ annotated-types==0.7.0
12
+ anyio==4.9.0
13
+ argon2-cffi==23.1.0
14
+ argon2-cffi-bindings==21.2.0
15
+ array_record==0.7.1
16
+ arviz==0.21.0
17
+ astropy==7.0.1
18
+ astropy-iers-data==0.2025.3.17.0.34.53
19
+ astunparse==1.6.3
20
+ atpublic==5.1
21
+ attrs==25.3.0
22
+ audioread==3.0.1
23
+ autograd==1.7.0
24
+ babel==2.17.0
25
+ backcall==0.2.0
26
+ beautifulsoup4==4.13.3
27
+ betterproto==2.0.0b6
28
+ bigframes==1.41.0
29
+ bigquery-magics==0.8.0
30
+ bitsandbytes==0.45.4
31
+ bleach==6.2.0
32
+ blinker==1.9.0
33
+ blis==1.2.0
34
+ blosc2==3.2.0
35
+ bokeh==3.6.3
36
+ boto3==1.37.20
37
+ botocore==1.37.20
38
+ Bottleneck==1.4.2
39
+ bqplot==0.12.44
40
+ branca==0.8.1
41
+ CacheControl==0.14.2
42
+ cachetools==5.5.2
43
+ catalogue==2.0.10
44
+ certifi==2025.1.31
45
+ cffi==1.17.1
46
+ chardet==5.2.0
47
+ charset-normalizer==3.4.1
48
+ chex==0.1.89
49
+ clarabel==0.10.0
50
+ click==8.1.8
51
+ cloudpathlib==0.21.0
52
+ cloudpickle==3.1.1
53
+ cmake==3.31.6
54
+ cmdstanpy==1.2.5
55
+ colorama==0.4.6
56
+ colorcet==3.1.0
57
+ colorlover==0.3.0
58
+ colour==0.1.5
59
+ comm==0.2.2
60
+ community==1.0.0b1
61
+ confection==0.1.5
62
+ cons==0.4.6
63
+ contourpy==1.3.1
64
+ cramjam==2.9.1
65
+ cryptography==43.0.3
66
+ cuda-python==12.6.0
67
+ cudf-cu12 @ https://pypi.nvidia.com/cudf-cu12/cudf_cu12-25.2.1-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
68
+ cudf-polars-cu12==24.12.0
69
+ cufflinks==0.17.3
70
+ cuml-cu12==25.2.1
71
+ cupy-cuda12x==13.3.0
72
+ cut-cross-entropy==25.1.1
73
+ cuvs-cu12==25.2.1
74
+ cvxopt==1.3.2
75
+ cvxpy==1.6.4
76
+ cycler==0.12.1
77
+ cyipopt==1.5.0
78
+ cymem==2.0.11
79
+ Cython==3.0.12
80
+ dask==2024.12.1
81
+ dask-cuda==25.2.0
82
+ dask-cudf-cu12==25.2.2
83
+ dask-expr==1.1.21
84
+ datascience==0.17.6
85
+ datasets==3.1.0
86
+ db-dtypes==1.4.2
87
+ dbus-python==1.2.18
88
+ debugpy==1.8.0
89
+ decorator==4.4.2
90
+ defusedxml==0.7.1
91
+ Deprecated==1.2.18
92
+ diffusers==0.32.2
93
+ dill==0.3.8
94
+ distributed==2024.12.1
95
+ distributed-ucxx-cu12==0.42.0
96
+ distro==1.9.0
97
+ dlib==19.24.2
98
+ dm-tree==0.1.9
99
+ docker==7.1.0
100
+ docker-pycreds==0.4.0
101
+ docstring_parser==0.16
102
+ docutils==0.21.2
103
+ dopamine_rl==4.1.2
104
+ duckdb==1.2.1
105
+ earthengine-api==1.5.7
106
+ easydict==1.13
107
+ editdistance==0.8.1
108
+ eerepr==0.1.1
109
+ einops==0.8.1
110
+ en_core_web_sm @ https://github.com/explosion/spacy-models/releases/download/en_core_web_sm-3.8.0/en_core_web_sm-3.8.0-py3-none-any.whl#sha256=1932429db727d4bff3deed6b34cfc05df17794f4a52eeb26cf8928f7c1a0fb85
111
+ entrypoints==0.4
112
+ et_xmlfile==2.0.0
113
+ etils==1.12.2
114
+ etuples==0.3.9
115
+ Farama-Notifications==0.0.4
116
+ fastai==2.7.19
117
+ fastapi==0.115.12
118
+ fastcore==1.7.29
119
+ fastdownload==0.0.7
120
+ fastjsonschema==2.21.1
121
+ fastprogress==1.0.3
122
+ fastrlock==0.8.3
123
+ filelock==3.18.0
124
+ firebase-admin==6.7.0
125
+ Flask==3.1.0
126
+ flatbuffers==25.2.10
127
+ flax==0.10.4
128
+ folium==0.19.5
129
+ fonttools==4.56.0
130
+ frozendict==2.4.6
131
+ frozenlist==1.5.0
132
+ fsspec==2024.9.0
133
+ future==1.0.0
134
+ gast==0.6.0
135
+ GDAL==3.6.4
136
+ gdown==5.2.0
137
+ geemap==0.35.3
138
+ geocoder==1.38.1
139
+ geographiclib==2.0
140
+ geopandas==1.0.1
141
+ geopy==2.4.1
142
+ gin-config==0.5.0
143
+ gitdb==4.0.12
144
+ GitPython==3.1.44
145
+ glob2==0.7
146
+ google==2.0.3
147
+ google-ai-generativelanguage==0.6.15
148
+ google-api-core==2.24.2
149
+ google-api-python-client==2.164.0
150
+ google-auth==2.38.0
151
+ google-auth-httplib2==0.2.0
152
+ google-auth-oauthlib==1.2.1
153
+ google-cloud-aiplatform==1.84.0
154
+ google-cloud-bigquery==3.29.0
155
+ google-cloud-bigquery-connection==1.18.2
156
+ google-cloud-bigquery-storage==2.29.1
157
+ google-cloud-bigtable==2.30.0
158
+ google-cloud-core==2.4.3
159
+ google-cloud-dataproc==5.18.1
160
+ google-cloud-datastore==2.20.2
161
+ google-cloud-firestore==2.20.1
162
+ google-cloud-functions==1.20.2
163
+ google-cloud-iam==2.18.3
164
+ google-cloud-language==2.17.1
165
+ google-cloud-pubsub==2.29.0
166
+ google-cloud-resource-manager==1.14.2
167
+ google-cloud-spanner==3.53.0
168
+ google-cloud-storage==2.19.0
169
+ google-cloud-translate==3.20.2
170
+ google-colab @ file:///colabtools/dist/google_colab-1.0.0.tar.gz
171
+ google-crc32c==1.7.0
172
+ google-genai==1.7.0
173
+ google-generativeai==0.8.4
174
+ google-pasta==0.2.0
175
+ google-resumable-media==2.7.2
176
+ google-spark-connect==0.5.2
177
+ googleapis-common-protos==1.69.2
178
+ googledrivedownloader==1.1.0
179
+ graphviz==0.20.3
180
+ greenlet==3.1.1
181
+ grpc-google-iam-v1==0.14.2
182
+ grpc-interceptor==0.15.4
183
+ grpcio==1.71.0
184
+ grpcio-status==1.71.0
185
+ grpclib==0.4.7
186
+ gspread==6.2.0
187
+ gspread-dataframe==4.0.0
188
+ gym==0.25.2
189
+ gym-notices==0.0.8
190
+ gymnasium==1.1.1
191
+ h11==0.14.0
192
+ h2==4.2.0
193
+ h5netcdf==1.6.1
194
+ h5py==3.13.0
195
+ hdbscan==0.8.40
196
+ hf_transfer==0.1.9
197
+ highspy==1.9.0
198
+ holidays==0.69
199
+ holoviews==1.20.2
200
+ hpack==4.1.0
201
+ html5lib==1.1
202
+ httpcore==1.0.7
203
+ httpimport==1.4.1
204
+ httplib2==0.22.0
205
+ httptools==0.6.4
206
+ httpx==0.28.1
207
+ huggingface-hub==0.27.1
208
+ humanize==4.12.1
209
+ hyperframe==6.1.0
210
+ hyperopt==0.2.7
211
+ ibis-framework==9.5.0
212
+ idna==3.10
213
+ imageio==2.37.0
214
+ imageio-ffmpeg==0.6.0
215
+ imagesize==1.4.1
216
+ imbalanced-learn==0.13.0
217
+ immutabledict==4.2.1
218
+ importlib_metadata==8.6.1
219
+ importlib_resources==6.5.2
220
+ imutils==0.5.4
221
+ inflect==7.5.0
222
+ iniconfig==2.1.0
223
+ intel-cmplr-lib-ur==2025.1.0
224
+ intel-openmp==2025.1.0
225
+ ipyevents==2.0.2
226
+ ipyfilechooser==0.6.0
227
+ ipykernel==6.29.5
228
+ ipyleaflet==0.19.2
229
+ ipyparallel==8.8.0
230
+ ipython==7.34.0
231
+ ipython-genutils==0.2.0
232
+ ipython-sql==0.5.0
233
+ ipytree==0.2.2
234
+ ipywidgets==7.7.1
235
+ itsdangerous==2.2.0
236
+ jax==0.5.2
237
+ jax-cuda12-pjrt==0.5.1
238
+ jax-cuda12-plugin==0.5.1
239
+ jaxlib==0.5.1
240
+ jedi==0.19.2
241
+ jeepney==0.7.1
242
+ jellyfish==1.1.0
243
+ jieba==0.42.1
244
+ Jinja2==3.1.4
245
+ jiter==0.9.0
246
+ jmespath==1.0.1
247
+ joblib==1.4.2
248
+ jsonpatch==1.33
249
+ jsonpickle==4.0.2
250
+ jsonpointer==3.0.0
251
+ jsonschema==4.23.0
252
+ jsonschema-specifications==2024.10.1
253
+ jupyter-client==6.1.12
254
+ jupyter-console==6.1.0
255
+ jupyter-leaflet==0.19.2
256
+ jupyter-server==1.16.0
257
+ jupyter_core==5.7.2
258
+ jupyterlab_pygments==0.3.0
259
+ jupyterlab_widgets==3.0.13
260
+ kaggle==1.7.4.2
261
+ kagglehub==0.3.10
262
+ keras==3.8.0
263
+ keras-hub==0.18.1
264
+ keras-nlp==0.18.1
265
+ keyring==23.5.0
266
+ kiwisolver==1.4.8
267
+ langchain==0.3.21
268
+ langchain-core==0.3.47
269
+ langchain-text-splitters==0.3.7
270
+ langcodes==3.5.0
271
+ langsmith==0.3.18
272
+ language_data==1.3.0
273
+ launchpadlib==1.10.16
274
+ lazr.restfulclient==0.14.4
275
+ lazr.uri==1.0.6
276
+ lazy_loader==0.4
277
+ libclang==18.1.1
278
+ libcudf-cu12==24.12.0
279
+ libcugraph-cu12==25.2.0
280
+ libcuml-cu12==25.2.1
281
+ libcuvs-cu12==25.2.1
282
+ libkvikio-cu12==24.12.1
283
+ libraft-cu12==25.2.0
284
+ librosa==0.11.0
285
+ libucx-cu12==1.18.0
286
+ libucxx-cu12==0.42.0
287
+ lightgbm==4.5.0
288
+ linkify-it-py==2.0.3
289
+ litserve==0.2.6
290
+ llvmlite==0.43.0
291
+ locket==1.0.0
292
+ logical-unification==0.4.6
293
+ lxml==5.3.0
294
+ Mako==1.1.3
295
+ marisa-trie==1.2.1
296
+ Markdown==3.7
297
+ markdown-it-py==3.0.0
298
+ MarkupSafe==3.0.2
299
+ matplotlib==3.9.2
300
+ matplotlib-inline==0.1.7
301
+ matplotlib-venn==1.1.2
302
+ mdit-py-plugins==0.4.2
303
+ mdurl==0.1.2
304
+ metaflow==2.10.0
305
+ metaflow-card-html==1.0.2
306
+ miniKanren==1.0.3
307
+ missingno==0.5.2
308
+ mistune==3.1.3
309
+ mizani==0.13.1
310
+ mkl==2025.0.1
311
+ ml-dtypes==0.4.1
312
+ mlxtend==0.23.4
313
+ more-itertools==10.6.0
314
+ moviepy==1.0.3
315
+ mpmath==1.3.0
316
+ msgpack==1.1.0
317
+ multidict==6.2.0
318
+ multipledispatch==1.0.0
319
+ multiprocess==0.70.16
320
+ multitasking==0.0.11
321
+ murmurhash==1.0.12
322
+ music21==9.3.0
323
+ namex==0.0.8
324
+ narwhals==1.31.0
325
+ natsort==8.4.0
326
+ nbclassic==1.2.0
327
+ nbclient==0.10.2
328
+ nbconvert==7.16.6
329
+ nbformat==5.10.4
330
+ ndindex==1.9.2
331
+ nest-asyncio==1.6.0
332
+ networkx==3.2.1
333
+ nibabel==5.3.2
334
+ nltk==3.9.1
335
+ notebook==6.5.7
336
+ notebook_shim==0.2.4
337
+ numba==0.60.0
338
+ numba-cuda==0.2.0
339
+ numexpr==2.10.2
340
+ numpy==1.26.4
341
+ nvidia-cublas-cu12==12.4.5.8
342
+ nvidia-cuda-cupti-cu12==12.4.127
343
+ nvidia-cuda-nvcc-cu12==12.5.82
344
+ nvidia-cuda-nvrtc-cu12==12.4.127
345
+ nvidia-cuda-runtime-cu12==12.4.127
346
+ nvidia-cudnn-cu12==9.1.0.70
347
+ nvidia-cufft-cu12==11.2.1.3
348
+ nvidia-curand-cu12==10.3.5.147
349
+ nvidia-cusolver-cu12==11.6.1.9
350
+ nvidia-cusparse-cu12==12.3.1.170
351
+ nvidia-cusparselt-cu12==0.6.2
352
+ nvidia-ml-py==12.570.86
353
+ nvidia-nccl-cu12==2.21.5
354
+ nvidia-nvcomp-cu12==4.1.0.6
355
+ nvidia-nvjitlink-cu12==12.4.127
356
+ nvidia-nvtx-cu12==12.4.127
357
+ nvtx==0.2.11
358
+ nx-cugraph-cu12 @ https://pypi.nvidia.com/nx-cugraph-cu12/nx_cugraph_cu12-25.2.0-py3-none-any.whl
359
+ oauth2client==4.1.3
360
+ oauthlib==3.2.2
361
+ openai==1.68.2
362
+ opencv-contrib-python==4.11.0.86
363
+ opencv-python==4.11.0.86
364
+ opencv-python-headless==4.11.0.86
365
+ openpyxl==3.1.5
366
+ opentelemetry-api==1.31.1
367
+ opentelemetry-sdk==1.31.1
368
+ opentelemetry-semantic-conventions==0.52b1
369
+ opt_einsum==3.4.0
370
+ optax==0.2.4
371
+ optree==0.14.1
372
+ orbax-checkpoint==0.11.10
373
+ orjson==3.10.15
374
+ osqp==0.6.7.post3
375
+ packaging==24.2
376
+ pandas==2.2.2
377
+ pandas-datareader==0.10.0
378
+ pandas-gbq==0.28.0
379
+ pandas-stubs==2.2.2.240909
380
+ pandocfilters==1.5.1
381
+ panel==1.6.1
382
+ param==2.2.0
383
+ parso==0.8.4
384
+ parsy==2.1
385
+ partd==1.4.2
386
+ pathlib==1.0.1
387
+ patsy==1.0.1
388
+ peewee==3.17.9
389
+ peft==0.14.0
390
+ pexpect==4.9.0
391
+ pickleshare==0.7.5
392
+ pillow==11.1.0
393
+ platformdirs==4.3.7
394
+ plotly==5.24.1
395
+ plotnine==0.14.5
396
+ pluggy==1.5.0
397
+ ply==3.11
398
+ polars==1.11.0
399
+ pooch==1.8.2
400
+ portpicker==1.5.2
401
+ preshed==3.0.9
402
+ prettytable==3.15.1
403
+ proglog==0.1.10
404
+ progressbar2==4.5.0
405
+ prometheus_client==0.21.1
406
+ promise==2.3
407
+ prompt_toolkit==3.0.50
408
+ propcache==0.3.0
409
+ prophet==1.1.6
410
+ proto-plus==1.26.1
411
+ protobuf==3.20.3
412
+ psutil==5.9.5
413
+ psycopg2==2.9.10
414
+ ptyprocess==0.7.0
415
+ py-cpuinfo==9.0.0
416
+ py4j==0.10.9.7
417
+ pyarrow==17.0.0
418
+ pyasn1==0.6.1
419
+ pyasn1_modules==0.4.1
420
+ pycairo==1.27.0
421
+ pycocotools==2.0.8
422
+ pycparser==2.22
423
+ pydantic==2.9.2
424
+ pydantic_core==2.23.4
425
+ pydata-google-auth==1.9.1
426
+ pydot==1.4.2
427
+ pydotplus==2.0.2
428
+ PyDrive==1.3.1
429
+ PyDrive2==1.21.3
430
+ pyerfa==2.0.1.5
431
+ pygame==2.6.1
432
+ pygit2==1.17.0
433
+ Pygments==2.18.0
434
+ PyGObject==3.42.0
435
+ PyJWT==2.10.1
436
+ pylibcudf-cu12==24.12.0
437
+ pylibcugraph-cu12==25.2.0
438
+ pylibraft-cu12==25.2.0
439
+ pymc==5.21.1
440
+ pymystem3==0.2.0
441
+ pynndescent==0.5.13
442
+ pynvjitlink-cu12==0.5.2
443
+ pynvml==12.0.0
444
+ pyogrio==0.10.0
445
+ Pyomo==6.8.2
446
+ PyOpenGL==3.1.9
447
+ pyOpenSSL==24.2.1
448
+ pyparsing==3.2.1
449
+ pyperclip==1.9.0
450
+ pyproj==3.7.1
451
+ pyshp==2.3.1
452
+ PySocks==1.7.1
453
+ pyspark==3.5.5
454
+ pytensor==2.28.3
455
+ pytest==8.3.3
456
+ python-apt==0.0.0
457
+ python-box==7.3.2
458
+ python-dateutil==2.8.2
459
+ python-dotenv==1.0.1
460
+ python-louvain==0.16
461
+ python-multipart==0.0.20
462
+ python-slugify==8.0.4
463
+ python-snappy==0.7.3
464
+ python-utils==3.9.1
465
+ pytz==2025.1
466
+ pyviz_comms==3.0.4
467
+ PyYAML==6.0.2
468
+ pyzmq==24.0.1
469
+ qdldl==0.1.7.post5
470
+ raft-dask-cu12==25.2.0
471
+ rapids-dask-dependency==25.2.0
472
+ ratelim==0.1.6
473
+ referencing==0.36.2
474
+ regex==2024.11.6
475
+ requests==2.32.3
476
+ requests-oauthlib==2.0.0
477
+ requests-toolbelt==1.0.0
478
+ requirements-parser==0.9.0
479
+ retrain_pipelines @ git+https://github.com/aurelienmorgan/retrain-pipelines.git@f00ea62b3234bdcae6c3b09fb0570cf5134c46e9#subdirectory=pkg_src
480
+ rich==13.9.4
481
+ rmm-cu12==24.12.0
482
+ roman-numerals-py==3.1.0
483
+ rpds-py==0.23.1
484
+ rpy2==3.5.17
485
+ rsa==4.9
486
+ s3transfer==0.11.4
487
+ safetensors==0.5.3
488
+ scikit-image==0.25.2
489
+ scikit-learn==1.6.1
490
+ scipy==1.14.1
491
+ scooby==0.10.0
492
+ scs==3.2.7.post2
493
+ seaborn==0.13.2
494
+ SecretStorage==3.3.1
495
+ Send2Trash==1.8.3
496
+ sentence-transformers==3.4.1
497
+ sentencepiece==0.2.0
498
+ sentry-sdk==2.24.0
499
+ setproctitle==1.3.5
500
+ shap==0.47.0
501
+ shapely==2.0.7
502
+ shellingham==1.5.4
503
+ shtab==1.7.1
504
+ simple-parsing==0.1.7
505
+ simplejson==3.20.1
506
+ simsimd==6.2.1
507
+ six==1.17.0
508
+ sklearn-compat==0.1.3
509
+ sklearn-pandas==2.2.0
510
+ slicer==0.0.8
511
+ smart-open==7.1.0
512
+ smmap==5.0.2
513
+ sniffio==1.3.1
514
+ snowballstemmer==2.2.0
515
+ sortedcontainers==2.4.0
516
+ soundfile==0.13.1
517
+ soupsieve==2.6
518
+ soxr==0.5.0.post1
519
+ spacy==3.8.4
520
+ spacy-legacy==3.0.12
521
+ spacy-loggers==1.0.5
522
+ spanner-graph-notebook==1.1.5
523
+ Sphinx==8.2.3
524
+ sphinxcontrib-applehelp==2.0.0
525
+ sphinxcontrib-devhelp==2.0.0
526
+ sphinxcontrib-htmlhelp==2.1.0
527
+ sphinxcontrib-jsmath==1.0.1
528
+ sphinxcontrib-qthelp==2.0.0
529
+ sphinxcontrib-serializinghtml==2.0.0
530
+ SQLAlchemy==2.0.39
531
+ sqlglot==25.20.2
532
+ sqlparse==0.5.3
533
+ srsly==2.5.1
534
+ stanio==0.5.1
535
+ starlette==0.46.1
536
+ statsmodels==0.14.4
537
+ stringzilla==3.12.3
538
+ sympy==1.13.1
539
+ tables==3.10.2
540
+ tabulate==0.9.0
541
+ tbb==2022.1.0
542
+ tblib==3.0.0
543
+ tcmlib==1.3.0
544
+ tenacity==9.0.0
545
+ tensorboard==2.18.0
546
+ tensorboard-data-server==0.7.2
547
+ tensorflow==2.18.0
548
+ tensorflow-datasets==4.9.8
549
+ tensorflow-hub==0.16.1
550
+ tensorflow-io-gcs-filesystem==0.37.1
551
+ tensorflow-metadata==1.16.1
552
+ tensorflow-probability==0.25.0
553
+ tensorflow-text==2.18.1
554
+ tensorstore==0.1.72
555
+ termcolor==2.5.0
556
+ terminado==0.18.1
557
+ text-unidecode==1.3
558
+ textblob==0.19.0
559
+ tf-slim==1.1.0
560
+ tf_keras==2.18.0
561
+ thinc==8.3.4
562
+ threadpoolctl==3.6.0
563
+ tifffile==2025.3.13
564
+ timm==1.0.15
565
+ tinycss2==1.4.0
566
+ tokenizers==0.20.3
567
+ toml==0.10.2
568
+ toolz==0.12.1
569
+ torch==2.5.0
570
+ torchsummary==1.5.1
571
+ torchvision==0.20.0
572
+ tornado==6.4.2
573
+ tqdm==4.67.1
574
+ traitlets==5.7.1
575
+ traittypes==0.2.1
576
+ transformers==4.46.2
577
+ treelite==4.4.1
578
+ treescope==0.1.9
579
+ triton==3.1.0
580
+ trl==0.12.0
581
+ tweepy==4.15.0
582
+ typeguard==4.4.2
583
+ typer==0.15.2
584
+ types-pytz==2025.1.0.20250318
585
+ types-setuptools==76.0.0.20250313
586
+ typing_extensions==4.12.2
587
+ tyro==0.9.17
588
+ tzdata==2025.1
589
+ tzlocal==5.3.1
590
+ uc-micro-py==1.0.3
591
+ ucx-py-cu12==0.42.0
592
+ ucxx-cu12==0.42.0
593
+ umap-learn==0.5.7
594
+ umf==0.10.0
595
+ unsloth @ git+https://github.com/unslothai/unsloth.git@3a1e7ef8299f3c96fa6e8de11fd0772af3cbc83f
596
+ unsloth_zoo==2024.11.4
597
+ uritemplate==4.1.1
598
+ urllib3==2.3.0
599
+ uvicorn==0.34.0
600
+ uvloop==0.21.0
601
+ vega-datasets==0.9.0
602
+ wadllib==1.3.6
603
+ wandb==0.19.8
604
+ wasabi==1.1.3
605
+ watchfiles==1.0.4
606
+ wcwidth==0.2.13
607
+ weasel==0.4.1
608
+ webcolors==24.11.1
609
+ webencodings==0.5.1
610
+ websocket-client==1.8.0
611
+ websockets==15.0.1
612
+ Werkzeug==3.1.3
613
+ widgetsnbextension==3.6.10
614
+ wordcloud==1.9.4
615
+ wrapt==1.17.2
616
+ xarray==2025.1.2
617
+ xarray-einstats==0.8.0
618
+ xformers==0.0.28.post2
619
+ xgboost==2.1.4
620
+ xlrd==2.0.1
621
+ xxhash==3.5.0
622
+ xyzservices==2025.1.0
623
+ yarl==1.18.3
624
+ yellowbrick==1.5
625
+ yfinance==0.2.55
626
+ zict==3.0.0
627
+ zipp==3.21.0
628
+ zstandard==0.23.0
v0.20_20250326_233239071_UTC/retraining_pipeline.py ADDED
@@ -0,0 +1,2220 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from unsloth import FastLanguageModel, \
3
+ is_bfloat16_supported, UnslothTrainer, \
4
+ UnslothTrainingArguments
5
+
6
+ import torch
7
+
8
+ import os
9
+ import sys
10
+
11
+ import gc
12
+ import json
13
+ import time
14
+ import shutil
15
+ import logging
16
+ import traceback
17
+ import subprocess
18
+ import importlib.util
19
+ from enum import Enum
20
+ from io import StringIO
21
+ from textwrap import dedent
22
+ from datetime import datetime
23
+ from contextlib import redirect_stdout
24
+
25
+ import numpy as np
26
+ import pandas as pd
27
+
28
+ import polars as pl
29
+ from polars.exceptions import ComputeError
30
+
31
+ import matplotlib
32
+ import matplotlib.pyplot as plt
33
+
34
+ from jinja2 import Environment, FileSystemLoader
35
+
36
+ from metaflow import FlowSpec, step, Parameter, JSONType, \
37
+ IncludeFile, current, metaflow_config as mf_config, \
38
+ resources, Flow, Task, card
39
+ from metaflow.current import Current
40
+ from metaflow.cards import Image, Table, Markdown, \
41
+ Artifact, get_cards
42
+
43
+ from datasets import load_dataset, Dataset, DatasetDict
44
+ from datasets.config import HF_DATASETS_CACHE, HF_CACHE_HOME
45
+ from huggingface_hub import list_repo_commits
46
+ from huggingface_hub.utils import \
47
+ disable_progress_bars as hf_hub_disable_progress_bars, \
48
+ enable_progress_bars as hf_hub_enable_progress_bars
49
+ from transformers import AutoTokenizer
50
+ from transformers.utils import logging as transformers_logging
51
+
52
+ from retrain_pipelines import __version__
53
+ from retrain_pipelines.dataset.hf_utils import get_lazy_df, \
54
+ get_column_info, iterable_dataset_multi_buffer_sampler, \
55
+ push_dataset_version_to_hub
56
+ from retrain_pipelines.dataset.tool_calls import \
57
+ get_unique_tools, count_tool_occurrences, \
58
+ plot_tools_occurences, column_words_stats, \
59
+ plot_words_count
60
+ from retrain_pipelines.utils.hf_utils import \
61
+ get_new_repo_minor_version, push_files_to_hub_repo_branch
62
+ from retrain_pipelines.utils import create_requirements
63
+
64
+
65
+ class LocalServeReadinessEnum(Enum):
66
+ """
67
+ tracking local-serve (infra-validation)
68
+ status using a "3+"-states enum :
69
+ - "-1" for "not applicable"
70
+ (i.e. "model version not blessed"),
71
+ - "0/1" bool for failure/success.
72
+ """
73
+ NOT_APPLICABLE = -1
74
+ FAILURE = 0
75
+ FAILURE_NO_DOCKER = 2
76
+ SUCCESS = 1
77
+
78
+
79
+ class UnslothFuncCallFlow(FlowSpec):
80
+ """
81
+ Training pipeline
82
+ """
83
+ # @see https://github.com/unslothai/unsloth/wiki
84
+
85
+ #--- flow parameters -------------------------------------------------------
86
+
87
+ RETRAIN_PIPELINE_TYPE = "mf_unsloth_func_call_litserve"
88
+ # in order to share the config across subprocesses
89
+ os.environ["retrain_pipeline_type"] = RETRAIN_PIPELINE_TYPE
90
+
91
+ hf_dataset = Parameter(
92
+ "hf_dataset",
93
+ help="dict with 'repo_id' and 'commit_hash' keys. " + \
94
+ "if 'commit_hash is None, falls back to latest version " +\
95
+ "of the dataset available in parquet format.\n" +
96
+ "Note that there are 3 required 'attributes' of type " + \
97
+ "str, list[str], list[str]",
98
+ type=JSONType,
99
+ default=dedent("""{
100
+ "repo_id": "Salesforce/xlam-function-calling-60k",
101
+ "config_name": "",
102
+ "commit_hash": "",
103
+ "attributes": {
104
+ "query_attr": "query",
105
+ "answers_attr": "answers",
106
+ "tools_attr": "tools"
107
+ }
108
+ }""").replace("'", '"').strip('"')
109
+ )
110
+
111
+ augmentation_rate = Parameter(
112
+ "augmentation_rate",
113
+ type=float,
114
+ default=.05,
115
+ help="proportion of records to be augmented "+\
116
+ "(x% of original dataset is created"+\
117
+ " as additional augmented datapoints), i.e. "+\
118
+ "truncated queries to serve as negative examples, "+\
119
+ "meaning they trigger no tool call "+\
120
+ "due to info incompleteness."
121
+ )
122
+
123
+ hf_enrich_dataset = Parameter(
124
+ "hf_enrich_dataset",
125
+ help="dict with 'repo_id', 'config_name' and 'commit_hash', "+\
126
+ "query_attribute' and 'query_attribute_handler' keys. "+\
127
+ "if 'commit_hash is None, falls back to latest version "+\
128
+ "of the dataset available in parquet format."+\
129
+ "'query_attribute' depicts the dataset attribute "+\
130
+ "from which 'queries' are to be sampled."+\
131
+ "'query_attribute_handler' serves for attributes "+\
132
+ "that have complex structure, "+\
133
+ "other than 'string' datatype.",
134
+ type=JSONType,
135
+ # @see https://huggingface.co/datasets/google-research-datasets/natural_questions
136
+ default=dedent("""{
137
+ "repo_id": "lighteval/natural_questions_clean",
138
+ "config_name": "",
139
+ "commit_hash": "",
140
+ "query_attribute": "question",
141
+ "query_attribute_handler": "lambda x: x"
142
+ }""").replace("'", '"').strip('"')
143
+ )
144
+
145
+ enrichment_rate = Parameter(
146
+ "enrichment_rate",
147
+ type=float,
148
+ default=.1,
149
+ help="proportion of records "+\
150
+ "to be added from the 'hf_enrich_dataset'"+\
151
+ "(x% of original dataset is sampled and"+\
152
+ " added as enriching datapoints), i.e. "+\
153
+ "queries to serve as negative examples, "+\
154
+ "due to their complete disconnexion "+\
155
+ "to tool calling situations."
156
+ )
157
+
158
+ dataset_repo_id = Parameter(
159
+ "dataset_repo_id",
160
+ type=str,
161
+ default="retrain-pipelines/func_calls",
162
+ help="The 'repo_id' to be used " + \
163
+ "for the Hugging Face dataset version push " + \
164
+ "(will be created at runtime" + \
165
+ " if doesn't already exist)."
166
+ )
167
+
168
+ hf_base_model = Parameter(
169
+ "hf_base_model",
170
+ help="dict with 'repo_id' and 'commit_hash' keys."+\
171
+ "if 'commit_hash is None, falls back "+\
172
+ "to latest available version of the model.",
173
+ type=JSONType,
174
+ default=dedent("""{
175
+ "repo_id": "unsloth/Qwen2.5-1.5B",
176
+ "commit_hash": ""
177
+ }""").replace("'", '"').strip('"')
178
+ )
179
+
180
+ cpt_training_args = Parameter(
181
+ "cpt_training_args",
182
+ help="dict with `TrainingArguments` params "+\
183
+ "for the CPT job.",
184
+ type=JSONType,
185
+ default=dedent("""{
186
+ "warmup_ratio": 0.1,
187
+ "num_train_epochs": 1
188
+ }""").replace("'", '"').strip('"')
189
+ )
190
+
191
+ sft_training_args = Parameter(
192
+ "sft_training_args",
193
+ help="dict with `TrainingArguments` params "+\
194
+ "for the SFT job.",
195
+ type=JSONType,
196
+ default=dedent("""{
197
+ "warmup_ratio": 0.1,
198
+ "num_train_epochs": 1
199
+ }""").replace("'", '"').strip('"')
200
+ )
201
+
202
+ model_repo_id = Parameter(
203
+ "model_repo_id",
204
+ type=str,
205
+ default="retrain-pipelines/function_caller",
206
+ help="The 'repo_id' to be used " + \
207
+ "for the Hugging Face model version push " + \
208
+ "(will be created at runtime" + \
209
+ " if doesn't already exist)."
210
+ )
211
+
212
+ default_pipeline_card_module_dir = \
213
+ os.path.dirname(
214
+ importlib.util.find_spec(
215
+ f"retrain_pipelines.pipeline_card."+
216
+ f"{RETRAIN_PIPELINE_TYPE}"
217
+ ).origin)
218
+ pipeline_card_artifacts_path = Parameter(
219
+ "pipeline_card_artifacts_path",
220
+ type=str,
221
+ default=default_pipeline_card_module_dir,
222
+ help="pipeline_card artifacts location "+\
223
+ "(i.e. dir hosting your optional " + \
224
+ " custom documentation files :" + \
225
+ " 'pipeline_card.py' and/or 'template.html'"+\
226
+ " and/or 'model_readme.py'"+\
227
+ " and/or 'model_readme_template.md'," +\
228
+ " and/or 'dataset_readme.py'"+\
229
+ " and/or 'dataset_readme_template.md' file), " +\
230
+ "if different from default."
231
+ )
232
+ @staticmethod
233
+ def copy_default_dataset_readme_module(
234
+ target_dir: str,
235
+ exists_ok: bool = False
236
+ ) -> None:
237
+ os.makedirs(target_dir, exist_ok=True)
238
+ if (
239
+ not exists_ok and
240
+ os.path.exists(os.path.join(target_dir, "dataset_readme.py"))
241
+ ):
242
+ print("File already exists. Skipping copy.")
243
+ else:
244
+ filefullname = os.path.join(
245
+ UnslothFuncCallFlow.default_pipeline_card_module_dir,
246
+ "dataset_readme.py"
247
+ )
248
+ shutil.copy(filefullname, target_dir)
249
+ print(filefullname)
250
+ @staticmethod
251
+ def copy_default_dataset_readme_template(
252
+ target_dir: str,
253
+ exists_ok: bool = False
254
+ ) -> None:
255
+ os.makedirs(target_dir, exist_ok=True)
256
+ if (
257
+ not exists_ok and
258
+ os.path.exists(os.path.join(target_dir,
259
+ "dataset_readme_template.md"))
260
+ ):
261
+ print("File already exists. Skipping copy.")
262
+ else:
263
+ filefullname = os.path.join(
264
+ UnslothFuncCallFlow.default_pipeline_card_module_dir,
265
+ "dataset_readme_template.md")
266
+ shutil.copy(filefullname, target_dir)
267
+ print(filefullname)
268
+ @staticmethod
269
+ def copy_default_model_readme_module(
270
+ target_dir: str,
271
+ exists_ok: bool = False
272
+ ) -> None:
273
+ os.makedirs(target_dir, exist_ok=True)
274
+ if (
275
+ not exists_ok and
276
+ os.path.exists(os.path.join(target_dir, "model_readme.py"))
277
+ ):
278
+ print("File already exists. Skipping copy.")
279
+ else:
280
+ filefullname = os.path.join(
281
+ UnslothFuncCallFlow.default_pipeline_card_module_dir,
282
+ "model_readme.py"
283
+ )
284
+ shutil.copy(filefullname, target_dir)
285
+ print(filefullname)
286
+ @staticmethod
287
+ def copy_default_model_readme_template(
288
+ target_dir: str,
289
+ exists_ok: bool = False
290
+ ) -> None:
291
+ os.makedirs(target_dir, exist_ok=True)
292
+ if (
293
+ not exists_ok and
294
+ os.path.exists(os.path.join(target_dir,
295
+ "model_readme_template.md"))
296
+ ):
297
+ print("File already exists. Skipping copy.")
298
+ else:
299
+ filefullname = os.path.join(
300
+ UnslothFuncCallFlow.default_pipeline_card_module_dir,
301
+ "model_readme_template.md")
302
+ shutil.copy(filefullname, target_dir)
303
+ print(filefullname)
304
+ @staticmethod
305
+ def copy_default_pipeline_card_module(
306
+ target_dir: str,
307
+ exists_ok: bool = False
308
+ ) -> None:
309
+ os.makedirs(target_dir, exist_ok=True)
310
+ if (
311
+ not exists_ok and
312
+ os.path.exists(os.path.join(target_dir, "pipeline_card.py"))
313
+ ):
314
+ print("File already exists. Skipping copy.")
315
+ else:
316
+ filefullname = os.path.join(
317
+ UnslothFuncCallFlow.default_pipeline_card_module_dir,
318
+ "pipeline_card.py"
319
+ )
320
+ shutil.copy(filefullname, target_dir)
321
+ print(filefullname)
322
+ @staticmethod
323
+ def copy_default_pipeline_card_html_template(
324
+ target_dir: str,
325
+ exists_ok: bool = False
326
+ ) -> None:
327
+ os.makedirs(target_dir, exist_ok=True)
328
+ if (
329
+ not exists_ok and
330
+ os.path.exists(os.path.join(target_dir, "template.html"))
331
+ ):
332
+ print("File already exists. Skipping copy.")
333
+ else:
334
+ filefullname = os.path.join(
335
+ UnslothFuncCallFlow.default_pipeline_card_module_dir,
336
+ "template.html")
337
+ shutil.copy(filefullname, target_dir)
338
+ print(filefullname)
339
+
340
+ del RETRAIN_PIPELINE_TYPE
341
+
342
+ #---------------------------------------------------------------------------
343
+
344
+ @step
345
+ def start(self):
346
+ print(f"{current.flow_name} - {current.run_id}")
347
+
348
+ # GPU availability
349
+ print(torch.cuda.get_device_name(0))
350
+ print(torch.__version__)
351
+ self.engine = "gpu" if torch.cuda.is_available() else "cpu"
352
+
353
+ # hf_dataset
354
+ hf_dataset_dict = \
355
+ get_lazy_df(
356
+ repo_id=self.hf_dataset["repo_id"],
357
+ commit_hash=self.hf_dataset["commit_hash"],
358
+ files_filter=(
359
+ self.hf_dataset['config_name']+"/.*\\.parquet"
360
+ if (
361
+ self.hf_dataset["config_name"] and
362
+ "" < self.hf_dataset["config_name"]
363
+ ) else ".*\\.parquet"
364
+ ),
365
+ hf_token=os.getenv("HF_TOKEN", None)
366
+ )
367
+ try:
368
+ print(hf_dataset_dict["repo_id"], ", ",
369
+ hf_dataset_dict["commit_hash"], " - ",
370
+ hf_dataset_dict["commit_datetime"], "\n",
371
+ hf_dataset_dict["lazy_df"].explain())
372
+ except ComputeError as ex:
373
+ if "HF_TOKEN" not in os.environ:
374
+ print("Does the Hugging Face-hosted dataset " +
375
+ "require authentication ?",
376
+ file=sys.stderr, flush=True)
377
+ raise ex
378
+ self.hf_dataset_dict = hf_dataset_dict
379
+
380
+ # hf_enrich_dataset
381
+ print(self.hf_enrich_dataset)
382
+ hf_enrich_dataset_dict = \
383
+ get_lazy_df(
384
+ repo_id=self.hf_enrich_dataset["repo_id"],
385
+ commit_hash=self.hf_enrich_dataset["commit_hash"],
386
+ files_filter=(
387
+ self.hf_enrich_dataset['config_name']+"/.*\\.parquet"
388
+ if (
389
+ self.hf_enrich_dataset["config_name"] and
390
+ "" < self.hf_enrich_dataset["config_name"]
391
+ ) else ".*\\.parquet"
392
+ ),
393
+ hf_token=os.getenv("HF_TOKEN", None)
394
+ )
395
+ print(' ; '.join(f"{k}: {hf_enrich_dataset_dict[k]}"
396
+ for k in ['commit_hash',
397
+ 'commit_datetime']))
398
+ self.hf_enrich_dataset_dict = hf_enrich_dataset_dict
399
+
400
+ # hf_base_model
401
+ hf_base_model_commits = list_repo_commits(
402
+ repo_id=self.hf_base_model["repo_id"],
403
+ revision=(
404
+ None if (rev_commit_hash:=self.hf_base_model["commit_hash"]) == ""
405
+ else rev_commit_hash
406
+ ),
407
+ repo_type="model",
408
+ token=os.getenv("HF_TOKEN", None))
409
+ self.hf_base_model_dict = {
410
+ "repo_id": self.hf_base_model["repo_id"],
411
+ "commit_hash": hf_base_model_commits[0].commit_id,
412
+ "commit_datetime": \
413
+ hf_base_model_commits[0].created_at
414
+ }
415
+
416
+ self.model_version_blessed = False
417
+ self.current_blessed_run = None
418
+ self.current_blessed_version_dict = None
419
+ current.run.remove_tag("model_version_blessed")
420
+
421
+ self.retrain_pipelines = f"retrain-pipelines {__version__}"
422
+ self.retrain_pipeline_type = os.environ["retrain_pipeline_type"]
423
+
424
+ self.serving_artifacts_local_folder = \
425
+ os.path.realpath(os.path.join(
426
+ os.path.dirname(__file__),
427
+ '..', '..', 'serving_artifacts',
428
+ os.path.sep.join(current.run.path_components)
429
+ ))
430
+
431
+ if not os.path.exists(self.serving_artifacts_local_folder):
432
+ os.makedirs(self.serving_artifacts_local_folder)
433
+
434
+ self.unsloth_dir = os.path.join(
435
+ self.serving_artifacts_local_folder,
436
+ "Unsloth"
437
+ )
438
+ print(f"unsloth_dir : {self.unsloth_dir}")
439
+ self.cpt_model_dir = os.path.join(
440
+ self.unsloth_dir, "cpt_model")
441
+ self.sft_model_dir = os.path.join(
442
+ self.unsloth_dir, "sft_model")
443
+
444
+ self.next(self.eda)
445
+
446
+
447
+ @step
448
+ def eda(self):
449
+ """
450
+ exploratory data analysis.
451
+ """
452
+
453
+ ############################
454
+ # features and label #
455
+ # basic counts #
456
+ ############################
457
+ self.records_count = self.hf_dataset_dict["lazy_df"] \
458
+ .select(pl.len()).collect(engine=self.engine).item()
459
+ self.data_schema = get_column_info(
460
+ self.hf_dataset_dict["lazy_df"], engine=self.engine)
461
+ ############################
462
+
463
+ ############################
464
+ # Answers #
465
+ # tools count #
466
+ ############################
467
+ struct_schema = pl.Struct([
468
+ pl.Field("name",
469
+ pl.String
470
+ ),
471
+ pl.Field("arguments",
472
+ pl.List(pl.String) # we retrieve list of args names
473
+ # (without assigned values)
474
+ )
475
+ ])
476
+ tool_answer_occurrences_df = \
477
+ count_tool_occurrences(
478
+ self.hf_dataset_dict["lazy_df"],
479
+ self.hf_dataset["attributes"]["answers_attr"],
480
+ struct_schema) \
481
+ .collect(engine=self.engine)
482
+ print(f"{tool_answer_occurrences_df['occurrences'].sum():,} " +
483
+ f"query/tool-calls pairs")
484
+ fig = plot_tools_occurences(tool_answer_occurrences_df,
485
+ title_prefix="Dataset answers - ")
486
+ self.answers_tools_count_fig = fig
487
+ ############################
488
+
489
+ ############################
490
+ # Query #
491
+ # words count #
492
+ ############################
493
+ queries_max_length = self.hf_dataset_dict["lazy_df"].select(
494
+ pl.col(
495
+ self.hf_dataset["attributes"]["query_attr"]
496
+ ).str.len_chars().max().alias("max_query_length")
497
+ ).collect(engine=self.engine)
498
+ print(f"longuest query counts " +
499
+ f"{queries_max_length['max_query_length'][0]:,} characters")
500
+
501
+ # queries length quartiles
502
+ self.query_words_stats = \
503
+ column_words_stats(
504
+ self.hf_dataset_dict["lazy_df"],
505
+ self.hf_dataset["attributes"]["query_attr"]
506
+ ).collect(engine=self.engine)
507
+ print(self.query_words_stats.to_pandas().to_string(index=False))
508
+ print("Two thirds of the records have a query with less than " +
509
+ f"{self.query_words_stats['q3'][0]} words.")
510
+
511
+ fig = plot_words_count(
512
+ self.hf_dataset_dict["lazy_df"],
513
+ column_name=self.hf_dataset["attributes"]["query_attr"],
514
+ engine=self.engine)
515
+ self.words_count_fig = fig
516
+ ############################
517
+
518
+ ############################
519
+ # hf_enrich_dataset #
520
+ # Query words count #
521
+ ############################
522
+ enrich_question_words_stats = \
523
+ column_words_stats(
524
+ self.hf_enrich_dataset_dict['lazy_df'],
525
+ self.hf_enrich_dataset["query_attribute"],
526
+ column_attr_handler=eval(
527
+ self.hf_enrich_dataset["query_attribute_handler"])
528
+ ).collect(engine=self.engine)
529
+ print(enrich_question_words_stats.to_pandas()
530
+ .to_string(index=False))
531
+ del enrich_question_words_stats
532
+ ############################
533
+
534
+ self.next(self.augment_data)
535
+
536
+
537
+ @step
538
+ def augment_data(self):
539
+ """
540
+ Add 'negative' examples, where
541
+ queries do not trigger any tool call.
542
+ To achieve that, we sample long user queries,
543
+ truncate at half words count, and
544
+ associate this to an empty list of tool-calls.
545
+ """
546
+ """
547
+ We only consider :
548
+ - records with longuest queries,
549
+ i.e. queries in the last quartile
550
+ of "queries with most word-counts"
551
+ (this is to avoid that 'truncated' queries
552
+ get really short)
553
+ - records with answers consisting
554
+ in a single tool-call
555
+ (in order to minimize the risk
556
+ that truncating actually gives
557
+ a valid answer with
558
+ one tool-call [or more])
559
+
560
+ Note on flow 'augmentation_rate' :
561
+ we add that many records (at most),
562
+ as quartiles size permits.
563
+ """
564
+
565
+ print("Sampling within the population with more than " +
566
+ str(self.query_words_stats['q3'][0]) +
567
+ " words (longest queries quartile) =>")
568
+
569
+ samples_count = \
570
+ int(self.records_count * self.augmentation_rate)
571
+ print(f"would represent {samples_count:,.0f} " +
572
+ f"records to be sampled")
573
+
574
+ eligible_records_df = \
575
+ self.hf_dataset_dict["lazy_df"].filter(
576
+ pl.col(
577
+ self.hf_dataset["attributes"]["query_attr"]
578
+ )
579
+ .str.extract_all(r"\w+")
580
+ .map_elements(
581
+ lambda arr: len(arr),
582
+ return_dtype=pl.Int16)
583
+ .gt(self.query_words_stats['q3'][0])
584
+ & pl.col("answers")
585
+ .map_elements(
586
+ lambda x: len(json.loads(x)) == 1
587
+ if isinstance(x, str)
588
+ else False,
589
+ return_dtype=pl.Boolean)
590
+ ) \
591
+ .collect(engine=self.engine)
592
+ eligible_records_count = \
593
+ eligible_records_df.select(pl.len())["len"][0]
594
+ print(f"eligible_records_count : " +
595
+ f"{eligible_records_count:,.0f}")
596
+ samples_count = min(samples_count, eligible_records_count)
597
+ self.actual_augmentation_rate = \
598
+ samples_count / self.records_count
599
+ print("actual augmentation rate : " +
600
+ f"{self.actual_augmentation_rate:.1%}")
601
+ sampled_records_df = eligible_records_df.sample(
602
+ n=samples_count
603
+ )
604
+
605
+ self.augmented_records_df = \
606
+ sampled_records_df.with_columns(
607
+ pl.col("query")
608
+ .map_elements(
609
+ lambda query:
610
+ " ".join(
611
+ query.split()[
612
+ :len(query.split()) // 2]),
613
+ return_dtype=pl.Utf8)
614
+ .alias("truncated_query")
615
+ ).select([
616
+ pl.col("truncated_query").alias("query"),
617
+ pl.lit("[]").alias("answers")
618
+ ])
619
+ print(self.augmented_records_df.height,
620
+ self.augmented_records_df.columns)
621
+
622
+ self.next(self.enrich_data)
623
+
624
+
625
+ @step
626
+ def enrich_data(self):
627
+ """
628
+ Further enrich our dataset with 'negative' records from
629
+ another dataset (can be general-purpose text dataset)
630
+ as specified by the the flow 'hf_enrich_dataset' argument.
631
+ """
632
+ """
633
+ Note : we here use the Hugging Face `datasets` library
634
+ in 'streaming' mode for records sampling.
635
+ """
636
+
637
+ hf_enrich_ds = load_dataset(
638
+ path=self.hf_enrich_dataset["repo_id"],
639
+ name=self.hf_enrich_dataset["config_name"],
640
+ revision=self.hf_enrich_dataset_dict["commit_hash"],
641
+ streaming=True)
642
+ print(hf_enrich_ds["train"])
643
+
644
+ samples_count = \
645
+ int(self.records_count * self.enrichment_rate)
646
+ print(f"Samplig {samples_count:,.0f} records")
647
+
648
+ query_attribute_handler = \
649
+ eval(self.hf_enrich_dataset["query_attribute_handler"])
650
+ samples_iterator = iterable_dataset_multi_buffer_sampler(
651
+ hf_enrich_ds["train"],
652
+ total_samples=samples_count,
653
+ attributes_selector=\
654
+ (lambda x:query_attribute_handler(
655
+ x[self.hf_enrich_dataset["query_attribute"]])),
656
+ buffer_size=3_000,
657
+ num_passes=3,
658
+ seed=None
659
+ )
660
+ # Capitalize and add end punctuation if missing
661
+ start_time = time.time()
662
+ print("Starting sample enriching records, " +
663
+ "this may take some time if the source dataset " +
664
+ "has a complex structure..")
665
+ samples_list = [
666
+ s.capitalize() + ("" if s[-1] in ".!?" else "?")
667
+ for s in samples_iterator]
668
+ elapsed_time = time.time() - start_time
669
+ print(f".. sampling completed " +
670
+ f"({int(elapsed_time // 3_600)}h:" +
671
+ f"{int((elapsed_time % 3_600) // 60)}m:" +
672
+ f"{int(elapsed_time % 60)}s).")
673
+ enriched_records_df = pl.DataFrame(
674
+ {"query": samples_list,
675
+ "answers": \
676
+ ["[]"] * \
677
+ len(samples_list)}
678
+ )
679
+ self.enriched_records_df = enriched_records_df
680
+
681
+ self.next(self.dataset_to_hub)
682
+
683
+
684
+ @step
685
+ def dataset_to_hub(self):
686
+ """
687
+ Push to hub dataset version
688
+ - continued pre-training dataset
689
+ - training and validation splits of the
690
+ augmented and enriched
691
+ supervised finetuning dataset
692
+ - readme with versioning info
693
+ """
694
+
695
+ #############################
696
+ # case of user-provided #
697
+ # documentation artifact(s) #
698
+ #############################
699
+ # note that user can provide either
700
+ # 'pipeline_card.py' or 'template.html'
701
+ # or 'dataset_readme.py'
702
+ # or 'dataset_readme_template.md'
703
+ # or 'model_readme.py'
704
+ # or 'model_readme_template.md'
705
+ # or any combination of those
706
+ # when specifying custom
707
+ # 'pipeline_card_artifacts_path'
708
+ if (
709
+ "dataset_readme_template.md" in
710
+ os.listdir(self.pipeline_card_artifacts_path)
711
+ ):
712
+ template_dir = self.pipeline_card_artifacts_path
713
+ else:
714
+ template_dir = os.path.dirname(
715
+ importlib.util.find_spec(
716
+ f"retrain_pipelines.pipeline_card."+
717
+ f"{os.getenv('retrain_pipeline_type')}"
718
+ ).origin)
719
+ print(f"template_dir : '{template_dir}'")
720
+ #############################
721
+ if "dataset_readme.py" in os.listdir(
722
+ self.pipeline_card_artifacts_path):
723
+ from retrain_pipelines.utils import \
724
+ get_get_dataset_readme_content
725
+ get_dataset_readme_content = \
726
+ get_get_dataset_readme_content(
727
+ self.pipeline_card_artifacts_path)
728
+ else:
729
+ from retrain_pipelines.pipeline_card import \
730
+ get_dataset_readme_content
731
+ #############################
732
+
733
+
734
+ #############################
735
+ # augmented & enriched #
736
+ # finetuning dataset #
737
+ #############################
738
+ merged_df = pl.concat([
739
+ # dataset
740
+ self.hf_dataset_dict["lazy_df"].select([
741
+ self.hf_dataset["attributes"]["query_attr"],
742
+ self.hf_dataset["attributes"]["answers_attr"]
743
+ ]).collect(engine=self.engine),
744
+ # truncated queries augmentation
745
+ self.augmented_records_df,
746
+ # enriching dataset
747
+ self.enriched_records_df
748
+ ]).sample(
749
+ # shuffling
750
+ fraction=1,
751
+ shuffle=True,
752
+ with_replacement=False
753
+ )
754
+ merged_df = merged_df.sample(fraction=1, shuffle=True)
755
+ merged_df.rechunk()
756
+ print(("merged_df", f"{merged_df.shape[0]:,.0F}",
757
+ merged_df.columns))
758
+
759
+ pandas_df = merged_df.to_pandas()
760
+ train_size = int(0.8 * len(pandas_df))
761
+ print(f"validation : {len(pandas_df) - train_size}")
762
+ sft_dataset = DatasetDict({
763
+ "train": Dataset.from_pandas(pandas_df[:train_size]),
764
+ "validation": Dataset.from_pandas(pandas_df[train_size:])
765
+ })
766
+ #############################
767
+
768
+ #############################
769
+ # continued pre-training #
770
+ # dataset #
771
+ #############################
772
+ struct_schema = pl.Struct([
773
+ pl.Field("name", pl.String),
774
+ pl.Field("description", pl.String),
775
+ pl.Field(
776
+ "parameters",
777
+ pl.String # Use String to allow
778
+ # for varying structures
779
+ # (different tools indeed having
780
+ # different sets of parameters
781
+ # i.e. different parameters counts,
782
+ # datatypes and names)
783
+ # so parsing must be tolerant.
784
+ )
785
+ ])
786
+ unique_tools_df = get_unique_tools(
787
+ self.hf_dataset_dict["lazy_df"],
788
+ tools_attr_name=\
789
+ self.hf_dataset["attributes"]["tools_attr"],
790
+ struct_schema=struct_schema
791
+ ).collect(engine=self.engine)
792
+ unique_tools_arrow_table = unique_tools_df.to_arrow()
793
+ self.unique_tools_dataset = \
794
+ Dataset(unique_tools_arrow_table)
795
+ print(self.unique_tools_dataset)
796
+ #############################
797
+
798
+ #############################
799
+ # DatasetDict #
800
+ # with multiple tables #
801
+ #############################
802
+ dataset_dict = DatasetDict({
803
+ "continued_pre_training": \
804
+ self.unique_tools_dataset,
805
+ "supervised_finetuning": sft_dataset
806
+ })
807
+ print(dataset_dict, flush=True)
808
+ #############################
809
+
810
+ #############################
811
+ # dataset README #
812
+ # from template #
813
+ #############################
814
+ commit_datetime = datetime.utcnow()
815
+ new_dataset_version_label = get_new_repo_minor_version(
816
+ repo_id=self.dataset_repo_id,
817
+ repo_type="dataset",
818
+ hf_token=os.getenv("HF_TOKEN", None))
819
+ readme_content = get_dataset_readme_content(
820
+ template_folder=template_dir,
821
+
822
+ hf_dataset_dict=self.hf_dataset_dict,
823
+ hf_enrich_dataset_dict=self.hf_enrich_dataset_dict,
824
+ dataset_dict=dataset_dict,
825
+
826
+ augmentation_rate=self.actual_augmentation_rate,
827
+ enrichment_rate=self.enrichment_rate,
828
+
829
+ version_label=new_dataset_version_label,
830
+ commit_datetime=commit_datetime,
831
+
832
+ mf_flow_name=current.flow_name,
833
+ mf_run_id=current.run.id,
834
+ engine=self.engine
835
+ )
836
+ #############################
837
+
838
+ dataset_commit_hash = push_dataset_version_to_hub(
839
+ repo_id=self.dataset_repo_id,
840
+ version_label=new_dataset_version_label,
841
+ timestamp_str=commit_datetime.strftime(
842
+ "%Y-%m-%d %H:%M:%S UTC"),
843
+ dataset_dict=dataset_dict,
844
+ dataset_readme_content=readme_content,
845
+ hf_token=os.getenv("HF_TOKEN", None)
846
+ )
847
+ if not dataset_commit_hash:
848
+ raise Exception(
849
+ "Failed to publish dataset version.")
850
+ print(f"https://huggingface.co/datasets/{self.dataset_repo_id}" +
851
+ f"/blob/{dataset_commit_hash}/README.md")
852
+ self.dataset_commit_dict = {
853
+ "repo_id": self.dataset_repo_id,
854
+ "commit_hash": dataset_commit_hash,
855
+ "version_label": new_dataset_version_label,
856
+ "commit_datetime": commit_datetime,
857
+ }
858
+
859
+ self.next(self.continued_pre_training)
860
+
861
+
862
+ @step
863
+ def continued_pre_training(self):
864
+ """
865
+ Gives the base model some additional intrinsic knowkledge
866
+ through continued pre-training.
867
+ See unsloth.ai/blog/contpretraining
868
+ """
869
+ from retrain_pipelines.model.hf_utils import \
870
+ plot_log_history
871
+
872
+ #######################################
873
+ # base-model and associated tokenizer #
874
+ # from Hub (or local cache) #
875
+ #######################################
876
+ self.max_seq_length = 2048
877
+ model, tokenizer = FastLanguageModel.from_pretrained(
878
+ model_name=self.hf_base_model_dict["repo_id"],
879
+ revision=self.hf_base_model_dict["commit_hash"],
880
+ max_seq_length=self.max_seq_length,
881
+ dtype=None,
882
+ load_in_4bit=False,
883
+ # case of a gated or private base-model
884
+ token=os.getenv("HF_TOKEN", None)
885
+ )
886
+ #######################################
887
+
888
+ #######################################
889
+ # dataset prompt_template mapping #
890
+ #######################################
891
+ tools_dataset = DatasetDict(
892
+ {"train": self.unique_tools_dataset})
893
+ print(tools_dataset)
894
+ tool_prompt_template = "tool: {}"
895
+ def formatting_prompts_func(tools_batch):
896
+ tools_batch = tools_batch["tool"]
897
+ outputs = []
898
+ for tool in tools_batch:
899
+ # Must add EOS_TOKEN,
900
+ # otherwise generation will go on forever!
901
+ text = tool_prompt_template.format(tool) + \
902
+ tokenizer.eos_token
903
+ outputs.append(text)
904
+ return { "tools" : outputs, }
905
+ cpt_dataset = tools_dataset["train"].map(
906
+ formatting_prompts_func, batched=True,)
907
+ #######################################
908
+
909
+ #######################################
910
+ # PEFT adapter #
911
+ # for continued pre-training #
912
+ #######################################
913
+ model = FastLanguageModel.get_peft_model(
914
+ model,
915
+ r = 128, # any number >0 ; 8, 16, 32, 64, 128, 256
916
+ target_modules = ["q_proj", "k_proj", "v_proj", "o_proj",
917
+ "gate_proj", "up_proj", "down_proj",
918
+ # Add for continued pretraining
919
+ "embed_tokens", "lm_head",],
920
+ lora_alpha = 32,
921
+ lora_dropout = 0, # Supports any, 0 is optimized
922
+ bias = "none", # Supports any, "none" is optimized
923
+ # True or "unsloth" for very long context
924
+ use_gradient_checkpointing = "unsloth",
925
+ use_rslora = True, # rank-stabilized LoRA
926
+ loftq_config = None, # LoftQ
927
+ #random_state = 3407,
928
+ )
929
+ #######################################
930
+
931
+ #######################################
932
+ # cpt_trainer #
933
+ #######################################
934
+ if (
935
+ "records_cap" in self.cpt_training_args and
936
+ self.cpt_training_args["records_cap"] is not None and
937
+ isinstance(self.cpt_training_args["records_cap"], int)
938
+ ):
939
+ cpt_dataset = cpt_dataset.take(
940
+ self.cpt_training_args["records_cap"])
941
+ print(f"cpt_dataset : {cpt_dataset}")
942
+
943
+ train_args = UnslothTrainingArguments(
944
+ # https://huggingface.co/docs/transformers/main_classes/trainer#transformers.TrainingArguments.save_strategy
945
+ per_device_train_batch_size=2,
946
+ gradient_accumulation_steps=8,
947
+
948
+ **{k: v for k, v in self.cpt_training_args.items()
949
+ if k != "records_cap"},
950
+
951
+ # 2 to 10x smaller learning rate
952
+ # for the embedding matrices
953
+ learning_rate=5e-5,
954
+ embedding_learning_rate=1e-5,
955
+
956
+ fp16=not is_bfloat16_supported(),
957
+ bf16=is_bfloat16_supported(),
958
+ logging_steps=1,
959
+ optim="adamw_8bit",
960
+ weight_decay=0.01,
961
+ lr_scheduler_type="linear",
962
+ #seed=3407,
963
+
964
+ output_dir=os.path.join(
965
+ self.unsloth_dir, "outputs", "cpt"),
966
+ save_total_limit = 2,
967
+
968
+ report_to="tensorboard",
969
+ logging_dir=os.path.join(
970
+ self.sft_model_dir,
971
+ "runs", "cpt")
972
+ )
973
+
974
+ self.cpt_traces_file_fullname = os.path.join(
975
+ self.unsloth_dir, "cpt_trainer_traces.txt")
976
+ print("Training started. " +
977
+ f"Check {self.cpt_traces_file_fullname} for live traces.",
978
+ flush=True)
979
+
980
+ trainer = UnslothTrainer(
981
+ model=model, tokenizer=tokenizer,
982
+ train_dataset=cpt_dataset,
983
+ dataset_text_field="tools",
984
+ max_seq_length=self.max_seq_length,
985
+ dataset_num_proc=2,
986
+ args=train_args,
987
+ )
988
+ #######################################
989
+
990
+ #######################################
991
+ # Show current memory stats #
992
+ #######################################
993
+ torch.cuda.ipc_collect()
994
+ torch.cuda.empty_cache()
995
+ _ = gc.collect()
996
+
997
+ gpu_stats = torch.cuda.get_device_properties(0)
998
+ self.start_gpu_memory = \
999
+ round(torch.cuda.max_memory_reserved()
1000
+ / 1024 / 1024 / 1024, 3)
1001
+ self.max_memory = \
1002
+ round(gpu_stats.total_memory
1003
+ / 1024 / 1024 / 1024, 3)
1004
+ print(f"GPU = {gpu_stats.name}. " +
1005
+ f"Max memory = {self.max_memory} GB.")
1006
+ print(f"{self.start_gpu_memory} GB of memory reserved.")
1007
+ #######################################
1008
+
1009
+ with open(self.cpt_traces_file_fullname, 'w') as f:
1010
+ with redirect_stdout(f):
1011
+ transformers_logging.set_verbosity_error()
1012
+ transformers_logging.disable_progress_bar()
1013
+ trainer_stats = trainer.train()
1014
+ transformers_logging.set_verbosity_info()
1015
+ transformers_logging.enable_progress_bar()
1016
+ print(f"{trainer_stats.metrics['train_runtime']} " +
1017
+ f"seconds used for training " +
1018
+ f"({round(trainer_stats.metrics['train_runtime']/60, 2)}" +
1019
+ f" minutes).")
1020
+
1021
+ self.cpt_log_history = trainer.state.log_history
1022
+ # print(self.cpt_log_history)
1023
+ self.cpt_log_history_fig = \
1024
+ plot_log_history(
1025
+ self.cpt_log_history,
1026
+ title="Continued pretraining loss"
1027
+ )
1028
+
1029
+ model.save_pretrained_merged(
1030
+ save_directory=self.cpt_model_dir,
1031
+ tokenizer=tokenizer,
1032
+ save_method="lora"
1033
+ )
1034
+ print(f"cpt_model_dir : {self.cpt_model_dir}\n")
1035
+
1036
+ self.next(self.supervised_finetuning)
1037
+
1038
+
1039
+ @step
1040
+ def supervised_finetuning(self):
1041
+ """
1042
+ Trains the model on tool-calling
1043
+ task specialization.
1044
+ """
1045
+ from retrain_pipelines.model.hf_utils import \
1046
+ plot_log_history
1047
+
1048
+ torch.cuda.ipc_collect()
1049
+ torch.cuda.empty_cache()
1050
+ _ = gc.collect()
1051
+
1052
+ model, tokenizer = FastLanguageModel.from_pretrained(
1053
+ model_name=self.cpt_model_dir,
1054
+ max_seq_length=self.max_seq_length,
1055
+ dtype=None,
1056
+ load_in_4bit=False,
1057
+ )
1058
+ # !!!! bug fix BEGIN !!!!
1059
+ # otherwise, 'embed_tokens' and 'lm_head'
1060
+ # trained during CPT are "ignored",
1061
+ # i.e. not saved after SFT
1062
+ # (note that, alternatively, we could also
1063
+ # do this fix after sft-training and
1064
+ # just before saving ;
1065
+ # which would be equivalent to
1066
+ # freezing embeddings during finetuning
1067
+ # for better pretrained knowledge retention)
1068
+ # @see https://www.reddit.com/r/unsloth/comments/1dtzcd6/fastlanguagemodelpatch_peft_model_changing/
1069
+ model.model.model.embed_tokens.modules_to_save.default.to(
1070
+ device="cuda:0",
1071
+ dtype=torch.float32,
1072
+ non_blocking=True)
1073
+ model.model.model.embed_tokens.modules_to_save.default \
1074
+ .requires_grad_(True)
1075
+ model.model.lm_head.modules_to_save.default.to(
1076
+ device="cuda:0",
1077
+ dtype=torch.float32,
1078
+ non_blocking=True)
1079
+ model.model.lm_head.modules_to_save.default \
1080
+ .requires_grad_(True)
1081
+ # !!!! bug fix END !!!!
1082
+
1083
+ #######################################
1084
+ # dataset prompt_template mapping #
1085
+ #######################################
1086
+ # download from Hub (or get from local cache)
1087
+ queries_dataset = load_dataset(
1088
+ path=self.dataset_commit_dict["repo_id"],
1089
+ name="supervised_finetuning",
1090
+ revision=self.dataset_commit_dict["commit_hash"],
1091
+ token=os.getenv("HF_TOKEN", None))
1092
+ print(f"HF_DATASETS_CACHE : {HF_DATASETS_CACHE}") # HF_CACHE_HOME
1093
+ self.sft_prompt_template = dedent("""
1094
+ You specialize in generating tool calls. Given a query, your task is to return a list of tool calls based on your knowledge of known tools.
1095
+
1096
+ Rules:
1097
+ 1. You can only use tools you know. Do not create new tools under any circumstances.
1098
+ 2. If a query does not match any known tool, return an empty list ([]).
1099
+ 3. If information is missing to use a known tool, do not attempt to use it.
1100
+ 4. Your response must always be a valid JSON array, and nothing else.
1101
+
1102
+ Be precise and do not guess.
1103
+
1104
+ # query:
1105
+ {}
1106
+ # response:
1107
+ {}
1108
+ """).strip()
1109
+ tokenizer.chat_template = self.sft_prompt_template
1110
+
1111
+ EOS_TOKEN = tokenizer.eos_token
1112
+ def formatting_prompts_func(records):
1113
+ query = records["query"]
1114
+ tools = records["answers"]
1115
+ outputs = []
1116
+ for query, tools in zip(query, tools):
1117
+ # Must add EOS_TOKEN,
1118
+ # otherwise your generation will go on forever
1119
+ text = self.sft_prompt_template.format(query, tools) \
1120
+ + EOS_TOKEN
1121
+ outputs.append(text)
1122
+ return { "text" : outputs, }
1123
+ sft_train_dataset = queries_dataset["train"].map(
1124
+ formatting_prompts_func, batched=True)
1125
+ sft_valid_dataset = queries_dataset["validation"].map(
1126
+ formatting_prompts_func, batched=True,)
1127
+ #######################################
1128
+
1129
+ #######################################
1130
+ # PEFT adapter #
1131
+ # for supervised finetuning #
1132
+ #######################################
1133
+ # for cases where CPT has been merged into overall model
1134
+ # otherwize, keep on training current LoRa adapter
1135
+ # model = FastLanguageModel.get_peft_model(
1136
+ # model,
1137
+ # r = 128, # any number >0 ; 8, 16, 32, 64, 128, 256
1138
+ # target_modules = ["q_proj", "k_proj", "v_proj", "o_proj",
1139
+ # "gate_proj", "up_proj", "down_proj"],
1140
+ # lora_alpha = 32,
1141
+ # lora_dropout = 0, # Supports any, but = 0 is optimized
1142
+ # bias = "none", # Supports any, but = "none" is optimized
1143
+ # # True or "unsloth" for very long context
1144
+ # use_gradient_checkpointing = "unsloth",
1145
+ # random_state = 3407,
1146
+ # use_rslora = True, # rank stabilized LoRA
1147
+ # loftq_config = None, # LoftQ
1148
+ # )
1149
+ #######################################
1150
+
1151
+ #######################################
1152
+ # sft_trainer #
1153
+ #######################################
1154
+ split = sft_train_dataset.train_test_split(
1155
+ test_size=1000,
1156
+ #seed=42
1157
+ )
1158
+ train_dataset = split['train']
1159
+ eval_dataset = split['test']
1160
+ if (
1161
+ "records_cap" in self.sft_training_args and
1162
+ self.sft_training_args["records_cap"] is not None and
1163
+ isinstance(self.sft_training_args["records_cap"], int)
1164
+ ):
1165
+ train_dataset = train_dataset.take(
1166
+ self.sft_training_args["records_cap"])
1167
+ eval_dataset = eval_dataset.take(
1168
+ self.sft_training_args["records_cap"])
1169
+ print(f"train_dataset : {train_dataset}")
1170
+ print(f"eval_dataset : {eval_dataset}")
1171
+
1172
+ train_args = UnslothTrainingArguments(
1173
+ per_device_train_batch_size=2,
1174
+ gradient_accumulation_steps=8,
1175
+
1176
+ **{k: v for k, v in self.sft_training_args.items()
1177
+ if k != "records_cap"},
1178
+
1179
+ per_device_eval_batch_size=2,
1180
+ eval_steps=200,
1181
+ eval_strategy="steps",
1182
+ do_eval=True,
1183
+
1184
+ learning_rate=5e-5,
1185
+ # embedding_learning_rate=1e-5, # Optionally here
1186
+
1187
+ fp16=not is_bfloat16_supported(),
1188
+ bf16=is_bfloat16_supported(),
1189
+
1190
+ optim="adamw_8bit",
1191
+ weight_decay=0.00,
1192
+ lr_scheduler_type="linear",
1193
+ #seed=3407,
1194
+
1195
+ output_dir=os.path.join(
1196
+ self.unsloth_dir, "outputs", "sft"),
1197
+ save_total_limit=2,
1198
+
1199
+ logging_steps=1,
1200
+ report_to="tensorboard",
1201
+ logging_dir=os.path.join(
1202
+ self.sft_model_dir,
1203
+ "runs", "sft")
1204
+ )
1205
+
1206
+ self.sft_traces_file_fullname = os.path.join(
1207
+ self.unsloth_dir, "sft_trainer_traces.txt")
1208
+ print("Training started. " +
1209
+ f"Check {self.sft_traces_file_fullname} for live traces.",
1210
+ flush=True)
1211
+
1212
+ trainer = UnslothTrainer(
1213
+ model=model, tokenizer=tokenizer,
1214
+ train_dataset=train_dataset,
1215
+ dataset_text_field="text",
1216
+ eval_dataset=eval_dataset,
1217
+ max_seq_length=self.max_seq_length,
1218
+ dataset_num_proc=8,
1219
+ args=train_args
1220
+ )
1221
+ trainer.can_return_loss = True
1222
+ #######################################
1223
+
1224
+ #######################################
1225
+ # Show current memory stats #
1226
+ #######################################
1227
+ torch.cuda.ipc_collect()
1228
+ torch.cuda.empty_cache()
1229
+ _ = gc.collect()
1230
+
1231
+ used_memory = \
1232
+ round(torch.cuda.max_memory_reserved()
1233
+ /1024/1024/1024, 3)
1234
+ used_memory_for_lora = \
1235
+ round(used_memory-self.start_gpu_memory, 3)
1236
+ used_percentage = \
1237
+ round(used_memory/self.max_memory*100, 3)
1238
+ lora_percentage = \
1239
+ round(used_memory_for_lora/self.max_memory*100,
1240
+ 3)
1241
+ print(f"Peak reserved memory = " +
1242
+ f"{used_memory} GB.")
1243
+ print(f"Peak reserved memory for " +
1244
+ f"training = {used_memory_for_lora} " +
1245
+ f"GB.")
1246
+ print(f"Peak reserved memory % of " +
1247
+ f"max memory = {used_percentage} %.")
1248
+ print(f"Peak reserved memory for training " +
1249
+ f"% of max memory = {lora_percentage} %.")
1250
+ #######################################
1251
+
1252
+ with open(self.sft_traces_file_fullname, 'w') as f:
1253
+ with redirect_stdout(f):
1254
+ hf_hub_disable_progress_bars()
1255
+ trainer_stats = trainer.train()
1256
+ hf_hub_enable_progress_bars()
1257
+ print(f"{trainer_stats.metrics['train_runtime']} " +
1258
+ f"seconds used for training " +
1259
+ f"({round(trainer_stats.metrics['train_runtime']/60, 2)}" +
1260
+ f" minutes).")
1261
+
1262
+ self.sft_log_history = trainer.state.log_history
1263
+ self.sft_log_history_fig = \
1264
+ plot_log_history(
1265
+ self.sft_log_history,
1266
+ title="Supervised finetuning loss"
1267
+ )
1268
+
1269
+ model.save_pretrained_merged(
1270
+ self.sft_model_dir, tokenizer,
1271
+ save_method = "lora"
1272
+ )
1273
+ print(f"sft_model_dir : {self.sft_model_dir}\n")
1274
+
1275
+ self.next(self.evaluate_model)
1276
+
1277
+
1278
+ @step
1279
+ def evaluate_model(self):
1280
+ """
1281
+ Batch inference on the SFT validation dataset.
1282
+ """
1283
+ from retrain_pipelines.model import \
1284
+ infer_validation, compute_counts_n_metrics, \
1285
+ plot_validation_completions
1286
+
1287
+ torch.cuda.ipc_collect()
1288
+ torch.cuda.empty_cache()
1289
+ _ = gc.collect()
1290
+
1291
+
1292
+ ######################################################
1293
+ # loading trained adapter #
1294
+ ######################################################
1295
+ # Unsloth [and hf transformers before it] #
1296
+ # (if loading both model & tokenizer at once #
1297
+ # same as we did in prior tasks, but now #
1298
+ # with tokenizer.chat_template being set #
1299
+ # in tokenizer.config) is forcing on us some kind of #
1300
+ # chat_template format hard-requirements. #
1301
+ ######################################################
1302
+ # load base from cache
1303
+ # (with base tokenizer, which we ignore)
1304
+ model, _ = FastLanguageModel.from_pretrained(
1305
+ model_name=self.hf_base_model_dict["repo_id"],
1306
+ revision=self.hf_base_model_dict["commit_hash"],
1307
+ max_seq_length=self.max_seq_length,
1308
+ dtype=None,
1309
+ load_in_4bit=False,
1310
+ # case of a gated or private base-model
1311
+ token=os.getenv("HF_TOKEN", None)
1312
+ )
1313
+ model = FastLanguageModel.for_inference(model)
1314
+ # load our CPT+SFT trained & locally-saved adapter
1315
+ model.load_adapter(peft_model_id=self.sft_model_dir)
1316
+ # Separately load our (potentially trained &)
1317
+ # locally-saved adapter-tokenizer
1318
+ # (loading it below via HF and not Unsloth)
1319
+ tokenizer = AutoTokenizer.from_pretrained(
1320
+ pretrained_model_name_or_path=self.sft_model_dir
1321
+ )
1322
+ ######################################################
1323
+
1324
+ ######################################################
1325
+ # validation dataset #
1326
+ ######################################################
1327
+ # download from Hub (or get from local cache)
1328
+ queries_dataset = load_dataset(
1329
+ path=self.dataset_commit_dict["repo_id"],
1330
+ name="supervised_finetuning",
1331
+ revision=self.dataset_commit_dict["commit_hash"],
1332
+ token=os.getenv("HF_TOKEN", None))
1333
+ if (
1334
+ "records_cap" in self.sft_training_args and
1335
+ self.sft_training_args["records_cap"] is not None and
1336
+ isinstance(self.sft_training_args["records_cap"], int)
1337
+ ):
1338
+ validation_data = queries_dataset["validation"].take(
1339
+ self.sft_training_args["records_cap"])
1340
+ else:
1341
+ validation_data = queries_dataset["validation"]
1342
+ print(validation_data, flush=True)
1343
+ ######################################################
1344
+
1345
+ self.max_new_tokens = 400
1346
+ start_time = time.time()
1347
+ validation_results = infer_validation(
1348
+ tokenizer=tokenizer,
1349
+ model=model,
1350
+ validation_data=validation_data,
1351
+ prompt_template=tokenizer.chat_template,
1352
+ batch_size=32, # 64,
1353
+ queries_attr_name=\
1354
+ self.hf_dataset["attributes"]["query_attr"],
1355
+ answers_attr_name=\
1356
+ self.hf_dataset["attributes"]["answers_attr"],
1357
+ max_new_tokens=self.max_new_tokens,
1358
+ device="cuda"
1359
+ )
1360
+ print("infer_validation - Elapsed time: " +
1361
+ f"{(time.time() - start_time):.2f} seconds")
1362
+ self.validation_results = validation_results # <= to artifacts store
1363
+
1364
+ eval_df = pl.LazyFrame(validation_results)
1365
+
1366
+ records = eval_df.with_columns(
1367
+ (pl.col("answer") == pl.col("completion")) \
1368
+ .alias("is_ground_truth_identical")
1369
+ ).collect() #engine=self.engine)
1370
+ print("perfect characters-match accuracy : " +
1371
+ str(records['is_ground_truth_identical'].mean()))
1372
+
1373
+ eval_metrics_df = compute_counts_n_metrics(
1374
+ eval_df, is_format_fault_tolerant=True)
1375
+ overall_metrics_df = eval_metrics_df.select([
1376
+ pl.col("precision").mean(),
1377
+ pl.col("recall").mean(),
1378
+ pl.col("f1").mean(),
1379
+ pl.col("jaccard").mean()
1380
+ ]).collect() #engine=self.engine)
1381
+ self.perf_metrics = overall_metrics_df.row(0, named=True)
1382
+ print(self.perf_metrics)
1383
+
1384
+ self.validation_completions_fig = \
1385
+ plot_validation_completions(
1386
+ eval_metrics_df, engine=self.engine)
1387
+
1388
+ del model
1389
+ del tokenizer
1390
+ torch.cuda.ipc_collect()
1391
+ torch.cuda.empty_cache()
1392
+ _ = gc.collect()
1393
+
1394
+ self.next(self.model_version_blessing)
1395
+
1396
+
1397
+ @step
1398
+ def model_version_blessing(self):
1399
+ """
1400
+ Comparing newly-retrained model version
1401
+ against best-performing predecessor.
1402
+ """
1403
+ """
1404
+ Note: for Hugging Face integrated pipelines,
1405
+ we compare against lastest commit of main branch
1406
+ of the model repository there.
1407
+ When it comes to local "mf_run_id" of the pipeline run
1408
+ having generated that best prior model version
1409
+ (retrieved from model card metadata from HF yaml section),
1410
+ we check against records of the herein ML-framework instance,
1411
+ as "prior best version" of the model here beign retrained
1412
+ may have been originated from another one
1413
+ than the one executing the current retraining
1414
+ (in which case, we simply don't includ a "local" hyperlink
1415
+ in the model version pipeline_cards that will be
1416
+ produced later in the herein pipeline run).
1417
+ """
1418
+ from retrain_pipelines.model.hf_utils import \
1419
+ current_blessed_model_version_dict
1420
+
1421
+ main_perf_metric_name = "jaccard"
1422
+
1423
+ current_blessed_version_dict = \
1424
+ current_blessed_model_version_dict(
1425
+ repo_id=self.model_repo_id,
1426
+ hf_token=os.getenv("HF_TOKEN", None)
1427
+ )
1428
+ print("current_blessed_version_dict : " +
1429
+ str(current_blessed_version_dict))
1430
+
1431
+ if current_blessed_version_dict is None:
1432
+ print("case 'no prior blessed model version found"
1433
+ " => blessing.'")
1434
+ self.model_version_blessed = True
1435
+
1436
+ elif (
1437
+ main_perf_metric_name in
1438
+ current_blessed_version_dict["perf_metrics"]
1439
+ ):
1440
+ current_blessed_run_id = \
1441
+ current_blessed_version_dict["mf_run_id"]
1442
+ print(f"current_blessed_run_id : {current_blessed_run_id}")
1443
+ current_blessed_metric_value = \
1444
+ current_blessed_version_dict[
1445
+ "perf_metrics"][main_perf_metric_name]
1446
+
1447
+ self.model_version_blessed = (
1448
+ self.perf_metrics[main_perf_metric_name] >=
1449
+ current_blessed_metric_value
1450
+ )
1451
+
1452
+ if not self.model_version_blessed:
1453
+ self.current_blessed_version_dict = \
1454
+ current_blessed_version_dict
1455
+ for run in Flow(self.__class__.__name__):
1456
+ if str(run.id) == current_blessed_run_id:
1457
+ run_steps = iter(run.steps())
1458
+ last_run_step = next(run_steps)
1459
+ last_task = next(iter(last_run_step.tasks()))
1460
+
1461
+ # tasks are listed backwards, so last task is first item :
1462
+ # Has the run seen task "pipeline_card" prior to last task
1463
+ # (meaning, "pipeline_card" completed successfully and
1464
+ # "run" has generated a sutom pipeline-card artifact) ?
1465
+ # If not, hyperlink generation will later fail.
1466
+ run_has_custom_card_artifact = False
1467
+ for step in run_steps:
1468
+ if "pipeline_card" == step.id:
1469
+ run_has_custom_card_artifact = True
1470
+ break
1471
+
1472
+ if not run_has_custom_card_artifact:
1473
+ print(
1474
+ f"Run #{current_blessed_run_id} " +
1475
+ "Doesn't seem to have successfully " +
1476
+ "generated a pipeline-card artifact.",
1477
+ file=sys.stderr, flush=True)
1478
+ break
1479
+ else:
1480
+ # further filtering on successful runs that are
1481
+ # retraining of a prior version of the same model
1482
+ # (to minimize the risk that this was obtained
1483
+ # on another ML-framework instance)
1484
+ if (
1485
+ # last_task.successful and
1486
+ # may have failed after the "pipeline_card" step
1487
+ # and been resumed
1488
+ hasattr(last_task.artifacts,
1489
+ 'model_version_blessed') and
1490
+ last_task.artifacts.model_version_blessed.data and
1491
+ hasattr(last_task.artifacts,
1492
+ 'model_repo_id') and
1493
+ last_task.artifacts.model_repo_id.data == \
1494
+ self.model_repo_id
1495
+ ):
1496
+ self.current_blessed_run = run
1497
+ break
1498
+
1499
+ if not self.current_blessed_run:
1500
+ print(
1501
+ "Couldn't find blessed run " +
1502
+ f"{current_blessed_run_id} !\n" +
1503
+ "It seems that prior blessed run was " +
1504
+ "executed on another ML framework instance.",
1505
+ file=sys.stderr, flush=True)
1506
+
1507
+ print("new : " +
1508
+ str(self.perf_metrics[main_perf_metric_name]) +
1509
+ " - previous best : " +
1510
+ str(current_blessed_metric_value) +
1511
+ " - model_version_blessing : " +
1512
+ str(self.model_version_blessed))
1513
+
1514
+ else:
1515
+ raise Exception(
1516
+ "Performance metric '" +
1517
+ main_perf_metric_name +
1518
+ "' can't be found in eval results " +
1519
+ "from blessed run " +
1520
+ str(current_blessed_version_dict[
1521
+ "mf_run_id"]) + " !")
1522
+
1523
+ # self.model_version_blessed = True ### DEBUG - DELETE ###
1524
+
1525
+ self.next(self.model_to_hub)
1526
+
1527
+
1528
+ @step
1529
+ def model_to_hub(self):
1530
+ """
1531
+ Push to hub model version, including
1532
+ readme with versioning info.
1533
+ """
1534
+
1535
+ #############################
1536
+ # case of user-provided #
1537
+ # documentation artifact(s) #
1538
+ #############################
1539
+ # note that user can provide either
1540
+ # 'pipeline_card.py' or 'template.html'
1541
+ # or 'dataset_readme.py'
1542
+ # or 'dataset_readme_template.md'
1543
+ # or 'model_readme.py'
1544
+ # or 'model_readme_template.md'
1545
+ # or any combination of those
1546
+ # when specifying custom
1547
+ # 'pipeline_card_artifacts_path'
1548
+ if (
1549
+ "model_readme_template.md" in
1550
+ os.listdir(self.pipeline_card_artifacts_path)
1551
+ ):
1552
+ template_dir = self.pipeline_card_artifacts_path
1553
+ else:
1554
+ template_dir = os.path.dirname(
1555
+ importlib.util.find_spec(
1556
+ f"retrain_pipelines.pipeline_card."+
1557
+ f"{os.getenv('retrain_pipeline_type')}"
1558
+ ).origin)
1559
+ print(f"template_dir : '{template_dir}'")
1560
+ #############################
1561
+ if "model_readme.py" in os.listdir(
1562
+ self.pipeline_card_artifacts_path):
1563
+ from retrain_pipelines.utils import \
1564
+ get_get_model_readme_content
1565
+ get_model_readme_content = \
1566
+ get_get_model_readme_content(
1567
+ self.pipeline_card_artifacts_path)
1568
+ else:
1569
+ from retrain_pipelines.pipeline_card import \
1570
+ get_model_readme_content
1571
+ #############################
1572
+ from retrain_pipelines.model.hf_utils import \
1573
+ push_model_version_to_hub
1574
+
1575
+ #############################
1576
+ # model README #
1577
+ # from template #
1578
+ #############################
1579
+ commit_datetime = datetime.utcnow()
1580
+ new_model_version_label = get_new_repo_minor_version(
1581
+ repo_id=self.model_repo_id,
1582
+ repo_type="model",
1583
+ hf_token=os.getenv("HF_TOKEN", None))
1584
+ readme_content = get_model_readme_content(
1585
+ template_folder=template_dir,
1586
+
1587
+ model_repo_id=self.model_repo_id,
1588
+
1589
+ base_model_dict=self.hf_base_model_dict,
1590
+ training_dataset_dict=self.dataset_commit_dict,
1591
+
1592
+ version_label=new_model_version_label,
1593
+ commit_datetime=commit_datetime,
1594
+ perf_metrics=self.perf_metrics,
1595
+
1596
+ mf_flow_name=current.flow_name,
1597
+ mf_run_id=current.run.id
1598
+ )
1599
+ #############################
1600
+
1601
+ print("Pushing model version to HF hub " +
1602
+ ("(blessed). " if self.model_version_blessed
1603
+ else "(not blessed). ") +
1604
+ "May take a while..",
1605
+ flush=True)
1606
+ model_commit_hash = push_model_version_to_hub(
1607
+ repo_id=self.model_repo_id,
1608
+ model_version_blessed=\
1609
+ self.model_version_blessed,
1610
+ version_label=new_model_version_label,
1611
+ timestamp_str=commit_datetime.strftime(
1612
+ "%Y-%m-%d %H:%M:%S UTC"),
1613
+ model_dir=self.sft_model_dir,
1614
+ model_readme_content=readme_content,
1615
+ hf_token=os.getenv("HF_TOKEN", None)
1616
+ )
1617
+ if not model_commit_hash:
1618
+ raise Exception(
1619
+ "Failed to publish model version.")
1620
+ print("Push of model version to HF hub completed.",
1621
+ flush=True)
1622
+ print(f"https://huggingface.co/{self.model_repo_id}" +
1623
+ f"/blob/{model_commit_hash}/README.md")
1624
+
1625
+ self.model_commit_dict = {
1626
+ "repo_id": self.model_repo_id,
1627
+ "commit_hash": model_commit_hash,
1628
+ "version_label": new_model_version_label,
1629
+ "commit_datetime": commit_datetime,
1630
+ }
1631
+
1632
+ self.next(self.infra_validator)
1633
+
1634
+
1635
+ @step
1636
+ def infra_validator(self):
1637
+ """
1638
+ If the trained model version is blessed,
1639
+ validate serving.
1640
+ """
1641
+ """
1642
+ Note that using isolated virtual env
1643
+ (using @conda task decorator)
1644
+ is advisable to not embark the whole
1645
+ pipeline dependencies into the local server.
1646
+ We don't for educational purpose,
1647
+ keep things "simple" to grasp
1648
+ as well as to avoid forcing conda
1649
+ (for instance miniconda) as
1650
+ a virtual environment management mean
1651
+ to the user.
1652
+ """
1653
+ """
1654
+ Note : We load base model from HF-cache
1655
+ (mounted as /huggingface_hub_cache
1656
+ docker volume) and adapter from local dir
1657
+ (mounted as /FuncCallAdater docker volume.
1658
+ """
1659
+
1660
+ self.local_serve_is_ready = LocalServeReadinessEnum.NOT_APPLICABLE
1661
+
1662
+ if self.model_version_blessed:
1663
+ from retrain_pipelines.utils.docker import \
1664
+ env_has_docker
1665
+
1666
+ if env_has_docker():
1667
+ model_module_dir = \
1668
+ os.path.dirname(
1669
+ importlib.util.find_spec(
1670
+ "retrain_pipelines.model." +
1671
+ os.getenv('retrain_pipeline_type')
1672
+ ).origin)
1673
+
1674
+ # server & data-model & server-config modules artifacts
1675
+ files_to_copy = [
1676
+ "litserve_server.py",
1677
+ "litserve_datamodel.py",
1678
+ "litserve_serverconfig.py",
1679
+ ".dockerignore" # docker context loading
1680
+ # at image-build time,
1681
+ # exclude model weights
1682
+ ]
1683
+ for filename in files_to_copy:
1684
+ shutil.copy(
1685
+ os.path.join(model_module_dir, "litserve",
1686
+ filename),
1687
+ os.path.join(self.serving_artifacts_local_folder,
1688
+ filename)
1689
+ )
1690
+
1691
+ # save dependencies as artifact
1692
+ create_requirements(self.serving_artifacts_local_folder,
1693
+ exclude=["cudf-polars-.*", "cuda-python",
1694
+ "nvidia-.*", "(py)?libcudf-.*",
1695
+ "nvtx", "rmm-.*", "litserve",
1696
+ ".*retrain-pipelines.*"]
1697
+ )
1698
+
1699
+ # server config yaml
1700
+ env = Environment(loader=FileSystemLoader(
1701
+ os.path.join(model_module_dir, "litserve")))
1702
+ template = env.get_template(
1703
+ "litserve_serverconfig_template.yaml")
1704
+ server_config_data = {
1705
+ "port": "8000",
1706
+ "max_seq_length": self.max_seq_length,
1707
+ "max_new_token": self.max_new_tokens,
1708
+ "base_model": {
1709
+ "repo_id": self.hf_base_model_dict["repo_id"],
1710
+ "revision": self.hf_base_model_dict["commit_hash"]
1711
+ },
1712
+ "adapters": [
1713
+ {
1714
+ "name": "func_caller",
1715
+ "path": "/FuncCallAdapter"
1716
+ }
1717
+ ]
1718
+ }
1719
+ server_config_yaml = template.render(server_config_data)
1720
+ print(server_config_yaml)
1721
+ with open(os.path.join(
1722
+ self.serving_artifacts_local_folder,
1723
+ "litserve_serverconfig.yaml"), 'w'
1724
+ ) as output_file:
1725
+ output_file.write(server_config_yaml)
1726
+
1727
+ # Dockerfile
1728
+ env = Environment(loader=FileSystemLoader(
1729
+ os.path.join(model_module_dir)))
1730
+ template = env.get_template(
1731
+ "Dockerfile.litserve_template")
1732
+ # Change CUDA version here from available list
1733
+ # @see https://hub.docker.com/r/nvidia/cuda/tags
1734
+ dockerfile_content = template.render(
1735
+ {"cuda_version": "12.0.0"})
1736
+ with open(os.path.join(
1737
+ self.serving_artifacts_local_folder,
1738
+ "Dockerfile.litserve"), 'w'
1739
+ ) as output_file:
1740
+ output_file.write(dockerfile_content)
1741
+
1742
+ os.environ["no_proxy"] = "localhost,127.0.0.1,0.0.0.0"
1743
+
1744
+ ############################################
1745
+ # actually deploy the inference service #
1746
+ ############################################
1747
+ start_time = time.time()
1748
+ from retrain_pipelines.utils.docker import \
1749
+ build_and_run_docker, print_container_log_tail, \
1750
+ cleanup_docker
1751
+ from retrain_pipelines.model.litserve import \
1752
+ endpoint_started, endpoint_is_ready
1753
+
1754
+ self.port = 8765
1755
+ HF_HUB_CACHE = os.path.realpath(os.path.expanduser(
1756
+ os.getenv(
1757
+ "HF_HUB_CACHE",
1758
+ os.path.join(os.getenv("HF_HOME",
1759
+ "~/.cache/huggingface"),
1760
+ "hub")
1761
+ )))
1762
+ print(f"HF_HUB_CACHE : {HF_HUB_CACHE}")
1763
+ image_name = container_name = "litserve-model"
1764
+
1765
+ serving_container = build_and_run_docker(
1766
+ image_name=image_name, image_tag="1.0",
1767
+ build_path=self.serving_artifacts_local_folder,
1768
+ dockerfile="Dockerfile.litserve",
1769
+ ports_publish_dict={'8000/tcp': self.port},
1770
+ env_vars_dict={
1771
+ "HF_HUB_CACHE": "/huggingface_hub_cache",
1772
+ "HF_TOKEN": os.getenv("HF_TOKEN")
1773
+ },
1774
+ volumes_dict={
1775
+ self.sft_model_dir:
1776
+ {"bind": "/FuncCallAdapter",
1777
+ "mode": "ro"},
1778
+ HF_HUB_CACHE:
1779
+ {"bind": "/huggingface_hub_cache",
1780
+ "mode": "ro"}
1781
+ }
1782
+ )
1783
+
1784
+ if not serving_container:
1785
+ print("failed spinning the LitServe container",
1786
+ file=sys.stderr)
1787
+ self.local_serve_is_ready = \
1788
+ LocalServeReadinessEnum.FAILURE
1789
+ try:
1790
+ cleanup_docker(
1791
+ container_name=container_name,
1792
+ image_name=f"{image_name}:1.0",
1793
+ no_pruning=True # for intermediate layers recycling
1794
+ # (during later re-runs)
1795
+ # to avoid long rebuild time
1796
+ # of exactly the same.
1797
+ )
1798
+ except Exception as cleanup_ex:
1799
+ # fail silently
1800
+ pass
1801
+ else:
1802
+ print("Awaiting endpoint launch..")
1803
+ start_time = time.time()
1804
+ if not endpoint_started(
1805
+ container_name, port=self.port, timeout=10*60
1806
+ ):
1807
+ print(
1808
+ f"The endpoint '{container_name}' " +
1809
+ f"did not start.")
1810
+ self.local_serve_is_ready = \
1811
+ LocalServeReadinessEnum.FAILURE
1812
+ # health check on the spun-up endpoint
1813
+ elif endpoint_is_ready(port=self.port):
1814
+ self.local_serve_is_ready = \
1815
+ LocalServeReadinessEnum.SUCCESS
1816
+ elapsed_time = time.time() - start_time
1817
+ print("deploy_local - Elapsed time: " +
1818
+ f"{elapsed_time:.2f} seconds")
1819
+ ############################################
1820
+ else:
1821
+ # env doesn't have docker
1822
+ self.local_serve_is_ready = \
1823
+ LocalServeReadinessEnum.FAILURE_NO_DOCKER
1824
+
1825
+ if LocalServeReadinessEnum.SUCCESS == self.local_serve_is_ready:
1826
+ from retrain_pipelines.model.litserve.litserve_datamodel \
1827
+ import Response
1828
+
1829
+ import requests
1830
+
1831
+ url = f"http://localhost:{self.port}/predict"
1832
+ headers = {"accept": "application/x-www-form-urlencoded"}
1833
+
1834
+ try:
1835
+ start_time = time.time()
1836
+ data = {
1837
+ "adapter_name": "func_caller",
1838
+ "queries": '["Hello.", "Is 49 a perfect square?"]'
1839
+ }
1840
+ print(f"inference test - data: {data}")
1841
+ response = requests.post(url, headers=headers, data=data)
1842
+ parsed_response = Response(**{"output": response.json()})
1843
+ elapsed_time = time.time() - start_time
1844
+ print("parsed_response ('func_caller' adapter ON) :" +
1845
+ str(parsed_response) +
1846
+ f"\t-\tElapsed time: {elapsed_time:.2f} seconds")
1847
+
1848
+ start_time = time.time()
1849
+ data = {
1850
+ "queries": '["Hello.", "Is 49 a perfect square?"]'
1851
+ }
1852
+ print(f"inference test - data: {data}")
1853
+ response = requests.post(url, headers=headers, data=data)
1854
+ parsed_response = Response(**{"output": response.json()})
1855
+ elapsed_time = time.time() - start_time
1856
+ print(f"parsed_response (no adapter) : {parsed_response}" +
1857
+ f"\t-\tElapsed time: {elapsed_time:.2f} seconds")
1858
+
1859
+ except Exception as ex:
1860
+ print(ex, file=sys.stderr)
1861
+ traceback.print_tb(ex.__traceback__, file=sys.stderr)
1862
+ self.local_serve_is_ready = \
1863
+ LocalServeReadinessEnum.FAILURE
1864
+ pass
1865
+
1866
+ try:
1867
+ cleanup_docker(
1868
+ container_name=container_name,
1869
+ image_name=f"{image_name}:1.0",
1870
+ no_pruning=True # for intermediate layers recycling
1871
+ # (during later re-runs)
1872
+ # to avoid long rebuild time
1873
+ # of exactly the same.
1874
+ )
1875
+ except Exception as cleanup_ex:
1876
+ # fail silently
1877
+ pass
1878
+
1879
+ self.next(self.pipeline_card)
1880
+
1881
+
1882
+ @card(id='default')
1883
+ @card(type='html', id='custom')
1884
+ @step
1885
+ def pipeline_card(self):
1886
+ import re
1887
+ import datetime
1888
+ import importlib.metadata
1889
+
1890
+ #############################
1891
+ # case of user-provided #
1892
+ # documentation artifact(s) #
1893
+ #############################
1894
+ # note that user can provide either
1895
+ # 'pipeline_card.py' or 'template.html'
1896
+ # or 'dataset_readme.py'
1897
+ # or 'dataset_readme_template.md'
1898
+ # or 'model_readme.py'
1899
+ # or 'model_readme_template.md'
1900
+ # or any combination of those
1901
+ # when specifying custom
1902
+ # 'pipeline_card_artifacts_path'
1903
+ if "template.html" in os.listdir(
1904
+ self.pipeline_card_artifacts_path
1905
+ ):
1906
+ template_dir = self.pipeline_card_artifacts_path
1907
+ else:
1908
+ template_dir = os.path.dirname(
1909
+ importlib.util.find_spec(
1910
+ f"retrain_pipelines.pipeline_card."+
1911
+ f"{os.getenv('retrain_pipeline_type')}"
1912
+ ).origin)
1913
+ #############################
1914
+ if "pipeline_card.py" in os.listdir(
1915
+ self.pipeline_card_artifacts_path
1916
+ ):
1917
+ from retrain_pipelines.utils import get_get_html
1918
+ get_html = \
1919
+ get_get_html(self.pipeline_card_artifacts_path)
1920
+ else:
1921
+ from retrain_pipelines.pipeline_card import \
1922
+ get_html
1923
+ from retrain_pipelines.pipeline_card.helpers import \
1924
+ mf_dag_svg
1925
+ #############################
1926
+
1927
+
1928
+ #############################
1929
+ ## "default" card ##
1930
+ #############################
1931
+ self.metadata = {
1932
+ "name": "TabNet Model",
1933
+ "version": "1.0",
1934
+ "retrain_pipelines": f"retrain-pipelines {__version__}",
1935
+ "retrain_pipeline_type": os.environ["retrain_pipeline_type"],
1936
+ "description": "A PyTorch TabNet model retrained",
1937
+ "authors": [current.username],
1938
+ "tags": ["classification", "tabnet"],
1939
+ "license": "MIT License",
1940
+ "data_augmentation": [
1941
+ {
1942
+ "name": "Augmentation",
1943
+ "description": "Truncating queries and " + \
1944
+ "associate those to " + \
1945
+ "no tool-call answers. " + \
1946
+ "Intent being to instruct on " + \
1947
+ "not hallucinating missing " + \
1948
+ "tool-calls parameters values."
1949
+ },
1950
+ {
1951
+ "name": "Enrichment",
1952
+ "description": "Addition of records " + \
1953
+ "from an external data-source. " + \
1954
+ "Here to instruct on no tool-call."
1955
+ }
1956
+ ],
1957
+ "references": [
1958
+ {
1959
+ "title": "Base model",
1960
+ "link": f"https://hf.co/{self.hf_base_model_dict['repo_id']}"
1961
+ },
1962
+ {
1963
+ "title": "Function-calling dataset",
1964
+ "link": f"https://hf.co/{self.hf_dataset_dict['repo_id']}"
1965
+ },
1966
+ {
1967
+ "title": "Data-enrichment dataset",
1968
+ "link": f"https://hf.co/{self.hf_enrich_dataset_dict['repo_id']}"
1969
+ },
1970
+ {
1971
+ "title": "Unsloth",
1972
+ "link": "https://unsloth.ai/blog/contpretraining"
1973
+ }
1974
+ ]
1975
+ }
1976
+
1977
+ current.card['default'].append(Markdown(
1978
+ "model_version_blessed : **%s**" % str(self.model_version_blessed)))
1979
+ current.card['default'].append(Artifact(
1980
+ {"model_version_blessed": self.model_version_blessed}))
1981
+
1982
+ current.card['default'].append(
1983
+ Image.from_matplotlib(self.sft_log_history_fig))
1984
+ current.card['default'].append(
1985
+ Image.from_matplotlib(self.validation_completions_fig))
1986
+ #############################
1987
+
1988
+ #############################
1989
+ ## html "custom" card ##
1990
+ #############################
1991
+ dt = datetime.datetime.now(tz=datetime.timezone.utc)
1992
+ formatted_dt = dt.strftime("%A %b %d %Y %I:%M:%S %p %Z")
1993
+ task_obj_python_cmd = f"metaflow.Task(" + \
1994
+ f"\"{current.pathspec}\", " + \
1995
+ f"attempt={str(current.retry_count)})"
1996
+ params={
1997
+ 'template_dir': template_dir,
1998
+ 'title': f"{current.flow_name}",
1999
+ "subtitle": f"(flow run # {len(list(current.run.parent.runs()))}," + \
2000
+ f" run_id: {str(current.run.id)} - {formatted_dt})",
2001
+
2002
+ # blessed status / current_blessed version
2003
+ 'model_version_blessed': self.model_version_blessed,
2004
+ 'current_blessed_version_label': (
2005
+ self.current_blessed_version_dict["version_label"]
2006
+ if self.current_blessed_version_dict
2007
+ else None
2008
+ ),
2009
+ 'current_blessed_commit_datetime': (
2010
+ self.current_blessed_version_dict["commit_datetime"]
2011
+ if self.current_blessed_version_dict
2012
+ else None
2013
+ ),
2014
+ 'current_blessed_model_commit_hash': (
2015
+ self.current_blessed_version_dict["commit_hash"]
2016
+ if self.current_blessed_version_dict
2017
+ else None
2018
+ ),
2019
+ 'current_blessed_run': self.current_blessed_run,
2020
+
2021
+ 'LocalServeReadinessEnum': LocalServeReadinessEnum,
2022
+ 'local_serve_is_ready': self.local_serve_is_ready,
2023
+ # EDA
2024
+ 'main_dataset_repo_id': self.hf_dataset['repo_id'],
2025
+ 'main_dataset_commit_hash': self.hf_dataset_dict['commit_hash'],
2026
+ 'main_dataset_commit_datetime': \
2027
+ self.hf_dataset_dict['commit_datetime'],
2028
+
2029
+ 'records_count': self.records_count,
2030
+ 'data_schema': self.data_schema,
2031
+ 'answers_tools_count_fig': self.answers_tools_count_fig,
2032
+ 'words_count_fig': self.words_count_fig,
2033
+
2034
+ # model training
2035
+ 'dataset_repo_id': self.dataset_repo_id,
2036
+ 'dataset_version_label': self.dataset_commit_dict["version_label"],
2037
+ 'dataset_commit_datetime': self.dataset_commit_dict["commit_datetime"],
2038
+ 'dataset_commit_hash': self.dataset_commit_dict["commit_hash"],
2039
+ 'dataset_augmentation_rate': self.actual_augmentation_rate,
2040
+ 'dataset_enrichment_rate': self.enrichment_rate,
2041
+
2042
+ 'model_repo_id': self.model_repo_id,
2043
+ 'model_version_label': self.model_commit_dict["version_label"],
2044
+ 'model_commit_datetime': self.model_commit_dict["commit_datetime"],
2045
+ 'model_commit_hash': self.model_commit_dict["commit_hash"],
2046
+
2047
+ 'cpt_log_history_fig': self.cpt_log_history_fig,
2048
+ 'sft_log_history_fig': self.sft_log_history_fig,
2049
+
2050
+ 'validation_completions_fig': self.validation_completions_fig,
2051
+
2052
+ 'pipeline_parameters_dict': {"cpt": self.cpt_training_args,
2053
+ "sft": self.sft_training_args},
2054
+
2055
+ 'metrics_dict': self.perf_metrics,
2056
+
2057
+ 'task_obj_python_cmd': task_obj_python_cmd,
2058
+ 'dag_svg': mf_dag_svg(self)
2059
+ }
2060
+ self.html = get_html(params)
2061
+ #############################
2062
+ current
2063
+ #############################
2064
+
2065
+ self.next(self.pipeline_to_hub)
2066
+
2067
+
2068
+ @step
2069
+ def pipeline_to_hub(self):
2070
+ """
2071
+ publish versioned source-code and pipeline-card
2072
+ for ths run on the Hugging Face Hub.
2073
+ """
2074
+
2075
+ model_commit_datetime = \
2076
+ self.model_commit_dict["commit_datetime"]
2077
+ timestamp_str = \
2078
+ "{:%Y%m%d_%H%M%S}".format(model_commit_datetime) + \
2079
+ "{:03d}".format(model_commit_datetime.microsecond//1000) + \
2080
+ "_UTC"
2081
+ subfolder_name = \
2082
+ "v" + self.model_commit_dict["version_label"] + \
2083
+ "_" + timestamp_str
2084
+ commit_datetime = datetime.utcnow()
2085
+
2086
+ ###############################
2087
+ # source-code #
2088
+ ###############################
2089
+ # We upload only herein file #
2090
+ # plus user-provided versions #
2091
+ # of the customizable ones #
2092
+ # (if any). #
2093
+ ###############################
2094
+ custom_source_files = [os.path.abspath(__file__)]
2095
+ if (
2096
+ self.pipeline_card_artifacts_path != \
2097
+ self.default_pipeline_card_module_dir
2098
+ ):
2099
+ candidate_source_files = [
2100
+ "pipeline_card.py",
2101
+ "template.html",
2102
+ "dataset_readme.py",
2103
+ "dataset_readme_template.md",
2104
+ "model_readme.py",
2105
+ "model_readme_template.md"
2106
+ ]
2107
+ for candidate_source_file in candidate_source_files:
2108
+ file_fullpath = os.path.join(
2109
+ self.pipeline_card_artifacts_path,
2110
+ candidate_source_file)
2111
+ if os.path.exists(file_fullpath):
2112
+ custom_source_files.append(file_fullpath)
2113
+
2114
+ source_code_commit_hash = \
2115
+ push_files_to_hub_repo_branch(
2116
+ repo_id=self.model_repo_id,
2117
+ branch_name="retrain-pipelines_source-code",
2118
+ file_fullnames=custom_source_files,
2119
+ include_requirements_txt=True,
2120
+ path_in_repo=subfolder_name,
2121
+ commit_message=\
2122
+ "source-code for model version " + \
2123
+ subfolder_name + \
2124
+ f"- retrain-pipelines {__version__}",
2125
+ repo_type="model",
2126
+ hf_token=os.getenv("HF_TOKEN", None)
2127
+ )
2128
+ print(source_code_commit_hash)
2129
+ self.source_code_commit_dict = {
2130
+ "repo_id": self.model_repo_id,
2131
+ "branch_name": "retrain-pipelines_source-code",
2132
+ "commit_datetime": commit_datetime,
2133
+ "commit_hash": source_code_commit_hash
2134
+ }
2135
+ ###############################
2136
+
2137
+ ###############################
2138
+ # pipeline-card #
2139
+ ###############################
2140
+ pipeline_card_fullname = None
2141
+ for run_step in current.run.steps():
2142
+ task = list(run_step.tasks())[0]
2143
+ task_name = task.path_components[2]
2144
+ if "pipeline_card" == task_name:
2145
+ pipeline_card = get_cards(
2146
+ task, id='custom', type='html')[0]
2147
+ pipeline_card_fullname = os.path.realpath(
2148
+ os.path.join(
2149
+ task.metadata_dict.get("ds-root", None),
2150
+ mf_config.CARD_SUFFIX, pipeline_card.path
2151
+ ))
2152
+ print(pipeline_card_fullname)
2153
+ break
2154
+ pipeline_card_commit_hash = \
2155
+ push_files_to_hub_repo_branch(
2156
+ repo_id=self.model_repo_id,
2157
+ branch_name="retrain-pipelines_pipeline-card",
2158
+ file_fullnames=[pipeline_card_fullname],
2159
+ path_in_repo=subfolder_name,
2160
+ commit_message=\
2161
+ "pipeline-card for model version " + \
2162
+ subfolder_name + \
2163
+ f"- retrain-pipelines {__version__}",
2164
+ repo_type="model",
2165
+ hf_token=os.getenv("HF_TOKEN", None)
2166
+ )
2167
+ print(pipeline_card_commit_hash)
2168
+ self.pipeline_card_commit_dict = {
2169
+ "repo_id": self.model_repo_id,
2170
+ "branch_name": "retrain-pipelines_pipeline-card",
2171
+ "commit_datetime": commit_datetime,
2172
+ "commit_hash": pipeline_card_commit_hash
2173
+ }
2174
+ ###############################
2175
+
2176
+ self.next(self.deploy)
2177
+
2178
+
2179
+ @step
2180
+ def deploy(self):
2181
+ """
2182
+ placeholder for the serving SDK deploy call
2183
+ (on the target production platform).
2184
+ Include any artifact you want,
2185
+ consider including the portable pipelione-card
2186
+ itself !
2187
+ """
2188
+
2189
+ if (
2190
+ self.model_version_blessed and
2191
+ (self.local_serve_is_ready == LocalServeReadinessEnum.SUCCESS)
2192
+ ):
2193
+ pass # your code here
2194
+
2195
+ self.next(self.load_test)
2196
+
2197
+
2198
+ @step
2199
+ def load_test(self):
2200
+ """
2201
+ placeholder
2202
+ """
2203
+
2204
+ if (
2205
+ self.model_version_blessed and
2206
+ (self.local_serve_is_ready == LocalServeReadinessEnum.SUCCESS)
2207
+ ):
2208
+ pass # your code here
2209
+
2210
+ self.next(self.end)
2211
+
2212
+
2213
+ @step
2214
+ def end(self):
2215
+ pass
2216
+
2217
+
2218
+ if __name__ == "__main__":
2219
+ UnslothFuncCallFlow()
2220
+