-
Notifications
You must be signed in to change notification settings - Fork 441
Expand file tree
/
Copy pathpyproject.toml
More file actions
589 lines (548 loc) · 31.1 KB
/
Copy pathpyproject.toml
File metadata and controls
589 lines (548 loc) · 31.1 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
[build-system]
requires = ["setuptools"]
build-backend = "setuptools.build_meta"
[tool.setuptools.packages.find]
include = ["skyrl*"]
[project]
name = "skyrl"
dynamic = ["version"]
description = "Unified API for training and inference"
readme = "README.md"
requires-python = ">=3.12,<3.15"
dependencies = [
"datasets>=4.0.0",
"pillow>=11.3.0",
"rich>=14.1.0",
"safetensors>=0.6.2",
# Floor tracks the transformers pin below: transformers 5.16.1 requires
# tokenizers>=0.23.1.
"tokenizers>=0.23.1",
"transformers>=5.6.1,<=5.16.1",
"typer>=0.17.4",
"peft==0.18.1",
"hf_transfer",
"cloudpathlib>=0.23.0",
"zstandard>=0.23.0",
"xxhash>=3.0.0",
]
[project.optional-dependencies]
gpu = [
"jax[cuda13]>=0.7.2; sys_platform == 'linux'",
]
tpu = [
# libtpu is published for Linux x86_64, not Linux aarch64.
"jax[tpu]>=0.7.2; sys_platform == 'linux' and platform_machine == 'x86_64'",
]
tinker = [
"tinker>=0.25.0,<0.26.0",
"fastapi[standard]",
"sqlmodel",
"sqlalchemy[asyncio]",
"aiosqlite",
"asyncpg",
"psycopg2-binary",
]
ray = [
"ray[default]==2.57.0",
]
aws = [
"cloudpathlib[s3]",
"s5cmd>=0.3.3; sys_platform == 'linux'",
]
gcp = [
"cloudpathlib[gs]",
]
azure = [
"cloudpathlib[azure]",
]
# The extras "jax", "fsdp", and "megatron" are the dependencies the
# engine needs for --backend="jax", --backend="fsdp", and --backend="megatron",
# respectively.
jax = [
"jax>=0.8,<1.0",
"flax>=0.12.10",
"optax>=0.2.5",
]
skyrl-train = [
"loguru",
"tqdm",
"ninja",
"tensorboard",
"func_timeout",
"hydra-core==1.3.4",
"accelerate",
"torchdata",
"omegaconf",
"ray==2.57.0",
"peft==0.18.1",
"debugpy==1.8.0",
"hf_transfer",
"wandb",
"datasets>=4.0.0",
"tensordict",
"jaxtyping",
"skyrl-gym",
"flash-attn==2.8.3; sys_platform == 'linux'",
"polars",
"s3fs",
"s5cmd>=0.3.3; sys_platform == 'linux'",
"gcsfs",
"fastapi",
"orjson>=3.11.9",
"pybase64>=1.4.2",
"uvicorn",
"vllm-router; sys_platform == 'linux'",
"pybind11",
"setuptools",
# The `nixl` shim provides that namespace and dispatches on `torch.version.cuda`,
# so with a cu13 torch it loads `nixl_cu13`. `nixl-cu13` ships the `nixl_cu13`
# module, but vLLM imports `nixl._api`. Its metadata hard-depends on `nixl-cu12`
# too; that variant is overridden out below (it would drag the CUDA-12 runtime
# into the image for a module the shim would never load).
"nixl; sys_platform == 'linux'",
]
fsdp = [
"skyrl[skyrl-train]",
"vllm==0.30.0; sys_platform == 'linux'",
"vllm-router; sys_platform == 'linux'",
# The `nixl` shim provides that namespace and dispatches on `torch.version.cuda`,
# so with a cu13 torch it loads `nixl_cu13`. `nixl-cu13` ships the `nixl_cu13`
# module, but vLLM imports `nixl._api`. Its metadata hard-depends on `nixl-cu12`
# too; that variant is overridden out below (it would drag the CUDA-12 runtime
# into the image for a module the shim would never load).
"nixl; sys_platform == 'linux'",
"flash-linear-attention; sys_platform == 'linux'",
"causal-conv1d==1.6.2.post1; sys_platform == 'linux'",
"flash-attn==2.8.3; sys_platform == 'linux'",
"torch==2.13.0; sys_platform == 'linux'",
"flashinfer-python==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"flashinfer-jit-cache==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"flashinfer-cubin==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"torchvision; sys_platform == 'linux'",
]
megatron = [
"skyrl[skyrl-train]",
"transformer-engine[pytorch]==2.16.0; sys_platform == 'linux'",
# Declared directly (not just via transformer-engine[pytorch]) so the
# [tool.uv.sources] index below is honored — uv ignores sources for
# transitive-only packages, which would otherwise fall back to the PyPI sdist.
"transformer-engine-torch==2.16.0; sys_platform == 'linux'",
# TE core CUDA lib (libtransformer_engine.so). Declared directly because the
# meta package only pulls a core via its `core*` extras, which we don't request
# — without it the meta package's sanity check fails with "Found empty
# `transformer-engine` meta package installed". cu13 is the only core Astral
# publishes on the cu130 index, and it is what `transformer-engine-torch`'s
# metadata requires there, so the two agree without an override.
"transformer-engine-cu13==2.16.0+cu.13.0; sys_platform == 'linux'",
"flash-attn==2.8.3; sys_platform == 'linux'",
"flash-linear-attention; sys_platform == 'linux'",
"causal-conv1d==1.6.2.post1; sys_platform == 'linux'",
"mamba-ssm==2.3.2.post1; sys_platform == 'linux'",
"vllm==0.30.0; sys_platform == 'linux'",
"vllm-router; sys_platform == 'linux'",
# The `nixl` shim provides that namespace and dispatches on `torch.version.cuda`,
# so with a cu13 torch it loads `nixl_cu13`. `nixl-cu13` ships the `nixl_cu13`
# module, but vLLM imports `nixl._api`. Its metadata hard-depends on `nixl-cu12`
# too; that variant is overridden out below (it would drag the CUDA-12 runtime
# into the image for a module the shim would never load).
"nixl; sys_platform == 'linux'",
"torch==2.13.0; sys_platform == 'linux'",
"flashinfer-python==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"torchvision; sys_platform == 'linux'",
"megatron-bridge; sys_platform == 'linux'",
"megatron-core; sys_platform == 'linux'",
"nvidia-resiliency-ext; sys_platform == 'linux'",
"flashinfer-jit-cache==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"flashinfer-cubin==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"nvidia-modelopt; sys_platform == 'linux'",
"fast-hadamard-transform; sys_platform == 'linux'",
]
# Opt-in FlashAttention 4. NOT part of `megatron`, so the default install runs FA2.
#
# The combined `flash-attn` wheel (see [tool.uv.sources]) always carries FA4's
# `flash_attn/cute/` tree -- that is what stops it colliding with flash-attn 2's
# stale copy. But both consumers gate on the *`flash-attn-4` distribution
# metadata*, not on the tree:
# * Transformer Engine: `get_pkg_version("flash-attn-4")` raising
# PackageNotFoundError leaves `flash_attn_func_v4 = None`, so TE never
# imports or dispatches to FA4.
# * megatron-core: `HAVE_FA4` is False for the same reason.
# So installing this companion is exactly the switch that turns FA4 on, and
# leaving it out keeps FA2 as the default with no shim and no runtime patching.
#
# Markers: `< '3.15'` matches the companion's own `Requires-Python`, and the
# `platform_machine` bound keeps it off Linux arches with no combined wheel
# (there `flash-attn` resolves to plain PyPI 2.8.3, which cannot satisfy the
# companion's hard dep on ==2.8.3+cu.13.0.torch.2.13.fa4b28.skyrl1). Dropping
# either bound makes resolution fail outright.
#
# FA4 4.0.0b28 is a beta, and TE disables it under context parallelism (TE <= 2.18;
# main lifts this for cp_comm_type p2p/all_gather/a2a). Note TE 2.16's arch gate
# only excludes < sm80, so FA4 is selected on sm86/sm87/sm89 (A10, L4, L40S, 4090)
# where its kernels cannot launch -- set SKYRL_DISABLE_FA4=1 there, or just do not
# install this extra.
fa4 = [
"skyrl[megatron]",
"flash-attn-4[cu13]==4.0.0b28; sys_platform == 'linux' and (platform_machine == 'x86_64' or platform_machine == 'aarch64') and python_version < '3.15'",
]
miniswe = [
"skyrl[skyrl-train]",
# NOTE (sumanthrh): Needs to be a commit after https://github.com/SWE-agent/mini-swe-agent/commit/4f5d445e99d13b5482478c23508bf2fbf7c0670c
"mini-swe-agent>=1.12.0",
"litellm",
]
harbor = [
"harbor[daytona,modal]",
]
# Mooncake distributed KV for vLLM: the store connector (MooncakeStoreConnector,
# KV-cache offloading) and P2P PD transfer (MooncakeConnector). Provides the python
# bindings and the `mooncake_master` binary. Off-the-shelf PyPI wheel; Linux-only.
# Mooncake currently has Linux wheels through CPython 3.13 only.
mooncake = [
"mooncake-transfer-engine-cuda13; sys_platform == 'linux' and python_version < '3.14'",
]
dev = [
"mkdocs",
"mkdocs-material",
"mkdocstrings[python]>=0.24.0",
"pymdown-extensions>=10.7",
"pytest",
"pytest-forked",
"pytest-asyncio",
"pre-commit",
"litellm",
"torch",
"ty",
"cloudpathlib[s3]",
"alembic",
"griffe2md",
]
[tool.setuptools]
include-package-data = true
# Runtime patch payloads (skyrl/backends/skyrl_train/patches/**/*.patch) are read
# from disk at runtime, so they must ship inside the package. `include-package-data`
# alone does not cover them
[tool.setuptools.package-data]
"*" = ["*.patch"]
[tool.setuptools.dynamic]
version = {attr = "skyrl.__version__"}
[project.scripts]
# The following is for supporting the skyrl-train dependency
[tool.uv]
# Resolve only the Linux architectures covered by the CUDA wheels, plus the
# supported macOS development target. CUDA-only extras cannot resolve on other
# Linux architectures.
required-environments = [
"sys_platform == 'linux' and (platform_machine == 'x86_64' or platform_machine == 'aarch64')",
"sys_platform == 'darwin' and platform_machine == 'arm64'",
]
# each backend should have separate dependencies that can potentially clash
# megatron also clashes with the jax dependency from gpu and tpu extras
conflicts = [
[
{ extra = "jax" },
{ extra = "megatron" },
{ extra = "fsdp" },
],
[
{ extra = "megatron" },
{ extra = "gpu" },
{ extra = "tpu" },
{ extra = "miniswe" },
]
]
# disable build isolation for megatron related dependencies
no-build-isolation-package = [
"nv-grouped-gemm",
]
# override unnecessary dependencies and pin versions to override Megatron-Bridge
# unpinned dependencies.
override-dependencies = [
"transformer-engine[pytorch]==2.16.0; sys_platform == 'linux'",
"transformers>=5.6.1,<=5.16.1; sys_platform == 'linux'",
"megatron-core>=0.16.0; sys_platform == 'linux'",
"ml_dtypes>=0.5.0; sys_platform == 'linux'",
# `nixl` hard-depends on both nixl-cu12 and nixl-cu13; drop the cu12 variant
# so it doesn't pull a second, unused CUDA runtime in alongside the cu13 one.
"nixl-cu12; sys_platform == 'never'",
# Megatron-Bridge pins flashinfer-python==0.6.8.post1, which conflicts with
# our pin, so all three flashinfer packages are held here. vLLM 0.30.0 needs
# >=0.6.16: it imports `flashinfer.autotuner.set_autotune_process_group`,
# which 0.6.14 does not have, and the engine fails to start without it.
# flashinfer hard-errors when an installed cubin's version differs from its
# own, so the three must move together.
# They must also be >=0.6.13: older flashinfer rejects the `layout_code` vLLM's
# allreduce+RMS fusion pass hands the MNNVL backend ("MNNVL AllReduce does
# not support quantization fusion").
#
# flashinfer-cubin is not published to PyPI past 0.6.13, so it (like
# flashinfer-jit-cache) comes from flashinfer's own index -- see the
# `flashinfer` index and [tool.uv.sources] below.
"flashinfer-python==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"flashinfer-jit-cache==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
"flashinfer-cubin==0.6.18.post1; sys_platform == 'linux' and platform_machine == 'x86_64'",
# mamba-ssm 2.3.2.post1 pins tilelang==0.1.8 and apache-tvm-ffi<=0.1.9, while
# vLLM 0.30.0 pins tilelang==0.1.12 and apache-tvm-ffi==0.1.11. Both are
# overridden here; these two MUST be bumped together.
#
# FA4 requires apache-tvm-ffi>=0.1.12, newer than either package's metadata
# allows. tilelang 0.1.9's own bound is loose enough (`>=0.1.2,~=0.1.0`) that
# uv happily resolves it against 0.1.12 -- and then `import tilelang` aborts
# with `AttributeError: attribute '__dict__' of 'type' objects is not writable`
# from a duplicate ffi field registration, taking `mamba_ssm` and `fla` (the
# default FLA_TILELANG path) down with it. Measured:
# tilelang 0.1.9 + ffi 0.1.10 -> import OK (pre-FA4 state)
# tilelang 0.1.9 + ffi 0.1.12 -> AttributeError
# tilelang 0.1.13 + ffi 0.1.12 -> import OK
# 0.1.13 is the first tilelang whose bound (`>=0.1.11,<0.1.13`) admits 0.1.12.
"tilelang==0.1.13; sys_platform == 'linux'",
"apache-tvm-ffi==0.1.12; sys_platform == 'linux'",
"nvidia-cutlass-dsl[cu13]==4.7.1; sys_platform == 'linux'",
"quack-kernels[cu13]==0.6.5; sys_platform == 'linux'",
# vLLM allows xgrammar>=0.2.1,<1.0.0, but 0.2.4/0.2.5 never published a
# cp312 linux x86_64 wheel (only cp310/cp311 + aarch64/macOS), so resolving
# to them breaks install on x86_64 py3.12. 0.2.3 has the wheel.
"xgrammar==0.2.3",
"torch==2.13.0; sys_platform == 'linux'",
# torchvision is NOT stable-ABI (unlike vLLM), and each release hard-pins one
# torch: 0.30.0 wants torch 2.13.0, 0.26.0 wants 2.11.0. vLLM 0.30.0 pulls
# torchvision==0.28.0, so without this the torch override above would resolve
# a torchvision that cannot import against the torch we actually install.
"torchvision==0.28.0; sys_platform == 'linux'",
]
[tool.uv.extra-build-dependencies]
flash-attn = [{requirement = "torch", match-runtime = true}]
fast-hadamard-transform = ["torch==2.13.0", "ninja"]
[tool.uv.extra-build-variables]
flash-attn = { FLASH_ATTENTION_SKIP_CUDA_BUILD = "TRUE"}
fast-hadamard-transform = { FAST_HADAMARD_TRANSFORM_FORCE_BUILD = "TRUE"}
# SkyRL only uses NVRx's async checkpointing. Its CUPTI straggler-detection extension
# only adds CUPTI's include dir and fails with `cuda.h: No such file or directory` unless
# the toolkit's headers are on the default include path, so skip building it.
nvidia-resiliency-ext = { STRAGGLER_DET_SKIP_CUPTI_EXT_BUILD = "1" }
[[tool.uv.index]]
name = "pytorch-cu130"
url = "https://download.pytorch.org/whl/cu130"
explicit = true
[[tool.uv.index]]
name = "pytorch-cpu"
url = "https://download.pytorch.org/whl/cpu"
explicit = true
[[tool.uv.index]]
name = "jax-tpu"
url = "https://storage.googleapis.com/jax-releases/libtpu_releases.html"
[[tool.uv.index]]
name = "flashinfer-cu130"
url = "https://flashinfer.ai/whl/cu130"
explicit = true
# flashinfer's CUDA-agnostic index. flashinfer-cubin is a pure-python wheel
# (py3-none-any) that stopped being published to PyPI after 0.6.13, so it is
# sourced here instead; the cu130 index above only carries flashinfer-jit-cache.
[[tool.uv.index]]
name = "flashinfer"
url = "https://flashinfer.ai/whl/"
explicit = true
# Astral's prebuilt GPU wheels (https://wheels.astral.sh). cu130 to match torch's
# CUDA variant. Wheels are published per (CUDA, torch) pair and carry that pair in
# a local version segment (e.g. `+cu.13.0.torch.2.11`), so the exact build must be
# pinned in the requirement — see the pins in the `fsdp` / `megatron` extras.
[[tool.uv.index]]
name = "astral-cu130"
url = "https://wheels.astral.sh/simple/cu130/"
explicit = true
[tool.uv.sources]
skyrl-gym = { path = "./skyrl-gym", editable = true }
# Match torch's CUDA variant (cu130).
flashinfer-jit-cache = { index = "flashinfer-cu130", marker = "sys_platform == 'linux'" }
flashinfer-cubin = { index = "flashinfer", marker = "sys_platform == 'linux'" }
# TEMPORARY: torch-2.13 builds of these four are not published anywhere yet --
# Astral's cu130 index stops at `torch.2.12` (mamba-ssm at `torch.2.11`), so the
# usual `+cu.13.0.torch.<ver>` pins cannot resolve. These are wheels we built
# ourselves against torch 2.13.0 + CUDA 13.0.88 and host on a GitHub release.
#
# https://github.com/NovaSky-AI/skyrl-wheels/releases/tag/cu13torch2.13
#
# They cover cp312/cp313/cp314 on both x86_64 and aarch64 -- the whole of this
# project's `requires-python = ">=3.12"` range, matching the coverage Astral gave us.
# Anything outside that matrix falls through to the PyPI sdist and builds from source.
#
# NOTE: `TORCH_CUDA_ARCH_LIST` does NOT control what these wheels contain -- four of
# the five packages ignore it and hardcode their own gencode lists (mamba-ssm's
# setup.py explicitly clears the variable; flash-attn reads `FLASH_ATTN_CUDA_ARCHS`,
# default `80;90;100;120`). Actual embedded cubins, per `cuobjdump --list-elf`:
# flash-attn sm_80 90 100 120
# mamba-ssm / causal-conv1d / fast-hadamard sm_75 80 87 90 100 103 110 120 121
# transformer-engine-torch none (kernels live in -cu13)
# So H100/GH200 (sm_90) and B200/GB200 (sm_100) are covered on both arches.
#
# TODO: DELETE THIS and go back to `{ index = "astral-cu130" }` with
# `+cu.13.0.torch.2.13` pins as soon as Astral publishes those builds. Check with:
# curl -sL https://wheels.astral.sh/simple/cu130/mamba-ssm/ | grep -o 'torch\.2\.[0-9]*'
#
# flash-attn must stay at exactly 2.8.3, NOT 2.8.3.post1: TE gates flash-attn on
# `max_version = 2.8.3` and 2.8.3.post1 > 2.8.3. TE >= 2.16 strips the local segment
# before comparing (`PkgVersion(...).public`).
# WARNING: every `flash-attn` entry below MUST stay a combined FA2+FA4 wheel.
# The stock 2.8.3 wheel ships a stale `flash_attn/cute/` built against the
# nvidia-cutlass-dsl 4.0/4.1 API, and megatron-core imports that tree
# unconditionally -- before its own `flash-attn-4` metadata check -- under a
# bare `except ImportError`. Against cutlass-dsl 4.6.x the stale tree raises
# `AttributeError: ThrMma`, which escapes that handler and takes down every
# `megatron.core.transformer.attention` import. SkyRL used to defuse this at
# runtime in `skyrl/_compat.py`; that shim is gone, so pointing any of these
# back at a plain wheel (e.g. the `astral-cu130` TODO above) reintroduces the
# crash. Rebuild combined wheels instead.
flash-attn = [
# Both architectures get the combined FA2+FA4 wheels: the plain
# cu13torch2.13 FA2 build ships a stale `flash_attn/cute/` tree (21 files,
# built against the nvidia-cutlass-dsl 4.0/4.1 API) that collides with real
# FA4 (52 files). This wheel is the same torch-2.13 FA2 extension with that
# tree replaced by FA4 4.0.0b28's, so one package owns `flash_attn.cute`
# outright and FA2 and FA4 are importable side by side. See `flash-attn-4`
# below for the companion.
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13-fa4b28-skyrl1/flash_attn-2.8.3%2Bcu.13.0.torch.2.13.fa4b28.skyrl1-cp312-cp312-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13-fa4b28-skyrl1/flash_attn-2.8.3%2Bcu.13.0.torch.2.13.fa4b28.skyrl1-cp313-cp313-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13-fa4b28-skyrl1/flash_attn-2.8.3%2Bcu.13.0.torch.2.13.fa4b28.skyrl1-cp314-cp314-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.14'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13-fa4b28-skyrl1/flash_attn-2.8.3%2Bcu.13.0.torch.2.13.fa4b28.skyrl1-cp312-cp312-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13-fa4b28-skyrl1/flash_attn-2.8.3%2Bcu.13.0.torch.2.13.fa4b28.skyrl1-cp313-cp313-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13-fa4b28-skyrl1/flash_attn-2.8.3%2Bcu.13.0.torch.2.13.fa4b28.skyrl1-cp314-cp314-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.14'" },
]
# Metadata-only companion: it carries no runtime files, declares the `cu13`
# extra, and hard-depends on the exact combined wheel above (which is what
# actually ships `flash_attn/cute/`). Both must come from the same release --
# the companion cannot find the combined wheel on PyPI. Transformer Engine
# gates FA4 on this distribution's metadata, so without it TE never uses FA4.
flash-attn-4 = [
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13-fa4b28-skyrl1/flash_attn_4-4.0.0b28%2Bskyrl1-py3-none-any.whl", marker = "sys_platform == 'linux' and (platform_machine == 'x86_64' or platform_machine == 'aarch64') and python_version < '3.15'" },
]
causal-conv1d = [
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/causal_conv1d-1.6.2.post1-cp312-cp312-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/causal_conv1d-1.6.2.post1-cp313-cp313-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/causal_conv1d-1.6.2.post1-cp314-cp314-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.14'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/causal_conv1d-1.6.2.post1-cp312-cp312-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/causal_conv1d-1.6.2.post1-cp313-cp313-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/causal_conv1d-1.6.2.post1-cp314-cp314-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.14'" },
]
mamba-ssm = [
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/mamba_ssm-2.3.2.post1-cp312-cp312-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/mamba_ssm-2.3.2.post1-cp313-cp313-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/mamba_ssm-2.3.2.post1-cp314-cp314-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.14'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/mamba_ssm-2.3.2.post1-cp312-cp312-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/mamba_ssm-2.3.2.post1-cp313-cp313-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/mamba_ssm-2.3.2.post1-cp314-cp314-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.14'" },
]
# The TE meta and core packages carry no torch coupling in their versions
# (2.16.0 / 2.16.0+cu.13.0), so they stay on Astral; only the torch bindings are
# rebuilt.
transformer-engine = { index = "astral-cu130", marker = "sys_platform == 'linux'" }
transformer-engine-cu13 = { index = "astral-cu130", marker = "sys_platform == 'linux'" }
transformer-engine-torch = [
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/transformer_engine_torch-2.16.0-cp312-cp312-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/transformer_engine_torch-2.16.0-cp313-cp313-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/transformer_engine_torch-2.16.0-cp314-cp314-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.14'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/transformer_engine_torch-2.16.0-cp312-cp312-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/transformer_engine_torch-2.16.0-cp313-cp313-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/transformer_engine_torch-2.16.0-cp314-cp314-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.14'" },
]
# CUDA torch on Linux, CPU torch on macOS (must match skyrl-train).
# Linux uses the cu130 index (torch 2.13 wheels are published there).
torch = [
{ index = "pytorch-cu130", marker = "sys_platform == 'linux'" },
{ index = "pytorch-cpu", marker = "sys_platform == 'darwin'" },
]
torchvision = [
{ index = "pytorch-cu130", marker = "sys_platform == 'linux'" },
{ index = "pytorch-cpu", marker = "sys_platform == 'darwin'" },
]
harbor = { git = "https://github.com/laude-institute/harbor", rev = "3de07a0e01f3368921766437fc7afece3ddec23d" }
# NVRx 0.6 wheels require glibc 2.39; its supported source build also works on Ubuntu 22.04.
nvidia-resiliency-ext = { git = "https://github.com/NVIDIA/nvidia-resiliency-ext", rev = "6c5f2a13c7688d92a7ac7ee6e464721eb8b7345d", marker = "sys_platform == 'linux'" }
# Megatron-Bridge upstream snapshot, with the matching Core revision below.
megatron-bridge = {git = "https://github.com/NVIDIA-NeMo/Megatron-Bridge", rev = "8e7077c6826d17eb4d4d54e6eb15c5a581eda4c0", marker = "sys_platform == 'linux'"}
# megatron-core must match the `3rdparty/Megatron-LM` submodule pin of the
# megatron-bridge rev above -- that is the only megatron-core commit Bridge's
# own CI validates against. Bump the two together, reading the submodule rev with
# `git ls-tree <bridge-rev> 3rdparty/Megatron-LM`.
#
# Upstream PRs not yet in these pins (GLM-5.3-Flash's KDA / mHC / k-pool DSA, and a few bug
# fixes) are carried in skyrl/backends/skyrl_train/patches/megatron/. When bumping either pin,
# follow skyrl/backends/skyrl_train/patches/megatron/README.md to retire whatever has landed.
megatron-core = {git = "https://github.com/NVIDIA/Megatron-LM", rev = "d476c21be396dc666d2e04fa8c199c3851a1551f", marker = "sys_platform == 'linux'"}
# NOTE (sumanthrh): This custom wheel of vllm-router includes a fix for /chat/completions endpoint: https://github.com/vllm-project/router/pull/162 on top of 0.1.14 as well as a new `sticky_least_loaded` policy: https://github.com/SumanthRH/router/tree/session-aware-lb
vllm-router = [
{ url = "https://github.com/SumanthRH/router/releases/download/0.1.14.post1/vllm_router-0.1.14.post1-cp38-abi3-manylinux_2_35_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64'" },
{ url = "https://github.com/SumanthRH/router/releases/download/0.1.14.post1/vllm_router-0.1.14.post1-cp38-abi3-manylinux_2_28_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64'" }
]
# Built against torch 2.13 + CUDA 13.0.88 alongside the other CUDA extensions above
# TODO (sumanthrh): Move towards official torch 13 + CUDA 13 wheels when available
fast-hadamard-transform = [
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/fast_hadamard_transform-1.1.0-cp312-cp312-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/fast_hadamard_transform-1.1.0-cp313-cp313-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/fast_hadamard_transform-1.1.0-cp314-cp314-linux_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine == 'x86_64' and python_version == '3.14'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/fast_hadamard_transform-1.1.0-cp312-cp312-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.12'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/fast_hadamard_transform-1.1.0-cp313-cp313-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.13'" },
{ url = "https://github.com/NovaSky-AI/skyrl-wheels/releases/download/cu13torch2.13/fast_hadamard_transform-1.1.0-cp314-cp314-linux_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and python_version == '3.14'" }
]
[tool.pytest.ini_options]
# Registering markers silences PytestUnknownMarkWarning and makes `-m` selection
# in ci/gpu_ci_run_*.sh explicit rather than relying on free-form names.
markers = [
"vllm: requires the vllm package (present under the fsdp/megatron extras)",
"megatron: requires the megatron extra",
"megatron_models: long-running Megatron model-parity tests",
"integrations: third-party integration tests with extra install steps",
"mooncake: requires the mooncake extra (mooncake-transfer-engine)",
]
[tool.black]
line-length = 120
include = '\.pyi?$'
extend-exclude = '''
/(
# directories
\.eggs
| \.git
| \.hg
| \.mypy_cache
| \.tox
| \.venv
| build
| dist
)/
'''
[tool.flake8]
max-line-length = 120
max-doc-length = 120
extend-ignore = [
# Default ignored errors by flake8
"E121", "E123", "E126", "E226", "E24", "E704",
# F401 module imported but unused
"F401",
# E203 whitespace before ':' (conflict with black)
"E203",
# E231 missing whitespace after ',' (conflict with black)
"E231",
# E501 line too long (conflict with black)
"E501",
# E741 do not use variables named 'l', 'O', or 'I'
"E741",
# W503 line break before binary operator (conflict with black)
"W503",
# W504 line break after binary operator (conflict with black)
"W504",
# W505 doc line too long (conflict with black)
"W505",
# W605 invalid escape sequence 'x' (conflict with latex within docs)
"W605",
]
[tool.ruff]
exclude = [
"skyrl-agent/",
"examples/",
]
[tool.ruff.lint]
extend-select = ["I"]
ignore = [
"F722" # Syntax error in annotation - ignored because this doesn't play well with jaxtyping
]
[tool.ruff.lint.isort]
known-first-party = ["skyrl", "skyrl_gym", "skycap"]