Compare commits
618 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8435586d14 | |||
| 65df6ab599 | |||
| 6f44dcbf68 | |||
| 2bd0875fd3 | |||
| 217d8a55a3 | |||
| c632d8dd80 | |||
| 07c4ce894d | |||
| 4123f0c3bf | |||
| 30e675906f | |||
| 5415d850af | |||
| 8e22ca1902 | |||
| bd80905120 | |||
| 352f1bd7c1 | |||
| 2914f3e2ef | |||
| 5e9c711724 | |||
| 8f8b8f95fd | |||
| 0eb4af11f9 | |||
| d2c4d1f630 | |||
| d1d6895782 | |||
| 5361ade1e5 | |||
| b2ca3a50af | |||
| 5b6fd360f2 | |||
| f033bb2cb8 | |||
| 68cef03c03 | |||
| b81c9b8e69 | |||
| ea4e49993d | |||
| df20382d21 | |||
| 5b7b4bb7cc | |||
| 4a052204fa | |||
| 4a610aee27 | |||
| f0a77cfad7 | |||
| d6df2e584d | |||
| 8fd5bd8ea1 | |||
| ee8d8ff828 | |||
| aa6ab2c654 | |||
| 9e6ed8f8b5 | |||
| fdf8ec957d | |||
| f195883f95 | |||
| ec5e62aadb | |||
| 7bfb0e7b15 | |||
| 9fcb436b1a | |||
| ee707b6626 | |||
| 0b44d5d5d3 | |||
| 77001e158d | |||
| e1d42667c8 | |||
| 398dabfc91 | |||
| 1e63f4b89c | |||
| 2955648c69 | |||
| 8c0f2dcc77 | |||
| d95c303a8b | |||
| eca8429228 | |||
| b260ed4d5a | |||
| e444cab8c7 | |||
| 04eabd9fc6 | |||
| 7789a86280 | |||
| bc809ec535 | |||
| 675767e197 | |||
| 9303317969 | |||
| 5ac5e8ae02 | |||
| bfbcf3c903 | |||
| 4a27e1cbe5 | |||
| 4e3e924e52 | |||
| 28376d57d7 | |||
| a6d1994646 | |||
| 97f846bf74 | |||
| 6fcbcd5913 | |||
| fce6e1fde9 | |||
| 46d8df5ee8 | |||
| a939ed97d9 | |||
| 1b8b0e05fc | |||
| 3abc413cb6 | |||
| d31618ceb2 | |||
| caa77f1dfe | |||
| ba8faa0fe8 | |||
| ae44285677 | |||
| 29b2e61cf3 | |||
| 0b40cb724a | |||
| d60e54fe5a | |||
| 9acaa7964a | |||
| bfd2f124fe | |||
| d75088d485 | |||
| 2f27d835e6 | |||
| 8ed1b4f46e | |||
| 812896373d | |||
| 298c8c611b | |||
| 5d2efe65ea | |||
| 526b9c73b4 | |||
| 578053b88b | |||
| f7bf3f7848 | |||
| 92d8a5dfea | |||
| a42eab12d7 | |||
| 01cd9d07e9 | |||
| 1cdb2c5914 | |||
| 2a99432616 | |||
| 705475bc1b | |||
| e4754e08ee | |||
| 07e075dd06 | |||
| 287cff5ded | |||
| 485d3eb2e1 | |||
| dea1a9529d | |||
| 39ecc2233a | |||
| 50ba96ed73 | |||
| 18cfc88a90 | |||
| da8256572e | |||
| 811ca1d6cc | |||
| 25fac86231 | |||
| daec1a2513 | |||
| 6876699912 | |||
| aeeb5b24f1 | |||
| 6107f66810 | |||
| 4bb3748fe5 | |||
| 138da36ee9 | |||
| 75942cb066 | |||
| fcd61966bc | |||
| fe33bfd250 | |||
| 8ab16e5c9e | |||
| 626849b963 | |||
| 374b3a29c5 | |||
| dd01656483 | |||
| 44e4501007 | |||
| 2ebebfae9d | |||
| 1b60f07540 | |||
| 5d40fa78cb | |||
| 8c0d20ac85 | |||
| f036ecdb2f | |||
| 920194ec0e | |||
| 0e48b372af | |||
| 36e7cfe22e | |||
| c25eb23cdc | |||
| db7e7151d0 | |||
| cd5649b1c2 | |||
| b5a693ebc9 | |||
| c1e6c95863 | |||
| 4c2de5614d | |||
| a58eb1189a | |||
| fd557ce49d | |||
| 634f438bca | |||
| 9fc5ec9c31 | |||
| 6143e04e4a | |||
| 3a01068d1f | |||
| 2bfda1cf89 | |||
| 23ff6c8553 | |||
| 3ebd587eb9 | |||
| d6518d29dd | |||
| c5447e39da | |||
| 4111c9261a | |||
| c4721693fa | |||
| 668a44d513 | |||
| 7cb50273ce | |||
| add6bf50f4 | |||
| 103bf4e87a | |||
| 8a4f5e852f | |||
| 49e80825a0 | |||
| ef298efb1a | |||
| 79e357bde4 | |||
| d3e15b4d6a | |||
| ff5411eecd | |||
| 530a0d9dd6 | |||
| cf5019a241 | |||
| 220fb45bc9 | |||
| 35338f0d67 | |||
| 827e4fdc56 | |||
| ebfa78a8b3 | |||
| 48d4a0b086 | |||
| fa8b230543 | |||
| 465a0233c8 | |||
| 41855545b7 | |||
| 56f84be193 | |||
| 584aa39d86 | |||
| 317cbda491 | |||
| 17ff78d832 | |||
| 6766e2f78c | |||
| f8159ba872 | |||
| f074c1cd33 | |||
| 3482dfde10 | |||
| 17dffd7e1d | |||
| 87ff4ba4e2 | |||
| 672a710662 | |||
| d3b3fa43c0 | |||
| 081b093e6c | |||
| a2a28b4fc9 | |||
| ac0f5f5615 | |||
| e037ddbba2 | |||
| 805f70eca8 | |||
| ae53bfe125 | |||
| 254fa161c4 | |||
| 3e28829410 | |||
| 462f23fe26 | |||
| ad26645c6d | |||
| cff633f309 | |||
| 29033ea20c | |||
| a85b4d7993 | |||
| 88d90755c2 | |||
| 79c026d2ca | |||
| ab6ca6a574 | |||
| 3700d9c3a1 | |||
| 7dd323d0b3 | |||
| fe72b2019d | |||
| d2888593e0 | |||
| a5c3cf7aba | |||
| c3e9e476fe | |||
| 06cba884c4 | |||
| 80255e4548 | |||
| deabe53444 | |||
| 73617f82ba | |||
| 1f33e3c245 | |||
| be511e7be0 | |||
| 94b709af23 | |||
| 313a537b31 | |||
| c89247e6ed | |||
| bbcb514caf | |||
| 349ae38c64 | |||
| a803cc64f0 | |||
| b78d8a6bc2 | |||
| e3099b42fa | |||
| 5248d7fe86 | |||
| 3fc35e6ea3 | |||
| b8fdd094d8 | |||
| fc1b15aaf7 | |||
| 66093a4968 | |||
| b71252ef98 | |||
| eee55e2d0c | |||
| 4479cb036e | |||
| 9a5f046448 | |||
| 392e09efbc | |||
| 74d4c3576c | |||
| b89d105599 | |||
| e7aa70bacf | |||
| 307d5e6056 | |||
| 7b6f7b690a | |||
| 4e7e02853c | |||
| 2e9ac2a4be | |||
| cfa09ae395 | |||
| 9b3163a4df | |||
| 58dc0cea12 | |||
| bac897a316 | |||
| 6606e926ea | |||
| 6aea2f6ce7 | |||
| abc6f44526 | |||
| e09180e8f9 | |||
| c351368305 | |||
| b899c2f1b7 | |||
| 6e875a1dc1 | |||
| b9bd98b850 | |||
| 9dc3b80115 | |||
| ec76bb7fd4 | |||
| bd723cc86b | |||
| c690432db6 | |||
| 163606fc99 | |||
| 6dfa1ce93b | |||
| b8bdbb97b5 | |||
| 6212bc25f3 | |||
| 2ce355d1ef | |||
| 67089804a4 | |||
| 17bd0d5def | |||
| ed3dd4a00d | |||
| 1966c8714c | |||
| 77392d321d | |||
| e6859a0da1 | |||
| a41979a06a | |||
| ee35412ed9 | |||
| 16756d9569 | |||
| 51787311ce | |||
| 167e1ae9f6 | |||
| 7901c3b39a | |||
| 805d1454c7 | |||
| 7255d65220 | |||
| 9fa477335c | |||
| f6215e2eb5 | |||
| f6119deac9 | |||
| 46f7214e7e | |||
| 45fb153b65 | |||
| be5fa1a333 | |||
| 2d1f094ad8 | |||
| dba4531fa3 | |||
| 972d646771 | |||
| 017ef879d7 | |||
| f965fd9083 | |||
| 54065e7621 | |||
| ea75bc4921 | |||
| 2ba3aa4d03 | |||
| 7969c1d3e5 | |||
| 49d97afb98 | |||
| 9684a190dc | |||
| 6c4d8a017d | |||
| 9f14cfcca6 | |||
| e52af54a60 | |||
| 80723d8e72 | |||
| 5b60310abc | |||
| 495de26769 | |||
| e507e1eee3 | |||
| b93d74c433 | |||
| 16a1dd6b27 | |||
| b5624a5123 | |||
| c36ac5afd9 | |||
| 7795f6c38b | |||
| ca8ab4cae3 | |||
| 05f350747e | |||
| fa450dce8e | |||
| 372eb1de6b | |||
| c49a999782 | |||
| 0b6985d1b0 | |||
| 11cd7928e6 | |||
| 9e0d9f168d | |||
| 7563176679 | |||
| 630a7d6b2a | |||
| ba6cc3ad64 | |||
| 3ce443d865 | |||
| 472cf4ce2c | |||
| 7efbda7dfa | |||
| bf793c7ae0 | |||
| 56eb814780 | |||
| 43112c6535 | |||
| 72d71712ad | |||
| fc97f79f5d | |||
| d2a0b57705 | |||
| 80f7014b54 | |||
| 6711ac0cb0 | |||
| a4d03e7f65 | |||
| 697f4abebd | |||
| bbb230b35b | |||
| 2dde63e6a6 | |||
| 5ed39ef77d | |||
| 02d489f682 | |||
| a36ccefeeb | |||
| 37744274dc | |||
| bf7b7df6ab | |||
| 8b22b9152a | |||
| 3ea8d246bc | |||
| 7f1de15c1b | |||
| 2b83bcaadd | |||
| f7ffc45f4f | |||
| 3f4616aeb8 | |||
| dd4495a62d | |||
| ab3d6f5d8c | |||
| 2d31b235d6 | |||
| 28699ca231 | |||
| aef7277702 | |||
| 90e3004919 | |||
| bc7db177b5 | |||
| 95d2bc188f | |||
| 894e17b58f | |||
| 1429ebec36 | |||
| 1d5ce35a1f | |||
| e8ee4f19f7 | |||
| dfb7a6691a | |||
| 358f6af60e | |||
| 0b47b57a73 | |||
| eed4853804 | |||
| 191c3459c3 | |||
| 808694991f | |||
| dfc8a1346b | |||
| 9622589746 | |||
| dcf3c7c9d0 | |||
| ae6e2f3672 | |||
| 0cedac72d7 | |||
| 9d42a6b8b0 | |||
| a748ed74ba | |||
| d1a1864a90 | |||
| 93cdde60d7 | |||
| 83c053794a | |||
| ddaa7f8d9a | |||
| 4900f4b924 | |||
| 94f01bd6ab | |||
| 8b47960b14 | |||
| 73672f1bc2 | |||
| bdd92de828 | |||
| e1bbe5e3ba | |||
| d43d3b4b38 | |||
| 62d401b7ad | |||
| 46c46356b0 | |||
| 9a4326d7cb | |||
| 7fbca400e5 | |||
| 1a5ef5c1dd | |||
| e0216ffec4 | |||
| b11caf61a6 | |||
| 1339372556 | |||
| e3b3f650e0 | |||
| 790f122ee7 | |||
| 2fb2fa96d2 | |||
| 56a8ff1114 | |||
| c8e8a8acea | |||
| 087a8fdd65 | |||
| c6137b73a1 | |||
| 21c071bb7e | |||
| 8ed4cd3c7a | |||
| 14c1f9d5a2 | |||
| b9a2decc76 | |||
| d27508ca98 | |||
| 2abdac190d | |||
| 5e7221e736 | |||
| bd565aaa90 | |||
| 11bdda2b0d | |||
| c746af2f76 | |||
| 49bf9cd9ec | |||
| 9864b8fbff | |||
| 82d8535048 | |||
| 5fa6582491 | |||
| 3f36754d64 | |||
| b0592efe1f | |||
| 5f3093d62a | |||
| bfb94a8750 | |||
| 25eda47a29 | |||
| a214474402 | |||
| 3b5d733861 | |||
| 158f5df28e | |||
| 40dc59205b | |||
| c1a43b6c3a | |||
| 4235731a0d | |||
| 22b80b38bf | |||
| 9f178accaf | |||
| 68a47d5087 | |||
| 4b36b25aad | |||
| a13e09fc6c | |||
| e63c340ebd | |||
| e62b7ca252 | |||
| f66d87745f | |||
| 52090e9dd5 | |||
| fdaf3f1777 | |||
| 087c7d350a | |||
| 2203070ef0 | |||
| a6078caa6c | |||
| f1a8072546 | |||
| 60f9955606 | |||
| 1d199b21c7 | |||
| 5f62ecb6f4 | |||
| 1948c2ba6d | |||
| 1e36e660b1 | |||
| 267b9936bb | |||
| 14714ded2b | |||
| c3f5cc7a39 | |||
| ee0fffd3a2 | |||
| 94d1cd780e | |||
| bbd5c2a30a | |||
| c6f4e6c283 | |||
| 8c504518bb | |||
| 337cce7fdc | |||
| 42c63d7a01 | |||
| 5ecd23792d | |||
| ad89e173e1 | |||
| feebaec24c | |||
| 4adf4e3ea4 | |||
| 0294eb5d32 | |||
| f9fe772e10 | |||
| 9f8a25ffdc | |||
| 3082ac0ee0 | |||
| 56e5c7ce62 | |||
| 21892cc6d3 | |||
| 34811cb454 | |||
| 003b8c6f83 | |||
| 63b6d42669 | |||
| 8c219175f5 | |||
| 931ddcfdcc | |||
| ce80b09a4a | |||
| f0546ca6c5 | |||
| 3a3bfeef31 | |||
| a733950b74 | |||
| 662fd90784 | |||
| 475c2adb91 | |||
| bffc7013f9 | |||
| 4cf8fb94e7 | |||
| ab185a5c5a | |||
| d581cbcd63 | |||
| 3459caa1de | |||
| f3fd58e72a | |||
| b3cb5e1337 | |||
| 3761c0f54c | |||
| d4334182be | |||
| 57c3c0525e | |||
| e356593f73 | |||
| 0e033831d5 | |||
| 0d721228d5 | |||
| e49b75b7d8 | |||
| eab691b1a1 | |||
| fd6494873d | |||
| 6cbfc1fee0 | |||
| b986ae132a | |||
| f24a47969e | |||
| a0bc1827d9 | |||
| f2869cea30 | |||
| 77cf447717 | |||
| 790ed3efb3 | |||
| 5ae7933d41 | |||
| 2ab977ed18 | |||
| 1eae9a34f0 | |||
| 4e7748b059 | |||
| 582f67cade | |||
| 9e23ba6b50 | |||
| 3f8a3ac0f1 | |||
| e0b55ab057 | |||
| 421f2773c7 | |||
| f717f9982f | |||
| 44dbfde0b4 | |||
| 713511902d | |||
| 80531c9c28 | |||
| 37daf2104f | |||
| 9afdd4570c | |||
| 55fbe66fe7 | |||
| 3794c97c1e | |||
| 848623766d | |||
| 4cd09ec900 | |||
| 3ed5e1e5b5 | |||
| 3f372ff7b3 | |||
| c453c41fd2 | |||
| 5c8ac61af6 | |||
| a02e1b91d9 | |||
| 496e793f0b | |||
| 80d306ff54 | |||
| 5f67bfe137 | |||
| f8c45b6ca8 | |||
| 01955aead7 | |||
| 0a9e3d75f2 | |||
| a3b2db18fa | |||
| 268bd77ce6 | |||
| ab6ea3c131 | |||
| d16538da96 | |||
| 8ce40a0410 | |||
| a9c7dbef22 | |||
| e69d24f4a8 | |||
| 955a0cc9a3 | |||
| 3966db6d2a | |||
| 584600d72e | |||
| 955524658d | |||
| 4d5e133a06 | |||
| df2a159b00 | |||
| 675fc86727 | |||
| c16b3a21b6 | |||
| aab976558b | |||
| 91c85aef7e | |||
| a1c36b55a0 | |||
| 6700878f64 | |||
| 56fa8d6881 | |||
| b0f28423b2 | |||
| e28fb8cb6b | |||
| fae0fba3d7 | |||
| 0decbabfbe | |||
| af7a6aa2cc | |||
| ae4e992771 | |||
| 8abe85ad91 | |||
| 22454adedb | |||
| 483c518d74 | |||
| 948506f3b6 | |||
| 34437dd6f5 | |||
| 951fa685b5 | |||
| 2b12e29f32 | |||
| 0e04363f4c | |||
| e91187b491 | |||
| c4b829dbe7 | |||
| 5c274703fe | |||
| 55284f8394 | |||
| 895bffc5b6 | |||
| 8e06fe6902 | |||
| 7d8dccd2b0 | |||
| 89a887d835 | |||
| d31090e9ee | |||
| c6298a96fd | |||
| fcb2a0811e | |||
| bdf6a8f223 | |||
| 8c673c241e | |||
| b7d2d6d6cb | |||
| 46a08d7272 | |||
| cca9e9d62f | |||
| 8aeb0ec1ba | |||
| 65ba916743 | |||
| d35a33dc14 | |||
| 418691e5a2 | |||
| 8b33ddc028 | |||
| bdc0b7e2a8 | |||
| 2adaddbf7c | |||
| 86becfdbff | |||
| 994384cb9b | |||
| 7f8395b941 | |||
| 495d4eba38 | |||
| 8dd4c5a1a3 | |||
| b1ae0b75c4 | |||
| 9f8ec4950f | |||
| c08da2ae37 | |||
| f032ffa319 | |||
| a4cf2fd5fd | |||
| 8a4ecbacf6 | |||
| dd337d456e | |||
| a0626bdea9 | |||
| 2489d068ba | |||
| f4814949cb | |||
| a0791e8b13 | |||
| 26d1df698d | |||
| 7bf418ea67 | |||
| d735fb27c4 | |||
| 347638f218 | |||
| a42839b7fb | |||
| e11036cf7b | |||
| 685eea70a6 | |||
| a1a4fe39c6 | |||
| a9d0c9237d | |||
| 1513b52a05 | |||
| 4ec1029577 | |||
| 63c133051d | |||
| 2316a8451e | |||
| 138ad0e487 | |||
| 504ef2c627 | |||
| a63197355c | |||
| 3eb725fade | |||
| 66bcfeba11 | |||
| a9208ab700 | |||
| ddc8997b8c | |||
| 0a92600a4c | |||
| ba10c845e1 | |||
| 7ad967daf7 | |||
| f6db2dc8ab | |||
| 5724f63cfc | |||
| bd6c62dd7c | |||
| b9595b8b0d | |||
| a54c6f39e8 | |||
| 6bd01959cd | |||
| d00cc3e4aa | |||
| bea013632e | |||
| 13bd48dbfc | |||
| 31309810f7 |
@@ -0,0 +1,232 @@
|
||||
---
|
||||
name: release
|
||||
description: Prepare and publish stable Agent Lightning releases through the repository's version bump, pull-request checks, merge, tag, PyPI trusted-publishing, and versioned-documentation workflows. Use when asked to plan, cut, verify, or explain a release; treat nightly TestPyPI builds as a separate path.
|
||||
---
|
||||
|
||||
# Release Agent Lightning
|
||||
|
||||
Merging a pull request does not publish a stable release. Stable publication is
|
||||
triggered only by pushing a `v*` tag to the canonical repository; the tagged
|
||||
commit is what gets tested, built, and uploaded. That same tag push also deploys
|
||||
versioned documentation and moves the public `stable` alias, so a release has two
|
||||
public side effects, not one.
|
||||
|
||||
## Establish the release state
|
||||
|
||||
1. Confirm the repository root, clean working tree, current branch, and remotes.
|
||||
2. Resolve the canonical `OWNER/REPO`, its default branch, and its permitted
|
||||
merge methods with
|
||||
`gh repo view OWNER/REPO --json nameWithOwner,defaultBranchRef,mergeCommitAllowed,rebaseMergeAllowed,squashMergeAllowed`.
|
||||
Identify the local remotes for that repository and the contributor fork by
|
||||
their URLs; do not assume particular remote names or merge settings.
|
||||
3. Inspect the release contract in:
|
||||
- `.github/workflows/pypi-release.yml`
|
||||
- `.github/workflows/docs.yml`
|
||||
- `.github/workflows/tests.yml`
|
||||
- `scripts/bump_version.sh`
|
||||
- `pyproject.toml`
|
||||
- `agentlightning/__init__.py`
|
||||
4. Confirm the canonical default branch is already green before branching from
|
||||
it. Resolve its current commit with
|
||||
`gh api repos/OWNER/REPO/commits/<default-branch>` and inspect that commit's
|
||||
check runs; a general recent-run listing can omit or mix commits. A release
|
||||
branch inherits every failure that main is carrying.
|
||||
5. Query the canonical repository's tags and compare them with the versions
|
||||
published at `https://pypi.org/pypi/agentlightning/json`. Confirm the target
|
||||
version exists in neither place, and stop for an explicit release decision
|
||||
when either of these holds:
|
||||
- A tag exists with no matching PyPI version. A published version is
|
||||
immutable, and its tag must never be reused or moved. A tag that never
|
||||
published is a different situation and still needs a human decision,
|
||||
informed by why it did not publish. See "Recovering a tag that never
|
||||
published" below.
|
||||
- The proposed bump would skip a version that was tagged but never published.
|
||||
6. Treat verified PyPI trusted-publisher configuration for the canonical
|
||||
repository and `pypi-release.yml` as a prerequisite. If it cannot be
|
||||
inspected directly, require confirmation from an authorized PyPI project
|
||||
owner before pushing the release tag.
|
||||
|
||||
## Prepare and merge the version pull request
|
||||
|
||||
Start a release branch from a freshly fetched canonical default branch, not
|
||||
from another feature branch. The branch name is only a recommendation:
|
||||
|
||||
```bash
|
||||
git fetch <canonical-remote> <default-branch>
|
||||
git switch -c chore/release-vX.Y.Z <canonical-remote>/<default-branch>
|
||||
scripts/bump_version.sh patch # or minor / major
|
||||
```
|
||||
|
||||
The script updates the project version with uv and then edits
|
||||
`agentlightning/__init__.py` separately. If it fails or is interrupted between
|
||||
those writes, only some of the three version files may be updated. Inspect
|
||||
`git diff` after any failure and restore or reconcile all three files before
|
||||
retrying; blindly rerunning a partial patch bump can advance the version twice.
|
||||
|
||||
The bump rewrites exactly three files. Confirm that with `git diff --stat`:
|
||||
|
||||
- `pyproject.toml`
|
||||
- the `agentlightning` entry in `uv.lock`
|
||||
- `agentlightning.__version__` in `agentlightning/__init__.py`
|
||||
|
||||
Other version strings in the tree, such as the FastAPI `version` in
|
||||
`agentlightning/server/app.py`, are deliberately outside the bump. Leave them
|
||||
alone; changing them is a separate pull request, not release work.
|
||||
|
||||
Review the version diff, but do not run the release tests or package build
|
||||
locally as a matter of course. `tests.yml` runs a broader test suite and the
|
||||
same package build on the pull request, covering the narrower tests and build
|
||||
that `pypi-release.yml` will run on the tag. The pull request's GitHub checks
|
||||
are therefore the verification gate. Reproduce a single failure locally only
|
||||
when the workflow logs are not enough to fix it.
|
||||
|
||||
Commit the version change, push it to the fork, and open the pull request with
|
||||
the GitHub CLI when those external actions are authorized:
|
||||
|
||||
```bash
|
||||
git commit -am "Bump version to X.Y.Z"
|
||||
git push -u <fork-remote> <release-branch>
|
||||
gh pr create --repo OWNER/REPO \
|
||||
--base <default-branch> \
|
||||
--head <fork-owner>:<release-branch> \
|
||||
--title "Bump version to X.Y.Z" \
|
||||
--body "Prepare the vX.Y.Z release."
|
||||
```
|
||||
|
||||
`gh pr create` refuses to run without `--title` and `--body` outside an
|
||||
interactive terminal, and every `gh` call needs `--repo OWNER/REPO` so it acts
|
||||
on the canonical repository rather than the fork.
|
||||
|
||||
Follow the pull request through its required checks with
|
||||
`gh pr checks <pr> --repo OWNER/REPO --watch`. If a check fails, take the run id
|
||||
from that output, inspect it with
|
||||
`gh run view <run-id> --repo OWNER/REPO --log-failed`, correct the source on the
|
||||
same branch, and resume watching. Once every required check has succeeded,
|
||||
extract the reviewed head and pass both a permitted merge-method flag from step
|
||||
2 and `--match-head-commit` to `gh pr merge`:
|
||||
|
||||
```bash
|
||||
HEAD_SHA="$(gh pr view <pr> --repo OWNER/REPO --json headRefOid --jq .headRefOid)"
|
||||
gh pr merge <pr> --repo OWNER/REPO <merge-method-flag> \
|
||||
--match-head-commit "$HEAD_SHA"
|
||||
```
|
||||
|
||||
Replace `<merge-method-flag>` with one permitted flag discovered in step 2:
|
||||
`--merge`, `--rebase`, or `--squash`.
|
||||
|
||||
Committing, pushing, opening the pull request, and merging are each distinct
|
||||
external actions and each requires authorization.
|
||||
|
||||
## Tag and publish the merged release
|
||||
|
||||
After the pull request merges, update the local default branch from the
|
||||
canonical repository, then confirm that the commit you are about to tag is the
|
||||
one this pull request produced and not a later commit that landed behind it:
|
||||
|
||||
```bash
|
||||
git switch <default-branch>
|
||||
git pull --ff-only <canonical-remote> <default-branch>
|
||||
gh pr view <pr> --repo OWNER/REPO --json mergeCommit
|
||||
git rev-parse HEAD
|
||||
```
|
||||
|
||||
If HEAD has moved past the merge commit, tag the merge commit explicitly instead
|
||||
of HEAD.
|
||||
|
||||
`pypi-release.yml` fails the release when the packaged version does not equal the
|
||||
tag without its leading `v`, or when it does not equal the runtime
|
||||
`__version__`. Check both before tagging:
|
||||
|
||||
```bash
|
||||
uv version --short
|
||||
grep '^__version__' agentlightning/__init__.py
|
||||
```
|
||||
|
||||
The workflow itself reads the runtime value as
|
||||
`python -c 'from agentlightning import __version__; print(__version__)'` from
|
||||
the repository root before its dependency-sync step, so Python resolves the
|
||||
checkout through the current working directory. Read the file directly for the
|
||||
local pre-tag check; `agentlightning/__init__.py` assigns `__version__` as a
|
||||
single literal, making that check independent of the active Python environment.
|
||||
|
||||
GitHub reads workflow files as they exist **at the tagged commit**, not at the
|
||||
tip of the default branch. Confirm that the commit being tagged actually
|
||||
contains `.github/workflows/pypi-release.yml` with its `v*` trigger; a commit
|
||||
that predates the workflow will never publish, however the tag is pushed.
|
||||
|
||||
Immediately query the canonical repository and PyPI again to ensure that
|
||||
`vX.Y.Z` is still absent. Then create an annotated tag on the release commit and
|
||||
push it to the canonical repository:
|
||||
|
||||
```bash
|
||||
git tag -a vX.Y.Z -m "vX.Y.Z" <release-commit>
|
||||
git push <canonical-remote> vX.Y.Z
|
||||
```
|
||||
|
||||
The tag push starts the production PyPI publication, so obtain explicit
|
||||
authorization immediately before it.
|
||||
|
||||
## Follow both tag-triggered workflows
|
||||
|
||||
One tag push starts two workflows, and both belong to the release:
|
||||
|
||||
- `PyPI Release` (`pypi-release.yml`) re-checks the version against the tag,
|
||||
runs the tests, builds the wheel and source distribution, and uploads them to
|
||||
PyPI through trusted publishing.
|
||||
- `Deploy Documentation` (`docs.yml`) runs
|
||||
`mike deploy --push --update-aliases X.Y.Z stable`, which publishes the
|
||||
versioned documentation and repoints the public `stable` alias at this
|
||||
release.
|
||||
|
||||
Look up each run by workflow and tag rather than selecting from an unfiltered
|
||||
recent-run list:
|
||||
|
||||
```bash
|
||||
gh run list --repo OWNER/REPO --workflow pypi-release.yml \
|
||||
--branch vX.Y.Z --event push --limit 1
|
||||
gh run list --repo OWNER/REPO --workflow docs.yml \
|
||||
--branch vX.Y.Z --event push --limit 1
|
||||
```
|
||||
|
||||
Confirm both runs have the expected tag commit, then follow them to a terminal
|
||||
result with `gh run watch <run-id> --repo OWNER/REPO --exit-status`. After
|
||||
`PyPI Release` succeeds, verify that PyPI exposes the exact version with both
|
||||
the expected wheel and source distribution. After `Deploy Documentation`
|
||||
succeeds, verify that the published site serves `X.Y.Z` and that `stable`
|
||||
resolves to it. A green PyPI job with a failed documentation job is a
|
||||
half-finished release: report both workflow URLs and both outcomes.
|
||||
|
||||
For a transient workflow failure, rerun only with authorization. For a source
|
||||
or workflow defect, do not move the public tag; prepare a corrective release
|
||||
version. A GitHub Release and release notes are optional, separate publication
|
||||
actions and must not be created unless requested.
|
||||
|
||||
## Recovering a tag that never published
|
||||
|
||||
Separate the mechanics from the policy before proposing a recovery.
|
||||
|
||||
The mechanics: pushing a tag that already exists and points at the same commit
|
||||
changes no ref, so it starts no workflow run. Creating a tag, moving one to a
|
||||
different commit, or deleting and recreating one does change the ref and does
|
||||
start a run. What that run executes is the workflow file at the tagged commit,
|
||||
so a tag on a commit from before `pypi-release.yml` existed starts no PyPI
|
||||
publication no matter how it is pushed. Run
|
||||
`git ls-tree --name-only <tag> .github/workflows/` before assuming a re-push
|
||||
would help.
|
||||
|
||||
The policy: never move or reuse a tag whose version is on PyPI. That version is
|
||||
immutable, so a re-run could only fail at upload, and consumers who already
|
||||
resolved the tag would silently get different code.
|
||||
|
||||
Between those, a tag that never published is a decision for a release owner,
|
||||
not a default action. Releasing the next version from a commit that carries the
|
||||
current workflow is usually simpler and always safer than resurrecting the old
|
||||
tag. Note that a non-publishing tag may still have had effects: `docs.yml` has
|
||||
carried the `v*` trigger for longer than `pypi-release.yml`, so an older tag can
|
||||
have deployed documentation and moved `stable` without ever reaching PyPI.
|
||||
|
||||
## Nightly distinction
|
||||
|
||||
`.github/workflows/pypi-nightly.yml` publishes timestamped `.dev` builds to
|
||||
TestPyPI on its schedule or by manual dispatch. It does not create a stable
|
||||
release and should not be substituted for the tag-driven process above.
|
||||
@@ -0,0 +1,4 @@
|
||||
interface:
|
||||
display_name: "Release"
|
||||
short_description: "Prepare and publish Agent Lightning releases"
|
||||
default_prompt: "Use $release to prepare and publish a new Agent Lightning release."
|
||||
@@ -0,0 +1,81 @@
|
||||
# Version control / editor state
|
||||
# Local Python environments and caches
|
||||
# Local-only runtime/deploy state
|
||||
.git/
|
||||
.gitignore
|
||||
.gitattributes
|
||||
.vscode/
|
||||
.idea/
|
||||
.claude/
|
||||
.DS_Store
|
||||
**/.DS_Store
|
||||
|
||||
# Local Python environments and caches
|
||||
.venv/
|
||||
.venv.bak/
|
||||
venv/
|
||||
env/
|
||||
ENV/
|
||||
__pycache__/
|
||||
**/__pycache__/
|
||||
*.py[codz]
|
||||
*.pyo
|
||||
*.pyd
|
||||
*.so
|
||||
*.egg-info/
|
||||
.eggs/
|
||||
dist/
|
||||
build/
|
||||
.pytest_cache/
|
||||
.ruff_cache/
|
||||
.mypy_cache/
|
||||
.pyright/
|
||||
.ipynb_checkpoints/
|
||||
**/.ipynb_checkpoints/
|
||||
.cache/
|
||||
|
||||
# Local-only runtime/deploy state
|
||||
.local/
|
||||
.env
|
||||
**/.env
|
||||
.env.local
|
||||
*.env.local
|
||||
.envrc
|
||||
tmp/
|
||||
node_modules/
|
||||
checkpoints/
|
||||
artifacts/
|
||||
logs/
|
||||
**/logs/
|
||||
*.log
|
||||
*-debug.log
|
||||
2026-*-debug.log
|
||||
|
||||
# Files not needed for runtime images
|
||||
tests/
|
||||
docs/
|
||||
dev/
|
||||
uv.lock
|
||||
|
||||
# Large example data and generated outputs
|
||||
examples/*/data/
|
||||
examples/*/outputs/
|
||||
examples/*/wandb/
|
||||
examples/*/mlruns/
|
||||
wandb/
|
||||
runs/
|
||||
outputs/
|
||||
mlruns/
|
||||
|
||||
# Archives and large packaged artifacts
|
||||
*.zip
|
||||
*.tar
|
||||
*.tar.gz
|
||||
*.tgz
|
||||
*.tar.bz2
|
||||
*.tar.xz
|
||||
*.7z
|
||||
agentlightning-main.zip
|
||||
|
||||
# Docs/dev generated artifacts
|
||||
docs/refactor_review/public/
|
||||
@@ -0,0 +1,11 @@
|
||||
version: 2
|
||||
updates:
|
||||
- package-ecosystem: "github-actions"
|
||||
directory: "/"
|
||||
groups:
|
||||
github-actions:
|
||||
patterns: ["*"]
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
cooldown:
|
||||
default-days: 7
|
||||
@@ -8,6 +8,10 @@ on:
|
||||
- 'v*'
|
||||
workflow_dispatch:
|
||||
|
||||
concurrency:
|
||||
group: docs-deploy
|
||||
cancel-in-progress: false
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pages: write
|
||||
@@ -17,18 +21,17 @@ jobs:
|
||||
deploy:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
./scripts/setup_stable.sh
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
- name: Sync dependencies
|
||||
run: uv sync --frozen --no-default-groups --group dev --group docs
|
||||
|
||||
- name: Configure Git
|
||||
run: |
|
||||
@@ -51,10 +54,11 @@ jobs:
|
||||
- name: Deploy versioned docs
|
||||
if: startsWith(github.ref, 'refs/tags/')
|
||||
run: |
|
||||
mike deploy --push --update-aliases ${{ steps.version.outputs.version }} stable
|
||||
uv run --locked --no-sync mike deploy --push --update-aliases ${{ steps.version.outputs.version }} stable
|
||||
|
||||
- name: Deploy dev docs
|
||||
if: github.ref == 'refs/heads/main'
|
||||
run: |
|
||||
mike deploy --push latest
|
||||
mike set-default --push latest
|
||||
uv run --locked --no-sync mike deploy --push latest
|
||||
# Always set stable to default
|
||||
uv run --locked --no-sync mike set-default --push stable
|
||||
|
||||
@@ -1,156 +0,0 @@
|
||||
name: GPU Test
|
||||
permissions:
|
||||
contents: read
|
||||
on:
|
||||
schedule:
|
||||
# Every day at 3 AM UTC+8
|
||||
- cron: '0 19 * * *'
|
||||
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
examples:
|
||||
runs-on: [self-hosted, linux, gpu]
|
||||
timeout-minutes: 60
|
||||
strategy:
|
||||
matrix:
|
||||
setup: [stable, latest]
|
||||
fail-fast: false
|
||||
container:
|
||||
image: ghcr.io/microsoft/agent-lightning/base:latest
|
||||
options: --gpus all --ipc=host --interactive --tty
|
||||
steps:
|
||||
- name: Check GPU status
|
||||
run: nvidia-smi
|
||||
- uses: actions/checkout@v4
|
||||
- name: Create a virtual environment
|
||||
run: python3 -m venv .venv
|
||||
- name: Install deps inside the container (${{ matrix.setup }})
|
||||
run: |
|
||||
. .venv/bin/activate
|
||||
./scripts/setup_${{ matrix.setup }}_gpu.sh
|
||||
- name: Freeze dependencies
|
||||
run: |
|
||||
. .venv/bin/activate
|
||||
which python
|
||||
which pip
|
||||
which uvx
|
||||
pip list | tee requirements-freeze.txt
|
||||
- name: Upload dependencies artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: dependencies-${{ matrix.setup }}
|
||||
path: requirements-freeze.txt
|
||||
compression-level: 0
|
||||
- name: Prepare Spider dataset
|
||||
run: |
|
||||
set -ex
|
||||
. .venv/bin/activate
|
||||
cd examples/spider
|
||||
gdown --fuzzy https://drive.google.com/file/d/1oi9J1jZP9TyM35L85CL3qeGWl2jqlnL6/view
|
||||
unzip -q spider-data.zip -d data
|
||||
rm spider-data.zip
|
||||
- name: Prepare Calc-X dataset
|
||||
run: |
|
||||
set -ex
|
||||
. .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
gdown --fuzzy https://drive.google.com/file/d/1FQMyKLLd6hP9dw9rfZn1EZOWNvKaDsqw/view
|
||||
unzip calc-x-data.zip -d data
|
||||
rm calc-x-data.zip
|
||||
- name: Spider sanity check
|
||||
run: |
|
||||
set -ex
|
||||
. .venv/bin/activate
|
||||
cd examples/spider
|
||||
python sql_agent.py --trainer.n-workers 1 --trainer.dev true --trainer.max-tasks 2
|
||||
env:
|
||||
VERL_API_BASE: http://localhost:9999/
|
||||
OPENAI_API_BASE: ${{ secrets.OPENAI_API_BASE }}
|
||||
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
|
||||
- name: Calc-X MCP sanity check
|
||||
run: |
|
||||
set -ex
|
||||
. .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
python tests/test_mcp_calculator.py
|
||||
env:
|
||||
OPENAI_API_BASE: ${{ secrets.OPENAI_API_BASE }}
|
||||
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
|
||||
- name: Calc-X sanity check
|
||||
run: |
|
||||
set -ex
|
||||
. .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
python calc_agent_dev.py
|
||||
env:
|
||||
OPENAI_API_BASE: ${{ secrets.OPENAI_API_BASE }}
|
||||
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
|
||||
|
||||
# Calc-X training suddenly works after running the sanity check.
|
||||
# And it has to be run before Spider training.
|
||||
# The client side used to hang in many of my attempts.
|
||||
# Don't ask why. Don't touch this.
|
||||
- name: Calc-X training
|
||||
run: |
|
||||
set -ex
|
||||
source .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
../../scripts/restart_ray.sh
|
||||
sleep 5
|
||||
PYTHONUNBUFFERED=1 python calc_agent.py &
|
||||
bash train_ci.sh
|
||||
pkill -f calc_agent.py && echo "SIGTERM sent to calc_agent.py" || echo "No calc_agent.py process found"
|
||||
while pgrep -f calc_agent.py; do
|
||||
echo "Waiting for calc_agent.py to finish..."
|
||||
sleep 5
|
||||
done
|
||||
echo "calc_agent.py has finished."
|
||||
sleep 10
|
||||
shell: bash
|
||||
env:
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
id: calc_x_train
|
||||
|
||||
- name: Validate Calc-X training
|
||||
run: |
|
||||
set -ex
|
||||
. .venv/bin/activate
|
||||
python scripts/validate_example_wandb.py ${{ steps.calc_x_train.outputs.project_name }} ${{ steps.calc_x_train.outputs.run_name }}
|
||||
env:
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
- name: Spider training
|
||||
run: |
|
||||
set -ex
|
||||
source .venv/bin/activate
|
||||
cd examples/spider
|
||||
../../scripts/restart_ray.sh
|
||||
sleep 5
|
||||
PYTHONUNBUFFERED=1 python sql_agent.py --trainer.n-workers 10 &
|
||||
bash train_ci.sh
|
||||
pkill -f sql_agent.py && echo "SIGTERM sent to sql_agent.py" || echo "No sql_agent.py process found"
|
||||
while pgrep -f sql_agent.py; do
|
||||
echo "Waiting for sql_agent.py to finish..."
|
||||
sleep 5
|
||||
done
|
||||
echo "sql_agent.py has finished."
|
||||
sleep 10
|
||||
shell: bash
|
||||
env:
|
||||
VERL_API_BASE: http://localhost:9991/
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
id: spider_train
|
||||
if: success() || failure()
|
||||
|
||||
- name: Validate Spider training
|
||||
run: |
|
||||
set -ex
|
||||
. .venv/bin/activate
|
||||
python scripts/validate_example_wandb.py ${{ steps.spider_train.outputs.project_name }} ${{ steps.spider_train.outputs.run_name }}
|
||||
env:
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
- name: Cleanup
|
||||
run: ./scripts/cleanup.sh
|
||||
if: success() || failure()
|
||||
@@ -2,58 +2,67 @@ name: PyPI Nightly Build
|
||||
|
||||
on:
|
||||
schedule:
|
||||
# Run daily at 6:00 AM UTC
|
||||
- cron: '0 6 * * *'
|
||||
workflow_dispatch: # Allow manual trigger
|
||||
# Run daily at 6:00 AM UTC+8.
|
||||
- cron: '0 22 * * *'
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: pypi-nightly
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
publish-test-pypi:
|
||||
name: Publish nightly package to TestPyPI
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
permissions:
|
||||
id-token: write # IMPORTANT: this permission is mandatory for trusted publishing
|
||||
contents: read
|
||||
id-token: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
- name: Install build dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -e .[dev]
|
||||
|
||||
- name: Get current version
|
||||
id: get_version
|
||||
run: |
|
||||
VERSION=$(grep '^version = ' pyproject.toml | sed 's/version = "\(.*\)"/\1/')
|
||||
echo "version=$VERSION" >> $GITHUB_OUTPUT
|
||||
echo "Current version: $VERSION"
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Create development version
|
||||
id: version
|
||||
shell: bash
|
||||
run: |
|
||||
# Create a dev version with timestamp
|
||||
TIMESTAMP=$(date +%Y%m%d%H%M%S)
|
||||
DEV_VERSION="${{ steps.get_version.outputs.version }}.dev$TIMESTAMP"
|
||||
echo "Creating dev version: $DEV_VERSION"
|
||||
./scripts/bump_version.sh "$DEV_VERSION"
|
||||
set -euo pipefail
|
||||
BASE_VERSION=$(uv version --short)
|
||||
TIMESTAMP=$(date -u +%Y%m%d%H%M%S)
|
||||
DEV_VERSION="${BASE_VERSION}.dev${TIMESTAMP}"
|
||||
uv version --frozen "${DEV_VERSION}"
|
||||
sed -i "s/^__version__ = \".*\"$/__version__ = \"${DEV_VERSION}\"/" agentlightning/__init__.py
|
||||
echo "version=${DEV_VERSION}" >> "${GITHUB_OUTPUT}"
|
||||
|
||||
- name: Verify version consistency
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
PACKAGE_VERSION=$(uv version --short)
|
||||
RUNTIME_VERSION=$(python -c 'from agentlightning import __version__; print(__version__)')
|
||||
if [[ "${PACKAGE_VERSION}" != "${RUNTIME_VERSION}" ]]; then
|
||||
echo "Package version ${PACKAGE_VERSION} does not match runtime version ${RUNTIME_VERSION}." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
hatch build
|
||||
run: uv build --no-sources
|
||||
|
||||
- name: Publish to Test PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
- name: Verify package contents
|
||||
run: |
|
||||
python -m tarfile -l dist/*.tar.gz
|
||||
python -m zipfile -l dist/*.whl
|
||||
|
||||
- name: Publish ${{ steps.version.outputs.version }} to TestPyPI
|
||||
uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
|
||||
with:
|
||||
repository-url: https://test.pypi.org/legacy/
|
||||
|
||||
- name: Test installation from Test PyPI
|
||||
run: |
|
||||
# Wait a bit for the package to be available
|
||||
sleep 30
|
||||
pip install --index-url https://test.pypi.org/simple/ --extra-index-url https://pypi.org/simple/ agentlightning
|
||||
python -c "import agentlightning; print('Package installed successfully')"
|
||||
|
||||
@@ -3,67 +3,64 @@ name: PyPI Release
|
||||
on:
|
||||
push:
|
||||
tags:
|
||||
- 'v*' # Trigger on version tags like v1.0.0, v1.2.3, etc.
|
||||
workflow_dispatch: # Allow manual trigger
|
||||
- 'v*'
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
check-version:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
outputs:
|
||||
version: ${{ steps.get_version.outputs.version }}
|
||||
tag_version: ${{ steps.get_tag.outputs.tag_version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Get version from pyproject.toml
|
||||
id: get_version
|
||||
run: |
|
||||
VERSION=$(grep '^version = ' pyproject.toml | sed 's/version = "\(.*\)"/\1/')
|
||||
echo "version=$VERSION" >> $GITHUB_OUTPUT
|
||||
echo "Package version: $VERSION"
|
||||
|
||||
- name: Get tag version
|
||||
id: get_tag
|
||||
run: |
|
||||
TAG_VERSION=${GITHUB_REF#refs/tags/v}
|
||||
echo "tag_version=$TAG_VERSION" >> $GITHUB_OUTPUT
|
||||
echo "Tag version: $TAG_VERSION"
|
||||
|
||||
- name: Verify version matches tag
|
||||
run: |
|
||||
if [ "${{ steps.get_version.outputs.version }}" != "${{ steps.get_tag.outputs.tag_version }}" ]; then
|
||||
echo "Error: Version in pyproject.toml (${{ steps.get_version.outputs.version }}) does not match tag (${{ steps.get_tag.outputs.tag_version }})"
|
||||
exit 1
|
||||
fi
|
||||
echo "Version check passed!"
|
||||
|
||||
publish-pypi:
|
||||
needs: check-version
|
||||
name: Test, build, and publish package
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
permissions:
|
||||
id-token: write # IMPORTANT: this permission is mandatory for trusted publishing
|
||||
contents: read
|
||||
id-token: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Install build dependencies
|
||||
- name: Verify version matches tag
|
||||
shell: bash
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -e .[dev]
|
||||
set -euo pipefail
|
||||
PACKAGE_VERSION=$(uv version --short)
|
||||
TAG_VERSION="${GITHUB_REF_NAME#v}"
|
||||
if [[ "${PACKAGE_VERSION}" != "${TAG_VERSION}" ]]; then
|
||||
echo "Package version ${PACKAGE_VERSION} does not match tag ${GITHUB_REF_NAME}." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Verify version consistency
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
PACKAGE_VERSION=$(uv version --short)
|
||||
RUNTIME_VERSION=$(python -c 'from agentlightning import __version__; print(__version__)')
|
||||
if [[ "${PACKAGE_VERSION}" != "${RUNTIME_VERSION}" ]]; then
|
||||
echo "Package version ${PACKAGE_VERSION} does not match runtime version ${RUNTIME_VERSION}." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Sync test dependencies
|
||||
run: uv sync --frozen --no-default-groups --extra dev --group dev
|
||||
|
||||
- name: Run tests
|
||||
run: >-
|
||||
uv run --locked --no-sync pytest -v --durations=20
|
||||
tests/server
|
||||
tests/controller
|
||||
tests/test_package.py
|
||||
tests/examples/test_swe_smith_images.py
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
hatch build
|
||||
run: uv build --no-sources
|
||||
|
||||
- name: Verify package contents
|
||||
run: |
|
||||
@@ -71,11 +68,4 @@ jobs:
|
||||
python -m zipfile -l dist/*.whl
|
||||
|
||||
- name: Publish to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
|
||||
- name: Test installation from PyPI
|
||||
run: |
|
||||
# Wait a bit for the package to be available
|
||||
sleep 30
|
||||
pip install --index-url https://test.pypi.org/simple/ --extra-index-url https://pypi.org/simple/ agentlightning
|
||||
python -c "import agentlightning; print('Package installed successfully')"
|
||||
uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
name: Validate Agent Skills
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'skills/**'
|
||||
- '.agents/skills/**'
|
||||
- '.github/workflows/skills.yml'
|
||||
pull_request:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'skills/**'
|
||||
- '.agents/skills/**'
|
||||
- '.github/workflows/skills.yml'
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
validate:
|
||||
name: Validate Agent Skills and Claude plugin formats
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
- name: Validate Agent Skills format
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
# Published skills live in skills/ and repository-local agent skills in
|
||||
# .agents/skills/. Both use the <root>/<name>/SKILL.md layout.
|
||||
mapfile -t SKILLS < <(
|
||||
find skills .agents/skills -mindepth 2 -maxdepth 2 -name SKILL.md -printf '%h\n' | sort
|
||||
)
|
||||
if [ "${#SKILLS[@]}" -eq 0 ]; then
|
||||
echo "No SKILL.md found under skills/ or .agents/skills/." >&2
|
||||
exit 1
|
||||
fi
|
||||
for SKILL in "${SKILLS[@]}"; do
|
||||
echo "::group::${SKILL}"
|
||||
uvx --from 'skills-ref==0.1.1' agentskills validate "${SKILL}"
|
||||
echo "::endgroup::"
|
||||
done
|
||||
- uses: actions/setup-node@249970729cb0ef3589644e2896645e5dc5ba9c38 # v6.5.0
|
||||
with:
|
||||
node-version: '22'
|
||||
- name: Validate Claude Code plugins
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
mapfile -t PLUGINS < <(
|
||||
find skills .agents/skills -mindepth 1 -maxdepth 1 -type d \
|
||||
-exec test -f '{}/.claude-plugin/plugin.json' \; -print | sort
|
||||
)
|
||||
if [ "${#PLUGINS[@]}" -eq 0 ]; then
|
||||
echo "No Claude Code plugin found under skills/ or .agents/skills/." >&2
|
||||
exit 1
|
||||
fi
|
||||
for PLUGIN in "${PLUGINS[@]}"; do
|
||||
echo "::group::${PLUGIN}"
|
||||
npx --yes @anthropic-ai/claude-code@2.1.218 plugin validate "${PLUGIN}"
|
||||
echo "::endgroup::"
|
||||
done
|
||||
@@ -1,99 +1,117 @@
|
||||
name: CPU Test
|
||||
name: Test
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [ main ]
|
||||
branches: [main]
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
branches: [main]
|
||||
workflow_dispatch:
|
||||
|
||||
schedule:
|
||||
# Every day at noon and midnight
|
||||
- cron: '0 0,12 * * *'
|
||||
concurrency:
|
||||
group: test-${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
|
||||
lint:
|
||||
name: Lint with Black
|
||||
name: Lint Python and repository files
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/setup-python@v4
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
- name: Install dependencies
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
- name: Sync lint dependencies
|
||||
run: uv sync --frozen --no-default-groups --extra dev --group dev
|
||||
- name: Run pre-commit checks
|
||||
run: uv run --locked --no-sync pre-commit run --all-files --show-diff-on-failure
|
||||
- name: Run Ruff
|
||||
run: uv run --locked --no-sync ruff check .
|
||||
- name: Check Ruff formatting
|
||||
run: uv run --locked --no-sync ruff format --check .
|
||||
- name: Check Python headers
|
||||
run: uv run --locked --no-sync python scripts/check_headers.py
|
||||
|
||||
typecheck:
|
||||
name: Type-check Python
|
||||
runs-on: ubuntu-latest
|
||||
# Split from `lint` because the verl-cpu group pulls torch and friends.
|
||||
# Pyright needs them to check agentlightning/verl and tests/verl; the fast
|
||||
# checks above should not wait on that download.
|
||||
timeout-minutes: 20
|
||||
steps:
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
- name: Sync type-check dependencies
|
||||
run: uv sync --frozen --no-default-groups --extra dev --group dev --group verl-cpu
|
||||
- name: Run Pyright
|
||||
run: uv run --locked --no-sync pyright
|
||||
|
||||
test:
|
||||
name: Run tests
|
||||
runs-on: ubuntu-latest
|
||||
# verl-cpu is needed for tests/verl; the rest of the suite only needs dev.
|
||||
timeout-minutes: 20
|
||||
steps:
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
- name: Sync test dependencies
|
||||
run: uv sync --frozen --no-default-groups --extra dev --group dev --group verl-cpu
|
||||
- name: Run tests
|
||||
run: uv run --locked --no-sync pytest -v --durations=20 tests
|
||||
|
||||
package:
|
||||
name: Build package
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
- name: Build package
|
||||
run: uv build --no-sources
|
||||
- name: Verify package contents
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -e .[dev]
|
||||
- name: Run Black
|
||||
run: |
|
||||
black --check --diff --line-length=120 .
|
||||
python -m tarfile -l dist/*.tar.gz
|
||||
python -m zipfile -l dist/*.whl
|
||||
|
||||
docs:
|
||||
name: Build documentation
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6.1.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- uses: actions/setup-python@v4
|
||||
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
|
||||
with:
|
||||
python-version: '3.12'
|
||||
- name: Install documentation dependencies
|
||||
run: |
|
||||
./scripts/setup_stable.sh
|
||||
- name: Set source commit for docs
|
||||
run: |
|
||||
echo "SOURCE_COMMIT=${{ github.sha }}" >> $GITHUB_ENV
|
||||
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
enable-cache: true
|
||||
- name: Sync documentation dependencies
|
||||
run: uv sync --frozen --no-default-groups --group docs
|
||||
- name: Build documentation
|
||||
run: |
|
||||
mkdocs build --strict
|
||||
- name: Upload docs artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: documentation-site
|
||||
path: site/
|
||||
compression-level: 6
|
||||
|
||||
test:
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- python-version: '3.10'
|
||||
setup-script: 'stable'
|
||||
- python-version: '3.12'
|
||||
setup-script: 'latest'
|
||||
- python-version: '3.12'
|
||||
setup-script: 'stable'
|
||||
fail-fast: false
|
||||
|
||||
name: Test with Python ${{ matrix.python-version }} (${{ matrix.setup-script }})
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
./scripts/setup_${{ matrix.setup-script }}.sh
|
||||
- name: Freeze dependencies
|
||||
run: |
|
||||
pip list | tee requirements-freeze-${{ matrix.python-version }}-${{ matrix.setup-script }}.txt
|
||||
- name: Upload dependencies artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: dependencies-python-${{ matrix.python-version }}-${{ matrix.setup-script }}
|
||||
path: requirements-freeze-${{ matrix.python-version }}-${{ matrix.setup-script }}.txt
|
||||
compression-level: 0
|
||||
- name: Run tests
|
||||
run: |
|
||||
pytest -v tests
|
||||
env:
|
||||
PYTEST_ADDOPTS: "--color=yes"
|
||||
SOURCE_COMMIT: ${{ github.sha }}
|
||||
run: uv run --locked --no-sync mkdocs build --strict
|
||||
|
||||
@@ -1,15 +1,9 @@
|
||||
# Agentlightning specific files
|
||||
verl_old
|
||||
meta-llama/**
|
||||
debug/*.png
|
||||
requirements-freeze*.txt
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
|
||||
# C extensions
|
||||
# Distribution / packaging
|
||||
__pycache__/
|
||||
*.py[codz]
|
||||
*$py.class
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
@@ -33,11 +27,6 @@ share/python-wheels/
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
@@ -52,26 +41,29 @@ htmlcov/
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
*.py,cover
|
||||
*.py.cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
cover/
|
||||
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
# Django stuff
|
||||
*.log
|
||||
!examples/math-poc/reference_output.log
|
||||
!examples/math-poc/reference_output_vllm.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
db.sqlite3-journal
|
||||
|
||||
# Flask stuff:
|
||||
# Flask stuff
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
# Scrapy stuff
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
@@ -83,46 +75,54 @@ target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
.ipynb_checkpoints/
|
||||
*.ipynb_checkpoints
|
||||
|
||||
# IPython
|
||||
profile_default/
|
||||
ipython_config.py
|
||||
|
||||
# pyenv
|
||||
# For a library or package, you might want to ignore these files since the code is
|
||||
# intended to run in multiple environments; otherwise, check them in:
|
||||
# For a library or package, you might want to ignore these files since the code is
|
||||
# intended to run in multiple environments; otherwise, check them in:
|
||||
# .python-version
|
||||
|
||||
# pipenv
|
||||
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
||||
# install all needed dependencies.
|
||||
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||
# having no cross-platform support, pipenv may install dependencies that do not work, or not
|
||||
# install all needed dependencies.
|
||||
#Pipfile.lock
|
||||
|
||||
# UV
|
||||
# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
#uv.lock
|
||||
|
||||
# poetry
|
||||
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
||||
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
||||
#poetry.lock
|
||||
#poetry.toml
|
||||
|
||||
# pdm
|
||||
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
||||
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
||||
# pdm recommends including project-wide configuration in pdm.toml, but excluding .pdm-python.
|
||||
# https://pdm-project.org/en/latest/usage/project/#working-with-version-control
|
||||
#pdm.lock
|
||||
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
|
||||
# in version control.
|
||||
# https://pdm.fming.dev/latest/usage/project/#working-with-version-control
|
||||
.pdm.toml
|
||||
#pdm.toml
|
||||
.pdm-python
|
||||
.pdm-build/
|
||||
|
||||
# pixi
|
||||
# Similar to Pipfile.lock, it is generally recommended to include pixi.lock in version control.
|
||||
#pixi.lock
|
||||
# Pixi creates a virtual environment in the .pixi directory, just like venv module creates one.
|
||||
.pixi
|
||||
|
||||
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
|
||||
__pypackages__/
|
||||
|
||||
@@ -133,14 +133,23 @@ celerybeat.pid
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
# Environments and local settings
|
||||
.env
|
||||
**/.env
|
||||
.env.local
|
||||
*.env.local
|
||||
.envrc
|
||||
.local/
|
||||
.venv
|
||||
.venv/
|
||||
.venv.bak/
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
.claude/
|
||||
auto_docs/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
@@ -152,52 +161,87 @@ venv.bak/
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
# mypy / pyright / pyre / pytype
|
||||
.mypy_cache/
|
||||
.dmypy.json
|
||||
dmypy.json
|
||||
|
||||
# Pyre type checker
|
||||
.pyright/
|
||||
.pyre/
|
||||
|
||||
# pytype static type analyzer
|
||||
.pytype/
|
||||
|
||||
# Cython debug symbols
|
||||
cython_debug/
|
||||
|
||||
# MacOS
|
||||
.DS_Store
|
||||
|
||||
# PyCharm
|
||||
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
|
||||
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
||||
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
||||
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
||||
#.idea/
|
||||
|
||||
# Abstra
|
||||
# Abstra is an AI-powered process automation framework.
|
||||
# Ignore directories containing user credentials, local state, and settings.
|
||||
# Learn more at https://abstra.io/docs
|
||||
.abstra/
|
||||
|
||||
# Visual Studio Code
|
||||
# Visual Studio Code specific template is maintained in a separate VisualStudioCode.gitignore
|
||||
# that can be found at https://github.com/github/gitignore/blob/main/Global/VisualStudioCode.gitignore
|
||||
# and can be added to the global gitignore or merged into this file. However, if you prefer,
|
||||
# you could uncomment the following to ignore the enitre vscode folder
|
||||
.vscode/
|
||||
|
||||
# Ruff stuff:
|
||||
# Ruff stuff
|
||||
.ruff_cache/
|
||||
|
||||
# PyPI configuration file
|
||||
.pypirc
|
||||
|
||||
# Cursor
|
||||
# Cursor is an AI-powered code editor. `.cursorignore` specifies files/directories to
|
||||
# exclude from AI features like autocomplete and code analysis. Recommended for sensitive data
|
||||
# refer to https://docs.cursor.com/context/ignore-files
|
||||
# Cursor ignore files can contain local/sensitive context selection.
|
||||
.cursorignore
|
||||
.cursorindexingignore
|
||||
|
||||
# Marimo
|
||||
marimo/_static/
|
||||
marimo/_lsp/
|
||||
__marimo__/
|
||||
|
||||
# Runtime logs and generated outputs
|
||||
logs/
|
||||
examples/*/logs/
|
||||
.vscode/examples/math-poc/logs/
|
||||
artifacts/
|
||||
publicartifacts/
|
||||
checkpoints/
|
||||
wandb/
|
||||
runs/
|
||||
outputs/
|
||||
mlruns/
|
||||
*.ckpt
|
||||
*.pt
|
||||
*.pth
|
||||
*.bin
|
||||
*.safetensors
|
||||
|
||||
# Data not maintained in the repo
|
||||
examples/calc_x/data/
|
||||
!examples/calc_x/data/sample.jsonl
|
||||
examples/calc_x/logs/
|
||||
|
||||
# Site/build output
|
||||
site/
|
||||
public/
|
||||
|
||||
# Archives / packaged artifacts
|
||||
*.zip
|
||||
*.tar
|
||||
*.tar.gz
|
||||
*.tgz
|
||||
*.tar.bz2
|
||||
*.tar.xz
|
||||
*.7z
|
||||
|
||||
# Editor / OS
|
||||
.vscode/
|
||||
.idea/
|
||||
.DS_Store
|
||||
**/.DS_Store
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# Local example datasets
|
||||
examples/**/*dataset*.jsonl
|
||||
examples/**/subset*.jsonl
|
||||
examples/**/verified_*.jsonl
|
||||
|
||||
# SWE-smith rollout stats
|
||||
examples/swe_smith/rollout_stats.json
|
||||
|
||||
@@ -1,8 +1,16 @@
|
||||
exclude: ^(\.agents/|examples/llm-in-sandbox/vendor/)
|
||||
|
||||
repos:
|
||||
- repo: https://github.com/psf/black
|
||||
rev: 25.1.0
|
||||
- repo: https://github.com/pre-commit/pre-commit-hooks
|
||||
rev: v6.0.0
|
||||
hooks:
|
||||
- id: black
|
||||
pass_filenames: false
|
||||
always_run: true
|
||||
args: ["--line-length=120", "."]
|
||||
- id: end-of-file-fixer
|
||||
- id: trailing-whitespace
|
||||
- id: check-yaml
|
||||
exclude: ^(mkdocs\.yml|examples/calc_x/job-template\.yaml|examples/llm-in-sandbox/job-template\.yaml|examples/swe_smith/job-template-openai\.yaml)$
|
||||
- id: check-toml
|
||||
- id: check-added-large-files
|
||||
args: ["--maxkb=1024"]
|
||||
exclude: ^uv\.lock$
|
||||
- id: check-shebang-scripts-are-executable
|
||||
- id: detect-private-key
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
3.12
|
||||
@@ -16,4 +16,4 @@ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
|
||||
THE SOFTWARE.
|
||||
THE SOFTWARE.
|
||||
|
||||
@@ -1,158 +1,130 @@
|
||||

|
||||
<p align="center">
|
||||
<img src="docs/images/agl-v1.0.svg" alt="Agent Lightning v1.0" width="500">
|
||||
</p>
|
||||
|
||||
# Agent Lightning⚡
|
||||
<p align="center"><em>3,500-Line Lightweight Agentic RL Framework for Training Agents with Real Harnesses!</em></p>
|
||||
|
||||
[](https://github.com/microsoft/agent-lightning/actions/workflows/tests.yml)
|
||||
[](https://github.com/microsoft/agent-lightning/actions/workflows/examples.yml)
|
||||
[](https://badge.fury.io/py/agentlightning)
|
||||
[](LICENSE)
|
||||
<p align="center">
|
||||
<a href="https://microsoft.github.io/agent-lightning/stable/">Documentation</a> · <a href="https://arxiv.org/pdf/2608.17528">Technical Report</a> · <a href="LICENSE">MIT License</a>
|
||||
</p>
|
||||
|
||||
**The absolute trainer to light up AI agents.**
|
||||
> Agent Lightning was completely refactored in v1.0. For legacy releases earlier than v1.0, see [this branch](https://github.com/microsoft/agent-lightning/tree/v0.x).
|
||||
|
||||
## ⚡ Core Features
|
||||
## ⚡ Key Features
|
||||
|
||||
- Turn your agent into an optimizable beast with **ZERO CODE CHANGE** (almost)! 💤
|
||||
- Build with **ANY** agent framework (LangChain, OpenAI Agent SDK, AutoGen, CrewAI, ...); or even WITHOUT agent framework (Python OpenAI). You name it! 🤖
|
||||
- **Selectively** optimize one or more agents in a multi-agent system. 🎯
|
||||
- Embraces Reinforcement Learning, Automatic Prompt Optimization and more **algorithms**. 🤗
|
||||
- 🪶 **~3,500 lines of code:** We treat simplicity as the first principle.
|
||||
- 🧩 **Train with real agent harnesses:** Agents interact with the model through the Agent Lightning v1.0 proxy with **ZERO changes**, while keeping tools, context, control flow, and environments in the loop.
|
||||
- ☸️ **Native Kubernetes support:** Run agents directly as Kubernetes Jobs without relying on external sandbox services.
|
||||
- 💻 **Full coding agent training example:** Using only **6K training samples**, an end-to-end Qwen3.5-9B workflow improves SWE-bench Verified from **41.8% to 56.4%**, a gain of **14.6 percentage points**. We release the full pipeline, including data cleaning, reward-hacking prevention, and training scripts.
|
||||
|
||||

|
||||
## ⚡ Installation
|
||||
|
||||
## ⚡ Resources
|
||||
The following is an example installation on a CUDA 13.0 machine:
|
||||
|
||||
```bash
|
||||
cd <this-repo>
|
||||
uv sync
|
||||
bash scripts/setup_verl.sh 0.8.0 cu130
|
||||
```
|
||||
|
||||
See the [Installation Guide](https://microsoft.github.io/agent-lightning/stable/00-installation/) for details.
|
||||
|
||||
|
||||
## ⚡ Architecture
|
||||
|
||||
<p align="center">
|
||||
<img src="docs/images/architecture.jpg" alt="Agent Lightning v1.0 architecture" width="800">
|
||||
</p>
|
||||
|
||||
Agent Lightning v1.0 keeps the training architecture simple with three lightweight components:
|
||||
|
||||
- **Trainer:** Runs `verl` and vLLM, builds training samples, and updates the policy.
|
||||
- **API Gateway:** Proxies model requests and captures training data.
|
||||
- **Rollout Controller:** Runs agents locally or as Kubernetes Jobs.
|
||||
|
||||
The Trainer creates rollouts, the Controller launches agents, and the Gateway turns interactions into training data, while agents continue to run with their real harnesses.
|
||||
|
||||
## ⚡ Results
|
||||
|
||||
We evaluate Agent Lightning v1.0 across several practical training domains, including Search R1, LLM-in-Sandbox, and Coding Agent. Pure RL delivers substantial improvements across all three domains, as shown below.
|
||||
|
||||
<p align="center">
|
||||
<img src="docs/images/benchmark-comparison.jpg" alt="Agent Lightning v1.0 benchmark comparison" width="600">
|
||||
</p>
|
||||
|
||||
## ⚡ Documentation
|
||||
|
||||
| Section | Content |
|
||||
|---------|---------|
|
||||
| [Installation](https://microsoft.github.io/agent-lightning/stable/00-installation/) | Base environment and `verl` GPU stack |
|
||||
| [Quick Start](https://microsoft.github.io/agent-lightning/stable/01-quick-start/) | Local first run and end-to-end flow |
|
||||
| [Basics](https://microsoft.github.io/agent-lightning/stable/05-basics/) | Components, rollouts, events, and trajectories |
|
||||
| [Trainer Configuration](https://microsoft.github.io/agent-lightning/stable/20-trainer-configuration/) | `verl` integration and trace aggregation |
|
||||
| [API Gateway Configuration](https://microsoft.github.io/agent-lightning/stable/25-api-gateway-configuration/) | Gateway and model proxy settings |
|
||||
| [Controller Configuration](https://microsoft.github.io/agent-lightning/stable/30-controller-configuration/) | Local and Kubernetes runners |
|
||||
| [Asynchronous Training](https://microsoft.github.io/agent-lightning/stable/35-asynchronous-training/) | Collocated async collection and pause/drain |
|
||||
|
||||
## ⚡ Examples
|
||||
|
||||
| Example | Description |
|
||||
|---|---|
|
||||
| [Calc-X](https://microsoft.github.io/agent-lightning/stable/50-example-calc-x/) | POC math reasoning example with AutoGen and MCP calculator tools, requiring only one GPU. |
|
||||
| [GSM8K](https://microsoft.github.io/agent-lightning/stable/55-example-gsm8k/) | POC grade-school math reasoning example. |
|
||||
| [ScienceWorld](https://microsoft.github.io/agent-lightning/stable/60-example-science-world/) | Interactive science tasks in a text-based environment. |
|
||||
| [Search-R1](https://microsoft.github.io/agent-lightning/stable/65-example-search-r1/) | Multi-turn retrieval and reasoning agent. |
|
||||
| [LLM-in-Sandbox](https://microsoft.github.io/agent-lightning/stable/70-example-llm-in-sandbox/) | General agent with computer and code execution tools. |
|
||||
| [Coding Agent](https://microsoft.github.io/agent-lightning/stable/75-example-coding-agent/) | Coding agent trained with repository tests. |
|
||||
|
||||
## ⚡ Articles
|
||||
|
||||
- 8/19/2026 [Agent Lightning v1.0: Towards Harnessed Agentic RL](https://arxiv.org/abs/2608.17528) technical report.
|
||||
- 12/17/2025 [Adopting the Trajectory Level Aggregation for Faster Training](https://agent-lightning.github.io/posts/trajectory_level_aggregation/) Agent-lightning blog.
|
||||
- 11/4/2025 [Tuning ANY AI agent with Tinker ✕ Agent-lightning](https://medium.com/@yugez/tuning-any-ai-agent-with-tinker-agent-lightning-part-1-1d8c9a397f0e) Medium. See also [Part 2](https://medium.com/@yugez/tuning-any-ai-agent-with-tinker-agent-lightning-part-2-332c5437f0dc).
|
||||
- 10/22/2025 [No More Retokenization Drift: Returning Token IDs via the OpenAI Compatible API Matters in Agent RL](https://blog.vllm.ai/2025/10/22/agent-lightning.html) vLLM blog. See also [Zhihu writeup](https://zhuanlan.zhihu.com/p/1965067274642785725).
|
||||
- 8/11/2025 [Training AI Agents to Write and Self-correct SQL with Reinforcement Learning](https://medium.com/@yugez/training-ai-agents-to-write-and-self-correct-sql-with-reinforcement-learning-571ed31281ad) Medium.
|
||||
- 8/5/2025 [Agent Lightning: Train ANY AI Agents with Reinforcement Learning](https://arxiv.org/abs/2508.03680) arXiv paper.
|
||||
- 7/26/2025 [We discovered an approach to train any AI agent with RL, with (almost) zero code changes.](https://www.reddit.com/r/LocalLLaMA/comments/1m9m670/we_discovered_an_approach_to_train_any_ai_agent/) Reddit.
|
||||
- 6/6/2025 [Agent Lightning - Microsoft Research](https://www.microsoft.com/en-us/research/project/agent-lightning/) Project page.
|
||||
|
||||
## ⚡ Installation
|
||||
## ⚡ Community Projects
|
||||
|
||||
First, let's get your environment set up. We'll be using `/path/to/agentlightning` to refer to the directory containing this README file.
|
||||
|
||||
### 1. Set Up Your Environment
|
||||
|
||||
We strongly recommend creating a new virtual environment to avoid conflicts with other packages. You can use either `conda` or `venv`. **Python 3.10 or later** is recommended.
|
||||
|
||||
### 2. Install Core Training Dependencies (Optional)
|
||||
|
||||
If you are running RL with Agent-Lightning, the next step is to install the essential packages: `PyTorch`, `FlashAttention`, `vLLM` and `VERL`. The following versions and installation order have been tested and are confirmed to work.
|
||||
|
||||
```bash
|
||||
pip install torch==2.7.0 torchvision==0.22.0 torchaudio==2.7.0 --index-url https://download.pytorch.org/whl/cu128
|
||||
pip install flash-attn --no-build-isolation
|
||||
pip install vllm==0.9.2
|
||||
pip install verl==0.5.0
|
||||
```
|
||||
|
||||
See `scripts/setup_stable_gpu.sh` for a full installation script.
|
||||
|
||||
### 3. Install Agent Lightning
|
||||
|
||||
Now, you're ready to install Agent Lightning itself.
|
||||
|
||||
```bash
|
||||
pip install agentlightning
|
||||
```
|
||||
|
||||
### 4. Install Agent Frameworks (Optional)
|
||||
|
||||
If you plan to use other agent frameworks, you can install them with the following commands. If you don't need these, feel free to skip this step.
|
||||
We recommend doing this as the final step to avoid dependency versions being overwritten by mistake.
|
||||
|
||||
```bash
|
||||
# AutoGen (Recommended to install first)
|
||||
pip install "autogen-agentchat" "autogen-ext[openai]"
|
||||
|
||||
# LiteLLM
|
||||
pip install "litellm[proxy]"
|
||||
|
||||
# MCP
|
||||
pip install mcp
|
||||
|
||||
# UV
|
||||
pip install uv
|
||||
|
||||
# OpenAI Agents
|
||||
pip install openai-agents
|
||||
|
||||
# LangChain
|
||||
pip install langgraph "langchain[openai]" langchain-community langchain-text-splitters
|
||||
|
||||
# SQL-related dependencies
|
||||
pip install sqlparse nltk
|
||||
```
|
||||
|
||||
Don't worry if dependency conflicts arise during this step. Follow the installation order above and the conflicts generally do not matter.
|
||||
|
||||
## ⚡ Examples
|
||||
|
||||
For more detailed examples, please see the `examples` folder:
|
||||
|
||||
1. [calc_x](examples/calc_x): An agent built with AutoGen with calculator tool use, trained on Calc-X dataset with Reinforcement Learning.
|
||||
2. [spider](examples/spider): A write-check-rewrite looped agent with LangGraph with SQL execution; selectively optimize write and rewrite on Spider dataset with Reinforcement Learning.
|
||||
3. [apo](examples/apo): An example to customize an optimization algorithm: Automatic Prompt Optimization.
|
||||
|
||||
## ⚡ Important Caveats
|
||||
|
||||
1. **AgentOps Integration**: Agent Lightning uses [AgentOps](https://github.com/AgentOps-AI/agentops) for agent tracking by default. If you're already using AgentOps in your own code, you'll need to disable our managed AgentOps client by modifying the `tracer` parameter of trainer.
|
||||
2. **Debugging Traces**: If you encounter issues with tracing, you can visualize the trace tree using `tracer.last_trace().visualize("tree_graph")`. Please note that this API is experimental and may change in future releases.
|
||||
3. **Launching the Server and Agents**: Currently, the training server and agent clients must be launched in separate processes. You can open two terminal windows or run one of them in the background. The launching order generally doesn't matter.
|
||||
4. **Environment Variables**: The environment variables and working directory at the time of `ray init` are important. If you run into "file not found" errors, try restarting Ray from your current working directory.
|
||||
5. **Handling Timeouts**: The training server may hang if samples fail or time out on the agent side. To prevent this, we recommend setting limits on the prompt and response lengths, as this is the most common cause of failures.
|
||||
6. **VERL Failures**: Save checkpoints frequently, as VERL with vLLM may sometimes experience out-of-memory issues. If you encounter a VERL failure, you can resume training from the last checkpoint.
|
||||
|
||||
## ⚡ Architecture
|
||||
|
||||
Currently, Agent Lightning is built around a **training server** and one or multiple **agents**.
|
||||
|
||||
* The **server** manages the training data, prepares samples for the agents, and provides the LLM endpoint.
|
||||
* **Agents** retrieve samples from the server, process them (which may involve interacting with the LLM), and send the results back. These results, or "trajectories," are lists of prompts and responses from the LLM.
|
||||
* The **server** then collects these trajectories and computes the losses to optimize the language models.
|
||||
|
||||

|
||||
|
||||
## ⚡ Development Instructions
|
||||
|
||||
Install with development dependencies:
|
||||
|
||||
```
|
||||
git clone https://github.com/microsoft/agent-lightning
|
||||
cd agent-lightning
|
||||
pip install -e .[dev]
|
||||
```
|
||||
|
||||
Please run pre-commit hooks before checking in code:
|
||||
|
||||
```
|
||||
pre-commit install
|
||||
pre-commit run --all-files --show-diff-on-failure --color=always
|
||||
```
|
||||
|
||||
Serve documentation locally:
|
||||
|
||||
```bash
|
||||
mkdocs serve
|
||||
```
|
||||
- [DeepWerewolf](https://github.com/af-74413592/DeepWerewolf) — A case study of agent RL training for the Chinese Werewolf game built with AgentScope and Agent Lightning.
|
||||
- [AgentFlow](https://agentflow.stanford.edu/) — A modular multi-agent framework that combines planner, executor, verifier, and generator agents with the Flow-GRPO algorithm to tackle long-horizon, sparse-reward tasks.
|
||||
- [Youtu-Agent](https://github.com/TencentCloudADP/Youtu-agent) — Youtu-Agent lets you build and train your agent with ease. Built with [a modified branch](https://github.com/microsoft/agent-lightning/tree/contrib/youtu-agent-lightning) of Agent Lightning, Youtu-Agent has verified up to 128 GPUs RL training on maths/code and search capabilities with steady convergence. Also check [the recipe](https://github.com/TencentCloudADP/youtu-agent/tree/rl/agl) and their blog [*Stop Wrestling with Your Agent RL: How Youtu-Agent Achieved Stable, 128-GPU Scaling Without Breaking a Sweat*](https://spotted-coconut-df8.notion.site/Stop-Wrestling-with-Your-Agent-RL-How-Youtu-Agent-Achieved-Stable-128-GPU-Scaling-Without-Breaking-2ca5e8f089ba80539a98c582b65e0233).
|
||||
|
||||
## ⚡ Citation
|
||||
|
||||
If you find Agent Lightning useful in your research or projects, please cite our paper:
|
||||
If you use Agent Lightning v1.0 in your research or projects, please cite the technical report:
|
||||
|
||||
```bibtex
|
||||
@misc{he2026agentlightningv10harnessed,
|
||||
title={Agent Lightning v1.0: Towards Harnessed Agentic RL},
|
||||
author={Zhiyuan He and Siwei Zhang and Zhiwen Zhou and Yuqing Yang and Yu Kang and Yuge Zhang and Luna K. Qiu and Tin Yan Tsui and Jiahang Xu and Chong Luo},
|
||||
year={2026},
|
||||
eprint={2608.17528},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.AI},
|
||||
url={https://arxiv.org/abs/2608.17528},
|
||||
}
|
||||
```
|
||||
|
||||
For the original Agent Lightning paper, please use:
|
||||
|
||||
```bibtex
|
||||
@misc{luo2025agentlightningtrainai,
|
||||
title={Agent Lightning: Train ANY AI Agents with Reinforcement Learning},
|
||||
title={Agent Lightning: Train ANY AI Agents with Reinforcement Learning},
|
||||
author={Xufang Luo and Yuge Zhang and Zhiyuan He and Zilong Wang and Siyun Zhao and Dongsheng Li and Luna K. Qiu and Yuqing Yang},
|
||||
year={2025},
|
||||
eprint={2508.03680},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.AI},
|
||||
url={https://arxiv.org/abs/2508.03680},
|
||||
url={https://arxiv.org/abs/2508.03680},
|
||||
}
|
||||
```
|
||||
|
||||
## ⚡ Contributing
|
||||
|
||||
This project welcomes contributions and suggestions. Most contributions require you to agree to a Contributor License Agreement (CLA) declaring that you have the right to, and actually do, grant us the rights to use your contribution. For details, visit https://cla.opensource.microsoft.com.
|
||||
This project welcomes contributions and suggestions. Start by reading the [Contributing Guide](docs/community/contributing.md) for recommended contribution points, environment setup, branching conventions, and pull request expectations. Most contributions require you to agree to a Contributor License Agreement (CLA) declaring that you have the right to, and actually do, grant us the rights to use your contribution. For details, visit https://cla.opensource.microsoft.com.
|
||||
|
||||
When you submit a pull request, a CLA bot will automatically determine whether you need to provide a CLA and decorate the PR appropriately (e.g., status check, comment). Simply follow the instructions provided by the bot. You will only need to do this once across all repos using our CLA.
|
||||
|
||||
@@ -166,6 +138,7 @@ This project may contain trademarks or logos for projects, products, or services
|
||||
|
||||
This project has been evaluated and certified to comply with the Microsoft Responsible AI Standard. The team will continue to monitor and maintain the repository, addressing any severe issues, including potential harms, if they arise.
|
||||
|
||||
|
||||
## ⚡ License
|
||||
|
||||
This project is licensed under the MIT License. See the [LICENSE](LICENSE) file for details.
|
||||
Agent Lightning v1.0 is released under the [MIT License](LICENSE).
|
||||
|
||||
@@ -11,4 +11,4 @@ For security reporting information, locations, contact information, and policies
|
||||
please review the latest guidance for Microsoft repositories at
|
||||
[https://aka.ms/SECURITY.md](https://aka.ms/SECURITY.md).
|
||||
|
||||
<!-- END MICROSOFT SECURITY.MD BLOCK -->
|
||||
<!-- END MICROSOFT SECURITY.MD BLOCK -->
|
||||
|
||||
@@ -1,10 +1,5 @@
|
||||
__version__ = "0.1.2"
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from .client import AgentLightningClient, DevTaskLoader
|
||||
from .config import lightning_cli
|
||||
from .litagent import LitAgent
|
||||
from .logging import configure_logger
|
||||
from .reward import reward
|
||||
from .server import AgentLightningServer
|
||||
from .trainer import Trainer
|
||||
from .types import *
|
||||
"""Agent Lightning."""
|
||||
|
||||
__version__ = "1.0.1"
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
import time
|
||||
from agentlightning.instrumentation.agentops import AgentOpsServerManager
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Start AgentOps server")
|
||||
parser.add_argument("--daemon", action="store_true", help="Run server as a daemon")
|
||||
parser.add_argument("--port", type=int, default=8002, help="Port to run the server on")
|
||||
args = parser.parse_args()
|
||||
|
||||
manager = AgentOpsServerManager(daemon=args.daemon, port=args.port)
|
||||
try:
|
||||
manager.start()
|
||||
# Wait forever
|
||||
while True:
|
||||
time.sleep(1)
|
||||
except KeyboardInterrupt:
|
||||
manager.stop()
|
||||
@@ -1,10 +0,0 @@
|
||||
from typing import List
|
||||
|
||||
from vllm.entrypoints.cli.main import main
|
||||
|
||||
from agentlightning.instrumentation.vllm import instrument_vllm
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
instrument_vllm()
|
||||
main()
|
||||
@@ -1,365 +1,82 @@
|
||||
import asyncio
|
||||
import logging
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Thin httpx clients for Agent Lightning."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import urllib.parse
|
||||
from typing import Any, Dict, Optional, List, Union
|
||||
from typing import Any
|
||||
|
||||
import aiohttp
|
||||
import requests
|
||||
|
||||
from .types import Rollout, Task, TaskInput, TaskIfAny, ResourcesUpdate, NamedResources
|
||||
import httpx
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
def _headers_with_key(headers: httpx.Headers | dict[str, str] | None, key: str | None) -> dict[str, str]:
|
||||
merged = dict(headers or {})
|
||||
if key:
|
||||
merged["Authorization"] = f"Bearer {key}"
|
||||
return merged
|
||||
|
||||
|
||||
class AgentLightningClient:
|
||||
"""
|
||||
Client for interacting with a version-aware Agent Lightning Server.
|
||||
|
||||
This client handles polling for tasks, fetching specific versions of resources
|
||||
(like model configurations), and posting completed rollouts back to the server.
|
||||
It provides both synchronous and asynchronous methods for these operations and
|
||||
includes a cache for resources.
|
||||
"""
|
||||
|
||||
_next_task_uri = "/task"
|
||||
_resources_uri = "/resources"
|
||||
_latest_resources_uri = "/resources/latest"
|
||||
_report_rollout_uri = "/rollout"
|
||||
|
||||
def __init__(self, endpoint: str, poll_interval: float = 5.0, timeout: float = 10.0):
|
||||
"""Initializes the AgentLightningClient.
|
||||
|
||||
Args:
|
||||
endpoint: The root URL of the Agent Lightning server.
|
||||
poll_interval: The interval in seconds to wait between polling for new tasks.
|
||||
timeout: The timeout in seconds for HTTP requests.
|
||||
"""
|
||||
self.endpoint = endpoint
|
||||
self.task_count = 0
|
||||
self.poll_interval = poll_interval
|
||||
self.timeout = timeout
|
||||
self._resource_cache: Dict[str, ResourcesUpdate] = {} # TODO: mechanism to evict cache
|
||||
self._default_headers = {"X-AgentLightning-Client": "true"}
|
||||
|
||||
async def _request_json_async(self, url: str) -> Optional[Dict[str, Any]]:
|
||||
"""Makes an async GET request to the specified URL and returns the JSON response.
|
||||
|
||||
Args:
|
||||
url: The URL to request.
|
||||
|
||||
Returns:
|
||||
The JSON response as a dictionary or None if the request fails.
|
||||
"""
|
||||
timeout = aiohttp.ClientTimeout(total=self.timeout)
|
||||
async with aiohttp.ClientSession(timeout=timeout) as session:
|
||||
try:
|
||||
async with session.get(url, headers=self._default_headers) as resp:
|
||||
resp.raise_for_status()
|
||||
return await resp.json()
|
||||
except Exception as e:
|
||||
logger.debug(f"Async GET request failed for {url}: {e}")
|
||||
return None
|
||||
|
||||
async def _post_json_async(self, url: str, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""Makes an async POST request with a JSON payload.
|
||||
|
||||
Args:
|
||||
url: The URL to post to.
|
||||
payload: The dictionary data to send as JSON.
|
||||
|
||||
Returns:
|
||||
The JSON response as a dictionary or None if the request fails.
|
||||
"""
|
||||
timeout = aiohttp.ClientTimeout(total=self.timeout)
|
||||
async with aiohttp.ClientSession(timeout=timeout) as session:
|
||||
try:
|
||||
async with session.post(url, json=payload, headers=self._default_headers) as resp:
|
||||
resp.raise_for_status()
|
||||
return await resp.json()
|
||||
except Exception as e:
|
||||
logger.debug(f"Async POST request failed for {url}: {e}")
|
||||
return None
|
||||
|
||||
async def poll_next_task_async(self) -> Task:
|
||||
"""Polls the server asynchronously for the next task until one is available.
|
||||
|
||||
Returns:
|
||||
A Task object containing the task details.
|
||||
"""
|
||||
url = urllib.parse.urljoin(self.endpoint, self._next_task_uri)
|
||||
while True:
|
||||
response = await self._request_json_async(url)
|
||||
if response:
|
||||
task_if_any = TaskIfAny.model_validate(response)
|
||||
if task_if_any.is_available and task_if_any.task:
|
||||
self.task_count += 1
|
||||
logger.info(f"[Task {self.task_count} Received] ID: {task_if_any.task.rollout_id}")
|
||||
return task_if_any.task
|
||||
logger.debug(f"No task available yet. Retrying in {self.poll_interval} seconds...")
|
||||
await asyncio.sleep(self.poll_interval)
|
||||
|
||||
async def get_resources_by_id_async(self, resource_id: str) -> Optional[ResourcesUpdate]:
|
||||
"""Fetches a specific version of resources by its ID, using a cache.
|
||||
|
||||
Args:
|
||||
resource_id: The ID of the resources to fetch, usually from a Task's metadata.
|
||||
|
||||
Returns:
|
||||
A ResourcesUpdate object containing the versioned resources, or None if not found.
|
||||
"""
|
||||
if resource_id in self._resource_cache:
|
||||
logger.debug(f"Found resources '{resource_id}' in cache.")
|
||||
return self._resource_cache[resource_id]
|
||||
|
||||
url = urllib.parse.urljoin(self.endpoint, f"{self._resources_uri}/{resource_id}")
|
||||
response = await self._request_json_async(url)
|
||||
if response:
|
||||
resources_update = ResourcesUpdate.model_validate(response)
|
||||
self._resource_cache[resource_id] = resources_update
|
||||
logger.info(f"Fetched and cached resources for ID: {resource_id}")
|
||||
return resources_update
|
||||
return None
|
||||
|
||||
async def get_latest_resources_async(self) -> Optional[ResourcesUpdate]:
|
||||
"""Fetches the latest available resources from the server.
|
||||
|
||||
Returns:
|
||||
A ResourcesUpdate object containing the latest resources.
|
||||
"""
|
||||
url = urllib.parse.urljoin(self.endpoint, self._latest_resources_uri)
|
||||
response = await self._request_json_async(url)
|
||||
if response:
|
||||
resources_update = ResourcesUpdate.model_validate(response)
|
||||
# Cache this result as well
|
||||
self._resource_cache[resources_update.resources_id] = resources_update
|
||||
return resources_update
|
||||
return None
|
||||
|
||||
async def post_rollout_async(self, rollout: Rollout) -> Optional[Dict[str, Any]]:
|
||||
"""Posts a completed rollout to the server asynchronously.
|
||||
|
||||
Args:
|
||||
rollout: A Rollout object containing the results of a task.
|
||||
|
||||
Returns:
|
||||
The server's JSON response as a dictionary.
|
||||
"""
|
||||
url = urllib.parse.urljoin(self.endpoint, self._report_rollout_uri)
|
||||
payload = rollout.model_dump(mode="json")
|
||||
return await self._post_json_async(url, payload)
|
||||
|
||||
def _request_json(self, url: str) -> Optional[Dict[str, Any]]:
|
||||
"""Makes a sync GET request to the specified URL and returns the JSON response.
|
||||
|
||||
Args:
|
||||
url: The URL to request.
|
||||
|
||||
Returns:
|
||||
The JSON response as a dictionary or None if the request fails.
|
||||
"""
|
||||
try:
|
||||
response = requests.get(url, timeout=self.timeout, headers=self._default_headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
logger.debug(f"Sync GET request failed for {url}: {e}")
|
||||
return None
|
||||
|
||||
def _post_json(self, url: str, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""Makes a sync POST request with a JSON payload.
|
||||
|
||||
Args:
|
||||
url: The URL to post to.
|
||||
payload: The dictionary data to send as JSON.
|
||||
|
||||
Returns:
|
||||
The JSON response as a dictionary or None if the request fails.
|
||||
"""
|
||||
try:
|
||||
response = requests.post(url, json=payload, timeout=self.timeout, headers=self._default_headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
logger.debug(f"Sync POST request failed for {url}: {e}")
|
||||
return None
|
||||
|
||||
def poll_next_task(self) -> Task:
|
||||
"""Polls the server synchronously for the next task until one is available.
|
||||
|
||||
Returns:
|
||||
A Task object containing the task details, including the required `resources_id`.
|
||||
"""
|
||||
url = urllib.parse.urljoin(self.endpoint, self._next_task_uri)
|
||||
while True:
|
||||
response = self._request_json(url)
|
||||
if response:
|
||||
task_if_any = TaskIfAny.model_validate(response)
|
||||
if task_if_any.is_available and task_if_any.task:
|
||||
self.task_count += 1
|
||||
logger.info(f"[Task {self.task_count} Received] ID: {task_if_any.task.rollout_id}")
|
||||
return task_if_any.task
|
||||
logger.debug(f"No task available yet. Retrying in {self.poll_interval} seconds...")
|
||||
time.sleep(self.poll_interval)
|
||||
|
||||
def get_resources_by_id(self, resource_id: str) -> Optional[ResourcesUpdate]:
|
||||
"""Fetches a specific version of resources by its ID synchronously, using a cache.
|
||||
|
||||
Args:
|
||||
resource_id: The ID of the resources to fetch, usually from a Task's metadata.
|
||||
|
||||
Returns:
|
||||
A ResourcesUpdate object containing the versioned resources, or None if not found.
|
||||
"""
|
||||
if resource_id in self._resource_cache:
|
||||
logger.debug(f"Found resources '{resource_id}' in cache.")
|
||||
return self._resource_cache[resource_id]
|
||||
|
||||
url = urllib.parse.urljoin(self.endpoint, f"{self._resources_uri}/{resource_id}")
|
||||
response = self._request_json(url)
|
||||
if response:
|
||||
resources_update = ResourcesUpdate.model_validate(response)
|
||||
self._resource_cache[resource_id] = resources_update
|
||||
logger.info(f"Fetched and cached resources for ID: {resource_id}")
|
||||
return resources_update
|
||||
return None
|
||||
|
||||
def get_latest_resources(self) -> Optional[ResourcesUpdate]:
|
||||
"""Fetches the latest available resources from the server synchronously.
|
||||
|
||||
Returns:
|
||||
A ResourcesUpdate object containing the latest resources.
|
||||
"""
|
||||
url = urllib.parse.urljoin(self.endpoint, self._latest_resources_uri)
|
||||
response = self._request_json(url)
|
||||
if response:
|
||||
resources_update = ResourcesUpdate.model_validate(response)
|
||||
self._resource_cache[resources_update.resources_id] = resources_update
|
||||
return resources_update
|
||||
return None
|
||||
|
||||
def post_rollout(self, rollout: Rollout) -> Optional[Dict[str, Any]]:
|
||||
"""Posts a completed rollout to the server synchronously.
|
||||
|
||||
Args:
|
||||
rollout: A Rollout object containing the results of a task.
|
||||
|
||||
Returns:
|
||||
The server's JSON response as a dictionary.
|
||||
"""
|
||||
url = urllib.parse.urljoin(self.endpoint, self._report_rollout_uri)
|
||||
payload = rollout.model_dump(mode="json")
|
||||
return self._post_json(url, payload)
|
||||
|
||||
|
||||
class DevTaskLoader(AgentLightningClient):
|
||||
"""A local task manager for development that provides sample tasks and resources.
|
||||
|
||||
This client mocks the server APIs by maintaining a local queue of tasks and resources
|
||||
within the same process. It's designed for development, testing, and scenarios where
|
||||
a full Agent Lightning server is not needed.
|
||||
|
||||
The DevTaskLoader overrides the polling and resource fetching methods to return data
|
||||
from local collections instead of making HTTP requests to a remote server.
|
||||
"""
|
||||
class AgentLightningAsyncClient(httpx.AsyncClient):
|
||||
"""Async httpx client with optional bearer key."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tasks: Union[List[TaskInput], List[Task]],
|
||||
resources: Union[NamedResources, ResourcesUpdate],
|
||||
*,
|
||||
key: str | None = None,
|
||||
headers: httpx.Headers | dict[str, str] | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Initializes the DevTaskLoader with pre-defined tasks and resources.
|
||||
) -> None:
|
||||
super().__init__(
|
||||
headers=_headers_with_key(headers, key),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
Args:
|
||||
tasks: Either a List of TaskInput objects or a List of Task objects.
|
||||
resources: Either NamedResources or ResourcesUpdate object.
|
||||
**kwargs: Additional arguments passed to the parent AgentLightningClient.
|
||||
|
||||
class AgentLightningSyncClient(httpx.Client):
|
||||
"""Sync httpx client with optional bearer key."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
key: str | None = None,
|
||||
headers: httpx.Headers | dict[str, str] | None = None,
|
||||
max_retries: int = 10,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
self.max_retries = max_retries
|
||||
super().__init__(
|
||||
headers=_headers_with_key(headers, key),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def get(self, *args: Any, **kwargs: Any) -> httpx.Response: # type: ignore[override]
|
||||
last_exc: Exception | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
return super().get(*args, **kwargs)
|
||||
except Exception as exc:
|
||||
last_exc = exc
|
||||
print(f"GET failed (attempt {attempt + 1}/{self.max_retries + 1}): {exc}")
|
||||
assert last_exc is not None
|
||||
raise last_exc
|
||||
|
||||
def post_with_retry(self, *args: Any, **kwargs: Any) -> httpx.Response:
|
||||
"""POST with retry + backoff, raising on non-2xx. Only for idempotent endpoints.
|
||||
|
||||
Retries both transport errors and error status codes, so a transient 5xx
|
||||
is retried too. Callers get an already status-checked response back.
|
||||
"""
|
||||
super().__init__(endpoint="local://", **kwargs)
|
||||
self._tasks = tasks.copy()
|
||||
if len(self._tasks) == 0:
|
||||
raise ValueError("DevTaskLoader requires at least one task to be provided.")
|
||||
|
||||
# Check if tasks are mixture of TaskInput and Task
|
||||
if any(isinstance(task, Task) for task in self._tasks):
|
||||
if not all(isinstance(task, Task) for task in self._tasks):
|
||||
raise ValueError("All tasks must be either Task or TaskInput objects.")
|
||||
|
||||
self._task_index = 0
|
||||
|
||||
if isinstance(resources, ResourcesUpdate):
|
||||
self._resources_update = resources
|
||||
else:
|
||||
self._resources_update = ResourcesUpdate(resources_id="local", resources=resources)
|
||||
|
||||
# Store rollouts posted back to the loader for easy debugging of local runs
|
||||
self._rollouts: List[Rollout] = []
|
||||
|
||||
@property
|
||||
def rollouts(self) -> List[Rollout]:
|
||||
"""Return rollouts that have been posted back to the loader."""
|
||||
return self._rollouts
|
||||
|
||||
def poll_next_task(self) -> Task:
|
||||
"""Returns the next task from the local queue.
|
||||
|
||||
If tasks are TaskInput objects, assembles them into Task objects.
|
||||
If tasks are already Task objects, returns them directly.
|
||||
|
||||
Returns:
|
||||
The next Task object from the local task list.
|
||||
"""
|
||||
if self._task_index >= len(self._tasks):
|
||||
self._task_index = 0
|
||||
|
||||
task_or_input = self._tasks[self._task_index]
|
||||
|
||||
if isinstance(task_or_input, Task):
|
||||
task = task_or_input
|
||||
else:
|
||||
rollout_id = f"local_task_{self._task_index + 1:03d}"
|
||||
task = Task(
|
||||
rollout_id=rollout_id,
|
||||
input=task_or_input,
|
||||
resources_id=self._resources_update.resources_id,
|
||||
create_time=time.time(),
|
||||
)
|
||||
|
||||
self._task_index += 1
|
||||
self.task_count += 1
|
||||
logger.info(f"[Task {self.task_count} Received] Task ID: {task.rollout_id}")
|
||||
return task
|
||||
|
||||
def get_resources_by_id(self, resource_id: str) -> Optional[ResourcesUpdate]:
|
||||
logger.debug(f"DevTaskLoader checking resources for ID: {resource_id}")
|
||||
if resource_id != self._resources_update.resources_id:
|
||||
raise ValueError(
|
||||
f"Resource ID '{resource_id}' not found. Only '{self._resources_update.resources_id}' is available."
|
||||
)
|
||||
return self._resources_update
|
||||
|
||||
def get_latest_resources(self) -> Optional[ResourcesUpdate]:
|
||||
logger.debug("DevTaskLoader returning latest resources.")
|
||||
return self._resources_update
|
||||
|
||||
def post_rollout(self, rollout: Rollout) -> Optional[Dict[str, Any]]:
|
||||
logger.debug(f"DevTaskLoader received rollout for task: {rollout.rollout_id}")
|
||||
self._rollouts.append(rollout)
|
||||
return {"status": "received", "rollout_id": rollout.rollout_id}
|
||||
|
||||
async def poll_next_task_async(self) -> Task:
|
||||
return self.poll_next_task()
|
||||
|
||||
async def get_resources_by_id_async(self, resource_id: str) -> Optional[ResourcesUpdate]:
|
||||
return self.get_resources_by_id(resource_id)
|
||||
|
||||
async def get_latest_resources_async(self) -> Optional[ResourcesUpdate]:
|
||||
return self.get_latest_resources()
|
||||
|
||||
async def post_rollout_async(self, rollout: Rollout) -> Optional[Dict[str, Any]]:
|
||||
return self.post_rollout(rollout)
|
||||
|
||||
def __repr__(self):
|
||||
return f"DevTaskLoader(num_tasks={len(self._tasks)}, resources={self._resources_update.resources})"
|
||||
last_exc: Exception | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
response = super().post(*args, **kwargs)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except Exception as exc:
|
||||
last_exc = exc
|
||||
print(f"POST failed (attempt {attempt + 1}/{self.max_retries + 1}): {exc}")
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(min(2 ** (attempt + 1), 30))
|
||||
assert last_exc is not None
|
||||
raise last_exc
|
||||
|
||||
@@ -1,336 +0,0 @@
|
||||
"""
|
||||
This file is not carefully reviewed.
|
||||
It might contain unintentional bugs and issues.
|
||||
Please always review the parsed construction arguments before using them.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import inspect
|
||||
import logging
|
||||
from typing import (
|
||||
Any,
|
||||
List,
|
||||
Type,
|
||||
TypeVar,
|
||||
Union,
|
||||
_GenericAlias, # type: ignore
|
||||
get_origin,
|
||||
get_args,
|
||||
Tuple,
|
||||
Callable,
|
||||
overload,
|
||||
Dict,
|
||||
get_type_hints,
|
||||
)
|
||||
|
||||
CliConfigurable = Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# TypeVars for precise return type hinting with overloads
|
||||
_C = TypeVar("_C", bound=CliConfigurable)
|
||||
_C1 = TypeVar("_C1", bound=CliConfigurable)
|
||||
_C2 = TypeVar("_C2", bound=CliConfigurable)
|
||||
_C3 = TypeVar("_C3", bound=CliConfigurable)
|
||||
_C4 = TypeVar("_C4", bound=CliConfigurable)
|
||||
|
||||
|
||||
# Custom type for CLI arguments that can be string or None
|
||||
def nullable_str(value: str) -> str | None:
|
||||
"""Converts specific string values (case-insensitive) to None, otherwise returns the string."""
|
||||
if value.lower() in ["none", "null", "~", "nil"]: # Define keywords for None
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
def nullable_int(value: str) -> int | None:
|
||||
"""Converts specific string values (case-insensitive) to None, otherwise returns the integer."""
|
||||
if value.lower() in ["none", "null", "~", "nil"]: # Define keywords for None
|
||||
return None
|
||||
try:
|
||||
return int(value)
|
||||
except ValueError:
|
||||
raise argparse.ArgumentTypeError(f"Invalid integer value: '{value}'")
|
||||
|
||||
|
||||
def nullable_float(value: str) -> float | None:
|
||||
"""Converts specific string values (case-insensitive) to None, otherwise returns the float."""
|
||||
if value.lower() in ["none", "null", "~", "nil"]: # Define keywords for None
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except ValueError:
|
||||
raise argparse.ArgumentTypeError(f"Invalid float value: '{value}'")
|
||||
|
||||
|
||||
def _str_to_bool(v: str) -> bool:
|
||||
"""Converts common string representations of bool to Python bool (case-insensitive)."""
|
||||
if isinstance(v, bool): # Allow passing bools directly if used programmatically
|
||||
return v
|
||||
lowered_v = v.lower()
|
||||
if lowered_v in ("yes", "true", "t", "y", "1"):
|
||||
return True
|
||||
elif lowered_v in ("no", "false", "f", "n", "0"):
|
||||
return False
|
||||
else:
|
||||
raise argparse.ArgumentTypeError(f"Boolean value expected (e.g., 'true', 'false', 'yes', 'no'), got '{v}'")
|
||||
|
||||
|
||||
def _get_param_type_details(param_annotation: Any) -> Tuple[Any, bool, bool]:
|
||||
"""
|
||||
Determines the core type, if it's Optional, and if it's a List.
|
||||
Returns: (core_type, is_optional, is_list)
|
||||
- For Optional[T]: (T, True, is_list_status_of_T)
|
||||
- For List[T]: (List[T], is_optional_status_of_List, True)
|
||||
- For Optional[List[T]]: (List[T], True, True)
|
||||
"""
|
||||
is_optional = False
|
||||
is_list = False
|
||||
current_type = param_annotation
|
||||
|
||||
# Check for outer Optional
|
||||
origin = get_origin(current_type)
|
||||
if origin is Union:
|
||||
union_args = get_args(current_type)
|
||||
if len(union_args) == 2 and type(None) in union_args:
|
||||
is_optional = True
|
||||
current_type = next(arg for arg in union_args if arg is not type(None)) # Unwrap Optional
|
||||
|
||||
# Check if the (potentially unwrapped) type is a List
|
||||
origin = get_origin(current_type) # Re-check origin after potential unwrap
|
||||
if origin is list or (isinstance(current_type, _GenericAlias) and current_type.__origin__ is list):
|
||||
is_list = True
|
||||
|
||||
return current_type, is_optional, is_list
|
||||
|
||||
|
||||
def _determine_argparse_type(param_type: Any) -> Callable[[str], Any]:
|
||||
"""Determines the type for argparse based on parameter type details."""
|
||||
core_type, is_optional, _ = _get_param_type_details(param_type)
|
||||
if core_type is str and is_optional:
|
||||
return nullable_str # Special handling for Optional[str]
|
||||
elif core_type is int and is_optional:
|
||||
return nullable_int
|
||||
elif core_type is float and is_optional:
|
||||
return nullable_float
|
||||
elif core_type is bool:
|
||||
return _str_to_bool # Special handling for bool
|
||||
elif core_type in (int, float, str):
|
||||
return core_type
|
||||
return str # Default to str if no specific type is provided (including empty)
|
||||
|
||||
|
||||
def _determine_argparse_type_and_nargs(
|
||||
core_param_type: Any, is_param_list: bool # The type after unwrapping an outer Optional
|
||||
) -> Dict[str, Any]:
|
||||
"""Determines the 'type' and 'nargs' for argparse based on parameter type details."""
|
||||
kwargs: Dict[str, Any] = {}
|
||||
|
||||
if is_param_list:
|
||||
kwargs["nargs"] = "*" # Allows zero or more arguments for lists
|
||||
list_item_annotations = get_args(core_param_type) # For List[T], core_param_type is List[T]
|
||||
|
||||
if list_item_annotations and list_item_annotations[0] is not Any:
|
||||
item_ann = list_item_annotations[0]
|
||||
# Check if the list item itself is, e.g., Optional[str] or bool
|
||||
kwargs["type"] = _determine_argparse_type(item_ann)
|
||||
else:
|
||||
kwargs["type"] = str
|
||||
else: # Not a list
|
||||
kwargs["type"] = _determine_argparse_type(core_param_type)
|
||||
return kwargs
|
||||
|
||||
|
||||
def _build_help_string(cls_name: str, param_name: str, core_type: Any, is_optional: bool, is_list: bool) -> str:
|
||||
"""Constructs a descriptive help string for a CLI argument."""
|
||||
type_display_name = "Any"
|
||||
if core_type is not inspect.Parameter.empty:
|
||||
type_display_name = getattr(core_type, "__name__", str(core_type))
|
||||
|
||||
if is_list:
|
||||
list_item_args = get_args(core_type) # core_type is List[T] here
|
||||
item_name = "Any"
|
||||
if list_item_args and list_item_args[0] is not Any:
|
||||
inner_item_core_type, inner_item_optional, _ = _get_param_type_details(list_item_args[0])
|
||||
item_name = getattr(inner_item_core_type, "__name__", str(inner_item_core_type))
|
||||
if inner_item_optional: # e.g. List[Optional[str]]
|
||||
item_name = f"Optional[{item_name}]"
|
||||
type_display_name = f"List[{item_name}]"
|
||||
|
||||
full_type_display = f"Optional[{type_display_name}]" if is_optional and not is_list else type_display_name
|
||||
if is_optional and is_list: # e.g. Optional[List[str]]
|
||||
full_type_display = f"Optional[{type_display_name}]"
|
||||
|
||||
help_str = f"For {cls_name}: '{param_name}'. Inferred type: {full_type_display}."
|
||||
return help_str
|
||||
|
||||
|
||||
def _add_argument_for_parameter(
|
||||
parser: argparse.ArgumentParser,
|
||||
cls: Type[CliConfigurable],
|
||||
param_name: str,
|
||||
param_obj: inspect.Parameter,
|
||||
dest_name: str,
|
||||
resolved_param_annotation: Any = None,
|
||||
) -> None:
|
||||
"""Configures and adds a single CLI argument for an __init__ parameter."""
|
||||
if resolved_param_annotation is None:
|
||||
param_type_annotation = param_obj.annotation
|
||||
else:
|
||||
param_type_annotation = resolved_param_annotation
|
||||
|
||||
# core_type is the main type (e.g., int, str, List[str]), after unwrapping the outermost Optional.
|
||||
# is_overall_optional indicates if the parameter itself can be None (e.g. param: Optional[T] = None)
|
||||
# is_list indicates if core_type is a List.
|
||||
core_type, is_overall_optional, is_list = _get_param_type_details(param_type_annotation)
|
||||
|
||||
has_init_default = param_obj.default is not inspect.Parameter.empty
|
||||
init_default_value = param_obj.default if has_init_default else None
|
||||
|
||||
argparse_kwargs = _determine_argparse_type_and_nargs(core_type if is_list else param_type_annotation, is_list)
|
||||
|
||||
if has_init_default:
|
||||
argparse_kwargs["default"] = init_default_value
|
||||
elif is_overall_optional: # Parameter is Optional (e.g. Optional[int]) and no explicit default in __init__
|
||||
argparse_kwargs["default"] = None # So, if not provided on CLI, it becomes None.
|
||||
|
||||
argparse_kwargs["help"] = _build_help_string(cls.__name__, param_name, core_type, is_overall_optional, is_list)
|
||||
|
||||
if not has_init_default and not is_overall_optional: # Required if no __init__ default AND not Optional
|
||||
argparse_kwargs["required"] = True
|
||||
if "default" in argparse_kwargs: # Should not happen if logic is correct
|
||||
del argparse_kwargs["default"]
|
||||
|
||||
cli_arg_name = f"--{cls.__name__.lower()}.{param_name.replace('_', '-')}"
|
||||
parser.add_argument(cli_arg_name, dest=dest_name, **argparse_kwargs)
|
||||
|
||||
|
||||
def _add_arguments_for_class(
|
||||
parser: argparse.ArgumentParser,
|
||||
cls: Type[CliConfigurable],
|
||||
class_arg_configs_maps: Dict[Type[CliConfigurable], Dict[str, str]], # Maps cls to {param_name: dest_name}
|
||||
) -> None:
|
||||
"""Adds all relevant CLI arguments for a given class by processing its __init__ parameters."""
|
||||
cls_name_lower = cls.__name__.lower()
|
||||
sig = inspect.signature(cls.__init__)
|
||||
|
||||
try:
|
||||
# Resolve string annotations to actual types using get_type_hints.
|
||||
# For methods, get_type_hints automatically uses obj.__globals__ for globalns.
|
||||
resolved_hints = get_type_hints(cls.__init__)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Could not resolve type hints for {cls.__name__}.__init__ using get_type_hints: {e}. "
|
||||
f"CLI argument parsing for this class might be based on string annotations, "
|
||||
"which could be unreliable for complex types."
|
||||
)
|
||||
resolved_hints = {} # Fallback to an empty dict if resolution fails
|
||||
|
||||
if cls not in class_arg_configs_maps: # Ensure the class entry exists
|
||||
class_arg_configs_maps[cls] = {}
|
||||
|
||||
for param_name, param_obj in sig.parameters.items():
|
||||
if param_name == "self": # Skip 'self'
|
||||
continue
|
||||
|
||||
dest_name = f"{cls_name_lower}_{param_name}" # Unique destination for argparse
|
||||
class_arg_configs_maps[cls][param_name] = dest_name # Store mapping for later instantiation
|
||||
|
||||
# Use the resolved hint if available, otherwise fallback to param_obj.annotation (which might be a string)
|
||||
actual_param_annotation = resolved_hints.get(param_name, param_obj.annotation)
|
||||
_add_argument_for_parameter(parser, cls, param_name, param_obj, dest_name, actual_param_annotation)
|
||||
|
||||
|
||||
def _create_argument_parser() -> argparse.ArgumentParser:
|
||||
"""Creates and returns the main ArgumentParser with default settings."""
|
||||
return argparse.ArgumentParser(
|
||||
description="CLI configurator for application components.",
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter, # Automatically shows default values in help
|
||||
)
|
||||
|
||||
|
||||
def _instantiate_classes(
|
||||
parsed_args: argparse.Namespace,
|
||||
classes: Tuple[Type[CliConfigurable], ...],
|
||||
class_arg_configs_maps: Dict[Type[CliConfigurable], Dict[str, str]],
|
||||
) -> Tuple[CliConfigurable, ...]:
|
||||
"""Instantiates classes using the parsed CLI arguments and the stored mappings."""
|
||||
instances_list: List[CliConfigurable] = []
|
||||
for cls in classes:
|
||||
constructor_args: Dict[str, Any] = {}
|
||||
# Get the {__init__ param_name: argparse_dest_name} map for the current class
|
||||
param_to_dest_map = class_arg_configs_maps.get(cls, {})
|
||||
|
||||
sig = inspect.signature(cls.__init__)
|
||||
for param_name_in_sig, _ in sig.parameters.items():
|
||||
if param_name_in_sig == "self":
|
||||
continue
|
||||
|
||||
dest_name_for_arg = param_to_dest_map.get(param_name_in_sig)
|
||||
if dest_name_for_arg and hasattr(parsed_args, dest_name_for_arg):
|
||||
value = getattr(parsed_args, dest_name_for_arg)
|
||||
constructor_args[param_name_in_sig] = value
|
||||
# If an argument was required by argparse, parse_args() would have exited if missing.
|
||||
# If not required and not provided, its default value (set by argparse) is used.
|
||||
|
||||
try:
|
||||
logger.info("Instantiating %s with args: %s", cls.__name__, constructor_args)
|
||||
instances_list.append(cls(**constructor_args))
|
||||
except Exception as e:
|
||||
parsed_args_for_cls = {
|
||||
k: getattr(parsed_args, v) for k, v in param_to_dest_map.items() if hasattr(parsed_args, v)
|
||||
}
|
||||
logger.error(
|
||||
f"Error instantiating {cls.__name__} with resolved args {constructor_args}. "
|
||||
f"Parsed args for class: "
|
||||
f"{parsed_args_for_cls}. "
|
||||
f"Error: {e}"
|
||||
)
|
||||
raise
|
||||
|
||||
return tuple(instances_list)
|
||||
|
||||
|
||||
@overload
|
||||
def lightning_cli(cls1: Type[_C1]) -> _C1: ...
|
||||
@overload
|
||||
def lightning_cli(cls1: Type[_C1], cls2: Type[_C2]) -> Tuple[_C1, _C2]: ...
|
||||
@overload
|
||||
def lightning_cli(cls1: Type[_C1], cls2: Type[_C2], cls3: Type[_C3]) -> Tuple[_C1, _C2, _C3]: ...
|
||||
@overload
|
||||
def lightning_cli(cls1: Type[_C1], cls2: Type[_C2], cls3: Type[_C3], cls4: Type[_C4]) -> Tuple[_C1, _C2, _C3, _C4]: ...
|
||||
@overload # Fallback for more than 4 or a dynamic number of classes
|
||||
def lightning_cli(*classes: Type[CliConfigurable]) -> Tuple[CliConfigurable, ...]: ...
|
||||
|
||||
|
||||
def lightning_cli(*classes: Type[CliConfigurable]) -> CliConfigurable | Tuple[CliConfigurable, ...]:
|
||||
"""
|
||||
Parses command-line arguments to configure and instantiate provided CliConfigurable classes.
|
||||
|
||||
Args:
|
||||
*classes: One or more classes that inherit from CliConfigurable. Each class's
|
||||
__init__ parameters will be exposed as command-line arguments.
|
||||
|
||||
Returns:
|
||||
A tuple of instantiated objects, corresponding to the input classes in order.
|
||||
"""
|
||||
if not classes:
|
||||
return tuple() # Return an empty tuple if no classes are provided
|
||||
|
||||
parser = _create_argument_parser()
|
||||
|
||||
# This map will store {cls: {init_param_name: argparse_dest_name}}
|
||||
class_arg_configs_maps: Dict[Type[CliConfigurable], Dict[str, str]] = {}
|
||||
|
||||
for cls in classes:
|
||||
_add_arguments_for_class(parser, cls, class_arg_configs_maps)
|
||||
|
||||
parsed_args = parser.parse_args() # Uses sys.argv[1:] by default
|
||||
|
||||
# Correctly handle single class case for return type matching overloads
|
||||
instances = _instantiate_classes(parsed_args, classes, class_arg_configs_maps)
|
||||
if len(classes) == 1:
|
||||
return instances[0]
|
||||
return instances
|
||||
@@ -0,0 +1,3 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Hydra configuration package for Agent Lightning."""
|
||||
@@ -0,0 +1,19 @@
|
||||
runner_type: k8s # k8s | local
|
||||
|
||||
agl_server:
|
||||
url: http://localhost:8080
|
||||
# Optional external URL for agent pods. If unset, falls back to agl_server.url.
|
||||
# Example for minikube docker driver:
|
||||
# agent_url: http://host.minikube.internal:8080
|
||||
agent_url: null
|
||||
key: ""
|
||||
|
||||
k8s_runner:
|
||||
namespace: default
|
||||
ttl_after_finished: 1200
|
||||
max_jobs_per_minute: 100
|
||||
poll_interval: 5
|
||||
|
||||
local_runner:
|
||||
maximum_size: 50
|
||||
poll_interval: 10
|
||||
@@ -0,0 +1,10 @@
|
||||
host: 0.0.0.0
|
||||
port: 8080
|
||||
key: ""
|
||||
default_proxy:
|
||||
model_name: "Qwen/Qwen2.5-7B-Instruct"
|
||||
include_log_probs: True
|
||||
train:
|
||||
temperature: 1
|
||||
val:
|
||||
temperature: 0.7
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
@@ -0,0 +1,52 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Hydra entrypoint for the Agent Lightning controller."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import signal
|
||||
|
||||
import hydra
|
||||
from omegaconf import DictConfig
|
||||
|
||||
from agentlightning.client import AgentLightningAsyncClient
|
||||
|
||||
|
||||
async def _run_controller(config: DictConfig) -> None:
|
||||
async with AgentLightningAsyncClient(
|
||||
base_url=str(config.agl_server.url),
|
||||
key=str(config.agl_server.key or "") or None,
|
||||
) as api:
|
||||
if config.runner_type == "k8s":
|
||||
try:
|
||||
from agentlightning.controller.k8s_reconciler import K8sReconciler
|
||||
except ImportError:
|
||||
raise RuntimeError("kr8s unavailable - install agentlightning[controller]") from None
|
||||
|
||||
reconciler = K8sReconciler(api=api, config=config)
|
||||
elif config.runner_type == "local":
|
||||
from agentlightning.controller.local_reconciler import LocalReconciler
|
||||
|
||||
reconciler = LocalReconciler(
|
||||
api=api,
|
||||
config=config,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"unknown runner_type: {config.runner_type}")
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
with contextlib.suppress(NotImplementedError):
|
||||
loop.add_signal_handler(sig, reconciler.stop)
|
||||
await reconciler.run()
|
||||
|
||||
|
||||
@hydra.main(version_base=None, config_path="../config", config_name="controller")
|
||||
def main(config: DictConfig) -> None:
|
||||
asyncio.run(_run_controller(config))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,362 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""K8s controller reconciler — manages rollout lifecycle via K8s Jobs.
|
||||
|
||||
Two concurrent tasks:
|
||||
1. periodic_reconcile() — poll queuing rollouts, create Jobs, expire stale
|
||||
2. watch_jobs() — react to Job completions/failures, update rollout status
|
||||
|
||||
Uses AgentLightningAsyncClient for store access and kr8s for K8s API.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from collections import deque
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
import kr8s
|
||||
import kr8s.asyncio
|
||||
import structlog
|
||||
import yaml
|
||||
from jinja2 import Environment
|
||||
from kr8s.asyncio import objects as k8s_objects
|
||||
from omegaconf import DictConfig
|
||||
|
||||
from agentlightning.client import AgentLightningAsyncClient
|
||||
from agentlightning.schemas import DEFAULT_ATTEMPT_ID, Rollout, RolloutPatch, RolloutState, RolloutStatusPatch
|
||||
|
||||
log = structlog.get_logger()
|
||||
|
||||
MANAGED_BY_SELECTOR = "app.kubernetes.io/managed-by=agentlightning"
|
||||
JOB_CREATION_WINDOW_SECONDS = 60
|
||||
|
||||
|
||||
def build_job_name(rollout_id: str) -> str:
|
||||
"""Deterministic Job name from rollout ID."""
|
||||
return f"agl-rollout-{rollout_id}"
|
||||
|
||||
|
||||
def build_job_spec(rollout: Rollout, controller_config: DictConfig) -> dict[str, Any]:
|
||||
"""Build a K8s Job manifest from the rollout's complete Jinja2 Job template."""
|
||||
template = rollout.config.k8s.job_template if rollout.config.k8s else None
|
||||
if not template:
|
||||
raise ValueError("invalid rollout config: missing config.k8s.job_template")
|
||||
|
||||
env = Environment()
|
||||
env.filters["yaml_escape"] = lambda value: json.dumps(str(value), ensure_ascii=True)
|
||||
rendered = env.from_string(template).render(
|
||||
job_name=build_job_name(rollout.rollout_id),
|
||||
input=rollout.input,
|
||||
)
|
||||
docs = [doc for doc in yaml.safe_load_all(rendered) if doc is not None]
|
||||
if len(docs) != 1:
|
||||
raise ValueError("invalid rollout config: config.k8s.job_template must render exactly one YAML document")
|
||||
|
||||
job = docs[0]
|
||||
if not isinstance(job, dict) or job.get("kind") != "Job":
|
||||
raise ValueError("invalid rollout config: config.k8s.job_template must render a Kubernetes Job")
|
||||
|
||||
metadata = job.setdefault("metadata", {})
|
||||
metadata["name"] = build_job_name(rollout.rollout_id)
|
||||
metadata["namespace"] = controller_config.k8s_runner.namespace
|
||||
labels = metadata.setdefault("labels", {})
|
||||
labels["app.kubernetes.io/managed-by"] = "agentlightning"
|
||||
labels["agentlightning/rollout-id"] = rollout.rollout_id
|
||||
labels["agentlightning/attempt-id"] = DEFAULT_ATTEMPT_ID
|
||||
|
||||
spec = job.setdefault("spec", {})
|
||||
spec["backoffLimit"] = 0
|
||||
spec["ttlSecondsAfterFinished"] = controller_config.k8s_runner.ttl_after_finished
|
||||
if rollout.config.timeout_seconds:
|
||||
spec["activeDeadlineSeconds"] = rollout.config.timeout_seconds
|
||||
pod_spec = spec.setdefault("template", {}).setdefault("spec", {})
|
||||
pod_spec["restartPolicy"] = "Never"
|
||||
|
||||
mode = "train" if rollout.is_train else "val"
|
||||
agent_base_url = str(
|
||||
controller_config.agl_server.get("agent_url", None) or controller_config.agl_server.url
|
||||
).rstrip("/")
|
||||
agl_openai_base_url = (
|
||||
f"{agent_base_url}/proxy/rollout/{rollout.rollout_id}/attempt/{DEFAULT_ATTEMPT_ID}/mode/{mode}/openai/v1"
|
||||
)
|
||||
for container in pod_spec.get("containers", []):
|
||||
env = container.setdefault("env", [])
|
||||
for name, value in {
|
||||
"AGL_OPENAI_BASE_URL": agl_openai_base_url,
|
||||
"AGL_EVENT_URL": (
|
||||
f"{agent_base_url}/api/rollouts/{rollout.rollout_id}/attempt/{DEFAULT_ATTEMPT_ID}/events"
|
||||
),
|
||||
"AGL_KEY": str(controller_config.agl_server.key or ""),
|
||||
}.items():
|
||||
existing = next((item for item in env if item.get("name") == name), None)
|
||||
if existing is None:
|
||||
env.append({"name": name, "value": value})
|
||||
else:
|
||||
existing.clear()
|
||||
existing.update({"name": name, "value": value})
|
||||
return job
|
||||
|
||||
|
||||
class K8sReconciler:
|
||||
"""Main controller loop. Reconciles rollouts into K8s Jobs.
|
||||
|
||||
Args:
|
||||
api: AgentLightningAsyncClient for store access.
|
||||
config: Controller configuration.
|
||||
"""
|
||||
|
||||
def __init__(self, api: AgentLightningAsyncClient, config: DictConfig) -> None:
|
||||
self._api = api
|
||||
self._config = config
|
||||
self._runner_config = config.k8s_runner
|
||||
self._namespace = str(self._runner_config.namespace)
|
||||
self._k8s_api: Any | None = None
|
||||
self._stop = asyncio.Event()
|
||||
self._job_creation_timestamps: deque[float] = deque()
|
||||
|
||||
async def _get_k8s_api(self) -> Any:
|
||||
if self._k8s_api is None:
|
||||
self._k8s_api = await kr8s.asyncio.api()
|
||||
return self._k8s_api
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Start both reconcile loops. Blocks until stop() is called."""
|
||||
log.info(
|
||||
"Controller starting",
|
||||
namespace=self._namespace,
|
||||
poll_interval=self._runner_config.poll_interval,
|
||||
)
|
||||
try:
|
||||
await asyncio.gather(
|
||||
self._periodic_reconcile_loop(),
|
||||
self._watch_jobs_loop(),
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
log.info("Controller stopped")
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Signal the controller to stop."""
|
||||
self._stop.set()
|
||||
|
||||
# --- Periodic reconcile ---
|
||||
|
||||
async def _periodic_reconcile_loop(self) -> None:
|
||||
"""Poll queuing rollouts and reconcile."""
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
await self._reconcile_once()
|
||||
except Exception:
|
||||
log.exception("Periodic reconcile error")
|
||||
# Sleep with cancellation support.
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), timeout=self._runner_config.poll_interval)
|
||||
return # stop was set
|
||||
except TimeoutError:
|
||||
pass
|
||||
|
||||
async def _reconcile_once(self) -> None:
|
||||
"""One reconcile cycle: align queuing/running rollouts with K8s Jobs."""
|
||||
rollouts = await self._query_rollouts(state_in=[RolloutState.QUEUING, RolloutState.RUNNING], limit=500)
|
||||
api = await self._get_k8s_api()
|
||||
jobs = [
|
||||
cast(k8s_objects.Job, job).raw
|
||||
async for job in k8s_objects.Job.async_list(
|
||||
namespace=self._namespace,
|
||||
label_selector=MANAGED_BY_SELECTOR,
|
||||
api=api,
|
||||
)
|
||||
]
|
||||
jobs_by_name = {job.get("metadata", {}).get("name", ""): job for job in jobs}
|
||||
|
||||
for rollout in rollouts:
|
||||
job_name = rollout.status.k8s_job_name or build_job_name(rollout.rollout_id)
|
||||
job = jobs_by_name.get(job_name)
|
||||
|
||||
if job is None:
|
||||
if rollout.status.state == RolloutState.QUEUING:
|
||||
await self._create_job(rollout)
|
||||
continue
|
||||
log.warning("Orphaned running rollout — Job gone", rollout_id=rollout.rollout_id, job_name=job_name)
|
||||
await self._patch_status(rollout.rollout_id, state=RolloutState.FAILED, error_message="Job disappeared")
|
||||
continue
|
||||
|
||||
attempt_id = (
|
||||
job.get("metadata", {}).get("labels", {}).get("agentlightning/attempt-id") or DEFAULT_ATTEMPT_ID
|
||||
)
|
||||
|
||||
job_status = job.get("status", {})
|
||||
state = None
|
||||
error_message = None
|
||||
for condition in job_status.get("conditions", []):
|
||||
if condition.get("status") != "True":
|
||||
continue
|
||||
if condition.get("type") == "Complete":
|
||||
state = RolloutState.SUCCEEDED
|
||||
break
|
||||
if condition.get("type") == "Failed":
|
||||
reason = condition.get("reason", "Unknown")
|
||||
message = condition.get("message", "")
|
||||
error_message = f"Job failed: {reason}"
|
||||
if message:
|
||||
error_message += f" — {message}"
|
||||
state = RolloutState.FAILED
|
||||
break
|
||||
|
||||
if state is None and job_status.get("succeeded", 0) > 0:
|
||||
state = RolloutState.SUCCEEDED
|
||||
elif state is None and job_status.get("failed", 0) > 0:
|
||||
state = RolloutState.FAILED
|
||||
error_message = "Job failed"
|
||||
|
||||
if state is None:
|
||||
if rollout.status.state == RolloutState.QUEUING:
|
||||
await self._patch_status(
|
||||
rollout.rollout_id,
|
||||
state=RolloutState.RUNNING,
|
||||
k8s_job_name=job_name,
|
||||
last_attempt_id=attempt_id,
|
||||
)
|
||||
continue
|
||||
|
||||
if rollout.status.state == RolloutState.QUEUING and state == RolloutState.SUCCEEDED:
|
||||
patched = await self._patch_status(
|
||||
rollout.rollout_id,
|
||||
state=RolloutState.RUNNING,
|
||||
k8s_job_name=job_name,
|
||||
last_attempt_id=attempt_id,
|
||||
)
|
||||
if not patched:
|
||||
continue
|
||||
await self._patch_status(
|
||||
rollout.rollout_id,
|
||||
state=state,
|
||||
k8s_job_name=job_name,
|
||||
last_attempt_id=attempt_id,
|
||||
error_message=error_message,
|
||||
)
|
||||
|
||||
async def _create_job(self, rollout: Rollout) -> None:
|
||||
"""Create a K8s Job for a queuing rollout without changing rollout state."""
|
||||
job_name = build_job_name(rollout.rollout_id)
|
||||
now = time.monotonic()
|
||||
window_start = now - JOB_CREATION_WINDOW_SECONDS
|
||||
while self._job_creation_timestamps and self._job_creation_timestamps[0] <= window_start:
|
||||
self._job_creation_timestamps.popleft()
|
||||
if len(self._job_creation_timestamps) >= self._runner_config.max_jobs_per_minute:
|
||||
log.info(
|
||||
"Job creation rate limit reached — deferring queued rollouts",
|
||||
rollout_id=rollout.rollout_id,
|
||||
jobs_in_last_minute=len(self._job_creation_timestamps),
|
||||
max_jobs_per_minute=self._runner_config.max_jobs_per_minute,
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
manifest = build_job_spec(rollout, self._config)
|
||||
attempt_id = manifest["metadata"]["labels"]["agentlightning/attempt-id"]
|
||||
api = await self._get_k8s_api()
|
||||
job = k8s_objects.Job(manifest, api=api)
|
||||
await job.async_create()
|
||||
self._job_creation_timestamps.append(time.monotonic())
|
||||
log.info("Job created", rollout_id=rollout.rollout_id, job_name=job_name, attempt_id=attempt_id)
|
||||
except Exception as exc:
|
||||
error_str = str(exc)
|
||||
lower_error = error_str.lower()
|
||||
if "422" in lower_error or "unprocessable" in lower_error or "invalid" in lower_error:
|
||||
log.error("Invalid Job spec — marking failed", rollout_id=rollout.rollout_id, error=error_str)
|
||||
await self._patch_status(
|
||||
rollout.rollout_id,
|
||||
state=RolloutState.FAILED,
|
||||
error_message=f"Invalid Job spec: {error_str}",
|
||||
)
|
||||
else:
|
||||
log.warning("Job creation failed — will retry", rollout_id=rollout.rollout_id, error=error_str)
|
||||
|
||||
# --- Watch Jobs ---
|
||||
|
||||
async def _watch_jobs_loop(self) -> None:
|
||||
"""Watch K8s Job events and react to completions/failures."""
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
watcher = kr8s.asyncio.watch(
|
||||
"jobs",
|
||||
namespace=self._namespace,
|
||||
label_selector=MANAGED_BY_SELECTOR,
|
||||
api=await self._get_k8s_api(),
|
||||
)
|
||||
async for event_type, obj in watcher:
|
||||
if self._stop.is_set():
|
||||
return
|
||||
if event_type in ("MODIFIED", "ADDED"):
|
||||
await self._handle_job_event(obj.raw)
|
||||
except Exception:
|
||||
log.exception("Watch error — restarting watch")
|
||||
await asyncio.sleep(5)
|
||||
|
||||
async def _handle_job_event(self, job: dict[str, Any]) -> None:
|
||||
"""Process a Job event — check conditions, update rollout status."""
|
||||
labels = job.get("metadata", {}).get("labels", {})
|
||||
rollout_id = labels.get("agentlightning/rollout-id")
|
||||
if not rollout_id:
|
||||
return
|
||||
attempt_id = labels.get("agentlightning/attempt-id") or DEFAULT_ATTEMPT_ID
|
||||
|
||||
conditions = job.get("status", {}).get("conditions", [])
|
||||
if not conditions:
|
||||
return
|
||||
|
||||
for condition in conditions:
|
||||
cond_type = condition.get("type", "")
|
||||
cond_status = condition.get("status", "")
|
||||
if cond_status != "True":
|
||||
continue
|
||||
|
||||
if cond_type == "Complete":
|
||||
log.info("Job completed", rollout_id=rollout_id, last_attempt_id=attempt_id)
|
||||
await self._patch_status(rollout_id, state=RolloutState.SUCCEEDED, last_attempt_id=attempt_id)
|
||||
return
|
||||
elif cond_type == "Failed":
|
||||
reason = condition.get("reason", "Unknown")
|
||||
message = condition.get("message", "")
|
||||
error_msg = f"Job failed: {reason}"
|
||||
if message:
|
||||
error_msg += f" — {message}"
|
||||
log.info("Job failed", rollout_id=rollout_id, last_attempt_id=attempt_id, reason=reason)
|
||||
await self._patch_status(
|
||||
rollout_id,
|
||||
state=RolloutState.FAILED,
|
||||
last_attempt_id=attempt_id,
|
||||
error_message=error_msg,
|
||||
)
|
||||
return
|
||||
|
||||
async def _query_rollouts(
|
||||
self,
|
||||
*,
|
||||
state_in: list[RolloutState],
|
||||
limit: int = 50,
|
||||
) -> list[Rollout]:
|
||||
params = httpx.QueryParams()
|
||||
for state in state_in:
|
||||
params = params.add("state_in", state.value)
|
||||
params = params.add("limit", limit)
|
||||
response = await self._api.get("/api/rollouts", params=params)
|
||||
response.raise_for_status()
|
||||
return [Rollout.model_validate(item) for item in response.json()]
|
||||
|
||||
async def _patch_status(self, rollout_id: str, **status: Any) -> bool:
|
||||
try:
|
||||
patch = RolloutPatch(status=RolloutStatusPatch.model_validate(status))
|
||||
response = await self._api.patch(
|
||||
f"/api/rollouts/{rollout_id}",
|
||||
json=patch.model_dump(mode="json", exclude_unset=True),
|
||||
)
|
||||
response.raise_for_status()
|
||||
return True
|
||||
except Exception as exc:
|
||||
log.warning("Failed to patch rollout", rollout_id=rollout_id, error=str(exc))
|
||||
return False
|
||||
@@ -0,0 +1,283 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Local reconciler that runs rollouts as short-lived Python subprocesses."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import importlib
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from dataclasses import dataclass
|
||||
|
||||
import httpx
|
||||
import structlog
|
||||
from omegaconf import DictConfig
|
||||
|
||||
from agentlightning.client import AgentLightningAsyncClient
|
||||
from agentlightning.schemas import DEFAULT_ATTEMPT_ID, Rollout, RolloutPatch, RolloutState, RolloutStatusPatch
|
||||
|
||||
log = structlog.get_logger()
|
||||
|
||||
_SHUTDOWN_WAIT_TIMEOUT = 5.0
|
||||
|
||||
|
||||
def _run_local_reconciler_worker(agent_class_path: str) -> int:
|
||||
try:
|
||||
if ":" in agent_class_path:
|
||||
module_name, class_name = agent_class_path.split(":", 1)
|
||||
else:
|
||||
module_name, class_name = agent_class_path.rsplit(".", 1)
|
||||
loaded = getattr(importlib.import_module(module_name), class_name)
|
||||
if not isinstance(loaded, type):
|
||||
raise TypeError(f"{agent_class_path} is not a class")
|
||||
result = loaded().run()
|
||||
if inspect.isawaitable(result):
|
||||
asyncio.run(result) # type: ignore[arg-type]
|
||||
return 0
|
||||
except Exception:
|
||||
traceback.print_exc()
|
||||
return 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class Proc:
|
||||
"""In-flight local subprocess."""
|
||||
|
||||
attempt_id: str
|
||||
proc: asyncio.subprocess.Process
|
||||
spawned_at: float
|
||||
killed: bool = False
|
||||
|
||||
|
||||
def _build_env_from_map(task_input: object, env_map: dict[str, str]) -> dict[str, str]:
|
||||
env: dict[str, str] = {}
|
||||
for name, path in env_map.items():
|
||||
value = _resolve_input_path(task_input, path)
|
||||
if isinstance(value, str):
|
||||
env[name] = value
|
||||
continue
|
||||
try:
|
||||
env[name] = json.dumps(value, ensure_ascii=False)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"local.env_map.{name} value is not JSON serializable") from exc
|
||||
return env
|
||||
|
||||
|
||||
def _resolve_input_path(task_input: object, path: str) -> object:
|
||||
if path == "input":
|
||||
return task_input
|
||||
if not path.startswith("input."):
|
||||
return path
|
||||
|
||||
value = task_input
|
||||
for part in path.split(".")[1:]:
|
||||
if isinstance(value, dict) and part in value:
|
||||
value = value[part]
|
||||
elif isinstance(value, list) and part.isdigit() and int(part) < len(value):
|
||||
value = value[int(part)]
|
||||
else:
|
||||
raise ValueError(f"local.env_map path not found: {path}")
|
||||
return value
|
||||
|
||||
|
||||
class LocalReconciler:
|
||||
"""Local-mode rollout reconciler."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api: AgentLightningAsyncClient,
|
||||
config: DictConfig,
|
||||
) -> None:
|
||||
assert config.runner_type == "local"
|
||||
self._api = api
|
||||
self._config = config
|
||||
self._runner_config = config.local_runner
|
||||
self._pool_size = int(self._runner_config.maximum_size)
|
||||
self._tick_interval = float(self._runner_config.poll_interval)
|
||||
self._rid_to_proc: dict[str, Proc] = {}
|
||||
self._stop = asyncio.Event()
|
||||
|
||||
async def run(self) -> None:
|
||||
log.info("LocalReconciler starting", pool_size=self._pool_size, tick=self._tick_interval)
|
||||
try:
|
||||
await self._reconcile_loop()
|
||||
finally:
|
||||
await self._shutdown()
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stop.set()
|
||||
|
||||
async def _reconcile_loop(self) -> None:
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
await self._reconcile_once()
|
||||
except Exception:
|
||||
log.exception("Local reconcile error")
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), timeout=self._tick_interval)
|
||||
break
|
||||
except TimeoutError:
|
||||
pass
|
||||
|
||||
async def _reconcile_once(self) -> None:
|
||||
params = httpx.QueryParams()
|
||||
params = params.add("state_in", RolloutState.QUEUING.value)
|
||||
params = params.add("state_in", RolloutState.RUNNING.value)
|
||||
params = params.add("limit", 50)
|
||||
response = await self._api.get("/api/rollouts", params=params)
|
||||
response.raise_for_status()
|
||||
rollouts = [Rollout.model_validate(item) for item in response.json()]
|
||||
rollouts_by_id = {rollout.rollout_id: rollout for rollout in rollouts}
|
||||
live_count = sum(1 for item in self._rid_to_proc.values() if item.proc.returncode is None)
|
||||
|
||||
for rollout in rollouts:
|
||||
item = self._rid_to_proc.get(rollout.rollout_id)
|
||||
|
||||
if item is None:
|
||||
if rollout.status.state == RolloutState.QUEUING and live_count < self._pool_size:
|
||||
if await self._spawn_for(rollout):
|
||||
live_count += 1
|
||||
elif rollout.status.state == RolloutState.RUNNING:
|
||||
await self._patch(rollout.rollout_id, RolloutState.FAILED, "local subprocess is not running")
|
||||
continue
|
||||
|
||||
if item.proc.returncode is None:
|
||||
if rollout.status.state == RolloutState.QUEUING:
|
||||
await self._patch(rollout.rollout_id, RolloutState.RUNNING, last_attempt_id=item.attempt_id)
|
||||
continue
|
||||
|
||||
await self._finish_proc(rollout, item)
|
||||
|
||||
now = time.monotonic()
|
||||
for rollout_id, item in list(self._rid_to_proc.items()):
|
||||
if item.proc.returncode is not None:
|
||||
continue
|
||||
rollout = rollouts_by_id.get(rollout_id)
|
||||
timeout = float(rollout.config.timeout_seconds) if rollout and rollout.config.timeout_seconds else None
|
||||
if (
|
||||
timeout is not None
|
||||
and (now - item.spawned_at) > timeout
|
||||
and await self._kill_process_group(rollout_id, item)
|
||||
):
|
||||
await self._patch(rollout_id, RolloutState.FAILED, "local subprocess timed out")
|
||||
|
||||
async def _finish_proc(self, rollout: Rollout, item: Proc) -> bool:
|
||||
if rollout.status.state == RolloutState.QUEUING:
|
||||
patched = await self._patch(rollout.rollout_id, RolloutState.RUNNING, last_attempt_id=item.attempt_id)
|
||||
if not patched:
|
||||
return False
|
||||
if item.proc.returncode == 0:
|
||||
return await self._patch(rollout.rollout_id, RolloutState.SUCCEEDED, last_attempt_id=item.attempt_id)
|
||||
return await self._patch(
|
||||
rollout.rollout_id,
|
||||
RolloutState.FAILED,
|
||||
f"subprocess exited with code {item.proc.returncode}",
|
||||
)
|
||||
|
||||
async def _kill_process_group(self, rollout_id: str, item: Proc) -> bool:
|
||||
"""SIGKILL the worker process group and wait for exit."""
|
||||
if item.proc.returncode is not None:
|
||||
return True
|
||||
if not item.killed:
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
os.killpg(item.proc.pid, signal.SIGKILL)
|
||||
item.killed = True
|
||||
log.info("SIGKILL sent to subprocess group", rollout_id=rollout_id, pid=item.proc.pid)
|
||||
try:
|
||||
await asyncio.wait_for(item.proc.wait(), timeout=_SHUTDOWN_WAIT_TIMEOUT)
|
||||
return True
|
||||
except TimeoutError:
|
||||
log.warning("Subprocess did not exit after SIGKILL within 5s", rollout_id=rollout_id, pid=item.proc.pid)
|
||||
return False
|
||||
|
||||
async def _spawn_for(self, rollout: Rollout) -> bool:
|
||||
"""Spawn one local subprocess for a rollout."""
|
||||
try:
|
||||
attempt_id = DEFAULT_ATTEMPT_ID
|
||||
if rollout.config.local is None or not rollout.config.local.agent_class:
|
||||
raise ValueError("invalid rollout config: missing config.local.agent_class")
|
||||
agent_class = rollout.config.local.agent_class
|
||||
mode = "train" if rollout.is_train else "val"
|
||||
env = {
|
||||
**os.environ,
|
||||
"AGL_KEY": str(self._config.agl_server.key or ""),
|
||||
"AGL_OPENAI_BASE_URL": (
|
||||
f"{self._config.agl_server.url}/proxy/rollout/{rollout.rollout_id}"
|
||||
f"/attempt/{attempt_id}/mode/{mode}/openai/v1"
|
||||
),
|
||||
"AGL_EVENT_URL": (
|
||||
f"{self._config.agl_server.url}/api/rollouts/{rollout.rollout_id}/attempt/{attempt_id}/events"
|
||||
),
|
||||
}
|
||||
env.update(_build_env_from_map(rollout.input, rollout.config.local.env_map))
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
sys.executable,
|
||||
"-c",
|
||||
(
|
||||
"import sys; "
|
||||
"from agentlightning.controller.local_reconciler import _run_local_reconciler_worker; "
|
||||
"sys.exit(_run_local_reconciler_worker(sys.argv[1]))"
|
||||
),
|
||||
agent_class,
|
||||
stdin=asyncio.subprocess.DEVNULL,
|
||||
stdout=None,
|
||||
stderr=None,
|
||||
env=env,
|
||||
start_new_session=True,
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception("Spawn failed", rollout_id=rollout.rollout_id)
|
||||
await self._patch(rollout.rollout_id, RolloutState.FAILED, f"local subprocess spawn failed: {e}")
|
||||
return False
|
||||
|
||||
self._rid_to_proc[rollout.rollout_id] = Proc(
|
||||
attempt_id=attempt_id,
|
||||
proc=proc,
|
||||
spawned_at=time.monotonic(),
|
||||
)
|
||||
await self._patch(rollout.rollout_id, RolloutState.RUNNING, last_attempt_id=attempt_id)
|
||||
log.info("Spawned rollout subprocess", rollout_id=rollout.rollout_id, attempt_id=attempt_id, pid=proc.pid)
|
||||
return True
|
||||
|
||||
async def _shutdown(self) -> None:
|
||||
"""Kill live subprocesses and mark them failed."""
|
||||
try:
|
||||
await self._reconcile_once()
|
||||
except Exception:
|
||||
log.exception("Final reconcile during shutdown failed")
|
||||
|
||||
for rollout_id, item in list(self._rid_to_proc.items()):
|
||||
if item.proc.returncode is None and await self._kill_process_group(rollout_id, item):
|
||||
await self._patch(rollout_id, RolloutState.FAILED, "local controller shutdown")
|
||||
|
||||
async def _patch(
|
||||
self,
|
||||
rollout_id: str,
|
||||
state: RolloutState,
|
||||
error_message: str | None = None,
|
||||
*,
|
||||
last_attempt_id: str | None = None,
|
||||
) -> bool:
|
||||
status = RolloutStatusPatch(state=state)
|
||||
if error_message is not None:
|
||||
status.error_message = error_message
|
||||
if last_attempt_id is not None:
|
||||
status.last_attempt_id = last_attempt_id
|
||||
patch = RolloutPatch(status=status)
|
||||
try:
|
||||
response = await self._api.patch(
|
||||
f"/api/rollouts/{rollout_id}",
|
||||
json=patch.model_dump(mode="json", exclude_unset=True),
|
||||
)
|
||||
response.raise_for_status()
|
||||
return True
|
||||
except Exception as e:
|
||||
log.warning("Failed to patch rollout", rollout_id=rollout_id, error=str(e))
|
||||
return False
|
||||
@@ -0,0 +1,63 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Rollout lifecycle hooks used by enqueue and fit flows."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
from agentlightning.schemas import RolloutCreate
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentlightning.schemas import Rollout
|
||||
|
||||
|
||||
class TraceWriter(Protocol):
|
||||
def add_event(self, rollout_id: str, attempt_id: str, event_type: str, data: dict[str, Any]) -> Any: ...
|
||||
|
||||
|
||||
class RolloutHooks:
|
||||
"""Base class for synchronous rollout lifecycle hooks."""
|
||||
|
||||
def on_startup(self, store: Any | None = None) -> None:
|
||||
"""Initialize hook state once after startup."""
|
||||
|
||||
def on_enqueue(self, request: RolloutCreate) -> RolloutCreate:
|
||||
"""Transform a rollout request before it is persisted."""
|
||||
return request
|
||||
|
||||
def on_succeeded(self, rollout: Rollout, events: dict[str, list[Any]], store: TraceWriter) -> None:
|
||||
"""Run after a rollout transitions to SUCCEEDED."""
|
||||
|
||||
def on_failed(self, rollout: Rollout, store: TraceWriter) -> None:
|
||||
"""Run after a rollout transitions to FAILED."""
|
||||
|
||||
|
||||
def load_hooks(path: str) -> RolloutHooks:
|
||||
"""Load the single ``RolloutHooks`` subclass from a Python file."""
|
||||
import importlib.util
|
||||
import inspect
|
||||
from pathlib import Path
|
||||
|
||||
module_path = Path(path).resolve()
|
||||
if not module_path.exists():
|
||||
raise FileNotFoundError(f"Hooks module not found: {module_path}")
|
||||
|
||||
spec = importlib.util.spec_from_file_location("_agl_hooks", str(module_path))
|
||||
assert spec is not None and spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
hook_classes = [
|
||||
obj
|
||||
for _, obj in inspect.getmembers(module, inspect.isclass)
|
||||
if issubclass(obj, RolloutHooks) and obj is not RolloutHooks
|
||||
]
|
||||
|
||||
if len(hook_classes) == 0:
|
||||
raise ValueError(f"No RolloutHooks subclass found in {path}")
|
||||
if len(hook_classes) > 1:
|
||||
names = [cls.__name__ for cls in hook_classes]
|
||||
raise ValueError(f"Multiple RolloutHooks subclasses found in {path}: {names}")
|
||||
|
||||
return hook_classes[0]()
|
||||
@@ -1,109 +0,0 @@
|
||||
import warnings
|
||||
|
||||
AGENTOPS_INSTALLED = False
|
||||
AGENTOPS_LANGCHAIN_INSTALLED = False
|
||||
LITELLM_INSTALLED = False
|
||||
VLLM_INSTALLED = False
|
||||
|
||||
try:
|
||||
from . import agentops
|
||||
|
||||
AGENTOPS_INSTALLED = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
from . import litellm
|
||||
|
||||
LITELLM_INSTALLED = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# MAGIC! DO NOT TOUCH THIS!
|
||||
# vllm import will cause reward tracing function to fail and produce nothing.
|
||||
# try:
|
||||
# from . import vllm
|
||||
|
||||
# VLLM_INSTALLED = True
|
||||
# except ImportError:
|
||||
# pass
|
||||
|
||||
|
||||
try:
|
||||
from . import agentops_langchain
|
||||
|
||||
AGENTOPS_LANGCHAIN_INSTALLED = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
def instrument_all():
|
||||
if AGENTOPS_INSTALLED:
|
||||
from .agentops import instrument_agentops
|
||||
|
||||
instrument_agentops()
|
||||
else:
|
||||
warnings.warn("agentops is not installed. It's therefore not instrumented.")
|
||||
|
||||
if LITELLM_INSTALLED:
|
||||
from .litellm import instrument_litellm
|
||||
|
||||
instrument_litellm()
|
||||
else:
|
||||
warnings.warn("litellm is not installed. It's therefore not instrumented.")
|
||||
|
||||
if VLLM_INSTALLED:
|
||||
from .vllm import instrument_vllm
|
||||
|
||||
instrument_vllm()
|
||||
else:
|
||||
warnings.warn("vllm is not installed. It's therefore not instrumented.")
|
||||
|
||||
if AGENTOPS_LANGCHAIN_INSTALLED:
|
||||
from .agentops_langchain import instrument_agentops_langchain
|
||||
|
||||
instrument_agentops_langchain()
|
||||
else:
|
||||
warnings.warn("Agentops-langchain integration is not installed. It's therefore not instrumented.")
|
||||
|
||||
|
||||
def uninstrument_all():
|
||||
if AGENTOPS_INSTALLED:
|
||||
try:
|
||||
from .agentops import uninstrument_agentops
|
||||
|
||||
uninstrument_agentops()
|
||||
except ImportError:
|
||||
warnings.warn("agentops is installed but uninstrument_agentops could not be imported.")
|
||||
else:
|
||||
warnings.warn("agentops is not installed. It's therefore not uninstrumented.")
|
||||
|
||||
if LITELLM_INSTALLED:
|
||||
try:
|
||||
from .litellm import uninstrument_litellm
|
||||
|
||||
uninstrument_litellm()
|
||||
except ImportError:
|
||||
warnings.warn("litellm is installed but uninstrument_litellm could not be imported.")
|
||||
else:
|
||||
warnings.warn("litellm is not installed. It's therefore not uninstrumented.")
|
||||
|
||||
if VLLM_INSTALLED:
|
||||
try:
|
||||
from .vllm import uninstrument_vllm
|
||||
|
||||
uninstrument_vllm()
|
||||
except ImportError:
|
||||
warnings.warn("vllm is installed but uninstrument_vllm could not be imported.")
|
||||
else:
|
||||
warnings.warn("vllm is not installed. It's therefore not uninstrumented.")
|
||||
|
||||
if AGENTOPS_LANGCHAIN_INSTALLED:
|
||||
try:
|
||||
from .agentops_langchain import uninstrument_agentops_langchain
|
||||
|
||||
uninstrument_agentops_langchain()
|
||||
except ImportError:
|
||||
warnings.warn("agentops_langchain is installed but uninstrument_agentops_langchain could not be imported.")
|
||||
else:
|
||||
warnings.warn("Agentops-langchain integration is not installed. It's therefore not uninstrumented.")
|
||||
@@ -1,240 +0,0 @@
|
||||
import logging
|
||||
import multiprocessing
|
||||
import signal
|
||||
import socket
|
||||
import time
|
||||
|
||||
import flask
|
||||
import setproctitle
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Module-level storage for originals
|
||||
_original_handle_chat_attributes = None
|
||||
_original_handle_response = None
|
||||
|
||||
|
||||
def _patch_new_agentops():
|
||||
import agentops.instrumentation.providers.openai.wrappers.chat
|
||||
import agentops.instrumentation.providers.openai.stream_wrapper
|
||||
from agentops.instrumentation.providers.openai.wrappers.chat import handle_chat_attributes
|
||||
|
||||
global _original_handle_chat_attributes
|
||||
|
||||
if _original_handle_chat_attributes is not None:
|
||||
logger.warning("AgentOps already patched. Skipping.")
|
||||
return True
|
||||
|
||||
_original_handle_chat_attributes = handle_chat_attributes
|
||||
|
||||
def _handle_chat_attributes_with_tokens(args=None, kwargs=None, return_value=None, **kws):
|
||||
attributes = _original_handle_chat_attributes(args=args, kwargs=kwargs, return_value=return_value, **kws)
|
||||
if hasattr(return_value, "prompt_token_ids"):
|
||||
attributes["prompt_token_ids"] = list(return_value.prompt_token_ids)
|
||||
if hasattr(return_value, "response_token_ids"):
|
||||
attributes["response_token_ids"] = list(return_value.response_token_ids[0])
|
||||
|
||||
# For LiteLLM, response is a openai._legacy_response.LegacyAPIResponse
|
||||
if hasattr(return_value, "http_response") and hasattr(return_value.http_response, "json"):
|
||||
json_data = return_value.http_response.json()
|
||||
if isinstance(json_data, dict):
|
||||
if "prompt_token_ids" in json_data:
|
||||
attributes["prompt_token_ids"] = list(json_data["prompt_token_ids"])
|
||||
if "response_token_ids" in json_data:
|
||||
attributes["response_token_ids"] = list(json_data["response_token_ids"][0])
|
||||
|
||||
return attributes
|
||||
|
||||
agentops.instrumentation.providers.openai.wrappers.chat.handle_chat_attributes = _handle_chat_attributes_with_tokens
|
||||
agentops.instrumentation.providers.openai.stream_wrapper.handle_chat_attributes = (
|
||||
_handle_chat_attributes_with_tokens
|
||||
)
|
||||
logger.info("Patched newer version of agentops using handle_chat_attributes")
|
||||
return True
|
||||
|
||||
|
||||
def _unpatch_new_agentops():
|
||||
import agentops.instrumentation.providers.openai.wrappers.chat
|
||||
import agentops.instrumentation.providers.openai.stream_wrapper
|
||||
|
||||
global _original_handle_chat_attributes
|
||||
if _original_handle_chat_attributes is not None:
|
||||
agentops.instrumentation.providers.openai.wrappers.chat.handle_chat_attributes = (
|
||||
_original_handle_chat_attributes
|
||||
)
|
||||
agentops.instrumentation.providers.openai.stream_wrapper.handle_chat_attributes = (
|
||||
_original_handle_chat_attributes
|
||||
)
|
||||
_original_handle_chat_attributes = None
|
||||
logger.info("Unpatched newer version of agentops using handle_chat_attributes")
|
||||
|
||||
|
||||
def _patch_old_agentops():
|
||||
import opentelemetry.instrumentation.openai.shared.chat_wrappers
|
||||
from opentelemetry.instrumentation.openai.shared.chat_wrappers import _handle_response, dont_throw
|
||||
|
||||
global _original_handle_response
|
||||
_original_handle_response = _handle_response
|
||||
|
||||
@dont_throw
|
||||
def _handle_response_with_tokens(response, span, *args, **kwargs):
|
||||
_original_handle_response(response, span, *args, **kwargs)
|
||||
if hasattr(response, "prompt_token_ids"):
|
||||
span.set_attribute("prompt_token_ids", list(response.prompt_token_ids))
|
||||
if hasattr(response, "response_token_ids"):
|
||||
span.set_attribute("response_token_ids", list(response.response_token_ids[0]))
|
||||
|
||||
# For LiteLLM, response is a openai._legacy_response.LegacyAPIResponse
|
||||
if hasattr(response, "http_response") and hasattr(response.http_response, "json"):
|
||||
json_data = response.http_response.json()
|
||||
if isinstance(json_data, dict):
|
||||
if "prompt_token_ids" in json_data:
|
||||
span.set_attribute("prompt_token_ids", list(json_data["prompt_token_ids"]))
|
||||
if "response_token_ids" in json_data:
|
||||
span.set_attribute("response_token_ids", list(json_data["response_token_ids"][0]))
|
||||
|
||||
opentelemetry.instrumentation.openai.shared.chat_wrappers._handle_response = _handle_response_with_tokens
|
||||
logger.info("Patched earlier version of agentops using _handle_response")
|
||||
return True
|
||||
|
||||
|
||||
def _unpatch_old_agentops():
|
||||
import opentelemetry.instrumentation.openai.shared.chat_wrappers
|
||||
|
||||
global _original_handle_response
|
||||
if _original_handle_response is not None:
|
||||
opentelemetry.instrumentation.openai.shared.chat_wrappers._handle_response = _original_handle_response
|
||||
_original_handle_response = None
|
||||
logger.info("Unpatched earlier version of agentops using _handle_response")
|
||||
|
||||
|
||||
def instrument_agentops():
|
||||
"""
|
||||
Instrument agentops to capture token IDs.
|
||||
Automatically detects and uses the appropriate patching method based on the installed agentops version.
|
||||
"""
|
||||
# Try newest version first (tested for 0.4.16)
|
||||
try:
|
||||
return _patch_new_agentops()
|
||||
except ImportError as e:
|
||||
logger.debug(f"Couldn't patch newer version of agentops: {str(e)}")
|
||||
|
||||
# Note: 0.4.15 needs another patching method, but it's too shortlived to be worth handling separately.
|
||||
|
||||
# Try older version (tested for 0.4.13)
|
||||
try:
|
||||
return _patch_old_agentops()
|
||||
except ImportError as e:
|
||||
logger.warning(f"Couldn't patch older version of agentops: {str(e)}")
|
||||
logger.error("Failed to instrument agentops - neither patching method was successful")
|
||||
return False
|
||||
|
||||
|
||||
def uninstrument_agentops():
|
||||
try:
|
||||
_unpatch_new_agentops()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
_unpatch_old_agentops()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def agentops_local_server():
|
||||
"""
|
||||
Returns a Flask app that can be used to test agentops integration.
|
||||
This server provides endpoints for token fetching and a catch-all endpoint.
|
||||
"""
|
||||
app = flask.Flask(__name__)
|
||||
|
||||
@app.route("/v3/auth/token", methods=["POST"])
|
||||
def fetch_token():
|
||||
return {"token": "dummy", "project_id": "dummy"}
|
||||
|
||||
@app.route("/", defaults={"path": ""}, methods=["GET", "POST"])
|
||||
@app.route("/<path:path>", methods=["GET", "POST"])
|
||||
def catch_all(path):
|
||||
return {"path": path}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def _run_server(**kwargs):
|
||||
"""
|
||||
Internal function to run the Flask server.
|
||||
This is used to avoid issues with multiprocessing and Flask's reloader.
|
||||
"""
|
||||
signal.signal(signal.SIGINT, signal.SIG_IGN) # Ignore SIGINT in worker processes
|
||||
setproctitle.setproctitle(multiprocessing.current_process().name)
|
||||
app = agentops_local_server()
|
||||
app.run(**kwargs)
|
||||
|
||||
|
||||
class AgentOpsServerManager:
|
||||
def __init__(self, daemon: bool = True, port: int | None = None):
|
||||
self.server_process: multiprocessing.Process | None = None
|
||||
self.server_port = port
|
||||
self.daemon = daemon
|
||||
logger.info("AgentOpsServerManager initialized.")
|
||||
|
||||
def _find_available_port(self) -> int:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(("", 0))
|
||||
return s.getsockname()[1]
|
||||
|
||||
def start(self):
|
||||
if self.server_process and self.server_process.is_alive():
|
||||
logger.warning("AgentOps server process appears to be already running.")
|
||||
return
|
||||
|
||||
if self.server_port is None:
|
||||
self.server_port = self._find_available_port()
|
||||
|
||||
logger.info(f"Starting AgentOps local server on port {self.server_port}...")
|
||||
|
||||
self.server_process = multiprocessing.Process(
|
||||
target=_run_server,
|
||||
kwargs={"host": "127.0.0.1", "port": self.server_port, "use_reloader": False, "debug": False},
|
||||
daemon=self.daemon,
|
||||
name="AgentLightning-AgentOpsServer",
|
||||
)
|
||||
self.server_process.start()
|
||||
logger.info(
|
||||
f"AgentOps local server process (PID: {self.server_process.pid}) started, targeting port {self.server_port}."
|
||||
)
|
||||
time.sleep(0.5) # Brief wait for server to start up
|
||||
if not self.server_process.is_alive():
|
||||
logger.error(f"AgentOps local server failed to start or exited prematurely.")
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
if self.server_process and self.server_process.is_alive():
|
||||
return True
|
||||
return False
|
||||
|
||||
def stop(self):
|
||||
if self.is_alive():
|
||||
logger.info(f"Stopping AgentOps local server (PID: {self.server_process.pid})...")
|
||||
self.server_process.terminate() # Send SIGTERM
|
||||
self.server_process.join(timeout=5) # Wait for clean exit
|
||||
if self.server_process.is_alive():
|
||||
logger.warning(
|
||||
f"AgentOps server (PID: {self.server_process.pid}) did not terminate gracefully, killing..."
|
||||
)
|
||||
self.server_process.kill() # Force kill
|
||||
self.server_process.join(timeout=10) # Wait for kill
|
||||
self.server_process = None
|
||||
logger.info(f"AgentOps local server stopped.")
|
||||
else:
|
||||
logger.info("AgentOps local server was not running or already stopped.")
|
||||
|
||||
def get_port(self) -> int | None:
|
||||
# Check liveness again in case it died since start()
|
||||
if self.is_alive() and self.server_port is not None:
|
||||
return self.server_port
|
||||
# If called after server stopped or failed, port might be stale or None
|
||||
if self.server_port is not None and (self.server_process is None or not self.server_process.is_alive()):
|
||||
logger.warning(
|
||||
f"AgentOps server port {self.server_port} is stored, but server process is not alive. Returning stored port."
|
||||
)
|
||||
return self.server_port
|
||||
@@ -1,36 +0,0 @@
|
||||
from typing import Dict, Any
|
||||
from agentops.integration.callbacks.langchain import LangchainCallbackHandler
|
||||
from agentops import instrumentation
|
||||
|
||||
|
||||
original_on_chain_start = LangchainCallbackHandler.on_chain_start
|
||||
langgraph_entry = None
|
||||
|
||||
|
||||
def on_chain_start(self, serialized: Dict[str, Any], inputs: Dict[str, Any], **kwargs: Any) -> None:
|
||||
if "name" in kwargs:
|
||||
if serialized is None:
|
||||
serialized = {}
|
||||
serialized = serialized.copy()
|
||||
serialized["name"] = kwargs["name"]
|
||||
if "run_id" in kwargs:
|
||||
if serialized is None:
|
||||
serialized = {}
|
||||
serialized = serialized.copy()
|
||||
if "id" not in serialized:
|
||||
serialized["id"] = kwargs["run_id"]
|
||||
return original_on_chain_start(self, serialized, inputs, **kwargs)
|
||||
|
||||
|
||||
def instrument_agentops_langchain():
|
||||
global langgraph_entry
|
||||
langgraph_entry = instrumentation.AGENTIC_LIBRARIES.pop("langgraph", None)
|
||||
LangchainCallbackHandler.on_chain_start = on_chain_start
|
||||
|
||||
|
||||
def uninstrument_agentops_langchain():
|
||||
global langgraph_entry
|
||||
if langgraph_entry is not None:
|
||||
instrumentation.AGENTIC_LIBRARIES["langgraph"] = langgraph_entry
|
||||
langgraph_entry = None
|
||||
LangchainCallbackHandler.on_chain_start = original_on_chain_start
|
||||
@@ -1,26 +0,0 @@
|
||||
from typing import Optional, Any
|
||||
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
|
||||
# It's unclear whether or not this file is useful
|
||||
# It seems that LiteLLM owns its own telemetry from their own entrance
|
||||
# https://docs.litellm.ai/docs/observability/agentops_integration
|
||||
|
||||
original_set_attributes = OpenTelemetry.set_attributes
|
||||
|
||||
|
||||
def patched_set_attributes(self, span: Any, kwargs, response_obj: Optional[Any]):
|
||||
original_set_attributes(self, span, kwargs, response_obj)
|
||||
# Add custom attributes
|
||||
if response_obj.get("prompt_token_ids"):
|
||||
span.set_attribute("prompt_token_ids", list(response_obj.get("prompt_token_ids")))
|
||||
if response_obj.get("response_token_ids"):
|
||||
span.set_attribute("response_token_ids", list(response_obj.get("response_token_ids")[0]))
|
||||
|
||||
|
||||
def instrument_litellm():
|
||||
OpenTelemetry.set_attributes = patched_set_attributes
|
||||
|
||||
|
||||
def uninstrument_litellm():
|
||||
OpenTelemetry.set_attributes = original_set_attributes
|
||||
@@ -1,148 +0,0 @@
|
||||
# type: ignore
|
||||
|
||||
# https://github.com/volcengine/verl/blob/bd94bd61fe4193e56f2845dc794004afbef7f818/examples/ppo_trainer/naive_chat_scheduler.py
|
||||
# This file is part of VERL example. It should be included in the VERL package but it's not currently.
|
||||
|
||||
# Copyright 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import asyncio
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import torch
|
||||
from openai.types.chat.chat_completion import ChatCompletion
|
||||
from tensordict import TensorDict
|
||||
|
||||
from verl.protocol import DataProto
|
||||
from verl.workers.rollout.async_server import ChatCompletionScheduler
|
||||
|
||||
|
||||
class NaiveChatCompletionScheduler(ChatCompletionScheduler):
|
||||
"""
|
||||
A very naive implementation of ChatCompletionScheduler for demo purpose,
|
||||
only do single-turn chat completion.
|
||||
"""
|
||||
|
||||
async def generate_sequences(self, batch: DataProto, **sampling_params) -> DataProto:
|
||||
kwargs = dict(
|
||||
n=self.config.n,
|
||||
max_completion_tokens=self.config.response_length,
|
||||
temperature=self.config.temperature,
|
||||
top_p=self.config.top_p,
|
||||
)
|
||||
|
||||
do_sample = batch.meta_info.get("do_sample", True)
|
||||
is_validate = batch.meta_info.get("validate", False)
|
||||
if not do_sample or is_validate:
|
||||
kwargs["n"] = 1
|
||||
kwargs["temperature"] = 0
|
||||
|
||||
kwargs.update(sampling_params)
|
||||
print(f"[NaiveChatCompletionScheduler] generate_sequences sampling params: {kwargs}")
|
||||
|
||||
async def callback(completions: ChatCompletion, info: Dict[str, Any], exception: Exception):
|
||||
assert exception is None, f"exception: {exception}"
|
||||
conversation, batch_conversations, batch_index = (
|
||||
info["conversation"],
|
||||
info["batch_conversations"],
|
||||
info["batch_index"],
|
||||
)
|
||||
|
||||
conversations = []
|
||||
for choice in completions.choices:
|
||||
chat = conversation.copy()
|
||||
chat.append({"role": choice.message.role, "content": choice.message.content})
|
||||
conversations.append(chat)
|
||||
batch_conversations[batch_index] = conversations
|
||||
|
||||
# NOTE: we can call tools and resubmit chat completions here.
|
||||
# call_tools(completions, info)
|
||||
# await self.submit_chat_completions(callback2, ...)
|
||||
|
||||
# TODO: we may need to control max concurrent requests here, or it will harm prefix cache hit rate.
|
||||
tasks, batch_conversations = [], [None] * len(batch)
|
||||
for batch_index, conversation in enumerate(batch.non_tensor_batch["raw_prompt"]):
|
||||
# raw_prompt: [{"role": "user", "content": ""}, ["role": "assistant", "content"], ...]
|
||||
tasks.append(
|
||||
asyncio.create_task(
|
||||
self.submit_chat_completions(
|
||||
callback=callback,
|
||||
callback_additional_info={
|
||||
"batch_conversations": batch_conversations,
|
||||
"batch_index": batch_index,
|
||||
"conversation": list(conversation),
|
||||
},
|
||||
model=self.model_name,
|
||||
messages=conversation.tolist(),
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
)
|
||||
await asyncio.gather(*tasks)
|
||||
print("[NaiveChatCompletionScheduler] generate_sequences done")
|
||||
|
||||
return self._postprocess(batch, batch_conversations, kwargs["n"])
|
||||
|
||||
def _postprocess(
|
||||
self, batch: DataProto, batch_conversations: List[List[List[Dict[str, str]]]], n: int
|
||||
) -> DataProto:
|
||||
# NOTE: consistent with batch version of generate_sequences in vllm_rollout_spmd.py
|
||||
# prompts: left pad
|
||||
# responses: right pad
|
||||
# input_ids: prompt + response
|
||||
# attention_mask: [0,0,0,0,1,1,1,1, | 1,1,1,0,0,0,0,0]
|
||||
# position_ids: [0,0,0,0,0,1,2,3, | 4,5,6,7,8,9,10,11]
|
||||
|
||||
# prompts: [prompt] from input dataset
|
||||
prompts = [
|
||||
self.tokenizer.apply_chat_template(prompt, add_generation_prompt=True, tokenize=False)
|
||||
for prompt in batch.non_tensor_batch["raw_prompt"]
|
||||
]
|
||||
|
||||
# flatten batch_conversations if n > 1
|
||||
assert len(batch_conversations) == len(prompts)
|
||||
batch_conversations = [conversation for conversations in batch_conversations for conversation in conversations]
|
||||
assert len(batch_conversations) == len(prompts) * n
|
||||
|
||||
# sequences: [prompt + response]
|
||||
sequences = [
|
||||
self.tokenizer.apply_chat_template(conversation, add_generation_prompt=False, tokenize=False)
|
||||
for conversation in batch_conversations
|
||||
]
|
||||
|
||||
# responses: [response]
|
||||
# TODO: mask out tools calling tokens?
|
||||
responses = [sequence[len(prompts[i // n]) :] for i, sequence in enumerate(sequences)]
|
||||
|
||||
prompts = self.tokenizer(prompts, return_tensors="pt", padding="longest", padding_side="left")
|
||||
responses = self.tokenizer(responses, return_tensors="pt", padding="longest", padding_side="right")
|
||||
if n > 1:
|
||||
prompts["input_ids"] = prompts["input_ids"].repeat_interleave(n, dim=0)
|
||||
prompts["attention_mask"] = prompts["attention_mask"].repeat_interleave(n, dim=0)
|
||||
|
||||
input_ids = torch.cat([prompts["input_ids"], responses["input_ids"]], dim=1)
|
||||
attention_mask = torch.cat([prompts["attention_mask"], responses["attention_mask"]], dim=1)
|
||||
position_ids = (attention_mask.cumsum(dim=1) - 1) * attention_mask
|
||||
|
||||
batch = TensorDict(
|
||||
{
|
||||
"prompts": prompts["input_ids"],
|
||||
"responses": responses["input_ids"],
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"position_ids": position_ids,
|
||||
},
|
||||
batch_size=len(input_ids),
|
||||
)
|
||||
|
||||
return DataProto(batch=batch)
|
||||
@@ -1,69 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import warnings
|
||||
from typing import List
|
||||
|
||||
from vllm.entrypoints.openai.protocol import ChatCompletionResponse
|
||||
import vllm.entrypoints.openai.protocol
|
||||
from vllm.entrypoints.openai.serving_chat import OpenAIServingChat
|
||||
|
||||
|
||||
class ChatCompletionResponsePatched(ChatCompletionResponse):
|
||||
prompt_token_ids: List[int] | None = None
|
||||
response_token_ids: List[int] | None = None
|
||||
|
||||
|
||||
original_chat_completion_full_generator = OpenAIServingChat.chat_completion_full_generator
|
||||
|
||||
|
||||
async def chat_completion_full_generator(
|
||||
self,
|
||||
request,
|
||||
result_generator,
|
||||
request_id: str,
|
||||
model_name: str,
|
||||
conversation,
|
||||
tokenizer,
|
||||
request_metadata,
|
||||
):
|
||||
prompt_token_ids: List[int] | None = None
|
||||
response_token_ids: List[List[int]] | None = None
|
||||
|
||||
async def _generate_inceptor():
|
||||
nonlocal prompt_token_ids, response_token_ids
|
||||
async for res in result_generator:
|
||||
yield res
|
||||
prompt_token_ids = res.prompt_token_ids
|
||||
response_token_ids = [output.token_ids for output in res.outputs]
|
||||
|
||||
response = await original_chat_completion_full_generator(
|
||||
self,
|
||||
request,
|
||||
_generate_inceptor(),
|
||||
request_id,
|
||||
model_name,
|
||||
conversation,
|
||||
tokenizer,
|
||||
request_metadata,
|
||||
)
|
||||
response = response.model_copy(
|
||||
update={
|
||||
"prompt_token_ids": prompt_token_ids,
|
||||
"response_token_ids": response_token_ids,
|
||||
}
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
def instrument_vllm():
|
||||
if vllm.entrypoints.openai.protocol.ChatCompletionResponse is ChatCompletionResponsePatched:
|
||||
warnings.warn("vllm is already instrumented. Skip the instrumentation.")
|
||||
return
|
||||
|
||||
vllm.entrypoints.openai.protocol.ChatCompletionResponse = ChatCompletionResponsePatched
|
||||
OpenAIServingChat.chat_completion_full_generator = chat_completion_full_generator
|
||||
|
||||
|
||||
def uninstrument_vllm():
|
||||
OpenAIServingChat.chat_completion_full_generator = original_chat_completion_full_generator
|
||||
@@ -1,178 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import weakref
|
||||
from typing import Any, List, Dict, Union, Optional, TYPE_CHECKING
|
||||
|
||||
from .types import NamedResources, Rollout, Task, TaskInput, Triplet, RolloutRawResult
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .trainer import Trainer
|
||||
from .runner import AgentRunner
|
||||
from .tracer import BaseTracer
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LitAgent:
|
||||
"""Base class for the training and validation logic of an agent.
|
||||
|
||||
Developers should subclass this class and implement the rollout methods
|
||||
to define the agent's behavior for a single task. The agent's logic
|
||||
is completely decoupled from the server communication and training
|
||||
infrastructure.
|
||||
"""
|
||||
|
||||
def __init__(self, *, trained_agents: Optional[str] = None) -> None: # FIXME: str | None won't work for cli
|
||||
"""
|
||||
Initialize the LitAgent.
|
||||
|
||||
Args:
|
||||
trained_agents: Optional string representing the trained agents.
|
||||
This can be used to track which agents have been trained by this instance.
|
||||
"""
|
||||
self.trained_agents = trained_agents
|
||||
self._trainer_ref: weakref.ReferenceType[Trainer] | None = None
|
||||
self._runner_ref: weakref.ReferenceType[AgentRunner] | None = None
|
||||
|
||||
def set_trainer(self, trainer: Trainer) -> None:
|
||||
"""
|
||||
Set the trainer for this agent.
|
||||
|
||||
Args:
|
||||
trainer: The Trainer instance that will handle training and validation.
|
||||
"""
|
||||
self._trainer_ref = weakref.ref(trainer)
|
||||
|
||||
@property
|
||||
def trainer(self) -> Trainer:
|
||||
"""
|
||||
Get the trainer for this agent.
|
||||
|
||||
Returns:
|
||||
The Trainer instance associated with this agent.
|
||||
"""
|
||||
if self._trainer_ref is None:
|
||||
raise ValueError("Trainer has not been set for this agent.")
|
||||
trainer = self._trainer_ref()
|
||||
if trainer is None:
|
||||
raise ValueError("Trainer reference is no longer valid (object has been garbage collected).")
|
||||
return trainer
|
||||
|
||||
@property
|
||||
def tracer(self) -> BaseTracer:
|
||||
"""
|
||||
Get the tracer for this agent.
|
||||
|
||||
Returns:
|
||||
The BaseTracer instance associated with this agent.
|
||||
"""
|
||||
return self.trainer.tracer
|
||||
|
||||
def set_runner(self, runner: AgentRunner) -> None:
|
||||
"""
|
||||
Set the runner for this agent.
|
||||
|
||||
Args:
|
||||
runner: The AgentRunner instance that will handle the execution of rollouts.
|
||||
"""
|
||||
self._runner_ref = weakref.ref(runner)
|
||||
|
||||
@property
|
||||
def runner(self) -> AgentRunner:
|
||||
"""
|
||||
Get the runner for this agent.
|
||||
|
||||
Returns:
|
||||
The AgentRunner instance associated with this agent.
|
||||
"""
|
||||
if self._runner_ref is None:
|
||||
raise ValueError("Runner has not been set for this agent.")
|
||||
runner = self._runner_ref()
|
||||
if runner is None:
|
||||
raise ValueError("Runner reference is no longer valid (object has been garbage collected).")
|
||||
return runner
|
||||
|
||||
def training_rollout(self, task: TaskInput, rollout_id: str, resources: NamedResources) -> RolloutRawResult:
|
||||
"""Defines the agent's behavior for a single training task.
|
||||
|
||||
This method should contain the logic for how the agent processes an
|
||||
input, uses the provided resources (like LLMs or prompts), and
|
||||
produces a result.
|
||||
|
||||
Args:
|
||||
task: The task object received from the server, containing the
|
||||
input data and metadata.
|
||||
rollout_id: A unique identifier for the rollout, used for tracking
|
||||
and reporting purposes.
|
||||
resources: A dictionary of named resources (e.g., LLMs, prompt
|
||||
templates) for the agent to use.
|
||||
|
||||
Returns:
|
||||
The result of the rollout, which can be one of:
|
||||
- None. The tracing should be handled by the agent runner.
|
||||
- A float representing the final reward.
|
||||
- A list of `Triplet` objects for detailed, step-by-step feedback.
|
||||
- A list of `ReadableSpan` objects for OpenTelemetry tracing.
|
||||
- A list of dictionaries for any trace spans.
|
||||
- A complete `Rollout` object for full control over reporting.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement the `training_rollout` method.")
|
||||
|
||||
def validation_rollout(self, task: TaskInput, rollout_id: str, resources: NamedResources) -> RolloutRawResult:
|
||||
"""Defines the agent's behavior for a single validation task.
|
||||
|
||||
By default, this method redirects to `training_rollout`. Override it
|
||||
if the agent should behave differently during validation.
|
||||
|
||||
Args:
|
||||
task: The task object received from the server, containing the
|
||||
input data and metadata.
|
||||
rollout_id: A unique identifier for the validation rollout,
|
||||
used for tracking and reporting purposes.
|
||||
resources: A dictionary of named resources for the agent to use.
|
||||
|
||||
Returns:
|
||||
The result of the validation rollout. See `training_rollout` for
|
||||
possible return types.
|
||||
"""
|
||||
return self.training_rollout(task, rollout_id, resources)
|
||||
|
||||
async def training_rollout_async(
|
||||
self, task: TaskInput, rollout_id: str, resources: NamedResources
|
||||
) -> RolloutRawResult:
|
||||
"""Asynchronous version of `training_rollout`.
|
||||
|
||||
This method should be implemented by agents that perform asynchronous
|
||||
operations (e.g., non-blocking I/O, concurrent API calls).
|
||||
|
||||
Args:
|
||||
task: The task object received from the server.
|
||||
rollout_id: A unique identifier for the training rollout,
|
||||
used for tracking and reporting purposes.
|
||||
resources: A dictionary of named resources for the agent to use.
|
||||
|
||||
Returns:
|
||||
The result of the asynchronous training rollout.
|
||||
"""
|
||||
raise NotImplementedError("Async agents must implement the `training_rollout_async` method.")
|
||||
|
||||
async def validation_rollout_async(
|
||||
self, task: TaskInput, rollout_id: str, resources: NamedResources
|
||||
) -> RolloutRawResult:
|
||||
"""Asynchronous version of `validation_rollout`.
|
||||
|
||||
By default, this method redirects to `training_rollout_async`.
|
||||
Override it for different asynchronous validation behavior.
|
||||
|
||||
Args:
|
||||
task: The task object received from the server.
|
||||
rollout_id: A unique identifier for the validation rollout,
|
||||
used for tracking and reporting purposes.
|
||||
resources: A dictionary of named resources for the agent to use.
|
||||
|
||||
Returns:
|
||||
The result of the asynchronous validation rollout.
|
||||
"""
|
||||
return await self.training_rollout_async(task, rollout_id, resources)
|
||||
@@ -1,16 +0,0 @@
|
||||
import logging
|
||||
|
||||
|
||||
def configure_logger(level: int = logging.INFO, name: str = "agentlightning") -> logging.Logger:
|
||||
logger = logging.getLogger(name)
|
||||
logger.handlers.clear() # clear existing handlers
|
||||
|
||||
# log to stdout
|
||||
handler = logging.StreamHandler()
|
||||
handler.setLevel(level)
|
||||
formatter = logging.Formatter("%(asctime)s [%(levelname)s] (Process-%(process)d %(name)s) %(message)s")
|
||||
handler.setFormatter(formatter)
|
||||
logger.addHandler(handler)
|
||||
logger.setLevel(level)
|
||||
logger.propagate = False # prevent double logging
|
||||
return logger
|
||||
@@ -1,66 +0,0 @@
|
||||
import asyncio
|
||||
import inspect
|
||||
import warnings
|
||||
from typing import TypedDict, Optional
|
||||
|
||||
from agentops.sdk.decorators import operation
|
||||
|
||||
|
||||
class RewardSpanData(TypedDict):
|
||||
type: "reward"
|
||||
value: Optional[float]
|
||||
|
||||
|
||||
def reward(fn: callable) -> callable:
|
||||
"""
|
||||
A decorator to wrap a function that computes rewards.
|
||||
It will automatically handle the input and output of the function.
|
||||
"""
|
||||
|
||||
def wrap_result(result: Optional[float]) -> RewardSpanData:
|
||||
"""
|
||||
Wrap the result of the function in a dict.
|
||||
"""
|
||||
if result is None:
|
||||
return {"type": "reward", "value": None}
|
||||
if not isinstance(result, (float, int)):
|
||||
warnings.warn(f"Reward is ignored because it is not a number: {result}")
|
||||
return {"type": "reward", "value": None}
|
||||
return {"type": "reward", "value": float(result)}
|
||||
|
||||
# Check if the function is async
|
||||
is_async = asyncio.iscoroutinefunction(fn) or inspect.iscoroutinefunction(fn)
|
||||
|
||||
if is_async:
|
||||
|
||||
async def wrapper_async(*args, **kwargs):
|
||||
result: Optional[float] = None
|
||||
|
||||
@operation
|
||||
async def agentops_reward_operation() -> RewardSpanData:
|
||||
# The reward function we are interested in tracing
|
||||
# It takes zero inputs and return a formatted dict
|
||||
nonlocal result
|
||||
result = await fn(*args, **kwargs)
|
||||
return wrap_result(result)
|
||||
|
||||
await agentops_reward_operation()
|
||||
return result
|
||||
|
||||
return wrapper_async
|
||||
|
||||
else:
|
||||
|
||||
def wrapper(*args, **kwargs):
|
||||
result: Optional[float] = None
|
||||
|
||||
@operation
|
||||
def agentops_reward_operation() -> RewardSpanData:
|
||||
nonlocal result
|
||||
result = fn(*args, **kwargs)
|
||||
return wrap_result(result)
|
||||
|
||||
agentops_reward_operation()
|
||||
return result
|
||||
|
||||
return wrapper
|
||||
@@ -1,253 +0,0 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from contextlib import nullcontext
|
||||
from typing import List, Optional, Union, Dict, Any
|
||||
|
||||
import agentops
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from .client import AgentLightningClient
|
||||
from .litagent import LitAgent
|
||||
from .types import Rollout, Task, Triplet, RolloutRawResult
|
||||
from .types import ParallelWorkerBase
|
||||
from .tracer.base import BaseTracer
|
||||
from .tracer import TripletExporter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AgentRunner(ParallelWorkerBase):
|
||||
"""Manages the agent's execution loop and integrates with AgentOps.
|
||||
|
||||
This class orchestrates the interaction between the agent (`LitAgent`) and
|
||||
the server (`AgentLightningClient`). It handles polling for tasks, executing
|
||||
the agent's logic, and reporting results back to the server. If enabled,
|
||||
it will also automatically trace each rollout using AgentOps.
|
||||
|
||||
Attributes:
|
||||
agent: The `LitAgent` instance containing the agent's logic.
|
||||
client: The `AgentLightningClient` for server communication.
|
||||
tracer: The tracer instance for this runner/worker.
|
||||
worker_id: An optional identifier for the worker process.
|
||||
max_tasks: The maximum number of tasks to process before stopping.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
agent: LitAgent,
|
||||
client: AgentLightningClient,
|
||||
tracer: BaseTracer,
|
||||
triplet_exporter: TripletExporter,
|
||||
worker_id: Optional[int] = None,
|
||||
max_tasks: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.agent = agent
|
||||
self.client = client
|
||||
self.tracer = tracer
|
||||
self.triplet_exporter = triplet_exporter
|
||||
|
||||
# Worker-specific attributes
|
||||
self.worker_id = worker_id
|
||||
self.max_tasks = max_tasks
|
||||
|
||||
def _log_prefix(self, rollout_id: Optional[str] = None) -> str:
|
||||
"""Generates a standardized log prefix for the current worker."""
|
||||
if self.worker_id is not None:
|
||||
if rollout_id:
|
||||
return f"[Worker {self.worker_id} | Rollout {rollout_id}]"
|
||||
else:
|
||||
return f"[Worker {self.worker_id}]"
|
||||
if rollout_id:
|
||||
return f"[Rollout {rollout_id}]"
|
||||
return "[Default Worker]"
|
||||
|
||||
def _to_rollout_object(
|
||||
self,
|
||||
result: RolloutRawResult,
|
||||
rollout_id: str,
|
||||
) -> Rollout:
|
||||
"""Standardizes the agent's return value into a Rollout object.
|
||||
|
||||
Args:
|
||||
result: The output from the agent's rollout method.
|
||||
rollout_id: The unique identifier for the current task.
|
||||
|
||||
Returns:
|
||||
A standardized `Rollout` object for reporting to the server.
|
||||
"""
|
||||
trace: Any = None
|
||||
final_reward: Optional[float] = None
|
||||
triplets: Optional[List[Triplet]] = None
|
||||
trace_spans: Optional[List[ReadableSpan]] = None
|
||||
|
||||
# Handle different types of results from the agent
|
||||
# Case 1: result is a float (final reward)
|
||||
if isinstance(result, float):
|
||||
final_reward = result
|
||||
# Case 2: result is a list of Triplets
|
||||
if isinstance(result, list) and all(isinstance(t, Triplet) for t in result):
|
||||
triplets = result # type: ignore
|
||||
# Case 3: result is a list of ReadableSpan (OpenTelemetry spans)
|
||||
if isinstance(result, list) and all(isinstance(t, ReadableSpan) for t in result):
|
||||
trace_spans = result # type: ignore
|
||||
trace = [json.loads(readable_span.to_json()) for readable_span in trace_spans] # type: ignore
|
||||
# Case 4: result is a list of dict (trace JSON)
|
||||
if isinstance(result, list) and all(isinstance(t, dict) for t in result):
|
||||
trace = result
|
||||
# Case 5: result is a Rollout object
|
||||
if isinstance(result, Rollout):
|
||||
final_reward = result.final_reward
|
||||
triplets = result.triplets
|
||||
trace = result.trace
|
||||
|
||||
# If the agent has tracing enabled, use the tracer's last trace if not already set
|
||||
if self.tracer and (trace is None or trace_spans is None):
|
||||
spans = self.tracer.get_last_trace()
|
||||
if spans:
|
||||
trace = [json.loads(readable_span.to_json()) for readable_span in spans]
|
||||
trace_spans = spans
|
||||
|
||||
# Always extract triplets from the trace using TripletExporter
|
||||
if trace_spans:
|
||||
triplets = self.triplet_exporter.export(trace_spans)
|
||||
|
||||
# If the agent has triplets, use the last one for final reward if not set
|
||||
if triplets and triplets[-1].reward is not None and final_reward is None:
|
||||
final_reward = triplets[-1].reward
|
||||
|
||||
# Create the Rollout object with standardized fields
|
||||
result_dict: Dict[str, Any] = {
|
||||
"rollout_id": rollout_id,
|
||||
}
|
||||
if final_reward is not None:
|
||||
result_dict["final_reward"] = final_reward
|
||||
if triplets is not None:
|
||||
result_dict["triplets"] = triplets
|
||||
if trace is not None:
|
||||
result_dict["trace"] = trace
|
||||
|
||||
if isinstance(result, Rollout):
|
||||
return result.model_copy(update=result_dict)
|
||||
return Rollout(**result_dict)
|
||||
|
||||
def run(self) -> bool:
|
||||
"""Poll the task and rollout once synchronously."""
|
||||
self.agent.set_runner(self) # Ensure the agent has a reference to this runner
|
||||
|
||||
task = self.client.poll_next_task()
|
||||
if task is None:
|
||||
logger.info(f"{self._log_prefix()} Poll returned no task. Exiting.")
|
||||
return False
|
||||
rollout_id = task.rollout_id
|
||||
|
||||
resources_id = task.resources_id
|
||||
resources_update = None
|
||||
if resources_id:
|
||||
resources_update = self.client.get_resources_by_id(resources_id)
|
||||
else:
|
||||
logger.debug(f"{self._log_prefix(rollout_id)} No 'resources_id'. Fetching latest resources.")
|
||||
resources_update = self.client.get_latest_resources()
|
||||
if not resources_update:
|
||||
logger.error(f"{self._log_prefix(rollout_id)} Failed to fetch resources. Skipping.")
|
||||
return False
|
||||
|
||||
rollout_obj = Rollout(rollout_id=task.rollout_id) # Default empty rollout
|
||||
|
||||
try:
|
||||
with self.tracer.trace_context(name=f"rollout_{rollout_id}"):
|
||||
start_time = time.time()
|
||||
rollout_method = self.agent.training_rollout if task.mode == "train" else self.agent.validation_rollout
|
||||
# Pass the task input, not the whole task object
|
||||
result = rollout_method(task.input, task.rollout_id, resources_update.resources)
|
||||
rollout_obj = self._to_rollout_object(result, task.rollout_id)
|
||||
end_time = time.time()
|
||||
logger.info(
|
||||
f"{self._log_prefix(rollout_id)} Completed in "
|
||||
f"{end_time - start_time:.2f}s. Triplet length: "
|
||||
f"{len(rollout_obj.triplets) if rollout_obj.triplets is not None else 'N/A'}. "
|
||||
f"Reward: {rollout_obj.final_reward}"
|
||||
)
|
||||
|
||||
except Exception:
|
||||
logger.exception(f"{self._log_prefix(rollout_id)} Exception during rollout.")
|
||||
finally:
|
||||
self.client.post_rollout(rollout_obj)
|
||||
|
||||
return True
|
||||
|
||||
def iter(self) -> int:
|
||||
"""Executes the synchronous polling and rollout loop."""
|
||||
num_tasks_processed = 0
|
||||
logger.info(f"{self._log_prefix()} Started sync rollouts (max: {self.max_tasks or 'unlimited'}).")
|
||||
|
||||
while self.max_tasks is None or num_tasks_processed < self.max_tasks:
|
||||
if self.run():
|
||||
num_tasks_processed += 1
|
||||
|
||||
if num_tasks_processed % 10 == 0 or num_tasks_processed == 1:
|
||||
logger.info(f"{self._log_prefix()} Progress: {num_tasks_processed}/{self.max_tasks or 'unlimited'}")
|
||||
|
||||
logger.info(f"{self._log_prefix()} Finished sync rollouts. Processed {num_tasks_processed} tasks.")
|
||||
return num_tasks_processed
|
||||
|
||||
async def run_async(self) -> bool:
|
||||
"""Poll the task and rollout once."""
|
||||
self.agent.set_runner(self) # Ensure the agent has a reference to this runner
|
||||
|
||||
task = await self.client.poll_next_task_async()
|
||||
if task is None:
|
||||
logger.info(f"{self._log_prefix()} Poll returned no task. Exiting.")
|
||||
return False
|
||||
rollout_id = task.rollout_id
|
||||
|
||||
resources_id = task.resources_id
|
||||
resources_update = None
|
||||
if resources_id:
|
||||
resources_update = await self.client.get_resources_by_id_async(resources_id)
|
||||
else:
|
||||
logger.debug(f"{self._log_prefix(rollout_id)} No 'resources_id'. Fetching latest resources.")
|
||||
resources_update = await self.client.get_latest_resources_async()
|
||||
if not resources_update:
|
||||
logger.error(f"{self._log_prefix(rollout_id)} Failed to fetch resources. Skipping.")
|
||||
return False
|
||||
|
||||
rollout_obj = Rollout(rollout_id=task.rollout_id) # Default empty rollout
|
||||
|
||||
try:
|
||||
with self.tracer.trace_context(name=f"rollout_{rollout_id}"):
|
||||
start_time = time.time()
|
||||
rollout_method = (
|
||||
self.agent.training_rollout_async if task.mode == "train" else self.agent.validation_rollout_async
|
||||
)
|
||||
# Pass the task input, not the whole task object
|
||||
result = await rollout_method(task.input, task.rollout_id, resources_update.resources)
|
||||
rollout_obj = self._to_rollout_object(result, task.rollout_id)
|
||||
end_time = time.time()
|
||||
logger.info(
|
||||
f"{self._log_prefix(rollout_id)} Completed in "
|
||||
f"{end_time - start_time:.2f}s. Reward: {rollout_obj.final_reward}"
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(f"{self._log_prefix(rollout_id)} Exception during rollout.")
|
||||
finally:
|
||||
await self.client.post_rollout_async(rollout_obj)
|
||||
|
||||
return True
|
||||
|
||||
async def iter_async(self) -> int:
|
||||
"""Executes the asynchronous polling and rollout loop."""
|
||||
num_tasks_processed = 0
|
||||
logger.info(f"{self._log_prefix()} Started async rollouts (max: {self.max_tasks or 'unlimited'}).")
|
||||
|
||||
while self.max_tasks is None or num_tasks_processed < self.max_tasks:
|
||||
if await self.run_async():
|
||||
num_tasks_processed += 1
|
||||
|
||||
if num_tasks_processed % 10 == 0 or num_tasks_processed == 1:
|
||||
logger.info(f"{self._log_prefix()} Progress: {num_tasks_processed}/{self.max_tasks or 'unlimited'}")
|
||||
logger.info(f"{self._log_prefix()} Finished async rollouts. Processed {num_tasks_processed} tasks.")
|
||||
return num_tasks_processed
|
||||
@@ -0,0 +1,183 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Shared Pydantic schemas for Agent Lightning."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class Event(BaseModel):
|
||||
"""Single event in a trajectory.
|
||||
|
||||
Events are stored in insertion order per rollout. Position in the list
|
||||
is the identity — no separate event ID needed. Only two event types
|
||||
have well-known structure (model_request, reward). Everything else is
|
||||
opaque pass-through.
|
||||
"""
|
||||
|
||||
event_type: str # "model_request", "reward", or any user-defined string
|
||||
rollout_id: str
|
||||
attempt_id: str
|
||||
timestamp: float # assigned by store at write time
|
||||
data: dict[str, Any] # event-type-specific payload
|
||||
|
||||
|
||||
class EventCreate(BaseModel):
|
||||
"""Input for appending a user-defined event."""
|
||||
|
||||
event_type: str
|
||||
data: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class ModelRequestData(BaseModel):
|
||||
"""Well-known structure for event_type='model_request'.
|
||||
|
||||
Created automatically by the Gateway on every proxied LLM call.
|
||||
Not enforced by the Store — this is a documentation/validation helper.
|
||||
"""
|
||||
|
||||
model: str
|
||||
model_version: int | None = None # training step of the serving model
|
||||
request: dict[str, Any] # original request body (messages, temperature, etc.)
|
||||
adjusted_params: dict[str, Any] | None = None # only if param adjustment changed anything
|
||||
response: dict[str, Any] # full response body
|
||||
latency_ms: float | None = None
|
||||
http_status: int | None = None
|
||||
status: str = "ok" # "ok" or "error"
|
||||
retry_count: int = 0
|
||||
usage: dict[str, Any] | None = None
|
||||
finish_reason: str | None = None
|
||||
|
||||
|
||||
class RewardData(BaseModel):
|
||||
"""Well-known structure for event_type='reward'.
|
||||
|
||||
Reported by the environment, evaluator, or runner.
|
||||
Not enforced by the Store — this is a documentation/validation helper.
|
||||
"""
|
||||
|
||||
value: float # scalar reward (required)
|
||||
message: str | None = None # optional human-readable explanation
|
||||
source: str | None = None # e.g. "agent" for explicit evaluator output, "fallback" for system fill-in
|
||||
reason: str | None = None # optional machine-readable explanation
|
||||
|
||||
|
||||
class Model(BaseModel):
|
||||
"""A registered model inference endpoint. Keyed by (model, endpoint)."""
|
||||
|
||||
model: str
|
||||
endpoint: str
|
||||
version: int = 0
|
||||
|
||||
|
||||
class RolloutState(StrEnum):
|
||||
"""Rollout lifecycle state values. Terminal states are final — no transitions out."""
|
||||
|
||||
QUEUING = "queuing"
|
||||
RUNNING = "running"
|
||||
SUCCEEDED = "succeeded"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
# Valid state transitions (Store-enforced).
|
||||
VALID_TRANSITIONS: dict[RolloutState, set[RolloutState]] = {
|
||||
RolloutState.QUEUING: {RolloutState.RUNNING, RolloutState.FAILED},
|
||||
RolloutState.RUNNING: {RolloutState.SUCCEEDED, RolloutState.FAILED},
|
||||
# Terminal states — no transitions out.
|
||||
RolloutState.SUCCEEDED: set(),
|
||||
RolloutState.FAILED: set(),
|
||||
}
|
||||
|
||||
TERMINAL_STATES: frozenset[RolloutState] = frozenset(
|
||||
{
|
||||
RolloutState.SUCCEEDED,
|
||||
RolloutState.FAILED,
|
||||
}
|
||||
)
|
||||
|
||||
DEFAULT_ATTEMPT_ID = "0"
|
||||
|
||||
|
||||
class RolloutLocalConfig(BaseModel):
|
||||
"""Local runner config for a rollout."""
|
||||
|
||||
agent_class: str | None = None
|
||||
env_map: dict[str, str] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class RolloutK8sConfig(BaseModel):
|
||||
"""K8s runner config for a rollout."""
|
||||
|
||||
job_template: str | None = None
|
||||
|
||||
|
||||
class RolloutConfig(BaseModel):
|
||||
"""Controller-facing rollout config."""
|
||||
|
||||
timeout_seconds: int = 3600
|
||||
local: RolloutLocalConfig | None = None
|
||||
k8s: RolloutK8sConfig | None = None
|
||||
|
||||
|
||||
class RolloutMetadata(BaseModel):
|
||||
"""Algorithm-facing batch context."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
batch_idx: int | None = None
|
||||
sample_idx_in_batch: int | None = None
|
||||
|
||||
|
||||
class RolloutCreate(BaseModel):
|
||||
"""Input for creating a rollout."""
|
||||
|
||||
input: Any
|
||||
is_train: bool = True
|
||||
config: RolloutConfig | None = None
|
||||
metadata: RolloutMetadata | dict[str, Any] | None = None
|
||||
# A caller-supplied id makes rollout creation idempotent and safe to retry.
|
||||
rollout_id: str | None = None
|
||||
|
||||
|
||||
class RolloutLifecycleStatus(BaseModel):
|
||||
"""Controller-managed rollout lifecycle status."""
|
||||
|
||||
state: RolloutState = RolloutState.QUEUING
|
||||
k8s_job_name: str | None = None
|
||||
last_attempt_id: str | None = None
|
||||
error_message: str | None = None
|
||||
version: int = 1
|
||||
created_at: float
|
||||
updated_at: float
|
||||
|
||||
|
||||
class RolloutStatusPatch(BaseModel):
|
||||
"""Partial update for the nested rollout status object."""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
state: RolloutState | None = None
|
||||
k8s_job_name: str | None = None
|
||||
last_attempt_id: str | None = None
|
||||
error_message: str | None = None
|
||||
|
||||
|
||||
class RolloutPatch(BaseModel):
|
||||
"""Partial rollout update. Only nested status may be patched."""
|
||||
|
||||
status: RolloutStatusPatch | None = None
|
||||
|
||||
|
||||
class Rollout(BaseModel):
|
||||
"""Unit of work. Lifecycle managed by the K8s controller."""
|
||||
|
||||
rollout_id: str
|
||||
input: Any
|
||||
is_train: bool = True
|
||||
config: RolloutConfig
|
||||
metadata: RolloutMetadata = Field(default_factory=RolloutMetadata)
|
||||
status: RolloutLifecycleStatus
|
||||
@@ -1,353 +0,0 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
import threading
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, Dict, List, Optional, Literal
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, HTTPException, Path
|
||||
from pydantic import Field
|
||||
|
||||
from .types import (
|
||||
Rollout,
|
||||
Task,
|
||||
TaskIfAny,
|
||||
NamedResources,
|
||||
GenericResponse,
|
||||
ResourcesUpdate,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ServerDataStore:
|
||||
"""
|
||||
A centralized, thread-safe, async, in-memory data store for the server's state.
|
||||
This holds the task queue, versioned resources, and completed rollouts.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._task_queue: asyncio.Queue[Task] = asyncio.Queue()
|
||||
self._processing_tasks: Dict[str, Task] = {} # Currently processing tasks
|
||||
self._completed_rollouts: Dict[str, Rollout] = {}
|
||||
|
||||
# Store for versioned resources
|
||||
self._resource_versions: Dict[str, NamedResources] = {}
|
||||
self._latest_resources_id: Optional[str] = None
|
||||
|
||||
# Locks for thread-safe access
|
||||
self._results_lock = asyncio.Lock()
|
||||
self._resources_lock = asyncio.Lock()
|
||||
|
||||
async def add_task(
|
||||
self,
|
||||
sample: Any,
|
||||
mode: Literal["train", "val", "test"] | None = None,
|
||||
resources_id: str | None = None,
|
||||
metadata: Dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Adds a new task to the queue with specific metadata and returns its unique ID.
|
||||
"""
|
||||
rollout_id = f"rollout-{uuid.uuid4()}"
|
||||
task = Task(
|
||||
rollout_id=rollout_id,
|
||||
input=sample,
|
||||
mode=mode,
|
||||
resources_id=resources_id,
|
||||
create_time=time.time(),
|
||||
num_claims=0,
|
||||
metadata=metadata or {},
|
||||
)
|
||||
await self._task_queue.put(task)
|
||||
logger.info(f"Task queued: {rollout_id} (mode: {mode}, resources_id: {resources_id})")
|
||||
return rollout_id
|
||||
|
||||
async def get_next_task(self) -> Optional[Task]:
|
||||
"""
|
||||
Retrieves the next task from the queue without blocking.
|
||||
Returns None if the queue is empty.
|
||||
"""
|
||||
try:
|
||||
async with self._results_lock:
|
||||
task = self._task_queue.get_nowait()
|
||||
task = task.model_copy(
|
||||
update={
|
||||
"last_claim_time": time.time(),
|
||||
"num_claims": (task.num_claims or 0) + 1,
|
||||
}
|
||||
)
|
||||
self._processing_tasks[task.rollout_id] = task
|
||||
if task.num_claims == 1:
|
||||
logger.debug(f"Next task retrieved: {task.rollout_id}")
|
||||
else:
|
||||
logger.info(f"Task {task.rollout_id} re-claimed (attempt {task.num_claims})")
|
||||
return task
|
||||
except asyncio.QueueEmpty:
|
||||
return None
|
||||
|
||||
async def update_resources(self, update: ResourcesUpdate):
|
||||
"""
|
||||
Safely stores a new version of named resources and sets it as the latest.
|
||||
"""
|
||||
# TODO: evict old resources if necessary.
|
||||
async with self._resources_lock:
|
||||
self._resource_versions[update.resources_id] = update.resources
|
||||
self._latest_resources_id = update.resources_id
|
||||
logger.info(f"Resources updated. New version '{update.resources_id}' is now latest.")
|
||||
|
||||
async def get_resources_by_id(self, resources_id: str) -> Optional[ResourcesUpdate]:
|
||||
"""
|
||||
Safely retrieves a specific version of named resources by its ID.
|
||||
"""
|
||||
async with self._resources_lock:
|
||||
resources = self._resource_versions.get(resources_id)
|
||||
if resources:
|
||||
return ResourcesUpdate(resources_id=resources_id, resources=resources)
|
||||
return None
|
||||
|
||||
async def get_latest_resources(self) -> Optional[ResourcesUpdate]:
|
||||
"""
|
||||
Safely retrieves the latest version of named resources.
|
||||
"""
|
||||
if self._latest_resources_id:
|
||||
return await self.get_resources_by_id(self._latest_resources_id)
|
||||
return None
|
||||
|
||||
async def store_rollout(self, rollout: Rollout):
|
||||
"""
|
||||
Safely stores a completed rollout from a client.
|
||||
"""
|
||||
async with self._results_lock:
|
||||
self._processing_tasks.pop(rollout.rollout_id, None)
|
||||
self._completed_rollouts[rollout.rollout_id] = rollout
|
||||
logger.info(f"Rollout received and stored: {rollout.rollout_id}")
|
||||
|
||||
async def retrieve_rollout(self, rollout_id: str) -> Optional[Rollout]:
|
||||
"""
|
||||
Safely retrieves a single rollout by its ID, removing it from the store.
|
||||
"""
|
||||
async with self._results_lock:
|
||||
return self._completed_rollouts.pop(rollout_id, None)
|
||||
|
||||
async def retrieve_completed_rollouts(self) -> List[Rollout]:
|
||||
"""
|
||||
Retrieves all completed rollouts and clears the store.
|
||||
"""
|
||||
async with self._results_lock:
|
||||
rollouts = list(self._completed_rollouts.values())
|
||||
self._completed_rollouts.clear()
|
||||
return rollouts
|
||||
|
||||
def get_processing_tasks(self) -> Dict[str, Task]:
|
||||
"""Returns a copy of currently processing tasks for timeout checking."""
|
||||
return self._processing_tasks.copy()
|
||||
|
||||
async def requeue_task(self, task: Task):
|
||||
"""Requeues a task that has timed out and removes it from processing."""
|
||||
logger.warning(f"Requeuing task {task.rollout_id} after timeout (attempt {task.num_claims})")
|
||||
async with self._results_lock:
|
||||
# Remove from processing tasks
|
||||
self._processing_tasks.pop(task.rollout_id, None)
|
||||
self._task_queue.put_nowait(task)
|
||||
|
||||
|
||||
class AgentLightningServer:
|
||||
"""
|
||||
The main SDK class for developers to control the Agent Lightning Server.
|
||||
|
||||
This class manages the server lifecycle, task queueing, resources updates,
|
||||
and retrieval of results, providing a simple interface for the optimization logic.
|
||||
"""
|
||||
|
||||
def __init__(self, host: str = "127.0.0.1", port: int = 8000, task_timeout_seconds: float = 300.0):
|
||||
"""
|
||||
Initializes the server controller.
|
||||
|
||||
Args:
|
||||
host: The host to bind the server to.
|
||||
port: The port to bind the server to.
|
||||
task_timeout_seconds: Time in seconds after which a claimed task is considered stale and requeued.
|
||||
"""
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.endpoint = f"http://{host}:{port}"
|
||||
self._task_timeout_seconds = task_timeout_seconds
|
||||
|
||||
# Defer initialization and use event for cross-thread communication
|
||||
self._store: Optional[ServerDataStore] = None
|
||||
self.loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
self.startup_event = threading.Event()
|
||||
|
||||
# Create FastAPI app instance with a lifespan manager
|
||||
self._app = FastAPI(lifespan=self._lifespan)
|
||||
self._setup_routes()
|
||||
|
||||
self._uvicorn_config = uvicorn.Config(self._app, host=self.host, port=self.port, log_level="info")
|
||||
self._uvicorn_server = uvicorn.Server(self._uvicorn_config)
|
||||
|
||||
# --- ADDED: Lifespan context manager ---
|
||||
@asynccontextmanager
|
||||
async def _lifespan(self, app: FastAPI):
|
||||
"""
|
||||
Manages server startup and shutdown. This runs inside the server's event loop.
|
||||
"""
|
||||
logger.info("Server is starting up...")
|
||||
self.loop = asyncio.get_running_loop()
|
||||
self._store = ServerDataStore() # Initialize data store here
|
||||
self.startup_event.set() # Signal that the server is ready
|
||||
|
||||
yield
|
||||
|
||||
logger.info("Server is shutting down.")
|
||||
self._store = None
|
||||
self.startup_event.clear() # Clear the startup event
|
||||
self.loop = None
|
||||
|
||||
async def _check_and_requeue_stale_tasks(self):
|
||||
"""
|
||||
Check for stale tasks and requeue them. Called reactively during get_next_task.
|
||||
"""
|
||||
current_time = time.time()
|
||||
# Ensure store is initialized before checking
|
||||
if not self._store:
|
||||
return
|
||||
processing_tasks = self._store.get_processing_tasks()
|
||||
|
||||
for rollout_id, task in processing_tasks.items():
|
||||
if task.last_claim_time and current_time - task.last_claim_time > self._task_timeout_seconds:
|
||||
await self._store.requeue_task(task)
|
||||
logger.warning(
|
||||
f"Task {task.rollout_id} timed out after {self._task_timeout_seconds}s, requeued (attempt {task.num_claims})"
|
||||
)
|
||||
|
||||
def _setup_routes(self):
|
||||
"""Setup FastAPI routes."""
|
||||
|
||||
@self._app.get("/task", response_model=TaskIfAny)
|
||||
async def next_task() -> TaskIfAny:
|
||||
"""Endpoint for clients to poll for the next available task."""
|
||||
await self._check_and_requeue_stale_tasks()
|
||||
|
||||
if not self._store:
|
||||
return TaskIfAny(is_available=False)
|
||||
|
||||
task = await self._store.get_next_task()
|
||||
if task:
|
||||
logger.debug(f"Serving task {task.rollout_id} to a client.")
|
||||
return TaskIfAny(is_available=True, task=task)
|
||||
else:
|
||||
logger.debug("No task available for client.")
|
||||
return TaskIfAny(is_available=False)
|
||||
|
||||
@self._app.get("/resources/latest", response_model=ResourcesUpdate)
|
||||
async def fetch_latest_resources() -> ResourcesUpdate:
|
||||
"""Endpoint for clients to poll for the latest available resources."""
|
||||
if not self._store:
|
||||
raise HTTPException(status_code=503, detail="Server not fully initialized.")
|
||||
resources_update = await self._store.get_latest_resources()
|
||||
if not resources_update:
|
||||
raise HTTPException(status_code=404, detail="No resources have been set on the server.")
|
||||
logger.debug(f"Serving latest resources '{resources_update.resources_id}' to a client.")
|
||||
return resources_update
|
||||
|
||||
@self._app.get("/resources/{resource_id}", response_model=ResourcesUpdate)
|
||||
async def fetch_resources_by_id(
|
||||
resource_id: str = Path(..., description="The unique identifier for the resource version.")
|
||||
) -> ResourcesUpdate:
|
||||
"""Endpoint for clients to fetch a specific version of resources."""
|
||||
if not self._store:
|
||||
raise HTTPException(status_code=503, detail="Server not fully initialized.")
|
||||
resources_update = await self._store.get_resources_by_id(resource_id)
|
||||
if not resources_update:
|
||||
raise HTTPException(status_code=404, detail=f"Resource ID '{resource_id}' not found.")
|
||||
logger.debug(f"Serving resources for ID '{resource_id}' to a client.")
|
||||
return resources_update
|
||||
|
||||
@self._app.post("/rollout", response_model=GenericResponse)
|
||||
async def post_rollout(payload: Rollout) -> GenericResponse:
|
||||
"""Endpoint for clients to report a completed rollout."""
|
||||
if not self._store:
|
||||
raise HTTPException(status_code=503, detail="Server not fully initialized.")
|
||||
await self._store.store_rollout(payload)
|
||||
return GenericResponse(
|
||||
status="ok",
|
||||
message=f"Rollout {payload.rollout_id} received and stored.",
|
||||
)
|
||||
|
||||
async def start(self):
|
||||
"""Starts the FastAPI server in the background."""
|
||||
logger.info(f"Starting server at {self.endpoint}")
|
||||
asyncio.create_task(self._uvicorn_server.serve())
|
||||
await asyncio.sleep(1) # Allow time for server to start up.
|
||||
|
||||
async def stop(self):
|
||||
"""Gracefully stops the running FastAPI server."""
|
||||
if self._uvicorn_server.started:
|
||||
logger.info("Stopping server...")
|
||||
self._uvicorn_server.should_exit = True
|
||||
await asyncio.sleep(1) # Allow time for graceful shutdown.
|
||||
logger.info("Server stopped.")
|
||||
|
||||
async def run_forever(self):
|
||||
"""
|
||||
Runs the server indefinitely until stopped.
|
||||
This is useful when async start and stop methods do not work.
|
||||
"""
|
||||
await self._uvicorn_server.serve()
|
||||
|
||||
async def queue_task(
|
||||
self,
|
||||
sample: Any,
|
||||
mode: Literal["train", "val", "test"] | None = None,
|
||||
resources_id: str | None = None,
|
||||
metadata: Dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Adds a task to the queue for a client to process.
|
||||
"""
|
||||
if not self._store:
|
||||
raise RuntimeError("Store not initialized. The server may not be running.")
|
||||
return await self._store.add_task(sample, mode=mode, resources_id=resources_id, metadata=metadata)
|
||||
|
||||
async def update_resources(self, resources: NamedResources) -> str:
|
||||
"""
|
||||
Updates the resources, creating a new version and setting it as the latest.
|
||||
"""
|
||||
if not self._store:
|
||||
raise RuntimeError("Store not initialized. The server may not be running.")
|
||||
resources_id = f"res-{uuid.uuid4()}"
|
||||
update = ResourcesUpdate(resources_id=resources_id, resources=resources)
|
||||
await self._store.update_resources(update)
|
||||
return resources_id
|
||||
|
||||
async def get_completed_rollout(self, rollout_id: str) -> Optional[Rollout]:
|
||||
"""
|
||||
Retrieves a specific completed rollout by its ID.
|
||||
"""
|
||||
if not self._store:
|
||||
raise RuntimeError("Store not initialized. The server may not be running.")
|
||||
return await self._store.retrieve_rollout(rollout_id)
|
||||
|
||||
async def poll_completed_rollout(self, rollout_id: str, timeout: Optional[float] = None) -> Optional[Rollout]:
|
||||
"""
|
||||
Polls for a completed rollout by its ID, waiting up to `timeout` seconds.
|
||||
"""
|
||||
start_time = time.time()
|
||||
while True:
|
||||
rollout = await self.get_completed_rollout(rollout_id)
|
||||
if rollout:
|
||||
return rollout
|
||||
if timeout and (time.time() - start_time) >= timeout:
|
||||
return None
|
||||
await asyncio.sleep(1)
|
||||
|
||||
async def retrieve_completed_rollouts(self) -> List[Rollout]:
|
||||
"""
|
||||
Retrieves all available completed trajectories and clears the internal store.
|
||||
"""
|
||||
if not self._store:
|
||||
raise RuntimeError("Store not initialized. The server may not be running.")
|
||||
return await self._store.retrieve_completed_rollouts()
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
@@ -0,0 +1,27 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Hydra entrypoint for the Agent Lightning server."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hydra
|
||||
import uvicorn
|
||||
from omegaconf import DictConfig
|
||||
|
||||
from agentlightning.server.app import create_app
|
||||
|
||||
|
||||
@hydra.main(version_base=None, config_path="../config", config_name="server")
|
||||
def main(config: DictConfig) -> None:
|
||||
application = create_app(config)
|
||||
uvicorn.run(
|
||||
application,
|
||||
host=str(config.host),
|
||||
port=int(config.port),
|
||||
workers=1,
|
||||
timeout_keep_alive=120,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,93 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""FastAPI application — lifespan, mount routes, wire proxy."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
import structlog
|
||||
from fastapi import Depends, FastAPI, Request
|
||||
from fastapi.exceptions import HTTPException
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
|
||||
from agentlightning.server.proxy import ProxyPauseState, ProxyRouter
|
||||
from agentlightning.server.routes import events, models, proxy, rollouts
|
||||
|
||||
log = structlog.get_logger()
|
||||
|
||||
|
||||
def _server_config(config: Mapping[str, Any] | DictConfig | None) -> dict[str, Any]:
|
||||
if config is None:
|
||||
raise ValueError("server config is required")
|
||||
elif OmegaConf.is_config(config):
|
||||
raw = dict(cast(Any, OmegaConf.to_container(config, resolve=True)))
|
||||
else:
|
||||
raw = dict(config)
|
||||
|
||||
return raw
|
||||
|
||||
|
||||
def _build_auth_dependency(key: str):
|
||||
"""Return a dependency that validates the optional API key."""
|
||||
|
||||
async def verify_key(request: Request) -> None:
|
||||
if not key:
|
||||
return
|
||||
|
||||
auth_header = request.headers.get("authorization", "")
|
||||
if auth_header.startswith("Bearer ") and auth_header[7:] == key:
|
||||
return
|
||||
|
||||
if request.headers.get("x-api-key", "") == key:
|
||||
return
|
||||
|
||||
raise HTTPException(status_code=401, detail="Invalid or missing API key")
|
||||
|
||||
return verify_key
|
||||
|
||||
|
||||
def create_app(config: Mapping[str, Any] | DictConfig | None = None) -> FastAPI:
|
||||
"""Create and configure the FastAPI application."""
|
||||
server_config = _server_config(config)
|
||||
key = str(server_config["key"] or "")
|
||||
|
||||
if not key:
|
||||
log.warning("AGL_KEY not set — authentication disabled. Do not use in production.")
|
||||
|
||||
verify_key = _build_auth_dependency(key)
|
||||
|
||||
default_proxy = server_config["default_proxy"]
|
||||
log.info("Proxy config loaded", model_name=default_proxy["model_name"])
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
||||
app.state.proxy_pause_state = ProxyPauseState()
|
||||
|
||||
app.state.proxy_router = ProxyRouter(default_proxy)
|
||||
async with httpx.AsyncClient(timeout=httpx.Timeout(timeout=300.0)) as client:
|
||||
app.state.http_client = client
|
||||
yield
|
||||
|
||||
app = FastAPI(title="Agent Lightning", version="1.0.0", lifespan=lifespan)
|
||||
|
||||
# Health check — no auth.
|
||||
@app.get("/healthz")
|
||||
async def healthz() -> dict[str, str]:
|
||||
return {"status": "ok"}
|
||||
|
||||
# Store API routes — all require auth.
|
||||
app.include_router(rollouts.router, prefix="/api", dependencies=[Depends(verify_key)])
|
||||
app.include_router(events.router, prefix="/api", dependencies=[Depends(verify_key)])
|
||||
app.include_router(models.router, prefix="/api", dependencies=[Depends(verify_key)])
|
||||
|
||||
# Proxy routes (LLM proxy + event ingestion) — require agent-facing auth.
|
||||
app.include_router(proxy.router, dependencies=[Depends(verify_key)])
|
||||
|
||||
# Proxy management routes use the same server key as the rest of the API.
|
||||
app.include_router(proxy.management_router, dependencies=[Depends(verify_key)])
|
||||
|
||||
return app
|
||||
@@ -0,0 +1,255 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Server-side OpenAI chat-completions proxy."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
import structlog
|
||||
from fastapi import HTTPException, Response
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from agentlightning.schemas import Model
|
||||
from agentlightning.server.routes.events import record_event
|
||||
from agentlightning.server.store import _models
|
||||
|
||||
log = structlog.get_logger()
|
||||
|
||||
_UPSTREAM_MAX_ATTEMPTS = 6
|
||||
_RETRY_STATUS_CODES = {408, 409, 429}
|
||||
_RETRY_BACKOFF_BASE_SECONDS = 0.5
|
||||
_RETRY_BACKOFF_CAP_SECONDS = 8.0
|
||||
|
||||
|
||||
class NoServersError(Exception):
|
||||
def __init__(self, model: str) -> None:
|
||||
self.model = model
|
||||
super().__init__(f"No servers available for model '{model}'")
|
||||
|
||||
|
||||
class ProxyRouter:
|
||||
"""Selects the configured default model server and rewrites request params."""
|
||||
|
||||
def __init__(self, default_proxy: Mapping[str, Any]) -> None:
|
||||
self._model_name = str(default_proxy["model_name"])
|
||||
self._train_temperature = float(default_proxy["train"]["temperature"])
|
||||
self._val_temperature = float(default_proxy["val"]["temperature"])
|
||||
self._include_log_probs = bool(default_proxy.get("include_log_probs", True))
|
||||
|
||||
@property
|
||||
def model_name(self) -> str:
|
||||
return self._model_name
|
||||
|
||||
def select_server(self, model: str, rollout_id: str) -> Model:
|
||||
servers = _models.get(model, {})
|
||||
if not servers:
|
||||
raise NoServersError(model)
|
||||
# Stable ordering pins each rollout to one endpoint for prefix-cache reuse.
|
||||
pool = [servers[endpoint] for endpoint in sorted(servers)]
|
||||
digest = hashlib.sha256(rollout_id.encode("utf-8")).digest()
|
||||
index = int.from_bytes(digest[:8], "big") % len(pool)
|
||||
return pool[index]
|
||||
|
||||
def prepare_body(self, body: dict[str, Any], mode: str) -> dict[str, Any]:
|
||||
if mode == "train":
|
||||
prepared = {
|
||||
**body,
|
||||
"model": self._model_name,
|
||||
"temperature": self._train_temperature,
|
||||
"return_token_ids": True,
|
||||
}
|
||||
if self._include_log_probs:
|
||||
prepared["logprobs"] = True
|
||||
return prepared
|
||||
if mode == "val":
|
||||
prepared = {
|
||||
**body,
|
||||
"model": self._model_name,
|
||||
"temperature": self._val_temperature,
|
||||
"return_token_ids": True,
|
||||
}
|
||||
return prepared
|
||||
raise ValueError(f"Unsupported proxy mode: {mode}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProxyPauseState:
|
||||
paused: bool = False
|
||||
retry_after_seconds: int = 5
|
||||
reason: str | None = None
|
||||
inflight: int = 0
|
||||
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||
|
||||
|
||||
async def forward_request(
|
||||
*,
|
||||
client: httpx.AsyncClient,
|
||||
server: Model,
|
||||
body: dict[str, Any],
|
||||
upstream_path: str = "chat/completions",
|
||||
rollout_id: str,
|
||||
attempt_id: str,
|
||||
pause_state: ProxyPauseState | None = None,
|
||||
) -> Response:
|
||||
if pause_state is not None:
|
||||
async with pause_state.lock:
|
||||
if pause_state.paused:
|
||||
retry_after = pause_state.retry_after_seconds
|
||||
reason = pause_state.reason
|
||||
return Response(
|
||||
status_code=429,
|
||||
headers={"Retry-After": str(retry_after), "X-Agl-Paused": "true"},
|
||||
content=json.dumps({"error": "gateway paused", "reason": reason}),
|
||||
media_type="application/json",
|
||||
)
|
||||
pause_state.inflight += 1
|
||||
|
||||
try:
|
||||
if body.get("stream", False):
|
||||
raise HTTPException(status_code=400, detail="Streaming responses are not supported")
|
||||
|
||||
url = f"{server.endpoint.rstrip('/')}/{upstream_path}"
|
||||
log.debug("Proxying request", rollout_id=rollout_id, model=server.model, path=upstream_path)
|
||||
|
||||
started_at = time.perf_counter()
|
||||
response = await _send_upstream_with_retries(client=client, url=url, body=body)
|
||||
latency_ms = (time.perf_counter() - started_at) * 1000
|
||||
response_body = (
|
||||
response.json() if response.headers.get("content-type", "").startswith("application/json") else {}
|
||||
)
|
||||
|
||||
_capture_event(
|
||||
rollout_id=rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
request_body=body,
|
||||
response_body=response_body,
|
||||
server=server,
|
||||
latency_ms=latency_ms,
|
||||
http_status=response.status_code,
|
||||
status=_status_from_http_status(response.status_code),
|
||||
retry_count=int(response.extensions.get("agl_retry_count", 0)),
|
||||
)
|
||||
return JSONResponse(content=response_body, status_code=response.status_code)
|
||||
finally:
|
||||
if pause_state is not None:
|
||||
await _dec_inflight(pause_state)
|
||||
|
||||
|
||||
async def _send_upstream_with_retries(
|
||||
*,
|
||||
client: httpx.AsyncClient,
|
||||
url: str,
|
||||
body: dict[str, Any],
|
||||
) -> httpx.Response:
|
||||
for attempt_index in range(_UPSTREAM_MAX_ATTEMPTS):
|
||||
try:
|
||||
response = await client.post(url, json=body, headers={"content-type": "application/json"})
|
||||
except httpx.TimeoutException as exc:
|
||||
if attempt_index == _UPSTREAM_MAX_ATTEMPTS - 1:
|
||||
raise HTTPException(status_code=504, detail="Upstream model server timed out") from exc
|
||||
await _sleep_before_retry(url=url, attempt_index=attempt_index, reason="timeout")
|
||||
continue
|
||||
except httpx.TransportError as exc:
|
||||
if attempt_index == _UPSTREAM_MAX_ATTEMPTS - 1:
|
||||
raise HTTPException(status_code=502, detail="Upstream model server request failed") from exc
|
||||
await _sleep_before_retry(url=url, attempt_index=attempt_index, reason="transport error")
|
||||
continue
|
||||
|
||||
if not _is_retryable_status(response.status_code) or attempt_index == _UPSTREAM_MAX_ATTEMPTS - 1:
|
||||
response.extensions["agl_retry_count"] = attempt_index
|
||||
return response
|
||||
|
||||
await response.aclose()
|
||||
await _sleep_before_retry(
|
||||
url=url,
|
||||
attempt_index=attempt_index,
|
||||
reason=f"status {response.status_code}",
|
||||
)
|
||||
|
||||
raise HTTPException(status_code=502, detail="Upstream model server request failed")
|
||||
|
||||
|
||||
async def _sleep_before_retry(*, url: str, attempt_index: int, reason: str) -> None:
|
||||
delay = _retry_delay_seconds(attempt_index)
|
||||
log.warning(
|
||||
"Retrying upstream request",
|
||||
url=url,
|
||||
attempt=attempt_index + 1,
|
||||
max_attempts=_UPSTREAM_MAX_ATTEMPTS,
|
||||
delay_seconds=round(delay, 3),
|
||||
reason=reason,
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
|
||||
|
||||
def _is_retryable_status(status_code: int) -> bool:
|
||||
return status_code in _RETRY_STATUS_CODES or status_code >= 500
|
||||
|
||||
|
||||
def _retry_delay_seconds(attempt_index: int) -> float:
|
||||
delay = min(_RETRY_BACKOFF_BASE_SECONDS * (2**attempt_index), _RETRY_BACKOFF_CAP_SECONDS)
|
||||
return delay * random.uniform(0.75, 1.25)
|
||||
|
||||
|
||||
async def _dec_inflight(pause_state: ProxyPauseState) -> None:
|
||||
async with pause_state.lock:
|
||||
pause_state.inflight = max(0, pause_state.inflight - 1)
|
||||
|
||||
|
||||
def _capture_event(
|
||||
*,
|
||||
rollout_id: str,
|
||||
attempt_id: str,
|
||||
request_body: dict[str, Any],
|
||||
response_body: dict[str, Any],
|
||||
server: Model,
|
||||
latency_ms: float,
|
||||
http_status: int,
|
||||
status: str,
|
||||
retry_count: int,
|
||||
) -> None:
|
||||
record_event(
|
||||
rollout_id,
|
||||
attempt_id,
|
||||
"model_request",
|
||||
{
|
||||
"model": server.model,
|
||||
"model_version": server.version,
|
||||
"request": request_body,
|
||||
"response": response_body,
|
||||
"server": {"model": server.model, "endpoint": server.endpoint, "version": server.version},
|
||||
"latency_ms": latency_ms,
|
||||
"http_status": http_status,
|
||||
"status": status,
|
||||
"retry_count": retry_count,
|
||||
"usage": _extract_usage(response_body),
|
||||
"finish_reason": _extract_finish_reason(response_body),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _status_from_http_status(http_status: int) -> str:
|
||||
return "ok" if http_status < 400 else "error"
|
||||
|
||||
|
||||
def _extract_usage(response_body: dict[str, Any]) -> dict[str, Any] | None:
|
||||
usage = response_body.get("usage")
|
||||
return usage if isinstance(usage, dict) else None
|
||||
|
||||
|
||||
def _extract_finish_reason(response_body: dict[str, Any]) -> str | None:
|
||||
choices = response_body.get("choices")
|
||||
if isinstance(choices, list) and choices:
|
||||
reason = choices[0].get("finish_reason") if isinstance(choices[0], dict) else None
|
||||
if isinstance(reason, str) and reason:
|
||||
return reason
|
||||
return None
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
@@ -0,0 +1,206 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Event API routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Query
|
||||
from fastapi.exceptions import HTTPException
|
||||
|
||||
from agentlightning.schemas import DEFAULT_ATTEMPT_ID, Event, EventCreate
|
||||
from agentlightning.server.store import _events, _rollouts
|
||||
|
||||
router = APIRouter(tags=["events"])
|
||||
|
||||
|
||||
def _not_found(rollout_id: str) -> HTTPException:
|
||||
return HTTPException(status_code=404, detail=f"Rollout not found: {rollout_id}")
|
||||
|
||||
|
||||
def record_event(rollout_id: str, attempt_id: str, event_type: str, data: dict[str, Any]) -> Event:
|
||||
"""Append a single event for an existing rollout."""
|
||||
if rollout_id not in _rollouts:
|
||||
raise _not_found(rollout_id)
|
||||
|
||||
event = Event(
|
||||
event_type=event_type,
|
||||
rollout_id=rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
timestamp=time.time(),
|
||||
data=data,
|
||||
)
|
||||
|
||||
rid_events = _events[rollout_id]
|
||||
if attempt_id not in rid_events:
|
||||
rid_events[attempt_id] = []
|
||||
rid_events[attempt_id].append(event)
|
||||
return event
|
||||
|
||||
|
||||
def _query_events(
|
||||
rollout_id: str,
|
||||
*,
|
||||
event_type: str | None = None,
|
||||
) -> list[Event]:
|
||||
if rollout_id not in _rollouts:
|
||||
raise _not_found(rollout_id)
|
||||
|
||||
rollout = _rollouts[rollout_id]
|
||||
attempt_id = rollout.status.last_attempt_id or DEFAULT_ATTEMPT_ID
|
||||
rid_events = _events.get(rollout_id, {})
|
||||
events = rid_events.get(attempt_id, [])
|
||||
if event_type is not None:
|
||||
events = [event for event in events if event.event_type == event_type]
|
||||
|
||||
return events
|
||||
|
||||
|
||||
def _extract_choice_log_probs(choice: dict[str, Any]) -> list[float] | None:
|
||||
"""Extract chosen-token logprobs from a single choice.
|
||||
|
||||
Returns the per-token logprobs, or None when they are missing or unusable
|
||||
(no logprobs field, unrecognized schema, or any non-finite/non-float value).
|
||||
Never raises: a malformed response yields None so the triplet query stays a
|
||||
successful HTTP response and the training bridge drops the sample.
|
||||
"""
|
||||
lp = choice.get("logprobs")
|
||||
if not isinstance(lp, dict):
|
||||
return None
|
||||
|
||||
raw: list[Any]
|
||||
if isinstance(lp.get("content"), list):
|
||||
# OpenAI chat schema: logprobs.content -> [{"logprob": float, ...}, ...]
|
||||
raw = []
|
||||
for item in lp["content"]:
|
||||
if not isinstance(item, dict) or "logprob" not in item:
|
||||
return None
|
||||
raw.append(item["logprob"])
|
||||
elif isinstance(lp.get("token_logprobs"), list):
|
||||
# Completions schema: logprobs.token_logprobs -> [float, ...]
|
||||
raw = list(lp["token_logprobs"])
|
||||
else:
|
||||
return None
|
||||
|
||||
out: list[float] = []
|
||||
for v in raw:
|
||||
try:
|
||||
f = float(v)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if not math.isfinite(f):
|
||||
return None
|
||||
out.append(f)
|
||||
return out
|
||||
|
||||
|
||||
def _trim_model_request(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Extract prompt_token_ids and response_token_ids from a model_request event.
|
||||
|
||||
Non-streaming gateway responses use a dict shape with prompt_token_ids at
|
||||
top level for chat completions or per choice for completions, and token_ids
|
||||
per choice. Legacy raw-chunk format (list) is also supported for backward
|
||||
compatibility.
|
||||
"""
|
||||
resp = data.get("response")
|
||||
prompt_token_ids: list[int] = []
|
||||
response_token_ids: list[int] = []
|
||||
response_log_probs: list[float] | None = None
|
||||
|
||||
if isinstance(resp, dict):
|
||||
prompt_token_ids = resp.get("prompt_token_ids", [])
|
||||
choices = resp.get("choices", [])
|
||||
if choices:
|
||||
if not prompt_token_ids:
|
||||
prompt_token_ids = choices[0].get("prompt_token_ids", [])
|
||||
response_token_ids = choices[0].get("token_ids", [])
|
||||
response_log_probs = _extract_choice_log_probs(choices[0])
|
||||
elif isinstance(resp, list):
|
||||
# Legacy: raw SSE chunks (pre-assembly format, backward compat).
|
||||
for chunk in resp:
|
||||
if not prompt_token_ids and chunk.get("prompt_token_ids"):
|
||||
prompt_token_ids = chunk["prompt_token_ids"]
|
||||
choices = chunk.get("choices", [])
|
||||
if choices:
|
||||
tids = choices[0].get("token_ids")
|
||||
if tids:
|
||||
response_token_ids.extend(tids)
|
||||
|
||||
srv = data.get("server", {})
|
||||
trimmed = {
|
||||
"prompt_token_ids": prompt_token_ids,
|
||||
"response_token_ids": response_token_ids,
|
||||
"response_log_probs": response_log_probs,
|
||||
"server": {"model": srv.get("model"), "version": srv.get("version")},
|
||||
}
|
||||
for key in ("http_status", "status"):
|
||||
if key in data:
|
||||
trimmed[key] = data[key]
|
||||
if isinstance(resp, dict) and "error" in resp:
|
||||
trimmed["error"] = resp["error"]
|
||||
return trimmed
|
||||
|
||||
|
||||
def _trim_reward(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Keep only the scalar value from a reward event."""
|
||||
trimmed = {"value": data.get("value")}
|
||||
for key in ("source", "reason"):
|
||||
if key in data:
|
||||
trimmed[key] = data[key]
|
||||
return trimmed
|
||||
|
||||
|
||||
def _to_triplet_format(event: Event) -> Event:
|
||||
"""Trim event data for triplet consumption.
|
||||
|
||||
- model_request: extract prompt_token_ids + response_token_ids only
|
||||
- reward: keep only the scalar value
|
||||
- other event types: pass through unchanged
|
||||
"""
|
||||
if event.event_type == "model_request":
|
||||
trimmed = _trim_model_request(event.data)
|
||||
return event.model_copy(update={"data": trimmed})
|
||||
elif event.event_type == "reward":
|
||||
trimmed = _trim_reward(event.data)
|
||||
return event.model_copy(update={"data": trimmed})
|
||||
return event
|
||||
|
||||
|
||||
def _dedupe_model_requests_by_prompt_token_ids(events: list[Event]) -> list[Event]:
|
||||
"""Keep only the last model_request event for each prompt_token_ids key."""
|
||||
last_index_by_prompt: dict[tuple[Any, ...], int] = {}
|
||||
for index, event in enumerate(events):
|
||||
if event.event_type != "model_request":
|
||||
continue
|
||||
prompt_token_ids = event.data.get("prompt_token_ids", [])
|
||||
prompt_key = tuple(prompt_token_ids) if isinstance(prompt_token_ids, list) else ()
|
||||
last_index_by_prompt[prompt_key] = index
|
||||
|
||||
last_indexes = set(last_index_by_prompt.values())
|
||||
return [event for index, event in enumerate(events) if event.event_type != "model_request" or index in last_indexes]
|
||||
|
||||
|
||||
@router.post("/rollouts/{rollout_id}/attempt/{attempt_id}/events", response_model=Event)
|
||||
async def post_event(rollout_id: str, body: EventCreate, attempt_id: str) -> Event:
|
||||
"""Post an event for one rollout attempt."""
|
||||
return record_event(rollout_id, attempt_id, body.event_type, body.data)
|
||||
|
||||
|
||||
@router.get("/rollouts/{rollout_id}/events", response_model=list[Event])
|
||||
async def query_events(
|
||||
rollout_id: str,
|
||||
event_type: str | None = None,
|
||||
format: str | None = Query(None, description="Set to 'triplet' to trim events for RL training"),
|
||||
) -> list[Event]:
|
||||
"""Query events for the default rollout attempt."""
|
||||
events = _query_events(
|
||||
rollout_id=rollout_id,
|
||||
event_type=event_type,
|
||||
)
|
||||
if format == "triplet":
|
||||
events = [_to_triplet_format(e) for e in events]
|
||||
events = _dedupe_model_requests_by_prompt_token_ids(events)
|
||||
return events
|
||||
@@ -0,0 +1,31 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Model server API routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from agentlightning.schemas import Model
|
||||
from agentlightning.server.store import _models
|
||||
|
||||
router = APIRouter(tags=["models"])
|
||||
|
||||
|
||||
@router.post("/models", status_code=201, response_model=list[Model])
|
||||
async def register_models(body: list[Model]) -> list[Model]:
|
||||
"""Register model server(s). Upsert by (model, endpoint)."""
|
||||
results: list[Model] = []
|
||||
for req in body:
|
||||
if req.model not in _models:
|
||||
_models[req.model] = {}
|
||||
_models[req.model][req.endpoint] = req
|
||||
results.append(req)
|
||||
return results
|
||||
|
||||
|
||||
@router.delete("/models")
|
||||
async def delete_all_models() -> dict[str, str]:
|
||||
"""Remove all model servers."""
|
||||
_models.clear()
|
||||
return {"status": "ok"}
|
||||
@@ -0,0 +1,137 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Proxy forwarding and pause/drain management routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import structlog
|
||||
from fastapi import APIRouter, Request, Response
|
||||
from fastapi.exceptions import HTTPException
|
||||
from pydantic import BaseModel
|
||||
|
||||
from agentlightning.server.proxy import NoServersError, ProxyPauseState, ProxyRouter, forward_request
|
||||
from agentlightning.server.store import _rollouts
|
||||
|
||||
log = structlog.get_logger()
|
||||
|
||||
router = APIRouter(tags=["gateway"])
|
||||
management_router = APIRouter(tags=["gateway-management"], prefix="/proxy")
|
||||
|
||||
|
||||
def _get_pause_state(request: Request) -> ProxyPauseState:
|
||||
state: ProxyPauseState | None = getattr(request.app.state, "proxy_pause_state", None)
|
||||
if state is None:
|
||||
raise HTTPException(status_code=503, detail="Gateway pause state not configured")
|
||||
return state
|
||||
|
||||
|
||||
@router.post(
|
||||
"/proxy/rollout/{rollout_id}/attempt/{attempt_id}/mode/{mode}/openai/v1/{upstream_path:path}",
|
||||
)
|
||||
async def llm_proxy(rollout_id: str, attempt_id: str, mode: str, upstream_path: str, request: Request) -> Response:
|
||||
"""LLM reverse proxy — forwards to model server, captures events."""
|
||||
if mode not in {"train", "val"}:
|
||||
raise HTTPException(status_code=404, detail=f"Unsupported proxy mode: {mode}")
|
||||
if upstream_path not in {"chat/completions", "completions"}:
|
||||
raise HTTPException(status_code=404, detail=f"Unsupported upstream path: {upstream_path}")
|
||||
|
||||
# Validate rollout exists.
|
||||
if rollout_id not in _rollouts:
|
||||
raise HTTPException(status_code=404, detail=f"Rollout not found: {rollout_id}")
|
||||
|
||||
# Get gateway router and httpx client from app state.
|
||||
proxy_router: ProxyRouter | None = getattr(request.app.state, "proxy_router", None)
|
||||
http_client = getattr(request.app.state, "http_client", None)
|
||||
|
||||
if proxy_router is None or http_client is None:
|
||||
raise HTTPException(status_code=503, detail="Proxy not configured")
|
||||
|
||||
pause_state: ProxyPauseState | None = getattr(request.app.state, "proxy_pause_state", None)
|
||||
|
||||
# Read and parse request body.
|
||||
raw_body = await request.body()
|
||||
try:
|
||||
body = json.loads(raw_body) if raw_body else {}
|
||||
except json.JSONDecodeError:
|
||||
raise HTTPException(status_code=400, detail="Invalid JSON in request body") from None
|
||||
|
||||
# Select server.
|
||||
model_name = proxy_router.model_name
|
||||
try:
|
||||
server = proxy_router.select_server(model_name, rollout_id)
|
||||
except NoServersError:
|
||||
raise HTTPException(status_code=503, detail=f"No servers available for model '{model_name}'") from None
|
||||
|
||||
prepared_body = proxy_router.prepare_body(body, mode)
|
||||
|
||||
# Server endpoint includes the OpenAI base path (e.g., "http://vllm:8000/v1").
|
||||
return await forward_request(
|
||||
client=http_client,
|
||||
server=server,
|
||||
body=prepared_body,
|
||||
upstream_path=upstream_path,
|
||||
rollout_id=rollout_id,
|
||||
attempt_id=attempt_id,
|
||||
pause_state=pause_state,
|
||||
)
|
||||
|
||||
|
||||
# --- Management routes ------------------------------------------------------
|
||||
|
||||
|
||||
class PauseRequest(BaseModel):
|
||||
retry_after_seconds: int = 5
|
||||
reason: str | None = None
|
||||
|
||||
|
||||
class PauseStateResponse(BaseModel):
|
||||
paused: bool
|
||||
retry_after_seconds: int
|
||||
reason: str | None
|
||||
inflight: int
|
||||
|
||||
|
||||
@management_router.post("/pause", response_model=PauseStateResponse)
|
||||
async def pause_proxy(body: PauseRequest, request: Request) -> PauseStateResponse:
|
||||
"""Pause new proxy forwarding requests while existing in-flight requests drain."""
|
||||
state = _get_pause_state(request)
|
||||
async with state.lock:
|
||||
state.paused = True
|
||||
state.retry_after_seconds = body.retry_after_seconds
|
||||
state.reason = body.reason
|
||||
return PauseStateResponse(
|
||||
paused=state.paused,
|
||||
retry_after_seconds=state.retry_after_seconds,
|
||||
reason=state.reason,
|
||||
inflight=state.inflight,
|
||||
)
|
||||
|
||||
|
||||
@management_router.post("/resume", response_model=PauseStateResponse)
|
||||
async def resume_proxy(request: Request) -> PauseStateResponse:
|
||||
"""Resume proxy forwarding after a pause."""
|
||||
state = _get_pause_state(request)
|
||||
async with state.lock:
|
||||
state.paused = False
|
||||
state.reason = None
|
||||
return PauseStateResponse(
|
||||
paused=state.paused,
|
||||
retry_after_seconds=state.retry_after_seconds,
|
||||
reason=state.reason,
|
||||
inflight=state.inflight,
|
||||
)
|
||||
|
||||
|
||||
@management_router.get("/state", response_model=PauseStateResponse)
|
||||
async def proxy_state(request: Request) -> PauseStateResponse:
|
||||
"""Return the proxy pause state and in-flight request count."""
|
||||
state = _get_pause_state(request)
|
||||
async with state.lock:
|
||||
return PauseStateResponse(
|
||||
paused=state.paused,
|
||||
retry_after_seconds=state.retry_after_seconds,
|
||||
reason=state.reason,
|
||||
inflight=state.inflight,
|
||||
)
|
||||
@@ -0,0 +1,220 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Rollout API routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Query
|
||||
from fastapi.exceptions import HTTPException
|
||||
from pydantic import BaseModel
|
||||
|
||||
from agentlightning.schemas import (
|
||||
TERMINAL_STATES,
|
||||
VALID_TRANSITIONS,
|
||||
Rollout,
|
||||
RolloutConfig,
|
||||
RolloutCreate,
|
||||
RolloutLifecycleStatus,
|
||||
RolloutMetadata,
|
||||
RolloutPatch,
|
||||
RolloutState,
|
||||
)
|
||||
from agentlightning.server.store import _events, _rollouts, _terminal_order
|
||||
|
||||
router = APIRouter(tags=["rollouts"])
|
||||
|
||||
|
||||
class RolloutDetail(BaseModel):
|
||||
"""Rollout with attempt list."""
|
||||
|
||||
rollout: Rollout
|
||||
attempts: list[str]
|
||||
|
||||
|
||||
class TerminalRolloutItem(BaseModel):
|
||||
"""Lightweight projection of a terminal rollout (no input/config payload)."""
|
||||
|
||||
rollout_id: str
|
||||
state: RolloutState
|
||||
data_id: str
|
||||
is_train: bool
|
||||
|
||||
|
||||
class TerminalRolloutsPage(BaseModel):
|
||||
"""A page of terminal rollouts plus the cursor to fetch the next page."""
|
||||
|
||||
items: list[TerminalRolloutItem]
|
||||
next_after: int
|
||||
total_terminal: int
|
||||
|
||||
|
||||
def _not_found(rollout_id: str) -> HTTPException:
|
||||
return HTTPException(status_code=404, detail=f"Rollout not found: {rollout_id}")
|
||||
|
||||
|
||||
def _invalid_transition(rollout_id: str, from_status: str, to_status: str) -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=409,
|
||||
detail=f"Rollout {rollout_id}: cannot transition {from_status} -> {to_status}",
|
||||
)
|
||||
|
||||
|
||||
def _get_rollout(rollout_id: str) -> Rollout:
|
||||
try:
|
||||
return _rollouts[rollout_id]
|
||||
except KeyError:
|
||||
raise _not_found(rollout_id) from None
|
||||
|
||||
|
||||
def _metadata_from_request(req: RolloutCreate) -> RolloutMetadata:
|
||||
if isinstance(req.metadata, dict):
|
||||
return RolloutMetadata(**req.metadata)
|
||||
if req.metadata is not None:
|
||||
return req.metadata
|
||||
return RolloutMetadata()
|
||||
|
||||
|
||||
def _list_attempts(rollout_id: str) -> list[str]:
|
||||
if rollout_id not in _rollouts:
|
||||
raise _not_found(rollout_id)
|
||||
|
||||
rid_events = _events.get(rollout_id, {})
|
||||
if not rid_events:
|
||||
return []
|
||||
|
||||
return sorted(
|
||||
rid_events.keys(),
|
||||
key=lambda attempt_id: rid_events[attempt_id][0].timestamp if rid_events[attempt_id] else float("inf"),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/rollouts", status_code=201, response_model=list[Rollout])
|
||||
async def enqueue_rollouts(body: list[RolloutCreate]) -> list[Rollout]:
|
||||
"""Enqueue rollouts. Each item in the list is self-contained.
|
||||
|
||||
If a request carries a `rollout_id` that already exists, the existing
|
||||
rollout is returned unchanged (its events are left intact), making creation
|
||||
idempotent so callers can pre-assign ids and retry safely.
|
||||
"""
|
||||
results: list[Rollout] = []
|
||||
for req in body:
|
||||
if req.rollout_id is not None and req.rollout_id in _rollouts:
|
||||
results.append(_rollouts[req.rollout_id])
|
||||
continue
|
||||
now = time.time()
|
||||
rollout_id = req.rollout_id or uuid.uuid4().hex
|
||||
metadata = _metadata_from_request(req)
|
||||
rollout = Rollout(
|
||||
rollout_id=rollout_id,
|
||||
input=req.input,
|
||||
is_train=req.is_train,
|
||||
config=req.config or RolloutConfig(),
|
||||
metadata=metadata,
|
||||
status=RolloutLifecycleStatus(created_at=now, updated_at=now),
|
||||
)
|
||||
_rollouts[rollout_id] = rollout
|
||||
_events[rollout_id] = {}
|
||||
results.append(rollout)
|
||||
return results
|
||||
|
||||
|
||||
@router.get("/rollouts", response_model=list[Rollout])
|
||||
async def list_rollouts(
|
||||
state_in: Annotated[list[RolloutState], Query()],
|
||||
limit: int = 500,
|
||||
) -> list[Rollout]:
|
||||
"""List rollouts by lifecycle states."""
|
||||
states = set(state_in)
|
||||
matches = [rollout for rollout in _rollouts.values() if rollout.status.state in states]
|
||||
return matches[:limit]
|
||||
|
||||
|
||||
def _data_id_of(rollout: Rollout) -> str:
|
||||
inp = rollout.input
|
||||
if isinstance(inp, dict):
|
||||
return str(inp.get("data_id") or inp.get("instance_id") or "")
|
||||
return ""
|
||||
|
||||
|
||||
@router.get("/rollouts/terminal", response_model=TerminalRolloutsPage)
|
||||
async def list_terminal_rollouts(after: int = 0, limit: int = 1000) -> TerminalRolloutsPage:
|
||||
"""Cursor-paginate terminal rollouts in completion order (lightweight projection).
|
||||
|
||||
`after` is an index into the append-only completion log; pass back `next_after`
|
||||
to fetch only rollouts that completed since the last call. Out-of-order
|
||||
completions are never missed because the log is append-on-terminal-transition.
|
||||
Returns only id/state/data_id/is_train — fetch events per rollout for details.
|
||||
"""
|
||||
if after < 0:
|
||||
after = 0
|
||||
if limit < 1:
|
||||
limit = 1
|
||||
total = len(_terminal_order)
|
||||
slice_ids = _terminal_order[after : after + limit]
|
||||
items: list[TerminalRolloutItem] = []
|
||||
for rid in slice_ids:
|
||||
rollout = _rollouts.get(rid)
|
||||
if rollout is None:
|
||||
continue
|
||||
items.append(
|
||||
TerminalRolloutItem(
|
||||
rollout_id=rid,
|
||||
state=rollout.status.state,
|
||||
data_id=_data_id_of(rollout),
|
||||
is_train=rollout.is_train,
|
||||
)
|
||||
)
|
||||
return TerminalRolloutsPage(items=items, next_after=after + len(slice_ids), total_terminal=total)
|
||||
|
||||
|
||||
@router.get("/rollouts/{rollout_id}", response_model=RolloutDetail)
|
||||
async def get_rollout(rollout_id: str) -> RolloutDetail:
|
||||
"""Get a single rollout with its attempt list."""
|
||||
rollout = _get_rollout(rollout_id)
|
||||
attempts = _list_attempts(rollout_id)
|
||||
return RolloutDetail(rollout=rollout, attempts=attempts)
|
||||
|
||||
|
||||
@router.patch("/rollouts/{rollout_id}", response_model=Rollout)
|
||||
async def patch_rollout(rollout_id: str, body: RolloutPatch) -> Rollout:
|
||||
"""Patch the lifecycle status of a rollout."""
|
||||
rollout = _get_rollout(rollout_id)
|
||||
updates = body.status.model_dump(exclude_unset=True) if body.status is not None else {}
|
||||
|
||||
if not updates:
|
||||
return rollout
|
||||
|
||||
if "state" in updates:
|
||||
new_state = updates["state"]
|
||||
if new_state not in VALID_TRANSITIONS[rollout.status.state]:
|
||||
raise _invalid_transition(rollout_id, rollout.status.state, str(new_state))
|
||||
|
||||
updated_status = rollout.status.model_copy(
|
||||
update={
|
||||
**updates,
|
||||
"version": rollout.status.version + 1,
|
||||
"updated_at": time.time(),
|
||||
}
|
||||
)
|
||||
|
||||
updated = rollout.model_copy(
|
||||
update={
|
||||
"status": updated_status,
|
||||
}
|
||||
)
|
||||
_rollouts[rollout_id] = updated
|
||||
if "state" in updates and updated_status.state in TERMINAL_STATES:
|
||||
# One-way terminal transition (guarded above) => append exactly once.
|
||||
_terminal_order.append(rollout_id)
|
||||
return updated
|
||||
|
||||
|
||||
@router.delete("/rollouts/{rollout_id}", status_code=204)
|
||||
async def delete_rollout(rollout_id: str) -> None:
|
||||
"""Delete a rollout and its events. Idempotent: missing id is a no-op."""
|
||||
_rollouts.pop(rollout_id, None)
|
||||
_events.pop(rollout_id, None)
|
||||
@@ -0,0 +1,18 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""In-memory server state — single-threaded, no locks, plain dict/list.
|
||||
|
||||
Route handlers mutate these module-level dictionaries directly on the event loop
|
||||
thread. See docs/dev_guidelines.md § Concurrency Model.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from agentlightning.schemas import Event, Model, Rollout
|
||||
|
||||
_rollouts: dict[str, Rollout] = {}
|
||||
_events: dict[str, dict[str, list[Event]]] = {}
|
||||
_models: dict[str, dict[str, Model]] = {}
|
||||
|
||||
# Completion-ordered ids enable cursor pagination without rescanning rollouts.
|
||||
_terminal_order: list[str] = []
|
||||
@@ -1,3 +0,0 @@
|
||||
from .base import BaseTracer
|
||||
from .agentops import AgentOpsTracer
|
||||
from .triplet import TripletExporter
|
||||
@@ -1,240 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import List, Optional, TYPE_CHECKING
|
||||
|
||||
import agentops.sdk.core
|
||||
import agentops
|
||||
from agentops.sdk.core import TracingCore
|
||||
from agentops.sdk.processors import SpanProcessor
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
|
||||
from agentlightning.instrumentation.agentops import AgentOpsServerManager
|
||||
from agentlightning.instrumentation import instrument_all, uninstrument_all
|
||||
from .base import BaseTracer
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentops.integration.callbacks.langchain import LangchainCallbackHandler
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AgentOpsTracer(BaseTracer):
|
||||
"""Traces agent execution using AgentOps.
|
||||
|
||||
This tracer provides functionality to capture execution details using the
|
||||
AgentOps library. It manages the AgentOps client initialization, server setup,
|
||||
and integration with the OpenTelemetry tracing ecosystem.
|
||||
|
||||
Attributes:
|
||||
agentops_managed: Whether to automatically manage `agentops`.
|
||||
When set to true, tracer calls `agentops.init()`
|
||||
automatically and launches an agentops endpoint locally.
|
||||
If not, you are responsible for calling and using it
|
||||
before using the tracer.
|
||||
instrument_managed: Whether to automatically manage instrumentation.
|
||||
When set to false, you will manage the instrumentation
|
||||
yourself and the tracer might not work as expected.
|
||||
daemon: Whether the AgentOps server runs as a daemon process.
|
||||
Only applicable if `agentops_managed` is True.
|
||||
"""
|
||||
|
||||
def __init__(self, *, agentops_managed: bool = True, instrument_managed: bool = True, daemon: bool = True):
|
||||
super().__init__()
|
||||
self._lightning_span_processor: Optional[LightningSpanProcessor] = None
|
||||
self.agentops_managed = agentops_managed
|
||||
self.instrument_managed = instrument_managed
|
||||
self.daemon = daemon
|
||||
|
||||
self._agentops_server_manager = AgentOpsServerManager(self.daemon)
|
||||
self._agentops_server_port_val: Optional[int] = None
|
||||
|
||||
if not self.agentops_managed:
|
||||
logger.warning("agentops_managed=False. You are responsible for AgentOps setup.")
|
||||
if not self.instrument_managed:
|
||||
logger.warning("instrument_managed=False. You are responsible for all instrumentation.")
|
||||
|
||||
def __getstate__(self):
|
||||
state = self.__dict__.copy()
|
||||
state["_agentops_server_manager"] = None # Exclude the unpicklable server manager
|
||||
# _agentops_server_port_val (int) is inherently picklable and will be included.
|
||||
logger.debug(f"Getting state for pickling Trainer (PID {os.getpid()}). _agentops_server_manager excluded.")
|
||||
return state
|
||||
|
||||
def __setstate__(self, state):
|
||||
self.__dict__.update(state)
|
||||
# In child process, self._agentops_server_manager will be None.
|
||||
logger.debug(f"Setting state for unpickled Trainer (PID {os.getpid()}). _agentops_server_manager is None.")
|
||||
|
||||
def init(self, *args, **kwargs):
|
||||
if self.agentops_managed and self._agentops_server_manager:
|
||||
self._agentops_server_manager.start()
|
||||
self._agentops_server_port_val = self._agentops_server_manager.get_port()
|
||||
if self._agentops_server_port_val is None:
|
||||
if (
|
||||
self._agentops_server_manager.server_process is not None
|
||||
and self._agentops_server_manager.server_process.is_alive()
|
||||
):
|
||||
raise RuntimeError("AgentOps server started but port is None. Check server manager logic.")
|
||||
elif (
|
||||
self._agentops_server_port_val is None and self._agentops_server_manager.server_process is None
|
||||
): # Server failed to start
|
||||
raise RuntimeError("AgentOps server manager indicates server is not running and port is None.")
|
||||
|
||||
def teardown(self):
|
||||
if self.agentops_managed:
|
||||
self._agentops_server_manager.stop()
|
||||
logger.info("AgentOps server stopped.")
|
||||
|
||||
def instrument(self, worker_id: int):
|
||||
instrument_all()
|
||||
|
||||
def uninstrument(self, worker_id: int):
|
||||
uninstrument_all()
|
||||
|
||||
def init_worker(self, worker_id: int):
|
||||
super().init_worker(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] Setting up tracer...") # worker_id included in process name
|
||||
|
||||
if self.instrument_managed:
|
||||
self.instrument(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] Instrumentation applied.")
|
||||
|
||||
if self.agentops_managed:
|
||||
if self._agentops_server_port_val: # Use the stored, picklable port value
|
||||
base_url = f"http://localhost:{self._agentops_server_port_val}"
|
||||
env_vars_to_set = {
|
||||
"AGENTOPS_API_KEY": "dummy",
|
||||
"AGENTOPS_API_ENDPOINT": base_url,
|
||||
"AGENTOPS_APP_URL": f"{base_url}/notavailable",
|
||||
"AGENTOPS_EXPORTER_ENDPOINT": f"{base_url}/traces",
|
||||
}
|
||||
for key, value in env_vars_to_set.items():
|
||||
os.environ[key] = value
|
||||
logger.info(f"[Worker {worker_id}] Env var set: {key}={value}")
|
||||
else:
|
||||
logger.warning(
|
||||
f"[Worker {worker_id}] AgentOps managed, but local server port is not available. Client may not connect as expected."
|
||||
)
|
||||
|
||||
if not agentops.get_client().initialized:
|
||||
agentops.init()
|
||||
logger.info(f"[Worker {worker_id}] AgentOps client initialized.")
|
||||
else:
|
||||
logger.warning(f"[Worker {worker_id}] AgentOps client was already initialized.")
|
||||
|
||||
self._lightning_span_processor = LightningSpanProcessor()
|
||||
|
||||
try:
|
||||
# new versions
|
||||
instance = agentops.sdk.core.tracer
|
||||
instance.provider.add_span_processor(self._lightning_span_processor)
|
||||
except AttributeError:
|
||||
# old versions
|
||||
instance = TracingCore.get_instance()
|
||||
instance._provider.add_span_processor(self._lightning_span_processor)
|
||||
|
||||
def teardown_worker(self, worker_id: int) -> None:
|
||||
super().teardown_worker(worker_id)
|
||||
|
||||
if self.instrument_managed:
|
||||
self.uninstrument(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] Instrumentation removed.")
|
||||
|
||||
@contextmanager
|
||||
def trace_context(self, name: Optional[str] = None):
|
||||
"""
|
||||
Starts a new tracing context. This should be used as a context manager.
|
||||
|
||||
Args:
|
||||
name: Optional name for the tracing context.
|
||||
|
||||
Yields:
|
||||
The LightningSpanProcessor instance to collect spans.
|
||||
"""
|
||||
if not self._lightning_span_processor:
|
||||
raise RuntimeError("LightningSpanProcessor is not initialized. Call init_worker() first.")
|
||||
|
||||
with self._lightning_span_processor:
|
||||
yield self._lightning_span_processor
|
||||
|
||||
def get_last_trace(self) -> List[ReadableSpan]:
|
||||
"""
|
||||
Retrieves the raw list of captured spans from the most recent trace.
|
||||
|
||||
Returns:
|
||||
A list of OpenTelemetry `ReadableSpan` objects.
|
||||
"""
|
||||
if not self._lightning_span_processor:
|
||||
raise RuntimeError("LightningSpanProcessor is not initialized. Call init_worker() first.")
|
||||
return self._lightning_span_processor.spans()
|
||||
|
||||
def get_langchain_callback_handler(self, tags: List[str] | None = None) -> LangchainCallbackHandler:
|
||||
"""
|
||||
Get the Langchain callback handler for integrating with Langchain.
|
||||
|
||||
Args:
|
||||
tags: Optional list of tags to apply to the Langchain callback handler.
|
||||
|
||||
Returns:
|
||||
An instance of the Langchain callback handler.
|
||||
"""
|
||||
import agentops
|
||||
from agentops.integration.callbacks.langchain import LangchainCallbackHandler
|
||||
|
||||
tags = tags or []
|
||||
client_instance = agentops.get_client()
|
||||
api_key = None
|
||||
if client_instance.initialized:
|
||||
api_key = client_instance.config.api_key
|
||||
else:
|
||||
logger.warning(
|
||||
"AgentOps client not initialized when creating LangchainCallbackHandler. API key may be missing."
|
||||
)
|
||||
return LangchainCallbackHandler(api_key=api_key, tags=tags)
|
||||
|
||||
|
||||
class LightningSpanProcessor(SpanProcessor):
|
||||
|
||||
_spans: List[ReadableSpan] = []
|
||||
|
||||
def __enter__(self):
|
||||
self._last_trace = None
|
||||
self._spans = []
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
pass
|
||||
|
||||
def spans(self) -> List[ReadableSpan]:
|
||||
"""
|
||||
Get the list of spans collected by this processor.
|
||||
This is useful for debugging and testing purposes.
|
||||
|
||||
Returns:
|
||||
List of ReadableSpan objects collected during tracing.
|
||||
"""
|
||||
return self._spans
|
||||
|
||||
def on_end(self, span: ReadableSpan) -> None:
|
||||
"""
|
||||
Process a span when it ends.
|
||||
|
||||
Args:
|
||||
span: The span that has ended.
|
||||
"""
|
||||
# Skip if span is not sampled
|
||||
if not span.context or not span.context.trace_flags.sampled:
|
||||
return
|
||||
|
||||
self._spans.append(span)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
pass
|
||||
|
||||
def force_flush(self, timeout_millis: int = 30000) -> bool:
|
||||
return True
|
||||
@@ -1,95 +0,0 @@
|
||||
from contextlib import contextmanager
|
||||
from typing import Iterator, List, Optional, Callable, Any, Awaitable
|
||||
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from agentlightning.types import ParallelWorkerBase
|
||||
|
||||
|
||||
class BaseTracer(ParallelWorkerBase):
|
||||
"""
|
||||
An abstract base class for tracers.
|
||||
|
||||
This class defines a standard interface for tracing code execution,
|
||||
capturing the resulting spans, and providing them for analysis. It is
|
||||
designed to be backend-agnostic, allowing for different implementations
|
||||
(e.g., for AgentOps, OpenTelemetry, Docker, etc.).
|
||||
|
||||
The primary interaction pattern is through the `trace_context`
|
||||
context manager, which ensures that traces are properly started and captured,
|
||||
even in the case of exceptions.
|
||||
|
||||
A typical workflow:
|
||||
|
||||
```python
|
||||
tracer = YourTracerImplementation()
|
||||
|
||||
try:
|
||||
with tracer.trace_context(name="my_traced_task"):
|
||||
# ... code to be traced ...
|
||||
run_my_agent_logic()
|
||||
except Exception as e:
|
||||
print(f"An error occurred: {e}")
|
||||
|
||||
# Retrieve the trace data after the context block
|
||||
spans: list[ReadableSpan] = tracer.get_last_trace()
|
||||
|
||||
# Process the trace data
|
||||
if trace_tree:
|
||||
rl_triplets = TripletExporter().export(spans)
|
||||
# ... do something with the triplets
|
||||
```
|
||||
"""
|
||||
|
||||
@contextmanager
|
||||
def trace_context(self, name: Optional[str] = None) -> Iterator[Any]:
|
||||
"""
|
||||
Starts a new tracing context. This should be used as a context manager.
|
||||
|
||||
The implementation should handle the setup and teardown of the tracing
|
||||
for the enclosed code block. It must ensure that any spans generated
|
||||
within the `with` block are collected and made available via
|
||||
`get_last_trace`.
|
||||
|
||||
Args:
|
||||
name: The name for the root span of this trace context.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
def get_last_trace(self) -> List[ReadableSpan]:
|
||||
"""
|
||||
Retrieves the raw list of captured spans from the most recent trace.
|
||||
|
||||
Returns:
|
||||
A list of OpenTelemetry `ReadableSpan` objects.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
def trace_run(self, func: Callable, *args, **kwargs) -> Any:
|
||||
"""
|
||||
A convenience wrapper to trace the execution of a single synchronous function.
|
||||
|
||||
Args:
|
||||
func: The synchronous function to execute and trace.
|
||||
*args: Positional arguments to pass to the function.
|
||||
**kwargs: Keyword arguments to pass to the function.
|
||||
|
||||
Returns:
|
||||
The return value of the function.
|
||||
"""
|
||||
with self.trace_context(name=func.__name__):
|
||||
return func(*args, **kwargs)
|
||||
|
||||
async def trace_run_async(self, func: Callable[..., Awaitable], *args, **kwargs) -> Any:
|
||||
"""
|
||||
A convenience wrapper to trace the execution of a single asynchronous function.
|
||||
|
||||
Args:
|
||||
func: The asynchronous function to execute and trace.
|
||||
*args: Positional arguments to pass to the function.
|
||||
**kwargs: Keyword arguments to pass to the function.
|
||||
|
||||
Returns:
|
||||
The return value of the function.
|
||||
"""
|
||||
with self.trace_context(name=func.__name__):
|
||||
return await func(*args, **kwargs)
|
||||
@@ -1,366 +0,0 @@
|
||||
from contextlib import contextmanager
|
||||
from typing import Iterator, List, Optional, Any, Dict, Callable, Awaitable
|
||||
import logging
|
||||
import uuid
|
||||
import pickle
|
||||
import multiprocessing
|
||||
import asyncio
|
||||
import queue
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from .base import BaseTracer
|
||||
|
||||
from httpdbg.hooks.all import httprecord
|
||||
from httpdbg.records import HTTPRecords
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from opentelemetry.trace import StatusCode, SpanKind, Status
|
||||
from opentelemetry.trace.span import (
|
||||
SpanContext,
|
||||
TraceFlags,
|
||||
TraceState,
|
||||
)
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class HttpTracer(BaseTracer):
|
||||
"""
|
||||
A tracer implementation that captures HTTP requests using httpdbg.
|
||||
|
||||
This tracer hooks into the Python HTTP libraries and captures all
|
||||
HTTP requests and responses made during the traced code execution.
|
||||
The captured requests are converted to OpenTelemetry spans for
|
||||
compatibility with the rest of the tracing ecosystem.
|
||||
|
||||
Caution: The current implementation of HttpTracer is very fragile,
|
||||
and we do not recommend using it in production.
|
||||
It is primarily for demonstration and testing purposes.
|
||||
|
||||
Attributes:
|
||||
include_headers: Whether to include HTTP headers in the spans.
|
||||
Headers may contain sensitive information. Use with caution.
|
||||
include_body: Whether to include HTTP request and response bodies in the spans.
|
||||
Bodies may be large and contain sensitive information. Use with caution.
|
||||
include_agentlightning_requests: Whether to include requests initiated by AgentLightning itself.
|
||||
subprocess_mode: Whether to run trace_run and trace_run_async in subprocesses for isolation.
|
||||
subprocess_timeout: Timeout for subprocess execution in seconds.
|
||||
"""
|
||||
|
||||
AGENTLIGHTNING_HEADERS = {"x-agentlightning-client"}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
include_headers: bool = False,
|
||||
include_body: bool = False,
|
||||
include_agentlightning_requests: bool = False,
|
||||
subprocess_mode: bool = True,
|
||||
subprocess_timeout: float = 3600.0,
|
||||
):
|
||||
super().__init__()
|
||||
self._last_records = None
|
||||
self.include_headers = include_headers
|
||||
self.include_body = include_body
|
||||
self.include_agentlightning_requests = include_agentlightning_requests
|
||||
self.subprocess_mode = subprocess_mode
|
||||
self.subprocess_timeout = subprocess_timeout
|
||||
|
||||
def init_worker(self, worker_id: int):
|
||||
"""
|
||||
Initialize the tracer in a worker process.
|
||||
|
||||
Args:
|
||||
worker_id: The ID of the worker process.
|
||||
"""
|
||||
super().init_worker(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] HttpTracer initialized.")
|
||||
|
||||
@contextmanager
|
||||
def trace_context(self, name: Optional[str] = None) -> Iterator[HTTPRecords]:
|
||||
"""
|
||||
Starts a new HTTP tracing context. This should be used as a context manager.
|
||||
|
||||
Args:
|
||||
name: Optional name for the tracing context.
|
||||
|
||||
Yields:
|
||||
The HTTPRecords instance containing traced HTTP activities.
|
||||
"""
|
||||
records = HTTPRecords()
|
||||
with httprecord(records):
|
||||
self._last_records = records
|
||||
yield records
|
||||
|
||||
def get_last_trace(self) -> List[ReadableSpan]:
|
||||
"""
|
||||
Retrieves the raw list of captured spans from the most recent trace.
|
||||
|
||||
Returns:
|
||||
A list of OpenTelemetry `ReadableSpan` objects converted from HTTP records.
|
||||
"""
|
||||
if self._last_records is None:
|
||||
return []
|
||||
|
||||
return self._convert_to_spans(self._last_records)
|
||||
|
||||
def _convert_to_spans(self, records: HTTPRecords) -> List[ReadableSpan]:
|
||||
"""
|
||||
Convert HTTPRecords to OpenTelemetry spans.
|
||||
|
||||
Args:
|
||||
records: The HTTPRecords instance containing HTTP traces.
|
||||
|
||||
Returns:
|
||||
A list of ReadableSpan objects representing the HTTP activities.
|
||||
"""
|
||||
spans = []
|
||||
|
||||
# Create a trace ID that will be shared by all spans in this trace
|
||||
trace_id = int(uuid.uuid4().hex[:16], 16)
|
||||
|
||||
for record in records.requests.values():
|
||||
# Skip AgentLightning requests if include_agentlightning_requests is False
|
||||
should_skip = False
|
||||
if not self.include_agentlightning_requests and record.request and record.request.headers:
|
||||
for header in record.request.headers:
|
||||
if header.name.lower() in self.AGENTLIGHTNING_HEADERS and header.value.lower() == "true":
|
||||
should_skip = True
|
||||
break
|
||||
|
||||
if should_skip:
|
||||
continue
|
||||
|
||||
# Create a span ID for this specific HTTP request
|
||||
span_id = int(uuid.uuid4().hex[:8], 16)
|
||||
|
||||
# Create a span context
|
||||
span_context = SpanContext(
|
||||
trace_id=trace_id,
|
||||
span_id=span_id,
|
||||
is_remote=False,
|
||||
trace_flags=TraceFlags(TraceFlags.SAMPLED),
|
||||
trace_state=TraceState(),
|
||||
)
|
||||
|
||||
# Extract important information from the HTTP record
|
||||
method = record.method
|
||||
url = record.url
|
||||
parsed_url = urlparse(url)
|
||||
status_code = record.status_code
|
||||
|
||||
# Create attributes dictionary
|
||||
attributes: Dict[str, Any] = {
|
||||
"http.method": method,
|
||||
"http.url": url,
|
||||
"http.target": parsed_url.path,
|
||||
"http.host": parsed_url.netloc,
|
||||
}
|
||||
|
||||
if status_code is not None and status_code > 0:
|
||||
attributes["http.status_code"] = status_code
|
||||
|
||||
# Calculate duration - from begin time to last update
|
||||
duration = None
|
||||
if hasattr(record, "last_update") and record.last_update and record.tbegin:
|
||||
duration = (record.last_update - record.tbegin).total_seconds()
|
||||
attributes["http.duration_ms"] = duration * 1000 # Convert to ms
|
||||
|
||||
# Optionally include headers
|
||||
if self.include_headers and record.request and record.request.headers:
|
||||
for header in record.request.headers:
|
||||
header_name = header.name.lower()
|
||||
attributes[f"http.request.header.{header_name}"] = header.value
|
||||
|
||||
if self.include_headers and record.response and record.response.headers:
|
||||
for header in record.response.headers:
|
||||
header_name = header.name.lower()
|
||||
attributes[f"http.response.header.{header_name}"] = header.value
|
||||
|
||||
# Optionally include body - preserve complete content for analysis
|
||||
if self.include_body and record.request:
|
||||
body_content = record.request.content
|
||||
if body_content:
|
||||
# Store raw body content for later parsing/analysis
|
||||
attributes["http.request.body"] = body_content
|
||||
|
||||
if self.include_body and record.response:
|
||||
body_content = record.response.content
|
||||
if body_content:
|
||||
# Store raw body content for later parsing/analysis
|
||||
attributes["http.response.body"] = body_content
|
||||
|
||||
# Determine span status
|
||||
span_status = StatusCode.OK
|
||||
if status_code and status_code >= 400 or record.exception:
|
||||
span_status = StatusCode.ERROR
|
||||
|
||||
# Create start and end timestamps in nanoseconds
|
||||
# If we have duration, use it, otherwise default to current time - 1ms
|
||||
start_time_ns = int(record.tbegin.timestamp() * 1e9)
|
||||
if duration:
|
||||
end_time_ns = int((record.tbegin.timestamp() + duration) * 1e9)
|
||||
else:
|
||||
end_time_ns = int(record.last_update.timestamp() * 1e9)
|
||||
|
||||
span = ReadableSpan(
|
||||
name=f"HTTP {method} {url}",
|
||||
context=span_context,
|
||||
parent=None,
|
||||
kind=SpanKind.CLIENT,
|
||||
status=Status(span_status),
|
||||
start_time=start_time_ns,
|
||||
end_time=end_time_ns,
|
||||
attributes=attributes,
|
||||
events=[],
|
||||
links=[],
|
||||
resource=None,
|
||||
)
|
||||
|
||||
spans.append(span)
|
||||
|
||||
return spans
|
||||
|
||||
def trace_run(self, func: Callable, *args, **kwargs) -> Any:
|
||||
"""
|
||||
A convenience wrapper to trace the execution of a single synchronous function.
|
||||
|
||||
If subprocess_mode is enabled, the function will be executed in an isolated subprocess
|
||||
to prevent HTTP hooks from affecting the parent process.
|
||||
|
||||
Args:
|
||||
func: The synchronous function to execute and trace.
|
||||
*args: Positional arguments to pass to the function.
|
||||
**kwargs: Keyword arguments to pass to the function.
|
||||
|
||||
Returns:
|
||||
The return value of the function.
|
||||
"""
|
||||
if self.subprocess_mode:
|
||||
return self._trace_run_subprocess(func, args, kwargs)
|
||||
else:
|
||||
return super().trace_run(func, *args, **kwargs)
|
||||
|
||||
async def trace_run_async(self, func: Callable[..., Awaitable], *args, **kwargs) -> Any:
|
||||
"""
|
||||
A convenience wrapper to trace the execution of a single asynchronous function.
|
||||
|
||||
If subprocess_mode is enabled, the function will be executed in an isolated subprocess
|
||||
to prevent HTTP hooks from affecting the parent process.
|
||||
|
||||
Args:
|
||||
func: The asynchronous function to execute and trace.
|
||||
*args: Positional arguments to pass to the function.
|
||||
**kwargs: Keyword arguments to pass to the function.
|
||||
|
||||
Returns:
|
||||
The return value of the function.
|
||||
"""
|
||||
if self.subprocess_mode:
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(
|
||||
None, self._trace_run_subprocess, func, args, kwargs, True # True for async
|
||||
)
|
||||
else:
|
||||
return await super().trace_run_async(func, *args, **kwargs)
|
||||
|
||||
def _trace_run_subprocess(self, func: Callable, args=None, kwargs=None, is_async: bool = False) -> Any:
|
||||
"""
|
||||
Execute a function in a subprocess with HTTP tracing.
|
||||
|
||||
Args:
|
||||
func: The function to execute.
|
||||
args: Positional arguments to pass to the function.
|
||||
kwargs: Keyword arguments to pass to the function.
|
||||
is_async: Whether the function is asynchronous.
|
||||
|
||||
Returns:
|
||||
The return value of the function.
|
||||
"""
|
||||
if args is None:
|
||||
args = ()
|
||||
if kwargs is None:
|
||||
kwargs = {}
|
||||
|
||||
# Create a queue to receive results from the subprocess
|
||||
result_queue = multiprocessing.Queue()
|
||||
|
||||
# Create and start the subprocess
|
||||
process = multiprocessing.Process(
|
||||
target=self._subprocess_worker, args=(func, args, kwargs, result_queue, is_async)
|
||||
)
|
||||
process.start()
|
||||
|
||||
try:
|
||||
# Wait for the process to complete and get the result
|
||||
process.join(timeout=self.subprocess_timeout)
|
||||
result = result_queue.get_nowait()
|
||||
|
||||
if result["success"]:
|
||||
# Store the captured records for get_last_trace()
|
||||
self._last_records = result["records"]
|
||||
return result["return_value"]
|
||||
else:
|
||||
if "records" in result:
|
||||
self._last_records = result["records"]
|
||||
# Re-raise the exception that occurred in the subprocess
|
||||
raise result["exception"]
|
||||
|
||||
except multiprocessing.TimeoutError:
|
||||
process.terminate()
|
||||
process.join()
|
||||
raise TimeoutError(f"Subprocess execution timed out after {self.subprocess_timeout} seconds.")
|
||||
except queue.Empty:
|
||||
logger.error("Traced result is empty. This may indicate a timeout or an issue with the subprocess.")
|
||||
finally:
|
||||
if process.is_alive():
|
||||
process.terminate()
|
||||
process.join()
|
||||
|
||||
def _subprocess_worker(self, func: Callable, args, kwargs, result_queue: multiprocessing.Queue, is_async: bool):
|
||||
"""
|
||||
Worker function that runs in the subprocess to execute the traced function.
|
||||
|
||||
Args:
|
||||
func: The function to execute.
|
||||
args: Positional arguments.
|
||||
kwargs: Keyword arguments.
|
||||
result_queue: Queue to send results back to parent process.
|
||||
is_async: Whether the function is asynchronous.
|
||||
"""
|
||||
# Create a new tracer instance in the subprocess (without subprocess mode to avoid recursion)
|
||||
subprocess_tracer = HttpTracer(
|
||||
include_headers=self.include_headers,
|
||||
include_body=self.include_body,
|
||||
include_agentlightning_requests=self.include_agentlightning_requests,
|
||||
subprocess_mode=False, # Disable subprocess mode in the worker
|
||||
)
|
||||
|
||||
try:
|
||||
if is_async:
|
||||
# Run async function in new event loop
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
try:
|
||||
return_value = loop.run_until_complete(subprocess_tracer.trace_run_async(func, *args, **kwargs))
|
||||
finally:
|
||||
loop.close()
|
||||
else:
|
||||
# Run sync function
|
||||
return_value = subprocess_tracer.trace_run(func, *args, **kwargs)
|
||||
|
||||
# Get the captured records
|
||||
records = subprocess_tracer._last_records
|
||||
|
||||
# Send success result back to parent
|
||||
result_queue.put({"success": True, "return_value": return_value, "records": records})
|
||||
|
||||
except Exception as e:
|
||||
# Log the exception
|
||||
logger.exception(f"Error in subprocess worker in http tracer: {e}")
|
||||
|
||||
# Get the captured records even when there's an exception
|
||||
records = subprocess_tracer._last_records
|
||||
# Send error result back to parent
|
||||
result_queue.put({"success": False, "exception": e, "records": records})
|
||||
@@ -1,539 +0,0 @@
|
||||
import json
|
||||
import re
|
||||
from enum import Enum
|
||||
from typing import List, Dict, Tuple, Optional, Any
|
||||
|
||||
from pydantic import BaseModel
|
||||
from opentelemetry import trace as trace_api
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
from agentlightning.types import Triplet
|
||||
|
||||
|
||||
class Transition(BaseModel):
|
||||
"""
|
||||
Transition class representing one transition in a trajectory.
|
||||
State and action are a list of token IDs.
|
||||
"""
|
||||
|
||||
state: List[int]
|
||||
action: List[int]
|
||||
response_id: Optional[str]
|
||||
# action_logprobs: List[float]
|
||||
agent_name: str
|
||||
reward: Optional[float]
|
||||
|
||||
|
||||
class RewardMatchPolicy(str, Enum):
|
||||
"""How to find the reward for each transition from the trace.
|
||||
In all cases, the reward must have data `{"type": "reward", "value": <float>|None}`,
|
||||
as defined in `reward.py`.
|
||||
"""
|
||||
|
||||
FIRST_SIBLING = "first_sibling"
|
||||
"""Use the first sibling in the current trace subtree as the reward, except another LLM call match is found."""
|
||||
|
||||
FIRST_OCCURRENCE = "first_occurrence"
|
||||
"""Use the first occurrence of the reward (in start time order) that occur after the current LLM call match.
|
||||
"""
|
||||
|
||||
|
||||
class TraceTree:
|
||||
"""
|
||||
A trace item, along with its span and children.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
id: str,
|
||||
span: ReadableSpan,
|
||||
children: Optional[List["TraceTree"]] = None,
|
||||
):
|
||||
self.id = id
|
||||
self.span = span
|
||||
self.children = children or []
|
||||
|
||||
@property
|
||||
def start_time(self):
|
||||
return self.span.start_time
|
||||
|
||||
@property
|
||||
def end_time(self):
|
||||
return self.span.end_time
|
||||
|
||||
def find_id(self, id: str) -> "TraceTree | None":
|
||||
if self.id == id:
|
||||
return self
|
||||
for child in self.children:
|
||||
found = child.find_id(id)
|
||||
if found:
|
||||
return found
|
||||
return None
|
||||
|
||||
def add_child(self, child: "TraceTree") -> None:
|
||||
self.children.append(child)
|
||||
|
||||
def visualize(self, filename: str, interested_span_match: str | None = None) -> None:
|
||||
"""
|
||||
Visualize the trace tree using graphviz.
|
||||
For debugging purposes only.
|
||||
Use `interested_span_match` to filter the spans (and its ancesters) to be visualized.
|
||||
"""
|
||||
import graphviz
|
||||
|
||||
dot = graphviz.Digraph(comment="Trace Tree")
|
||||
|
||||
should_visit_cache = {}
|
||||
|
||||
def should_visit(node: "TraceTree") -> bool:
|
||||
if node.id in should_visit_cache:
|
||||
return should_visit_cache[node.id]
|
||||
if interested_span_match is not None:
|
||||
if re.search(interested_span_match, node.span.name):
|
||||
should_visit_cache[node.id] = True
|
||||
return True
|
||||
else:
|
||||
should_visit_cache[node.id] = False
|
||||
for child in node.children:
|
||||
if should_visit(child):
|
||||
should_visit_cache[node.id] = True
|
||||
|
||||
return should_visit_cache[node.id]
|
||||
else:
|
||||
return True
|
||||
|
||||
def visit(node: "TraceTree") -> bool:
|
||||
if not should_visit(node):
|
||||
return False
|
||||
agent_name = node.agent_name()
|
||||
vis_name = node.id[:8] + " (" + node.span.name + ")"
|
||||
if agent_name is not None:
|
||||
vis_name += " [" + agent_name + "]"
|
||||
dot.node(node.id, vis_name)
|
||||
for child in node.children:
|
||||
if visit(child):
|
||||
dot.edge(node.id, child.id)
|
||||
return True
|
||||
|
||||
visit(self)
|
||||
dot.render(filename, format="png", cleanup=True)
|
||||
|
||||
def names_tuple(self) -> Tuple[str, List[Any]]:
|
||||
"""Return the span name, and a list of children.
|
||||
Each child is also a tuple of span name and a list of children.
|
||||
Useful for debugging and testing.
|
||||
"""
|
||||
name = self.span.name
|
||||
agent_name = self.agent_name()
|
||||
if agent_name is not None:
|
||||
name += " [" + agent_name + "]"
|
||||
children_names = []
|
||||
for child in self.children:
|
||||
child_name, child_children = child.names_tuple()
|
||||
children_names.append((child_name, child_children))
|
||||
return name, children_names
|
||||
|
||||
def traverse(self) -> List["TraceTree"]:
|
||||
"""
|
||||
Traverse the trace tree and return a list of all spans.
|
||||
"""
|
||||
spans: List["TraceTree"] = [self]
|
||||
for child in self.children:
|
||||
spans.extend(child.traverse())
|
||||
return spans
|
||||
|
||||
def to_json(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"span": self.span.to_json(),
|
||||
"children": [child.to_json() for child in self.children],
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_spans(cls, spans: List[ReadableSpan]) -> "TraceTree":
|
||||
"""
|
||||
Create a TraceTree from a list of spans.
|
||||
All spans without parents found will be considered as candidate root spans.
|
||||
If multiple root spans are found, a virtual root span will be created as the parent of all root spans.
|
||||
"""
|
||||
|
||||
if not spans:
|
||||
raise ValueError("No spans provided to create TraceTree.")
|
||||
|
||||
# Process trace items in topological order
|
||||
id_to_span = {span.get_span_context().span_id: span for span in spans}
|
||||
|
||||
forward_graph: dict[int, list[int]] = {}
|
||||
root_ids: list[int] = []
|
||||
for span in spans:
|
||||
if span.parent is None:
|
||||
root_ids.append(span.get_span_context().span_id)
|
||||
else:
|
||||
if span.parent.span_id not in forward_graph:
|
||||
forward_graph[span.parent.span_id] = []
|
||||
forward_graph[span.parent.span_id].append(span.get_span_context().span_id)
|
||||
|
||||
# Diff between span with data and forward_graph keys
|
||||
# Sometimes the top-level session span is lost.
|
||||
unfound_roots = set(forward_graph.keys()) - set(id_to_span.keys())
|
||||
for unfound_root in unfound_roots:
|
||||
root_ids.append(unfound_root)
|
||||
|
||||
def visit(node_id):
|
||||
children: list[TraceTree] = []
|
||||
if node_id in forward_graph:
|
||||
for child_id in forward_graph[node_id]:
|
||||
children.append(visit(child_id))
|
||||
|
||||
if node_id not in id_to_span:
|
||||
assert len(children) > 0
|
||||
virtual_span = ReadableSpan(
|
||||
context=trace_api.SpanContext(
|
||||
trace_id=children[0].span.get_span_context().trace_id,
|
||||
span_id=node_id,
|
||||
is_remote=False,
|
||||
),
|
||||
name="virtual-node",
|
||||
kind=trace_api.SpanKind.INTERNAL,
|
||||
attributes={},
|
||||
start_time=min(child.start_time for child in children),
|
||||
end_time=max(child.end_time for child in children),
|
||||
)
|
||||
return cls(trace_api.format_span_id(node_id), virtual_span, children=children)
|
||||
else:
|
||||
return cls(
|
||||
trace_api.format_span_id(node_id),
|
||||
id_to_span[node_id],
|
||||
children=children,
|
||||
)
|
||||
|
||||
# Create a virtual root span if multiple root spans are found
|
||||
if len(root_ids) > 1:
|
||||
root_spans = [visit(root_id) for root_id in root_ids]
|
||||
virtual_root = TraceTree(
|
||||
id="virtual-root",
|
||||
span=ReadableSpan(
|
||||
context=trace_api.SpanContext(
|
||||
trace_id=root_spans[0].span.get_span_context().trace_id,
|
||||
span_id=0,
|
||||
is_remote=False,
|
||||
),
|
||||
name="virtual-root",
|
||||
kind=trace_api.SpanKind.INTERNAL,
|
||||
attributes={},
|
||||
start_time=root_spans[0].start_time,
|
||||
end_time=root_spans[-1].end_time,
|
||||
),
|
||||
children=root_spans,
|
||||
)
|
||||
return virtual_root
|
||||
elif len(root_ids) == 0:
|
||||
# No root spans found
|
||||
raise ValueError("No root spans found in the trace.")
|
||||
else:
|
||||
root_span = visit(root_ids[0])
|
||||
return root_span
|
||||
|
||||
def agent_name(self) -> Optional[str]:
|
||||
"""Return the name of agent span. Return the agent or None (not an agent at all).
|
||||
Extend this function to support more agent frameworks."""
|
||||
|
||||
# Case 1: OpenAI Agent SDK
|
||||
agent_name = self.span.attributes.get("agent.name")
|
||||
if agent_name is not None:
|
||||
return agent_name
|
||||
|
||||
# Case 2: Agentops decorator @agent
|
||||
is_agent = self.span.attributes.get("agentops.span.kind") == "agent"
|
||||
if is_agent:
|
||||
agent_name = self.span.attributes.get("operation.name")
|
||||
if agent_name is not None:
|
||||
return agent_name
|
||||
|
||||
# Case 3: Autogen team
|
||||
agent_name = self.span.attributes.get("recipient_agent_type")
|
||||
if agent_name is not None:
|
||||
return agent_name
|
||||
|
||||
# Case 4: LangGraph
|
||||
agent_name = self.span.attributes.get("langchain.chain.type")
|
||||
if agent_name is not None:
|
||||
return agent_name
|
||||
|
||||
def maybe_reward_dict(self) -> dict[str, Any]:
|
||||
for key in [
|
||||
"agentops.task.output", # newer versions of agentops
|
||||
"agentops.entity.output",
|
||||
]:
|
||||
output = self.span.attributes.get(key)
|
||||
if output:
|
||||
if isinstance(output, dict):
|
||||
return output
|
||||
elif isinstance(output, str):
|
||||
try:
|
||||
return json.loads(output)
|
||||
except json.JSONDecodeError:
|
||||
return {}
|
||||
return {}
|
||||
|
||||
def is_reward_span(self) -> bool:
|
||||
maybe_reward = self.maybe_reward_dict()
|
||||
return maybe_reward and maybe_reward.get("type") == "reward"
|
||||
|
||||
def find_llm_calls(
|
||||
self,
|
||||
*,
|
||||
llm_call_match: str,
|
||||
agent_match: Optional[str],
|
||||
within_matching_subtree: str | None = None,
|
||||
within_reward: Optional[bool] = None,
|
||||
within_llm_call: Optional[bool] = None,
|
||||
existing_llm_call_response_ids: Optional[set[str]] = None,
|
||||
) -> List[Tuple["TraceTree", str]]:
|
||||
"""Find all LLM calls in the trace tree.
|
||||
|
||||
The LLM call is defined as a span with type = request and name matching `llm_call_match`.
|
||||
If `agent_match` is not None, it must also reside in an agent span (type = agent) with name matched.
|
||||
|
||||
Return a list of traces and the agent names (why it's selected).
|
||||
"""
|
||||
llm_calls: List[Tuple[TraceTree, str]] = []
|
||||
|
||||
is_llm_call = True
|
||||
if within_matching_subtree is None or within_reward is True:
|
||||
# We must be in an interesting agent subtree, and not in a reward span.
|
||||
is_llm_call = False
|
||||
if re.search(llm_call_match, self.span.name) is None:
|
||||
# The span name does not match the LLM call match.
|
||||
is_llm_call = False
|
||||
if is_llm_call:
|
||||
# Check the response id
|
||||
response_id = self.span.attributes.get("gen_ai.response.id")
|
||||
if response_id is None and within_llm_call is True:
|
||||
is_llm_call = False
|
||||
if (
|
||||
response_id is not None
|
||||
and existing_llm_call_response_ids is not None
|
||||
and response_id in existing_llm_call_response_ids
|
||||
):
|
||||
is_llm_call = False
|
||||
|
||||
if is_llm_call:
|
||||
llm_calls.append((self, within_matching_subtree))
|
||||
existing_llm_call_response_ids = existing_llm_call_response_ids or set()
|
||||
if response_id is not None:
|
||||
existing_llm_call_response_ids.add(response_id)
|
||||
if within_llm_call is not None:
|
||||
within_llm_call = True
|
||||
|
||||
agent_name = self.agent_name()
|
||||
if agent_name is not None:
|
||||
if agent_match is None or re.search(agent_match, agent_name):
|
||||
within_matching_subtree = agent_name
|
||||
else:
|
||||
within_matching_subtree = None
|
||||
|
||||
if within_reward is not None and self.is_reward_span():
|
||||
within_reward = True
|
||||
|
||||
for child in self.children:
|
||||
llm_calls.extend(
|
||||
child.find_llm_calls(
|
||||
llm_call_match=llm_call_match,
|
||||
agent_match=agent_match,
|
||||
within_matching_subtree=within_matching_subtree,
|
||||
within_reward=within_reward,
|
||||
within_llm_call=within_llm_call,
|
||||
existing_llm_call_response_ids=existing_llm_call_response_ids,
|
||||
)
|
||||
)
|
||||
|
||||
return llm_calls
|
||||
|
||||
def repair_hierarchy(self) -> None:
|
||||
"""
|
||||
We find that sometimes the hierarchy is not correct, due to the way the spans are created.
|
||||
The spans within the agent frameworks (e.g., OpenAI Agent SDK) and spans within the LLM frameworks
|
||||
(e.g., Anthropic) are created in two systems.
|
||||
So the inner LLM completion span does not necessarily have an agent span as a parent.
|
||||
Rather they sometimes directly become children of the root span.
|
||||
This becomes a problem when we want to select the LLM completion span with agent as filter.
|
||||
To repair the hierarchy, for each children of the root span, we find a span over the whole tree,
|
||||
with duration covering the current span and being closest to the current span.
|
||||
|
||||
This function modifies the tree in place.
|
||||
"""
|
||||
nodes_to_repair = list(self.children)
|
||||
for repair_node in nodes_to_repair:
|
||||
if len(self.children) == 1:
|
||||
# If there is only one child, we don't need to repair the hierarchy.
|
||||
break
|
||||
# Find the closest parent span (but not the root itself)
|
||||
closest_parent = None
|
||||
closest_duration = float("inf")
|
||||
for node in self.traverse():
|
||||
if node.id == repair_node.id:
|
||||
continue
|
||||
if node is self:
|
||||
continue
|
||||
if node.start_time <= repair_node.start_time and node.end_time >= repair_node.end_time:
|
||||
duration_delta = node.end_time - repair_node.end_time + repair_node.start_time - node.start_time
|
||||
if duration_delta > 0 and duration_delta < closest_duration:
|
||||
closest_duration = duration_delta
|
||||
closest_parent = node
|
||||
|
||||
# Repair the hierarchy
|
||||
if closest_parent is not None:
|
||||
self.children.remove(repair_node)
|
||||
closest_parent.children.append(repair_node)
|
||||
|
||||
def match_rewards(self, reward_match: str, llm_calls: List["TraceTree"]) -> dict[str, Optional[float]]:
|
||||
"""Match the rewards to the LLM calls."""
|
||||
llm_call_ids = set([llm_call.id for llm_call in llm_calls])
|
||||
rewards: dict[str, Optional[float]] = {}
|
||||
|
||||
if reward_match == RewardMatchPolicy.FIRST_OCCURRENCE:
|
||||
time_sorted: List[TraceTree] = sorted(self.traverse(), key=lambda x: x.start_time)
|
||||
assign_to: List[Tuple[str, int]] = []
|
||||
for item in time_sorted:
|
||||
if item.id in llm_call_ids:
|
||||
assign_to.append((item.id, item.end_time))
|
||||
|
||||
# get reward
|
||||
agentops_output = item.maybe_reward_dict()
|
||||
if agentops_output and agentops_output.get("type") == "reward":
|
||||
for assign_to_id, assign_to_end_time in reversed(assign_to):
|
||||
# This reward happens before the end of the LLM call.
|
||||
if assign_to_end_time > item.start_time:
|
||||
continue
|
||||
# Ok, we found someone to assign to
|
||||
if assign_to_id in rewards:
|
||||
# If the reward is already set, skip
|
||||
continue
|
||||
rewards[assign_to_id] = agentops_output.get("value", None)
|
||||
break
|
||||
|
||||
elif reward_match == RewardMatchPolicy.FIRST_SIBLING:
|
||||
for item in self.traverse():
|
||||
assign_to: List[Tuple[str, int]] = []
|
||||
for child in item.children:
|
||||
if child.id in llm_call_ids:
|
||||
assign_to.append(child.id)
|
||||
|
||||
agentops_output = item.maybe_reward_dict()
|
||||
if agentops_output and agentops_output.get("type") == "reward":
|
||||
for assign_to_id, assign_to_end_time in reversed(assign_to):
|
||||
if assign_to_end_time > item.start_time:
|
||||
# This reward happens before the end of the LLM call.
|
||||
continue
|
||||
if assign_to_id in rewards:
|
||||
continue
|
||||
rewards[assign_to_id] = agentops_output.get("value", None)
|
||||
break
|
||||
|
||||
return rewards
|
||||
|
||||
def to_trajectory(
|
||||
self,
|
||||
llm_call_match: str = r"openai\.chat\.completion",
|
||||
agent_match: Optional[str] = None,
|
||||
exclude_llm_call_in_reward: bool = True,
|
||||
dedup_llm_call: bool = True,
|
||||
reward_match: RewardMatchPolicy = RewardMatchPolicy.FIRST_OCCURRENCE,
|
||||
final_reward: Optional[float] = None,
|
||||
) -> List[Triplet]:
|
||||
"""Convert the trace tree to a trajectory.
|
||||
|
||||
First, we find all the LLM calls (span type = request, `llm_call_match` matching the span name).
|
||||
If the agent match is set, we check, for each LLM call,
|
||||
if it resides in an agent (span type = agent, `agent_match` matching the span name).
|
||||
The above sets the basis for the trajectory, as we use the prompt token IDs and response token IDs for each LLM call,
|
||||
as the state and action of each transition.
|
||||
|
||||
Then, we find the reward for each transition.
|
||||
The reward is searched on the trace tree, after the LLM call,
|
||||
until the next LLM call or the end of the tree depending on the policy.
|
||||
It can be enforced to a sibling or the first occurrence in the time order, depending on the policy.
|
||||
If a reward is never found for a transition, it is set to None.
|
||||
"""
|
||||
# Find all LLM calls
|
||||
llm_calls = self.find_llm_calls(
|
||||
llm_call_match=llm_call_match,
|
||||
agent_match=agent_match,
|
||||
within_matching_subtree="*" if agent_match is None else None,
|
||||
within_reward=False if exclude_llm_call_in_reward else None,
|
||||
within_llm_call=False if dedup_llm_call else None,
|
||||
existing_llm_call_response_ids=set(),
|
||||
)
|
||||
id_transitions = [
|
||||
(
|
||||
llm_call.id,
|
||||
Triplet(
|
||||
prompt={"token_ids": llm_call.span.attributes.get("prompt_token_ids", [])},
|
||||
response={"token_ids": llm_call.span.attributes.get("response_token_ids", [])},
|
||||
reward=None,
|
||||
metadata=dict(
|
||||
response_id=llm_call.span.attributes.get(
|
||||
"gen_ai.response.id", None
|
||||
), # it works at least for OpenAI
|
||||
agent_name=agent_name,
|
||||
),
|
||||
),
|
||||
)
|
||||
for llm_call, agent_name in llm_calls
|
||||
]
|
||||
|
||||
rewards = self.match_rewards(reward_match, [call for call, _ in llm_calls])
|
||||
transitions = [
|
||||
transition.model_copy(update={"reward": rewards.get(id, None)}) for id, transition in id_transitions
|
||||
]
|
||||
if final_reward is not None and len(transitions) > 0:
|
||||
# Add the final reward to the last transition
|
||||
transitions[-1] = transitions[-1].model_copy(update={"reward": final_reward})
|
||||
return transitions
|
||||
|
||||
def __repr__(self):
|
||||
return (
|
||||
f"TraceTree(id={self.id}, span={self.span}, start_time={self.start_time}, "
|
||||
+ f"end_time={self.end_time}, children={self.children})"
|
||||
)
|
||||
|
||||
|
||||
class TripletExporter:
|
||||
"""
|
||||
A class to export triplet data from OpenTelemetry spans.
|
||||
|
||||
Attributes:
|
||||
repair_hierarchy: When `repair_hierarchy` is set to True, the trace will be repaired with the time information.
|
||||
See `TraceTree.repair_hierarchy` for more details.
|
||||
llm_call_match: Regular expression pattern to match LLM call span names.
|
||||
agent_match: Optional regular expression pattern to match agent span names. If None, all agents are matched.
|
||||
exclude_llm_call_in_reward: Whether to exclude LLM calls that occur within reward spans.
|
||||
reward_match: Policy for matching rewards to LLM calls.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
repair_hierarchy: bool = True,
|
||||
llm_call_match: str = r"openai\.chat\.completion",
|
||||
agent_match: Optional[str] = None,
|
||||
exclude_llm_call_in_reward: bool = True,
|
||||
reward_match: RewardMatchPolicy = RewardMatchPolicy.FIRST_OCCURRENCE,
|
||||
):
|
||||
self.repair_hierarchy = repair_hierarchy
|
||||
self.llm_call_match = llm_call_match
|
||||
self.agent_match = agent_match
|
||||
self.exclude_llm_call_in_reward = exclude_llm_call_in_reward
|
||||
self.reward_match = reward_match
|
||||
|
||||
def export(self, spans: List[ReadableSpan]) -> List[Triplet]:
|
||||
"""Convert OpenTelemetry spans to a list of Triplet objects."""
|
||||
trace_tree = TraceTree.from_spans(spans)
|
||||
if self.repair_hierarchy:
|
||||
trace_tree.repair_hierarchy()
|
||||
trajectory = trace_tree.to_trajectory(
|
||||
llm_call_match=self.llm_call_match,
|
||||
agent_match=self.agent_match,
|
||||
exclude_llm_call_in_reward=self.exclude_llm_call_in_reward,
|
||||
reward_match=self.reward_match,
|
||||
)
|
||||
return trajectory
|
||||
@@ -1,311 +0,0 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import multiprocessing
|
||||
import os
|
||||
import signal
|
||||
import time
|
||||
from typing import List, Optional, Union
|
||||
import importlib
|
||||
|
||||
import agentops
|
||||
|
||||
from .client import AgentLightningClient
|
||||
from .litagent import LitAgent
|
||||
from .runner import AgentRunner
|
||||
from .types import ParallelWorkerBase
|
||||
from .tracer.base import BaseTracer
|
||||
from .tracer.agentops import AgentOpsTracer
|
||||
from .tracer.triplet import TripletExporter
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Trainer(ParallelWorkerBase):
|
||||
"""Orchestrates the distributed execution of agent rollouts.
|
||||
|
||||
The Trainer is responsible for launching one or more worker processes
|
||||
that run the agent's execution loop. It manages multiprocessing,
|
||||
handles graceful shutdown, and serves as the main entry point for
|
||||
running a client-side agent fleet.
|
||||
|
||||
Attributes:
|
||||
dev: If True, rollouts are run against the dev endpoint provided in `fit`.
|
||||
n_workers: Number of agent workers (processes) to run in parallel.
|
||||
max_tasks: Maximum number of tasks to process per worker. If None,
|
||||
workers run until no more tasks are available.
|
||||
daemon: Whether worker processes should be daemons. Daemon processes
|
||||
are terminated automatically when the main process exits.
|
||||
tracer: A tracer instance, or a string pointing to the class full name or a dictionary with a 'type' key
|
||||
that specifies the class full name and other initialization parameters.
|
||||
If None, a default `AgentOpsTracer` will be created with the current settings.
|
||||
triplet_exporter: An instance of `TripletExporter` to export triplets from traces,
|
||||
or a dictionary with the initialization parameters for the exporter.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
dev: bool = False,
|
||||
n_workers: int = 1,
|
||||
max_tasks: Optional[int] = None,
|
||||
daemon: bool = True,
|
||||
tracer: Union[BaseTracer, str, dict, None] = None,
|
||||
triplet_exporter: Union[TripletExporter, dict, None] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.n_workers = n_workers
|
||||
self.max_tasks = max_tasks
|
||||
self.daemon = daemon
|
||||
self.dev = dev
|
||||
self._client: AgentLightningClient | None = None # Will be initialized in fit method
|
||||
|
||||
self.tracer = self._make_tracer(tracer)
|
||||
if isinstance(triplet_exporter, TripletExporter):
|
||||
self.triplet_exporter = triplet_exporter
|
||||
elif isinstance(triplet_exporter, dict):
|
||||
self.triplet_exporter = TripletExporter(**triplet_exporter)
|
||||
elif triplet_exporter is None:
|
||||
self.triplet_exporter = TripletExporter()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid triplet_exporter type: {type(triplet_exporter)}. Expected TripletExporter, dict, or None."
|
||||
)
|
||||
|
||||
if not self.daemon:
|
||||
logger.warning(
|
||||
"daemon=False. Worker processes are non-daemonic. "
|
||||
"The worker processes will NOT be terminated when the main process exits. "
|
||||
"The cleanup must be handled manually."
|
||||
)
|
||||
|
||||
def _make_tracer(self, tracer: Union[BaseTracer, str, dict, None]) -> BaseTracer:
|
||||
"""Creates a tracer instance based on the provided configuration."""
|
||||
if isinstance(tracer, BaseTracer):
|
||||
return tracer
|
||||
if isinstance(tracer, str):
|
||||
module_name, class_name = tracer.rsplit(".", 1)
|
||||
module = importlib.import_module(module_name)
|
||||
tracer_cls = getattr(module, class_name)
|
||||
return tracer_cls()
|
||||
if isinstance(tracer, dict):
|
||||
tracer_type = tracer.get("type")
|
||||
if tracer_type is None:
|
||||
raise ValueError("tracer dict must have a 'type' key with the class full name")
|
||||
module_name, class_name = tracer_type.rsplit(".", 1)
|
||||
module = importlib.import_module(module_name)
|
||||
tracer_cls = getattr(module, class_name)
|
||||
# Remove 'type' key and pass remaining keys as kwargs
|
||||
tracer_kwargs = {k: v for k, v in tracer.items() if k != "type"}
|
||||
return tracer_cls(**tracer_kwargs)
|
||||
if tracer is None:
|
||||
return AgentOpsTracer(agentops_managed=True, instrument_managed=True, daemon=self.daemon)
|
||||
raise ValueError(f"Invalid tracer type: {type(tracer)}. Expected BaseTracer, str, dict, or None.")
|
||||
|
||||
def init(self, backend: Union[str, AgentLightningClient]) -> None:
|
||||
logger.info(f"Initializing Trainer...")
|
||||
|
||||
self._init_client(backend)
|
||||
|
||||
self.tracer.init()
|
||||
|
||||
logger.info(f"Trainer main initialization complete.")
|
||||
|
||||
def teardown(self) -> None:
|
||||
logger.info(f"Cleaning up Trainer...")
|
||||
self.tracer.teardown()
|
||||
|
||||
self._client = None
|
||||
logger.info(f"Trainer main cleanup complete.")
|
||||
|
||||
def client(self) -> AgentLightningClient:
|
||||
"""Returns the AgentLightningClient instance."""
|
||||
if self._client is None:
|
||||
raise RuntimeError("AgentLightningClient has not been initialized. Call `init` first.")
|
||||
return self._client
|
||||
|
||||
def _init_client(self, backend: Union[str, AgentLightningClient]) -> AgentLightningClient:
|
||||
if self._client is None:
|
||||
if isinstance(backend, AgentLightningClient):
|
||||
logger.info("Using provided AgentLightningClient instance.")
|
||||
self._client = backend
|
||||
else:
|
||||
logger.info(f"Initializing AgentLightningClient with endpoint: {backend}")
|
||||
if not isinstance(backend, str):
|
||||
raise ValueError("backend must be a string URL or an AgentLightningClient instance.")
|
||||
if not backend.startswith("http://") and not backend.startswith("https://"):
|
||||
raise ValueError("backend must be a valid URL starting with http:// or https://")
|
||||
# Initialize the client with the provided backend URL
|
||||
self._client = AgentLightningClient(endpoint=backend)
|
||||
else:
|
||||
logger.warning("AgentLightningClient already initialized. Returning existing instance.")
|
||||
return self._client
|
||||
|
||||
def _worker_main_loop(self, agent: LitAgent, worker_id: int, is_async: bool):
|
||||
"""The main function for each worker process.
|
||||
|
||||
This function initializes the client and the loop, then starts the
|
||||
execution. It also configures process-specific settings like the
|
||||
process title and signal handling.
|
||||
|
||||
Args:
|
||||
agent: The `LitAgent` instance to run.
|
||||
worker_id: The unique ID for this worker.
|
||||
is_async: A boolean indicating if the async loop should be run.
|
||||
"""
|
||||
if self.n_workers > 1:
|
||||
import setproctitle
|
||||
|
||||
# Ignore Ctrl+C in worker processes; the main process handles it
|
||||
signal.signal(signal.SIGINT, signal.SIG_IGN)
|
||||
setproctitle.setproctitle(multiprocessing.current_process().name)
|
||||
|
||||
# Now we are in child processes, so we can safely set up the environment.
|
||||
agent.set_trainer(self)
|
||||
# TODO: this should be set elsewhere
|
||||
if agent.trained_agents:
|
||||
self.triplet_exporter.agent_match = agent.trained_agents
|
||||
self._initialize_worker_env(worker_id)
|
||||
|
||||
mode = "Async" if is_async else "Sync"
|
||||
logger.info(f"[Worker {worker_id}] {mode} worker process started.")
|
||||
|
||||
num_processed = 0
|
||||
|
||||
try:
|
||||
client = self.client()
|
||||
loop = AgentRunner(
|
||||
agent=agent,
|
||||
client=client,
|
||||
tracer=self.tracer,
|
||||
triplet_exporter=self.triplet_exporter,
|
||||
max_tasks=self.max_tasks,
|
||||
worker_id=worker_id,
|
||||
)
|
||||
loop.init_worker(worker_id)
|
||||
if is_async:
|
||||
num_processed = asyncio.run(loop.iter_async())
|
||||
else:
|
||||
num_processed = loop.iter()
|
||||
except Exception:
|
||||
logger.exception(f"[Worker {worker_id}] Unhandled exception in worker loop.")
|
||||
finally:
|
||||
self._teardown_worker_env(worker_id)
|
||||
|
||||
return num_processed
|
||||
|
||||
def _initialize_worker_env(self, worker_id: int):
|
||||
logger.info(f"[Worker {worker_id}] Setting up trainer environment...") # worker_id included in process name
|
||||
self.tracer.init_worker(worker_id)
|
||||
|
||||
def _teardown_worker_env(self, worker_id: int):
|
||||
logger.info(f"[Worker {worker_id}] Cleaning up trainer environment...")
|
||||
self.tracer.teardown_worker(worker_id)
|
||||
logger.info(f"[Worker {worker_id}] Environment cleanup complete.")
|
||||
|
||||
@staticmethod
|
||||
def kill_orphaned_processes() -> None:
|
||||
"""
|
||||
Kill any orphaned processes that may have been left behind by previous runs.
|
||||
This is useful for cleaning up after crashes or unexpected exits.
|
||||
"""
|
||||
import psutil
|
||||
|
||||
for proc in psutil.process_iter():
|
||||
# check whether the process name matches
|
||||
if proc.name().startswith("AgentLightning-"):
|
||||
proc.kill()
|
||||
|
||||
def fit(
|
||||
self,
|
||||
agent: LitAgent,
|
||||
backend: Union[str, AgentLightningClient],
|
||||
dev_backend: Union[str, AgentLightningClient, None] = None,
|
||||
):
|
||||
if self.dev:
|
||||
if dev_backend is None:
|
||||
raise ValueError("dev_backend must be provided when dev=True.")
|
||||
logger.warning(f"Running in dev mode. Using dev backend: {dev_backend}")
|
||||
self.init(dev_backend)
|
||||
else:
|
||||
logger.debug(f"Running in non-dev mode. Using backend: {backend}")
|
||||
self.init(backend)
|
||||
|
||||
processes: List[multiprocessing.Process] = []
|
||||
|
||||
# Determine if the agent is asynchronous.
|
||||
is_async = (
|
||||
hasattr(agent, "training_rollout_async")
|
||||
and agent.__class__.training_rollout_async is not LitAgent.training_rollout_async
|
||||
)
|
||||
|
||||
mode = "asynchronous" if is_async else "synchronous"
|
||||
|
||||
try:
|
||||
if self.n_workers == 1:
|
||||
logger.info(f"Running with n_workers=1 ({mode} in main process).")
|
||||
num_tasks = self._worker_main_loop(agent, 0, is_async)
|
||||
logger.info(f"Single worker mode finished. Tasks processed: {num_tasks}")
|
||||
else:
|
||||
logger.info(f"Running with n_workers={self.n_workers} ({mode} multiprocessing).")
|
||||
for i in range(self.n_workers):
|
||||
process_name = f"AgentLightning-Worker-{i}"
|
||||
p = multiprocessing.Process(
|
||||
target=self._worker_main_loop,
|
||||
args=(agent, i, is_async),
|
||||
daemon=self.daemon,
|
||||
name=process_name,
|
||||
)
|
||||
processes.append(p)
|
||||
logger.info(f"Starting worker process {i} (name: {process_name})...")
|
||||
p.start()
|
||||
|
||||
if self.daemon:
|
||||
for i, p in enumerate(processes):
|
||||
p.join() # Wait for the process to complete
|
||||
logger.info(
|
||||
f"Worker process {i} (name: {p.name}, PID: {p.pid}) joined with exit code {p.exitcode}."
|
||||
)
|
||||
if p.exitcode != 0:
|
||||
logger.warning(
|
||||
f"Worker process {i} (name: {p.name}, PID: {p.pid}) exited with non-zero code: {p.exitcode}."
|
||||
)
|
||||
|
||||
logger.info(f"All {self.n_workers} worker processes have completed.")
|
||||
else:
|
||||
logger.info("All worker processes started. Main process will not wait.")
|
||||
|
||||
# A hack to stop the main process from waiting for child processes to finish.
|
||||
time.sleep(1) # Give workers time to start
|
||||
import multiprocessing.process as multiprocessing_process
|
||||
|
||||
multiprocessing_process._children.clear() # type: ignore
|
||||
|
||||
except KeyboardInterrupt:
|
||||
if self.n_workers > 1 and len(processes) > 0:
|
||||
logger.info(f"KeyboardInterrupt received. Terminating workers...")
|
||||
for i, p in enumerate(processes):
|
||||
if p.is_alive():
|
||||
logger.info(f"Terminating worker {i} (name: {p.name}, PID: {p.pid})...")
|
||||
p.terminate()
|
||||
else:
|
||||
logger.info(
|
||||
f"Worker {i} (name: {p.name}, PID: {p.pid}) is not alive or has already terminated."
|
||||
)
|
||||
for i, p in enumerate(processes):
|
||||
if p.is_alive():
|
||||
p.join(timeout=10) # Give some time to terminate
|
||||
if p.is_alive(): # If still alive, kill
|
||||
logger.warning(
|
||||
f"Worker {i} (name: {p.name}, PID: {p.pid}) did not terminate gracefully, killing..."
|
||||
)
|
||||
p.kill()
|
||||
p.join(timeout=10) # Ensure it's reaped
|
||||
logger.info(f"Workers terminated or single worker interrupted.")
|
||||
except Exception as e:
|
||||
logger.exception(f"Unhandled exception in fit method.")
|
||||
finally:
|
||||
if self.daemon:
|
||||
self.teardown()
|
||||
else:
|
||||
logger.info("Main process exiting. Please use Trainer.kill_orphaned_processes() for cleanup.")
|
||||
@@ -1,204 +0,0 @@
|
||||
from typing import Any, Dict, List, Optional, Union, Literal, Annotated
|
||||
|
||||
from pydantic import BaseModel, Field, Discriminator
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
|
||||
__all__ = [
|
||||
"Triplet",
|
||||
"Rollout",
|
||||
"Task",
|
||||
"TaskInput",
|
||||
"TaskIfAny",
|
||||
"RolloutRawResult",
|
||||
"Resource",
|
||||
"LLM",
|
||||
"PromptTemplate",
|
||||
"ResourceUnion",
|
||||
"NamedResources",
|
||||
"ResourcesUpdate",
|
||||
"GenericResponse",
|
||||
"ParallelWorkerBase",
|
||||
]
|
||||
|
||||
|
||||
class Triplet(BaseModel):
|
||||
"""A standard structure for a single turn in a trajectory."""
|
||||
|
||||
prompt: Any
|
||||
response: Any
|
||||
reward: Optional[float] = None
|
||||
metadata: Dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class Rollout(BaseModel):
|
||||
"""The standard reporting object from client to server."""
|
||||
|
||||
rollout_id: str
|
||||
|
||||
# Primary, high-level feedback
|
||||
final_reward: Optional[float] = None
|
||||
|
||||
# Structured, sequential feedback for RL-style optimization
|
||||
triplets: Optional[List[Triplet]] = None
|
||||
|
||||
# Optional, rich-context data for deep analysis
|
||||
trace: Optional[List[Dict[str, Any]]] = Field(
|
||||
default=None,
|
||||
description="A list of spans that conform to the OpenTelemetry JSON format. "
|
||||
"Users of the opentelemetry-sdk can generate this by calling "
|
||||
"json.loads(readable_span.to_json()).",
|
||||
)
|
||||
logs: Optional[List[str]] = None
|
||||
|
||||
# A bucket for any other relevant information
|
||||
metadata: Dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
TaskInput = Any
|
||||
|
||||
|
||||
class Task(BaseModel):
|
||||
"""A task (rollout request) to be processed by the client agent."""
|
||||
|
||||
rollout_id: str
|
||||
input: TaskInput
|
||||
|
||||
mode: Optional[Literal["train", "val", "test"]] = None
|
||||
resources_id: Optional[str] = None
|
||||
|
||||
# Optional fields for tracking task lifecycle
|
||||
create_time: Optional[float] = None
|
||||
last_claim_time: Optional[float] = None
|
||||
num_claims: Optional[int] = None
|
||||
|
||||
# Allow additional metadata fields
|
||||
metadata: Dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class TaskIfAny(BaseModel):
|
||||
is_available: bool
|
||||
task: Optional[Task] = None
|
||||
|
||||
|
||||
RolloutRawResult = Union[None, float, List[Triplet], List[Dict[str, Any]], List[ReadableSpan], Rollout]
|
||||
|
||||
|
||||
class Resource(BaseModel):
|
||||
"""
|
||||
Base class for all tunable resources.
|
||||
"""
|
||||
|
||||
resource_type: Any
|
||||
|
||||
|
||||
class LLM(Resource):
|
||||
"""
|
||||
Provide an LLM endpoint and model name as a resource.
|
||||
|
||||
Attributes:
|
||||
endpoint (str): The URL of the LLM API endpoint.
|
||||
model (str): The identifier for the model to be used (e.g., 'gpt-4o').
|
||||
sampling_parameters (SamplingParameters): A dictionary of hyperparameters
|
||||
for model inference, such as temperature, top_p, etc.
|
||||
"""
|
||||
|
||||
resource_type: Literal["llm"] = "llm"
|
||||
endpoint: str
|
||||
model: str
|
||||
sampling_parameters: Dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class PromptTemplate(Resource):
|
||||
"""
|
||||
A prompt template as a resource.
|
||||
|
||||
Attributes:
|
||||
template (str): The template string. The format depends on the engine.
|
||||
engine (Literal['jinja', 'f-string', 'poml']): The templating engine
|
||||
to use for rendering the prompt. I imagine users can use their own
|
||||
customized engines, but algos can only well operate on a subset of them.
|
||||
"""
|
||||
|
||||
resource_type: Literal["prompt_template"] = "prompt_template"
|
||||
template: str
|
||||
engine: Literal["jinja", "f-string", "poml"]
|
||||
|
||||
|
||||
# Use discriminated union for proper deserialization
|
||||
ResourceUnion = Annotated[Union[LLM, PromptTemplate], Field(discriminator="resource_type")]
|
||||
NamedResources = Dict[str, ResourceUnion]
|
||||
"""
|
||||
A dictionary-like class to hold named resources.
|
||||
|
||||
Example:
|
||||
resources: NamedResources = {
|
||||
'main_llm': LLM(
|
||||
endpoint="http://localhost:8080",
|
||||
model="llama3",
|
||||
sampling_parameters={'temperature': 0.7, 'max_tokens': 100}
|
||||
),
|
||||
'system_prompt': PromptTemplate(
|
||||
template="You are a helpful assistant.",
|
||||
engine='f-string'
|
||||
)
|
||||
}
|
||||
"""
|
||||
|
||||
|
||||
class ResourcesUpdate(BaseModel):
|
||||
"""
|
||||
A resource update message to be sent from the server to clients.
|
||||
|
||||
This message contains a dictionary of resources that clients should use
|
||||
for subsequent tasks. It is used to update the resources available to
|
||||
clients dynamically.
|
||||
"""
|
||||
|
||||
resources_id: str
|
||||
resources: NamedResources
|
||||
|
||||
|
||||
class GenericResponse(BaseModel):
|
||||
"""
|
||||
A generic response message that can be used for various purposes.
|
||||
"""
|
||||
|
||||
status: str = "success"
|
||||
message: Optional[str] = None
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class ParallelWorkerBase:
|
||||
"""Base class for objects that can be parallelized across multiple worker processes.
|
||||
|
||||
This class defines the standard lifecycle for parallel processing:
|
||||
|
||||
Main Process:
|
||||
1. init() - Initialize the object in the main process
|
||||
2. spawn workers and call init_worker() in each worker
|
||||
3. run() - Execute the main workload in parallel across workers
|
||||
4. teardown_worker() - Clean up resources in each worker
|
||||
5. teardown() - Final cleanup in the main process
|
||||
|
||||
Subclasses should implement the run() method and optionally override
|
||||
the lifecycle methods for custom initialization and cleanup behavior.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the base class. This method can be overridden by subclasses."""
|
||||
self.worker_id: Optional[int] = None
|
||||
|
||||
def init(self, *args: Any, **kwargs: Any) -> None:
|
||||
pass
|
||||
|
||||
def init_worker(self, worker_id: int, *args: Any, **kwargs: Any) -> None:
|
||||
self.worker_id = worker_id
|
||||
|
||||
def run(self, *args: Any, **kwargs: Any) -> Any:
|
||||
pass
|
||||
|
||||
def teardown_worker(self, worker_id: int, *args: Any, **kwargs: Any) -> None:
|
||||
pass
|
||||
|
||||
def teardown(self, *args: Any, **kwargs: Any) -> None:
|
||||
pass
|
||||
@@ -1,3 +1,3 @@
|
||||
from .trainer import *
|
||||
from .daemon import *
|
||||
from .dataset import *
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""VERL integration for Agent Lightning."""
|
||||
|
||||
@@ -1,4 +0,0 @@
|
||||
from .entrypoint import main
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,543 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Rollout managers for Agent Lightning VERL training."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import traceback
|
||||
import uuid
|
||||
from collections import defaultdict
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
import numpy as np
|
||||
from httpx_retries import Retry, RetryTransport
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from agentlightning.client import AgentLightningSyncClient
|
||||
from agentlightning.schemas import (
|
||||
TERMINAL_STATES,
|
||||
Event,
|
||||
EventCreate,
|
||||
Model,
|
||||
Rollout,
|
||||
RolloutCreate,
|
||||
RolloutState,
|
||||
)
|
||||
|
||||
try:
|
||||
import torch
|
||||
except ImportError: # pragma: no cover - torch is optional outside VERL installs.
|
||||
torch = None
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentlightning.hooks import RolloutHooks
|
||||
|
||||
|
||||
class Triplet(BaseModel):
|
||||
"""Single prompt-response-reward turn."""
|
||||
|
||||
prompt: Any
|
||||
response: Any
|
||||
reward: float | None = None
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class EnqueuedRollout(BaseModel):
|
||||
"""Enqueued rollout request metadata."""
|
||||
|
||||
data_id: str
|
||||
rollout_id: str
|
||||
step: int
|
||||
sample_idx_in_step: int
|
||||
enqueue_time: float
|
||||
input: Any = None
|
||||
# Server timestamps expose pod queue time and completion time.
|
||||
running_at: float | None = None
|
||||
finished_at: float | None = None
|
||||
|
||||
|
||||
class CompletedRollout(BaseModel):
|
||||
"""Completed rollout result."""
|
||||
|
||||
rollout_id: str
|
||||
data_id: str
|
||||
step: int
|
||||
sample_idx_in_step: int
|
||||
enqueue_time: float
|
||||
input: Any = None
|
||||
running_at: float | None = None
|
||||
finished_at: float | None = None
|
||||
final_reward: float | None = None
|
||||
triplets: list[Triplet] | None = None
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
events: list[dict[str, Any]] = Field(default_factory=list)
|
||||
triplet_events: list[dict[str, Any]] = Field(default_factory=list)
|
||||
rollout_state: RolloutState | None = None
|
||||
error_message: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class _TraceEvent:
|
||||
rollout_id: str
|
||||
attempt_id: str
|
||||
event_type: str
|
||||
data: dict[str, Any]
|
||||
|
||||
|
||||
class _TraceEventHelper:
|
||||
"""Queues hook events before HTTP flush."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._queued: list[_TraceEvent] = []
|
||||
|
||||
def add_event(self, rollout_id: str, attempt_id: str, event_type: str, data: dict[str, Any]) -> None:
|
||||
self._queued.append(_TraceEvent(rollout_id=rollout_id, attempt_id=attempt_id, event_type=event_type, data=data))
|
||||
|
||||
def flush(self, manager: AglRolloutManagerBase) -> None:
|
||||
for event in self._queued:
|
||||
manager._post_event(
|
||||
event.rollout_id,
|
||||
event.attempt_id,
|
||||
EventCreate(event_type=event.event_type, data=event.data),
|
||||
)
|
||||
|
||||
|
||||
def _as_reward_value(value: Any) -> float | None:
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
if isinstance(value, int | float | np.number):
|
||||
return float(value)
|
||||
return None
|
||||
|
||||
|
||||
def _to_native(obj: Any) -> Any:
|
||||
"""Convert numpy/torch values for JSON serialization."""
|
||||
if isinstance(obj, np.ndarray):
|
||||
return _to_native(obj.tolist())
|
||||
if isinstance(obj, np.generic):
|
||||
return _to_native(obj.item())
|
||||
if isinstance(obj, Mapping):
|
||||
return {_to_native(key): _to_native(value) for key, value in obj.items()}
|
||||
if isinstance(obj, (list, tuple, set)):
|
||||
return [_to_native(item) for item in obj]
|
||||
if torch is not None and isinstance(obj, torch.Tensor):
|
||||
return obj.item() if obj.ndim == 0 else obj.tolist()
|
||||
return obj
|
||||
|
||||
|
||||
class AglRolloutManagerBase:
|
||||
"""Base manager for Agent Lightning rollout HTTP operations."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
agl_base_url: str,
|
||||
agl_key: str,
|
||||
model: str,
|
||||
step: int,
|
||||
train_rollout_n: int = 1,
|
||||
rollout_timeout_seconds: float = 1200.0,
|
||||
poll_interval_seconds: float = 1.0,
|
||||
hooks: RolloutHooks | None = None,
|
||||
local_agent_class: str | None = None,
|
||||
local_env_map: dict[str, str] | None = None,
|
||||
k8s_job_template_path: str | None = None,
|
||||
) -> None:
|
||||
self._model = model
|
||||
self._step = step
|
||||
self._train_rollout_n = train_rollout_n
|
||||
self._poll_interval_seconds = poll_interval_seconds
|
||||
self._hooks = hooks
|
||||
self._rollout_config: dict[str, Any] = {"timeout_seconds": int(rollout_timeout_seconds)}
|
||||
if local_agent_class:
|
||||
self._rollout_config["local"] = {
|
||||
"agent_class": local_agent_class,
|
||||
"env_map": local_env_map or {},
|
||||
}
|
||||
if k8s_job_template_path:
|
||||
self._rollout_config["k8s"] = {"job_template": Path(k8s_job_template_path).read_text()}
|
||||
|
||||
self.client = AgentLightningSyncClient(
|
||||
base_url=agl_base_url,
|
||||
key=agl_key,
|
||||
timeout=120.0,
|
||||
transport=RetryTransport(retry=Retry(total=10, allowed_methods=["GET"])),
|
||||
)
|
||||
|
||||
def register_model(self, server_addresses: list[str]) -> list[Model]:
|
||||
"""Register model server endpoints."""
|
||||
models: list[Model] = []
|
||||
for address in server_addresses:
|
||||
endpoint = address if address.startswith("http") else f"http://{address}/v1"
|
||||
models.append(Model(model=self._model, endpoint=endpoint))
|
||||
|
||||
# Model registration is idempotent, so transient failures are safe to retry.
|
||||
payload = [model.model_dump(mode="json") for model in models]
|
||||
response = self.client.post_with_retry("/api/models", json=payload)
|
||||
return [Model.model_validate(item) for item in response.json()]
|
||||
|
||||
def delete_model(self) -> dict[str, Any]:
|
||||
"""Delete registered model endpoints. Best-effort: ignore errors."""
|
||||
try:
|
||||
response = self.client.delete("/api/models")
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except Exception as exc:
|
||||
print(f"RolloutManager: failed to delete models: {exc}")
|
||||
return {}
|
||||
|
||||
def _get_rollout(self, rollout_id: str) -> Rollout:
|
||||
response = self.client.get(f"/api/rollouts/{rollout_id}")
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
item = payload["rollout"] if isinstance(payload, dict) and "rollout" in payload else payload
|
||||
return Rollout.model_validate(item)
|
||||
|
||||
def _delete_rollout(self, rollout_id: str) -> None:
|
||||
try:
|
||||
self.client.delete(f"/api/rollouts/{rollout_id}")
|
||||
except Exception as exc:
|
||||
print(f"RolloutManager: failed to delete rollout {rollout_id}: {exc}")
|
||||
|
||||
@staticmethod
|
||||
def _record_lifecycle_timestamps(enqueued_rollout: EnqueuedRollout, rollout: Rollout) -> None:
|
||||
"""Capture server-authoritative running/finished timestamps in place.
|
||||
|
||||
Pods are launched in CPU-limited batches, so a rollout can sit QUEUING
|
||||
well after it was submitted; status.updated_at at the queuing->running
|
||||
flip is the moment its pod actually started. We record it the first time
|
||||
we observe RUNNING (or, if we polled too slowly and skipped straight to a
|
||||
terminal state, the terminal updated_at) so running_at - enqueue_time
|
||||
reflects the real queue/startup wait.
|
||||
"""
|
||||
state = rollout.status.state
|
||||
updated_at = rollout.status.updated_at
|
||||
if enqueued_rollout.running_at is None and state in (
|
||||
RolloutState.RUNNING,
|
||||
RolloutState.SUCCEEDED,
|
||||
RolloutState.FAILED,
|
||||
):
|
||||
enqueued_rollout.running_at = updated_at
|
||||
if state in TERMINAL_STATES:
|
||||
enqueued_rollout.finished_at = updated_at
|
||||
|
||||
def _get_events(self, rollout_id: str, *, event_type: str | None = None, format: str | None = None) -> list[Event]:
|
||||
params = {
|
||||
key: value for key, value in {"event_type": event_type, "format": format}.items() if value is not None
|
||||
}
|
||||
response = self.client.get(f"/api/rollouts/{rollout_id}/events", params=params)
|
||||
response.raise_for_status()
|
||||
return [Event.model_validate(item) for item in response.json()]
|
||||
|
||||
def _post_event(self, rollout_id: str, attempt_id: str, event: EventCreate) -> Event:
|
||||
response = self.client.post(
|
||||
f"/api/rollouts/{rollout_id}/attempt/{attempt_id}/events",
|
||||
json=event.model_dump(mode="json"),
|
||||
)
|
||||
response.raise_for_status()
|
||||
return Event.model_validate(response.json())
|
||||
|
||||
def _create_rollouts(self, data: dict[str, Any], *, is_train: bool) -> list[EnqueuedRollout]:
|
||||
keys = list(data.keys())
|
||||
if not keys:
|
||||
return []
|
||||
|
||||
num_samples = len(data[keys[0]])
|
||||
rollouts_per_sample = self._train_rollout_n if is_train else 1
|
||||
rollout_requests: list[RolloutCreate] = []
|
||||
enqueued_rollouts: list[EnqueuedRollout] = []
|
||||
for sample_idx in range(num_samples):
|
||||
original = {key: _to_native(data[key][sample_idx]) for key in keys}
|
||||
data_id = str(uuid.uuid4())
|
||||
for _ in range(rollouts_per_sample):
|
||||
request = RolloutCreate(
|
||||
input=_to_native(original),
|
||||
is_train=is_train,
|
||||
config=cast(Any, self._rollout_config), # pydantic coerces the dict
|
||||
metadata={},
|
||||
)
|
||||
if self._hooks is not None:
|
||||
request = self._hooks.on_enqueue(request)
|
||||
# Assign the id after hooks so creation remains idempotent.
|
||||
rollout_id = uuid.uuid4().hex
|
||||
request = request.model_copy(update={"rollout_id": rollout_id})
|
||||
enqueued_rollouts.append(
|
||||
EnqueuedRollout(
|
||||
data_id=data_id,
|
||||
input=request.input,
|
||||
rollout_id=rollout_id,
|
||||
step=self._step,
|
||||
sample_idx_in_step=sample_idx,
|
||||
enqueue_time=time.time(),
|
||||
)
|
||||
)
|
||||
rollout_requests.append(request)
|
||||
|
||||
if not rollout_requests:
|
||||
return []
|
||||
|
||||
# Preassigned ids make batch creation safe to retry without duplicates.
|
||||
payload = [request.model_dump(mode="json", exclude_none=True) for request in rollout_requests]
|
||||
response = self.client.post_with_retry("/api/rollouts", json=payload)
|
||||
created = [Rollout.model_validate(item) for item in response.json()]
|
||||
assert len(created) == len(rollout_requests), (
|
||||
f"Agent Lightning returned {len(created)} rollouts, expected {len(rollout_requests)}"
|
||||
)
|
||||
return [
|
||||
enqueued_rollout.model_copy(update={"rollout_id": rollout.rollout_id})
|
||||
for enqueued_rollout, rollout in zip(enqueued_rollouts, created, strict=True)
|
||||
]
|
||||
|
||||
def _fetch_rollout_events(self, rollout_id: str) -> tuple[list[Event], list[Event]]:
|
||||
raw_events = self._get_events(rollout_id)
|
||||
triplet_events = self._get_events(rollout_id, format="triplet")
|
||||
return raw_events, triplet_events
|
||||
|
||||
@staticmethod
|
||||
def _events_by_attempt(raw_events: list[Event], fallback_attempt_id: str) -> dict[str, list[Event]]:
|
||||
grouped: dict[str, list[Event]] = defaultdict(list)
|
||||
for event in raw_events:
|
||||
grouped[event.attempt_id or fallback_attempt_id].append(event)
|
||||
if not grouped:
|
||||
grouped[fallback_attempt_id] = []
|
||||
return dict(grouped)
|
||||
|
||||
def _run_succeeded_hook(self, rollout: Rollout) -> None:
|
||||
if self._hooks is None:
|
||||
return
|
||||
attempt_id = rollout.status.last_attempt_id or "unknown"
|
||||
trace_event_helper = _TraceEventHelper()
|
||||
raw_events = self._get_events(rollout.rollout_id)
|
||||
events_by_attempt = self._events_by_attempt(raw_events, attempt_id)
|
||||
try:
|
||||
self._hooks.on_succeeded(rollout, events_by_attempt, trace_event_helper)
|
||||
trace_event_helper.flush(self)
|
||||
except Exception:
|
||||
traceback.print_exc()
|
||||
print(f"RolloutManager: on_succeeded hook failed for rollout {rollout.rollout_id}")
|
||||
|
||||
def _run_failed_hook(self, rollout: Rollout) -> None:
|
||||
if self._hooks is None:
|
||||
return
|
||||
trace_event_helper = _TraceEventHelper()
|
||||
try:
|
||||
self._hooks.on_failed(rollout, trace_event_helper)
|
||||
trace_event_helper.flush(self)
|
||||
except Exception:
|
||||
traceback.print_exc()
|
||||
print(f"RolloutManager: on_failed hook failed for rollout {rollout.rollout_id}")
|
||||
|
||||
def _build_completed_rollout(self, enqueued_rollout: EnqueuedRollout, rollout: Rollout) -> CompletedRollout:
|
||||
"""Fetch triplets and reward for a terminal rollout."""
|
||||
raw_events, triplet_events = self._fetch_rollout_events(enqueued_rollout.rollout_id)
|
||||
|
||||
triplets: list[Triplet] = []
|
||||
for event in triplet_events:
|
||||
if event.event_type != "model_request":
|
||||
continue
|
||||
data = event.data
|
||||
http_status = data.get("http_status")
|
||||
response_token_ids = data.get("response_token_ids", [])
|
||||
if data.get("status") == "error" or (isinstance(http_status, int) and http_status >= 400):
|
||||
continue
|
||||
if not response_token_ids:
|
||||
continue
|
||||
triplets.append(
|
||||
Triplet(
|
||||
prompt={"token_ids": data.get("prompt_token_ids", [])},
|
||||
response={
|
||||
"token_ids": response_token_ids,
|
||||
"log_probs": data.get("response_log_probs"),
|
||||
},
|
||||
reward=None,
|
||||
metadata={"server": data.get("server", {})},
|
||||
)
|
||||
)
|
||||
|
||||
final_reward: float | None = None
|
||||
reward_events = [event for event in triplet_events if event.event_type == "reward"]
|
||||
if reward_events:
|
||||
reward_data = reward_events[-1].data
|
||||
final_reward = _as_reward_value(reward_data.get("value"))
|
||||
|
||||
if triplets and final_reward is not None:
|
||||
triplets[-1] = triplets[-1].model_copy(update={"reward": final_reward})
|
||||
|
||||
metadata = rollout.metadata.model_dump()
|
||||
finished_at = enqueued_rollout.finished_at
|
||||
if finished_at is None:
|
||||
finished_at = rollout.status.updated_at
|
||||
return CompletedRollout(
|
||||
rollout_id=enqueued_rollout.rollout_id,
|
||||
data_id=enqueued_rollout.data_id,
|
||||
step=enqueued_rollout.step,
|
||||
sample_idx_in_step=enqueued_rollout.sample_idx_in_step,
|
||||
input=enqueued_rollout.input,
|
||||
enqueue_time=enqueued_rollout.enqueue_time,
|
||||
running_at=enqueued_rollout.running_at,
|
||||
finished_at=finished_at,
|
||||
final_reward=final_reward,
|
||||
triplets=triplets,
|
||||
metadata=metadata,
|
||||
events=[event.model_dump() for event in raw_events],
|
||||
triplet_events=[event.model_dump() for event in triplet_events],
|
||||
rollout_state=rollout.status.state,
|
||||
error_message=rollout.status.error_message,
|
||||
)
|
||||
|
||||
|
||||
class AglRolloutManager(AglRolloutManagerBase):
|
||||
def enqueue_and_wait_until_completed(
|
||||
self,
|
||||
data: dict[str, Any],
|
||||
*,
|
||||
is_train: bool,
|
||||
) -> list[CompletedRollout]:
|
||||
"""Create rollouts, wait for completion, and return results."""
|
||||
enqueued_rollouts = self._create_rollouts(data, is_train=is_train)
|
||||
pending_rollouts = list(enqueued_rollouts)
|
||||
completed_rollouts: list[CompletedRollout] = []
|
||||
num_deleted = 0
|
||||
num_succeeded = 0
|
||||
num_failed = 0
|
||||
|
||||
while len(completed_rollouts) < len(enqueued_rollouts):
|
||||
# Delete prior completions before polling to bound server-side state.
|
||||
for completed_rollout in completed_rollouts[num_deleted:]:
|
||||
self._delete_rollout(completed_rollout.rollout_id)
|
||||
num_deleted = len(completed_rollouts)
|
||||
|
||||
for enqueued_rollout in list(pending_rollouts):
|
||||
rollout_id = enqueued_rollout.rollout_id
|
||||
rollout = self._get_rollout(rollout_id)
|
||||
state = rollout.status.state
|
||||
self._record_lifecycle_timestamps(enqueued_rollout, rollout)
|
||||
if state not in TERMINAL_STATES:
|
||||
continue
|
||||
|
||||
pending_rollouts.remove(enqueued_rollout)
|
||||
if state == RolloutState.SUCCEEDED:
|
||||
num_succeeded += 1
|
||||
self._run_succeeded_hook(rollout)
|
||||
elif state == RolloutState.FAILED:
|
||||
num_failed += 1
|
||||
self._run_failed_hook(rollout)
|
||||
|
||||
completed_rollouts.append(self._build_completed_rollout(enqueued_rollout, rollout))
|
||||
|
||||
print(
|
||||
f"AglRolloutManager: completed={len(completed_rollouts)}/{len(enqueued_rollouts)} "
|
||||
f"succeeded={num_succeeded} failed={num_failed}"
|
||||
)
|
||||
|
||||
if pending_rollouts:
|
||||
time.sleep(self._poll_interval_seconds)
|
||||
|
||||
# Delete whatever completed in the final round.
|
||||
for completed_rollout in completed_rollouts[num_deleted:]:
|
||||
self._delete_rollout(completed_rollout.rollout_id)
|
||||
|
||||
return completed_rollouts
|
||||
|
||||
|
||||
class AglAsyncRolloutManager(AglRolloutManagerBase):
|
||||
"""Async rollout manager."""
|
||||
|
||||
def enqueue_and_wait_until_group_completed(
|
||||
self,
|
||||
data: dict[str, Any],
|
||||
carry_over_enqueued_rollouts: list[EnqueuedRollout],
|
||||
*,
|
||||
is_train: bool,
|
||||
target_finished_group_num: int,
|
||||
) -> tuple[list[CompletedRollout], list[EnqueuedRollout]]:
|
||||
"""Enqueue rollouts and wait for enough completed rollout groups."""
|
||||
assert is_train is True
|
||||
enqueued_rollouts = self._create_rollouts(data, is_train=True)
|
||||
active_rollouts = carry_over_enqueued_rollouts + enqueued_rollouts
|
||||
if not active_rollouts:
|
||||
return [], []
|
||||
|
||||
grouped_rollouts: dict[str, list[EnqueuedRollout]] = defaultdict(list)
|
||||
for enqueued_rollout in active_rollouts:
|
||||
grouped_rollouts[enqueued_rollout.data_id].append(enqueued_rollout)
|
||||
|
||||
for group in grouped_rollouts.values():
|
||||
assert len(group) == self._train_rollout_n
|
||||
|
||||
finished_rollout_ids: set[str] = set()
|
||||
terminal_rollouts: dict[str, Rollout] = {}
|
||||
completed_group_keys: set[str] = set()
|
||||
completed_rollouts: list[CompletedRollout] = []
|
||||
num_succeeded = 0
|
||||
num_failed = 0
|
||||
|
||||
while len(completed_group_keys) < target_finished_group_num:
|
||||
for data_id, group in grouped_rollouts.items():
|
||||
if data_id in completed_group_keys:
|
||||
continue
|
||||
|
||||
for enqueued_rollout in group:
|
||||
if enqueued_rollout.rollout_id in finished_rollout_ids:
|
||||
continue
|
||||
|
||||
rollout = self._get_rollout(enqueued_rollout.rollout_id)
|
||||
state = rollout.status.state
|
||||
self._record_lifecycle_timestamps(enqueued_rollout, rollout)
|
||||
if state not in TERMINAL_STATES:
|
||||
continue
|
||||
|
||||
finished_rollout_ids.add(enqueued_rollout.rollout_id)
|
||||
terminal_rollouts[enqueued_rollout.rollout_id] = rollout
|
||||
if state == RolloutState.SUCCEEDED:
|
||||
num_succeeded += 1
|
||||
self._run_succeeded_hook(rollout)
|
||||
elif state == RolloutState.FAILED:
|
||||
num_failed += 1
|
||||
self._run_failed_hook(rollout)
|
||||
|
||||
if all(enqueued_rollout.rollout_id in finished_rollout_ids for enqueued_rollout in group):
|
||||
completed_group_keys.add(data_id)
|
||||
completed_rollouts.extend(
|
||||
self._build_completed_rollout(
|
||||
enqueued_rollout,
|
||||
terminal_rollouts[enqueued_rollout.rollout_id],
|
||||
)
|
||||
for enqueued_rollout in group
|
||||
)
|
||||
# Free completed group state after reading it.
|
||||
for enqueued_rollout in group:
|
||||
self._delete_rollout(enqueued_rollout.rollout_id)
|
||||
if len(completed_group_keys) >= target_finished_group_num:
|
||||
break
|
||||
|
||||
print(
|
||||
f"AglAsyncRolloutManager: completed_groups={len(completed_group_keys)}/{target_finished_group_num} "
|
||||
f"finished_rollouts={len(finished_rollout_ids)}/{len(active_rollouts)} "
|
||||
f"succeeded={num_succeeded} failed={num_failed}"
|
||||
)
|
||||
|
||||
if len(completed_group_keys) < target_finished_group_num:
|
||||
time.sleep(self._poll_interval_seconds)
|
||||
|
||||
new_carry_over_rollouts = [
|
||||
enqueued_rollout
|
||||
for data_id, group in grouped_rollouts.items()
|
||||
if data_id not in completed_group_keys
|
||||
for enqueued_rollout in group
|
||||
]
|
||||
return completed_rollouts, new_carry_over_rollouts
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AglAsyncRolloutManager",
|
||||
"AglRolloutManager",
|
||||
"AglRolloutManagerBase",
|
||||
"CompletedRollout",
|
||||
"EnqueuedRollout",
|
||||
"Triplet",
|
||||
]
|
||||
@@ -1,41 +0,0 @@
|
||||
import ray
|
||||
from copy import deepcopy
|
||||
|
||||
from agentlightning.instrumentation.vllm import instrument_vllm, ChatCompletionResponsePatched
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse, StreamingResponse
|
||||
from vllm.entrypoints.openai.protocol import ChatCompletionRequest, ErrorResponse
|
||||
from verl.workers.rollout.vllm_rollout.vllm_async_server import AsyncvLLMServer
|
||||
|
||||
|
||||
def _unwrap_ray_remote(cls):
|
||||
if hasattr(cls, "__ray_actor_class__"):
|
||||
cls = cls.__ray_actor_class__
|
||||
return cls
|
||||
|
||||
|
||||
@ray.remote(num_cpus=1)
|
||||
class PatchedvLLMServer(_unwrap_ray_remote(AsyncvLLMServer)):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
instrument_vllm()
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.config = deepcopy(self.config)
|
||||
self.config.rollout.multi_turn.tool_config_path = "/dev/null"
|
||||
|
||||
async def chat_completion(self, raw_request: Request):
|
||||
"""OpenAI-compatible HTTP endpoint.
|
||||
|
||||
API reference: https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html
|
||||
"""
|
||||
request_json = await raw_request.json()
|
||||
request = ChatCompletionRequest(**request_json)
|
||||
generator = await self.openai_serving_chat.create_chat_completion(request, raw_request)
|
||||
|
||||
if isinstance(generator, ErrorResponse):
|
||||
return JSONResponse(content=generator.model_dump(), status_code=generator.code)
|
||||
if request.stream:
|
||||
return StreamingResponse(content=generator, media_type="text/event-stream")
|
||||
else:
|
||||
return JSONResponse(content=generator.model_dump())
|
||||
@@ -6,16 +6,36 @@ defaults:
|
||||
- ppo_trainer
|
||||
- _self_
|
||||
|
||||
algorithm:
|
||||
enable_rollout_level_advantage: true
|
||||
|
||||
agentlightning:
|
||||
port: 9999
|
||||
agl_base_url: http://localhost:8080
|
||||
agl_key: ""
|
||||
hooks: null
|
||||
rollout_timeout_seconds: 1800
|
||||
local:
|
||||
agent_class: null
|
||||
env_map: {}
|
||||
k8s:
|
||||
job_template_path: null
|
||||
reward_fillna_value: 0.0
|
||||
max_ppo_update_times: null
|
||||
trace_aggregator:
|
||||
level: trajectory # transition | trajectory
|
||||
trajectory_max_prompt_length: 2048
|
||||
trajectory_max_response_length: 8192
|
||||
async_rollout:
|
||||
enabled: false
|
||||
async_train_batch_size: null
|
||||
|
||||
data:
|
||||
filter_overlong_prompts: false
|
||||
|
||||
actor_rollout_ref:
|
||||
actor:
|
||||
calculate_entropy: true
|
||||
policy_loss:
|
||||
loss_mode: per_rollout_mean
|
||||
rollout:
|
||||
mode: async
|
||||
agent:
|
||||
custom_async_server:
|
||||
path: pkg://agentlightning.verl.async_server
|
||||
name: PatchedvLLMServer
|
||||
|
||||
@@ -1,516 +0,0 @@
|
||||
import asyncio
|
||||
import json
|
||||
import random
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
from agentlightning import LLM, AgentLightningServer, NamedResources, Rollout, configure_logger
|
||||
from flask import Flask, Response, abort, request
|
||||
from openai.types.chat.chat_completion import ChatCompletion
|
||||
from tensordict import TensorDict
|
||||
|
||||
from verl import DataProto
|
||||
|
||||
configure_logger()
|
||||
|
||||
|
||||
def get_left_padded_ids_and_attention_mask(ids: List[int], max_length: int, pad_token_id: int):
|
||||
"""
|
||||
Left-pad (or truncate) a sequence of token IDs to a fixed length,
|
||||
and build the corresponding attention mask.
|
||||
|
||||
Args:
|
||||
ids: the original list of token IDs.
|
||||
max_length: desired total length after padding/truncation.
|
||||
pad_token_id: ID to use for padding.
|
||||
|
||||
Returns:
|
||||
padded_ids (any): list of length == max_length.
|
||||
attention_mask (any): list of same length: 1 for non-pad tokens, 0 for pads.
|
||||
"""
|
||||
seq_len = len(ids)
|
||||
|
||||
if seq_len >= max_length:
|
||||
# too long → truncate from the left, keep the last max_length tokens
|
||||
trimmed = ids[-max_length:]
|
||||
attention_mask = [1] * max_length
|
||||
return trimmed, attention_mask
|
||||
|
||||
# too short → pad on the left
|
||||
pad_len = max_length - seq_len
|
||||
padded_ids = [pad_token_id] * pad_len + ids
|
||||
attention_mask = [0] * pad_len + [1] * seq_len
|
||||
return padded_ids, attention_mask
|
||||
|
||||
|
||||
def get_right_padded_ids_and_attention_mask(ids: List[int], max_length: int, pad_token_id: int):
|
||||
"""
|
||||
Right-pad (or truncate) a sequence of token IDs to a fixed length,
|
||||
and build the corresponding attention mask.
|
||||
|
||||
Args:
|
||||
ids: the original list of token IDs.
|
||||
max_length: desired total length after padding/truncation.
|
||||
pad_token_id: ID to use for padding.
|
||||
|
||||
Returns:
|
||||
padded_ids (any): list of length == max_length.
|
||||
attention_mask (any): list of same length: 1 for non-pad tokens, 0 for pads.
|
||||
"""
|
||||
seq_len = len(ids)
|
||||
|
||||
if seq_len >= max_length:
|
||||
# too long → truncate to the first max_length tokens
|
||||
trimmed = ids[:max_length]
|
||||
attention_mask = [1] * max_length
|
||||
return trimmed, attention_mask
|
||||
|
||||
# too short → pad on the right
|
||||
pad_len = max_length - seq_len
|
||||
padded_ids = ids + [pad_token_id] * pad_len
|
||||
attention_mask = [1] * seq_len + [0] * pad_len
|
||||
return padded_ids, attention_mask
|
||||
|
||||
|
||||
def _find_available_port() -> int:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(("", 0))
|
||||
return s.getsockname()[1]
|
||||
|
||||
|
||||
class AgentModeDaemon:
|
||||
"""
|
||||
AgentModeDaemon using the AgentLightningServer SDK.
|
||||
|
||||
This class manages the server lifecycle, task queueing, and results
|
||||
retrieval, while also running a proxy server for LLM requests. It maintains
|
||||
the original interface for compatibility with the RayPPOTrainer.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
port,
|
||||
train_rollout_n,
|
||||
train_information,
|
||||
tokenizer,
|
||||
mini_batch_size,
|
||||
pad_token_id,
|
||||
reward_fillna_value=0.0,
|
||||
llm_timeout_seconds=600.0,
|
||||
):
|
||||
# Server and Task Configuration
|
||||
self.server_port = port
|
||||
self.llm_timeout_seconds = llm_timeout_seconds
|
||||
self.server = AgentLightningServer(
|
||||
host="0.0.0.0", port=self.server_port, task_timeout_seconds=self.llm_timeout_seconds
|
||||
)
|
||||
self.proxy_port = _find_available_port() # Run proxy on a different port
|
||||
|
||||
# Training and Data Configuration
|
||||
self.train_rollout_n = train_rollout_n
|
||||
self.train_information = train_information
|
||||
self.mini_batch_size = mini_batch_size
|
||||
self.pad_token_id = pad_token_id
|
||||
self.tokenizer = tokenizer
|
||||
self.reward_fillna_value = reward_fillna_value
|
||||
|
||||
# Internal State
|
||||
self.backend_llm_server_addresses: List[str] = []
|
||||
self._total_tasks_queued = 0
|
||||
self._completed_rollouts: Dict[str, Rollout] = {}
|
||||
self._task_id_to_original_sample: Dict[str, Dict] = {}
|
||||
self._server_thread: Optional[threading.Thread] = None
|
||||
self._proxy_thread: Optional[threading.Thread] = None
|
||||
self.is_train = True
|
||||
|
||||
def _start_proxy_server(self):
|
||||
"""
|
||||
Initializes and runs a Flask-based proxy server in a separate thread.
|
||||
This proxy load-balances requests to the actual backend LLM servers.
|
||||
"""
|
||||
app = Flask(__name__)
|
||||
|
||||
num_requests = 0
|
||||
last_request_time = 0
|
||||
|
||||
@app.route("/v1/<path:path>", methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"])
|
||||
def proxy(path):
|
||||
if not self.backend_llm_server_addresses:
|
||||
abort(503, description="No backend LLM servers available.")
|
||||
|
||||
# Randomly choose a backend server for load balancing
|
||||
target_server = random.choice(self.backend_llm_server_addresses)
|
||||
target_url = f"http://{target_server}/v1/{path}"
|
||||
|
||||
# Copy client request headers, removing the Host header
|
||||
headers = {key: value for key, value in request.headers if key.lower() != "host"}
|
||||
|
||||
# Log the request for debugging
|
||||
nonlocal num_requests, last_request_time
|
||||
current_time = time.time()
|
||||
num_requests += 1
|
||||
if current_time - last_request_time > 60 or num_requests == 1 or num_requests % 100 == 0:
|
||||
print(f"Proxying {request.method} request to {target_server}. Request data: {request.get_data()}")
|
||||
last_request_time = current_time
|
||||
|
||||
try:
|
||||
# Forward the request to the target backend
|
||||
resp = requests.request(
|
||||
method=request.method,
|
||||
url=target_url,
|
||||
headers=headers,
|
||||
params=request.args,
|
||||
data=request.get_data(),
|
||||
cookies=request.cookies,
|
||||
allow_redirects=False,
|
||||
timeout=self.llm_timeout_seconds,
|
||||
)
|
||||
# Filter out hop-by-hop headers before returning the response
|
||||
excluded_headers = [
|
||||
"content-encoding",
|
||||
"content-length",
|
||||
"transfer-encoding",
|
||||
"connection",
|
||||
"keep-alive",
|
||||
"proxy-authenticate",
|
||||
"proxy-authorization",
|
||||
"te",
|
||||
"trailers",
|
||||
"upgrade",
|
||||
]
|
||||
response_headers = [
|
||||
(name, value) for name, value in resp.raw.headers.items() if name.lower() not in excluded_headers
|
||||
]
|
||||
if resp.status_code == 200:
|
||||
# NOTE: from Zhiyuan's code.
|
||||
# https://github.com/hzy46/verl_agent_mode/blob/2db65ea9858f645a914120357412a7540f8bd82d/verl/trainer/ppo/ray_trainer.py#L692-L711
|
||||
# request_json = json.loads(request.get_data().decode("utf-8"))
|
||||
response_json = json.loads(resp.content.decode("utf-8"))
|
||||
# response_message = ChatCompletion(**response_json).choices[0].message.model_dump(exclude_unset=True, exclude_none=True)
|
||||
# tool_schemas = request_json.get("tools", None)
|
||||
# prompt_ids = self.tokenizer.apply_chat_template(request_json["messages"], tools=tool_schemas, add_generation_prompt=True, tokenize=True)
|
||||
# full_ids = self.tokenizer.apply_chat_template(request_json["messages"] + [response_message], tools=tool_schemas, add_generation_prompt=False, tokenize=True)
|
||||
# TBD: response_ids sometimes ends with "<eos_id>\n", shall we keep the extra "\n"?
|
||||
# sometimes it has some differences with the hacky method in the end, but this should align with ToolCompletionCallback
|
||||
# response_ids = full_ids[len(prompt_ids):]
|
||||
|
||||
# NOTE (yuge): They are different. Don't know why.
|
||||
# assert response_json['prompt_token_ids'] == prompt_ids
|
||||
# patched_response_ids = response_json['response_token_ids'][0]
|
||||
# assert patched_response_ids == response_ids[:len(patched_response_ids)], f"{patched_response_ids} != {response_ids[:len(patched_response_ids)]}"
|
||||
# response_json['prompt_token_ids'] = prompt_ids
|
||||
# response_json['response_token_ids'] = [response_ids]
|
||||
replaced_return_content = json.dumps(response_json).encode("utf-8")
|
||||
return Response(replaced_return_content, status=resp.status_code, headers=response_headers)
|
||||
return Response(resp.content, resp.status_code, response_headers)
|
||||
except requests.exceptions.RequestException as e:
|
||||
abort(500, description=f"Error proxying request: {e}")
|
||||
|
||||
def run_app():
|
||||
app.run(host="0.0.0.0", port=self.proxy_port, threaded=True, debug=False)
|
||||
|
||||
self._proxy_thread = threading.Thread(target=run_app, daemon=True)
|
||||
self._proxy_thread.start()
|
||||
print(f"Proxy server running on port {self.proxy_port}")
|
||||
|
||||
def start(self):
|
||||
"""Starts the main AgentLightningServer and the proxy server."""
|
||||
|
||||
def run_server():
|
||||
"""Run the AgentLightningServer in a separate thread."""
|
||||
asyncio.run(self.server.run_forever())
|
||||
|
||||
self._server_thread = threading.Thread(target=run_server, daemon=True)
|
||||
self._server_thread.start()
|
||||
|
||||
# Wait for the server's internal startup event to be set.
|
||||
print("Waiting for AgentLightningServer to start...")
|
||||
is_ready = self.server.startup_event.wait(timeout=20.0) # Wait up to 20s
|
||||
if not is_ready:
|
||||
raise RuntimeError("AgentLightningServer failed to start within the timeout period.")
|
||||
|
||||
print(f"AgentLightningServer control plane running on port {self.server_port}")
|
||||
|
||||
self._start_proxy_server()
|
||||
|
||||
async def _async_set_up(self, data, server_addresses, is_train=True):
|
||||
"""Async helper to set up data and resources on the server."""
|
||||
self.clear_data_and_server()
|
||||
self.backend_llm_server_addresses = server_addresses
|
||||
self.is_train = is_train
|
||||
|
||||
# 1. Update resources on the server for clients to use
|
||||
llm_resource = LLM(
|
||||
endpoint=f"http://127.0.0.1:{self.proxy_port}/v1",
|
||||
model=self.train_information.get("model", "default-model"),
|
||||
sampling_parameters={"temperature": self.train_information.get("temperature", 0.7)},
|
||||
)
|
||||
resources: NamedResources = {"main_llm": llm_resource}
|
||||
resources_id = await self.server.update_resources(resources)
|
||||
|
||||
# 2. Queue tasks for agents to process
|
||||
keys = list(data.keys())
|
||||
num_samples = len(data[keys[0]])
|
||||
rollouts_per_sample = self.train_rollout_n if is_train else 1
|
||||
|
||||
for i in range(num_samples):
|
||||
data_id = str(uuid.uuid4())
|
||||
original_sample = {key: data[key][i] for key in keys}
|
||||
original_sample["data_id"] = data_id
|
||||
|
||||
# For training, each sample is rolled out multiple times
|
||||
for j in range(rollouts_per_sample):
|
||||
task_metadata = {"data_id": data_id, "is_train": is_train}
|
||||
|
||||
# Data ID is different from Rollout ID, as one data can have multiple rollouts.
|
||||
rollout_id = await self.server.queue_task(
|
||||
sample=original_sample,
|
||||
mode="train" if is_train else "val",
|
||||
resources_id=resources_id,
|
||||
metadata=task_metadata,
|
||||
)
|
||||
# Store original sample data to reconstruct batch information later
|
||||
self._task_id_to_original_sample[rollout_id] = original_sample
|
||||
self._total_tasks_queued += 1
|
||||
|
||||
def set_up_data_and_server(self, data, server_addresses, is_train=True):
|
||||
"""Synchronous wrapper for setting up data and server resources."""
|
||||
if not self.server.loop or not self.server.startup_event.is_set():
|
||||
raise RuntimeError("Server is not running or ready.")
|
||||
|
||||
coro = self._async_set_up(data, server_addresses, is_train)
|
||||
future = asyncio.run_coroutine_threadsafe(coro, self.server.loop)
|
||||
try:
|
||||
future.result(timeout=60) # Wait for completion with a timeout
|
||||
except Exception as e:
|
||||
print(f"Failed to set up data on server: {e}")
|
||||
raise
|
||||
|
||||
def _validate_data(self, rollout: Rollout):
|
||||
if rollout.final_reward is None:
|
||||
print(
|
||||
f"Warning: Reward is None for rollout {rollout.rollout_id}, will be auto-set to {self.reward_fillna_value}."
|
||||
)
|
||||
if rollout.triplets is None:
|
||||
print(f"Warning: Triplet is None for rollout {rollout.rollout_id}.")
|
||||
elif len(rollout.triplets) == 0:
|
||||
print(f"Warning: Length of triplets is 0 for rollout {rollout.rollout_id}.")
|
||||
elif any(not r.response.get("token_ids", []) for r in rollout.triplets):
|
||||
print(f"Warning: Rollout {rollout.rollout_id} contains empty response: {rollout.triplets}")
|
||||
elif any(not r.prompt.get("token_ids", []) for r in rollout.triplets):
|
||||
print(f"Warning: Rollout {rollout.rollout_id} contains empty prompt: {rollout.triplets}")
|
||||
|
||||
async def _async_run_until_finished(self, verbose=True):
|
||||
"""Async helper to wait for all tasks to complete."""
|
||||
while len(self._completed_rollouts) < self._total_tasks_queued:
|
||||
completed_batch = await self.server.retrieve_completed_rollouts()
|
||||
for rollout in completed_batch:
|
||||
self._validate_data(rollout)
|
||||
self._completed_rollouts[rollout.rollout_id] = rollout
|
||||
if verbose:
|
||||
print(f"Completed {len(self._completed_rollouts)}/{self._total_tasks_queued} tasks...")
|
||||
await asyncio.sleep(5)
|
||||
print("All tasks finished.")
|
||||
|
||||
def run_until_all_finished(self, verbose=True):
|
||||
"""Synchronously waits for all queued tasks to be completed and reported."""
|
||||
if self._total_tasks_queued == 0:
|
||||
print("Warning: No tasks were queued.")
|
||||
return
|
||||
|
||||
if not self.server.loop or not self.server.startup_event.is_set():
|
||||
raise RuntimeError("Server is not running or ready.")
|
||||
|
||||
coro = self._async_run_until_finished(verbose)
|
||||
future = asyncio.run_coroutine_threadsafe(coro, self.server.loop)
|
||||
try:
|
||||
future.result() # Wait indefinitely for all tasks to complete
|
||||
except Exception as e:
|
||||
print(f"Error while waiting for tasks to finish: {e}")
|
||||
raise
|
||||
|
||||
def get_test_metrics(self):
|
||||
"""Calculates and returns metrics for a validation run."""
|
||||
assert not self.is_train, "This method should only be called during validation."
|
||||
assert len(self._completed_rollouts) == self._total_tasks_queued
|
||||
|
||||
sample_stat_list = []
|
||||
for rollout_id, rollout in self._completed_rollouts.items():
|
||||
if not rollout.triplets:
|
||||
continue
|
||||
response_length_list = [len(triplet.response.get("token_ids", [])) for triplet in rollout.triplets]
|
||||
final_reward = self._fillna_reward(rollout)
|
||||
sample_stat_list.append(
|
||||
{
|
||||
"sum_response_length": np.sum(response_length_list),
|
||||
"mean_response_length": np.mean(response_length_list) if response_length_list else 0,
|
||||
"turn_count": len(rollout.triplets),
|
||||
"reward": final_reward,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"val/reward": np.mean([stat["reward"] for stat in sample_stat_list]),
|
||||
"val/mean_response_length": np.mean([stat["mean_response_length"] for stat in sample_stat_list]),
|
||||
"val/sum_response_length": np.mean([stat["sum_response_length"] for stat in sample_stat_list]),
|
||||
"val/turn_count": np.mean([stat["turn_count"] for stat in sample_stat_list]),
|
||||
}
|
||||
|
||||
def get_train_data_batch(self, max_prompt_length, max_response_length, device):
|
||||
"""
|
||||
Processes completed rollouts to generate a training data batch.
|
||||
|
||||
This function reconstructs the logic from the original AgentModeDaemon,
|
||||
using data retrieved from the new server architecture. It handles padding,
|
||||
truncation, and tensor creation for the PPO training loop.
|
||||
"""
|
||||
assert self.is_train, "This method should only be called during training."
|
||||
assert len(self._completed_rollouts) == self._total_tasks_queued
|
||||
|
||||
# 1. Reconstruct the `finished_id_to_sample_info` structure from completed rollouts
|
||||
finished_id_to_sample_info = {}
|
||||
for rollout_id, rollout in self._completed_rollouts.items():
|
||||
original_sample = self._task_id_to_original_sample[rollout_id]
|
||||
|
||||
if not rollout.triplets:
|
||||
continue
|
||||
|
||||
# The client should report triplets that contain prompt_ids and response_ids.
|
||||
# Example triplet.prompt: {"token_ids": [...]}
|
||||
# Example triplet.response: {"token_ids": [...]}
|
||||
trace_list = [
|
||||
{"prompt_ids": t.prompt.get("token_ids", []), "response_ids": t.response.get("token_ids", [])}
|
||||
for t in rollout.triplets
|
||||
]
|
||||
|
||||
final_reward = self._fillna_reward(rollout)
|
||||
info = {
|
||||
"reward": final_reward,
|
||||
"trace_list": trace_list,
|
||||
"data_id": original_sample["data_id"],
|
||||
}
|
||||
finished_id_to_sample_info[rollout_id] = info
|
||||
#
|
||||
# --- Data processing and tensor creation logic ---
|
||||
# Get all the reported data.
|
||||
# prompt_ids are left-padded.
|
||||
# response_ids are right-padded.
|
||||
# They are concatenated in the middle.
|
||||
# Discard handling:
|
||||
# - Those exceeding max_prompt_length will be marked for discard, but not
|
||||
# discarded here. They are only truncated and marked, to be discarded later.
|
||||
# This is for the correctness of the advantage calculation.
|
||||
# - The discard for the PPO mini-batch should also be handled this way.
|
||||
input_ids_list, input_attention_mask_list = [], []
|
||||
response_ids_list, response_attention_mask_list = [], []
|
||||
reward_list, data_id_list, rollout_id_list, turn_index_list, is_drop_list = [], [], [], [], []
|
||||
n_trunc_sample_because_of_response = 0
|
||||
|
||||
for rollout_id, sample_info in finished_id_to_sample_info.items():
|
||||
for turn_index, trace in enumerate(sample_info["trace_list"]):
|
||||
|
||||
reward_list.append(sample_info["reward"])
|
||||
prompt_ids, response_ids = trace["prompt_ids"], trace["response_ids"]
|
||||
|
||||
# Mark samples with prompts exceeding max_prompt_length to be dropped later
|
||||
if len(prompt_ids) > max_prompt_length:
|
||||
prompt_ids = prompt_ids[:max_prompt_length]
|
||||
is_drop_list.append(True)
|
||||
else:
|
||||
is_drop_list.append(False)
|
||||
|
||||
# Truncate responses that exceed max_response_length
|
||||
if len(response_ids) > max_response_length:
|
||||
response_ids = response_ids[:max_response_length]
|
||||
n_trunc_sample_because_of_response += 1
|
||||
|
||||
# Pad prompts to the left and responses to the right
|
||||
one_input_ids, one_input_attention_mask = get_left_padded_ids_and_attention_mask(
|
||||
prompt_ids, max_prompt_length, self.pad_token_id
|
||||
)
|
||||
one_response_ids, one_response_attention_mask = get_right_padded_ids_and_attention_mask(
|
||||
response_ids, max_response_length, self.pad_token_id
|
||||
)
|
||||
|
||||
input_ids_list.append(one_input_ids)
|
||||
input_attention_mask_list.append(one_input_attention_mask)
|
||||
response_ids_list.append(one_response_ids)
|
||||
response_attention_mask_list.append(one_response_attention_mask)
|
||||
data_id_list.append(sample_info["data_id"])
|
||||
rollout_id_list.append(rollout_id)
|
||||
turn_index_list.append(turn_index)
|
||||
|
||||
n_transition = len(input_ids_list)
|
||||
batch_input_ids = torch.LongTensor(input_ids_list).to(device)
|
||||
input_attention_mask = torch.LongTensor(input_attention_mask_list).to(device)
|
||||
batch_response_ids = torch.LongTensor(response_ids_list).to(device)
|
||||
response_attention_mask = torch.LongTensor(response_attention_mask_list).to(device)
|
||||
|
||||
# Concatenate prompts and responses to form the full sequence
|
||||
batch_seq = torch.cat([batch_input_ids, batch_response_ids], dim=-1)
|
||||
attention_mask = torch.cat([input_attention_mask, response_attention_mask], dim=-1)
|
||||
position_ids = torch.clamp(torch.cumsum(attention_mask, dim=-1) - 1, min=0)
|
||||
is_drop_mask = torch.BoolTensor(is_drop_list).to(device)
|
||||
scores = torch.tensor(reward_list, dtype=torch.bfloat16).to(device)
|
||||
|
||||
# Create token-level scores by placing the final reward at the last token position
|
||||
token_level_scores = torch.zeros_like(attention_mask, dtype=scores.dtype)
|
||||
# At the eos_mask_idx position of each sample, fill in the corresponding scores.
|
||||
# torch.arange(n_transition) generates [0,1,2,...,bsz-1] as indices for the batch dimension.
|
||||
eos_mask_idx = torch.argmax(position_ids * attention_mask, dim=-1) # (bsz,)
|
||||
token_level_scores[torch.arange(n_transition), eos_mask_idx] = scores
|
||||
# Only take the last response_length part of the sequence to get the token-level scores for the model's response part.
|
||||
token_level_scores = token_level_scores[:, -max_response_length:]
|
||||
|
||||
# Form the final batch using TensorDict
|
||||
batch = TensorDict(
|
||||
{
|
||||
"prompts": batch_input_ids,
|
||||
"responses": batch_response_ids,
|
||||
"input_ids": batch_seq, # here input_ids become the whole sentences
|
||||
"attention_mask": attention_mask,
|
||||
"position_ids": position_ids,
|
||||
"is_drop_mask": is_drop_mask,
|
||||
"token_level_scores": token_level_scores.contiguous(),
|
||||
},
|
||||
batch_size=n_transition,
|
||||
)
|
||||
data_proto = DataProto(batch=batch)
|
||||
|
||||
data_metrics = {
|
||||
"agent_mode/n_trunc_sample_because_of_response": n_trunc_sample_because_of_response,
|
||||
"agent_mode/n_sample_to_train": n_transition,
|
||||
}
|
||||
|
||||
# Add non-tensor data for advantage calculation and logging
|
||||
data_proto.non_tensor_batch["data_id_list"] = np.array(data_id_list)
|
||||
data_proto.non_tensor_batch["rollout_id_list"] = np.array(rollout_id_list)
|
||||
data_proto.non_tensor_batch["turn_index_list"] = np.array(turn_index_list)
|
||||
|
||||
return data_proto, data_metrics
|
||||
|
||||
def clear_data_and_server(self):
|
||||
"""Resets the internal state of the daemon for the next run."""
|
||||
self.backend_llm_server_addresses = []
|
||||
self._completed_rollouts.clear()
|
||||
self._task_id_to_original_sample.clear()
|
||||
self._total_tasks_queued = 0
|
||||
# For a true reset, the server's internal queues would also need clearing.
|
||||
# This implementation assumes that `set_up_data_and_server` is called
|
||||
# for each new run, effectively starting a fresh batch.
|
||||
|
||||
def _fillna_reward(self, rollout):
|
||||
if rollout.final_reward is None:
|
||||
if self.reward_fillna_value is not None:
|
||||
final_reward = self.reward_fillna_value
|
||||
else:
|
||||
raise ValueError(f"Reward is None for rollout {rollout.rollout_id}, please check the reward function.")
|
||||
else:
|
||||
final_reward = rollout.final_reward
|
||||
return final_reward
|
||||
@@ -1,20 +1,45 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
# type: ignore
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from datasets import Dataset as HuggingFaceDataset
|
||||
from verl.utils.dataset.rl_dataset import RLHFDataset
|
||||
|
||||
__all__ = [
|
||||
"LoadedDataset",
|
||||
]
|
||||
|
||||
class AgentDataset(RLHFDataset):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
class LoadedDataset(RLHFDataset):
|
||||
"""Dataset wrapper for pre-loaded in-memory data sequences.
|
||||
|
||||
Bypasses RLHFDataset's file-based initialization and directly sets
|
||||
``self.dataframe`` from the provided sequence.
|
||||
"""
|
||||
|
||||
def __init__(self, dataset: Sequence[Any]):
|
||||
# Skip file-based RLHFDataset initialization; only dataframe behavior is needed.
|
||||
dataset_copy = [dataset[i] for i in range(len(dataset))]
|
||||
self.dataframe = HuggingFaceDataset.from_list(dataset_copy)
|
||||
self.filter_overlong_prompts = False
|
||||
self.serialize_dataset = True # Tell __getstate__ to serialize inline
|
||||
self.original_data_files = None # Not file-backed
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dataframe)
|
||||
|
||||
def __getitem__(self, item):
|
||||
row_dict: dict = self.dataframe[item]
|
||||
|
||||
# add index for each prompt
|
||||
index = row_dict.get("extra_info", {}).get("index", 0)
|
||||
row_dict["index"] = index
|
||||
# Workaround for data proto. At least one tensor is needed.
|
||||
row_dict["fake_ids"] = torch.ones(1, dtype=torch.int)
|
||||
return row_dict
|
||||
|
||||
def _read_files_and_tokenize(self):
|
||||
pass
|
||||
|
||||
@@ -1,156 +1,143 @@
|
||||
import hydra
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""VERL entrypoint for Agent Lightning — wraps verl's PPO setup with a custom trainer.
|
||||
|
||||
Customizations:
|
||||
1. Use AgentLightningRayPPOTrainer (subclass of RayPPOTrainer) that drives rollouts
|
||||
through the Agent Lightning HTTP API instead of stock VERL agent loop workers.
|
||||
2. Support pre-loaded in-memory datasets.
|
||||
"""
|
||||
|
||||
# pyright: reportPrivateImportUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import socket
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, cast
|
||||
|
||||
import ray
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
from .dataset import AgentDataset
|
||||
from .trainer import AgentLightningTrainer
|
||||
from verl.trainer.ppo.reward import load_reward_manager
|
||||
from verl.trainer.main_ppo import create_rl_sampler
|
||||
from .dataset import LoadedDataset
|
||||
|
||||
__all__ = [
|
||||
"run_ppo",
|
||||
]
|
||||
|
||||
|
||||
@hydra.main(config_path="pkg://agentlightning/verl", config_name="config", version_base=None)
|
||||
def main(config):
|
||||
run_ppo(config)
|
||||
def run_ppo(
|
||||
config: Any,
|
||||
train_dataset: Sequence[Any],
|
||||
val_dataset: Sequence[Any],
|
||||
) -> None:
|
||||
"""Launch VERL PPO training with Agent Lightning agent orchestration.
|
||||
|
||||
Datasets must be passed as non-empty in-memory sequences.
|
||||
"""
|
||||
from verl.trainer.main_ppo import get_ppo_ray_runtime_env
|
||||
|
||||
assert train_dataset is not None and len(train_dataset) > 0, "train_dataset must be non-empty"
|
||||
assert val_dataset is not None and len(val_dataset) > 0, "val_dataset must be non-empty"
|
||||
|
||||
def run_ppo(config) -> None:
|
||||
if not ray.is_initialized():
|
||||
# this is for local ray cluster
|
||||
default_runtime_env = cast(dict[str, Any], get_ppo_ray_runtime_env())
|
||||
ray_init_config = OmegaConf.to_container(config.ray_kwargs.get("ray_init", OmegaConf.create({})), resolve=True)
|
||||
ray_init_kwargs: dict[str, Any] = (
|
||||
{str(key): value for key, value in ray_init_config.items()} if isinstance(ray_init_config, dict) else {}
|
||||
)
|
||||
runtime_env_config = ray_init_kwargs.pop("runtime_env", {})
|
||||
runtime_env_kwargs = dict(runtime_env_config) if isinstance(runtime_env_config, dict) else {}
|
||||
runtime_env = {**default_runtime_env, **runtime_env_kwargs}
|
||||
# Register the custom policy loss in each Ray actor process.
|
||||
runtime_env.setdefault(
|
||||
"worker_process_setup_hook",
|
||||
"agentlightning.verl.per_rollout_loss.register_in_worker",
|
||||
)
|
||||
_temp_dir = os.environ.get("RAY_TMPDIR")
|
||||
ray.init(
|
||||
runtime_env={
|
||||
"env_vars": {"TOKENIZERS_PARALLELISM": "true", "NCCL_DEBUG": "WARN", "VLLM_LOGGING_LEVEL": "WARN"}
|
||||
},
|
||||
num_cpus=config.ray_init.num_cpus,
|
||||
runtime_env=runtime_env,
|
||||
**({"_temp_dir": _temp_dir} if _temp_dir else {}),
|
||||
**ray_init_kwargs,
|
||||
)
|
||||
|
||||
runner = TaskRunner.remote()
|
||||
ray.get(runner.run.remote(config))
|
||||
train_ds = LoadedDataset(train_dataset)
|
||||
val_ds = LoadedDataset(val_dataset)
|
||||
|
||||
runner = cast(Any, _AglTaskRunner).remote()
|
||||
ray.get(runner.run.remote(config, train_ds, val_ds))
|
||||
|
||||
|
||||
@ray.remote(num_cpus=1) # please make sure main_task is not scheduled on head
|
||||
class TaskRunner:
|
||||
def run(self, config):
|
||||
# print initial config
|
||||
@ray.remote(num_cpus=1)
|
||||
class _AglTaskRunner:
|
||||
"""TaskRunner that extends verl's TaskRunner with pre-loaded dataset support."""
|
||||
|
||||
def __init__(self):
|
||||
from verl.trainer.main_ppo import TaskRunner
|
||||
|
||||
self._delegate = TaskRunner()
|
||||
|
||||
def run(self, config, train_dataset, val_dataset):
|
||||
from pprint import pprint
|
||||
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
from verl.trainer.main_ppo import (
|
||||
create_rl_sampler,
|
||||
need_critic,
|
||||
need_reference_policy,
|
||||
validate_config,
|
||||
)
|
||||
from verl.utils.dataset.rl_dataset import collate_fn
|
||||
from verl.utils.fs import copy_to_local
|
||||
from verl.utils.tokenizer import hf_processor, hf_tokenizer
|
||||
|
||||
pprint(OmegaConf.to_container(config, resolve=True)) # resolve=True will eval symbol values
|
||||
from agentlightning.verl.trainer import AgentLightningRayPPOTrainer
|
||||
|
||||
print(f"AglTaskRunner hostname: {socket.gethostname()}, PID: {os.getpid()}")
|
||||
pprint(OmegaConf.to_container(config, resolve=True))
|
||||
OmegaConf.resolve(config)
|
||||
|
||||
# download the checkpoint from hdfs
|
||||
local_path = copy_to_local(config.actor_rollout_ref.model.path)
|
||||
# Worker setup — delegated to verl's TaskRunner
|
||||
d = self._delegate
|
||||
actor_rollout_cls, ray_worker_group_cls = d.add_actor_rollout_worker(config)
|
||||
d.add_critic_worker(config)
|
||||
d.add_reward_model_resource_pool(config)
|
||||
d.add_ref_policy_worker(config, actor_rollout_cls)
|
||||
|
||||
# instantiate tokenizer
|
||||
from verl.utils import hf_processor, hf_tokenizer
|
||||
validate_config(
|
||||
config=config,
|
||||
use_reference_policy=need_reference_policy(config),
|
||||
use_critic=need_critic(config),
|
||||
)
|
||||
|
||||
local_path = copy_to_local(
|
||||
config.actor_rollout_ref.model.path,
|
||||
use_shm=config.actor_rollout_ref.model.get("use_shm", False),
|
||||
)
|
||||
trust_remote_code = config.data.get("trust_remote_code", False)
|
||||
tokenizer = hf_tokenizer(local_path, trust_remote_code=trust_remote_code)
|
||||
processor = hf_processor(local_path, use_fast=True) # used for multimodal LLM, could be none
|
||||
processor = hf_processor(local_path, trust_remote_code=trust_remote_code, use_fast=True)
|
||||
|
||||
# define worker classes
|
||||
if config.actor_rollout_ref.actor.strategy in ["fsdp", "fsdp2"]:
|
||||
assert config.critic.strategy in ["fsdp", "fsdp2"]
|
||||
from verl.single_controller.ray import RayWorkerGroup
|
||||
from verl.workers.fsdp_workers import ActorRolloutRefWorker, AsyncActorRolloutRefWorker, CriticWorker
|
||||
resource_pool_manager = d.init_resource_pool_mgr(config)
|
||||
|
||||
actor_rollout_cls = (
|
||||
AsyncActorRolloutRefWorker
|
||||
if config.actor_rollout_ref.rollout.mode == "async"
|
||||
else ActorRolloutRefWorker
|
||||
)
|
||||
ray_worker_group_cls = RayWorkerGroup
|
||||
assert train_dataset is not None and len(train_dataset) > 0, "train_dataset must be non-empty"
|
||||
assert val_dataset is not None and len(val_dataset) > 0, "val_dataset must be non-empty"
|
||||
|
||||
elif config.actor_rollout_ref.actor.strategy == "megatron":
|
||||
assert config.actor_rollout_ref.actor.strategy == config.critic.strategy
|
||||
from verl.single_controller.ray.megatron import NVMegatronRayWorkerGroup
|
||||
from verl.workers.megatron_workers import ActorRolloutRefWorker, CriticWorker
|
||||
|
||||
actor_rollout_cls = ActorRolloutRefWorker
|
||||
ray_worker_group_cls = NVMegatronRayWorkerGroup
|
||||
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
from verl.trainer.ppo.ray_trainer import ResourcePoolManager, Role
|
||||
|
||||
role_worker_mapping = {
|
||||
Role.ActorRollout: ray.remote(actor_rollout_cls),
|
||||
Role.Critic: ray.remote(CriticWorker),
|
||||
}
|
||||
|
||||
global_pool_id = "global_pool"
|
||||
resource_pool_spec = {
|
||||
global_pool_id: [config.trainer.n_gpus_per_node] * config.trainer.nnodes,
|
||||
}
|
||||
mapping = {
|
||||
Role.ActorRollout: global_pool_id,
|
||||
Role.Critic: global_pool_id,
|
||||
}
|
||||
|
||||
# we should adopt a multi-source reward function here
|
||||
# - for rule-based rm, we directly call a reward score
|
||||
# - for model-based rm, we call a model
|
||||
# - for code related prompt, we send to a sandbox if there are test cases
|
||||
# - finally, we combine all the rewards together
|
||||
# - The reward type depends on the tag of the data
|
||||
if config.reward_model.enable:
|
||||
if config.reward_model.strategy in ["fsdp", "fsdp2"]:
|
||||
from verl.workers.fsdp_workers import RewardModelWorker
|
||||
elif config.reward_model.strategy == "megatron":
|
||||
from verl.workers.megatron_workers import RewardModelWorker
|
||||
else:
|
||||
raise NotImplementedError
|
||||
role_worker_mapping[Role.RewardModel] = ray.remote(RewardModelWorker)
|
||||
mapping[Role.RewardModel] = global_pool_id
|
||||
|
||||
# use reference model
|
||||
if config.algorithm.use_kl_in_reward or config.actor_rollout_ref.actor.use_kl_loss:
|
||||
role_worker_mapping[Role.RefPolicy] = ray.remote(ActorRolloutRefWorker)
|
||||
mapping[Role.RefPolicy] = global_pool_id
|
||||
|
||||
reward_fn = load_reward_manager(
|
||||
config, tokenizer, num_examine=0, **config.reward_model.get("reward_kwargs", {})
|
||||
)
|
||||
val_reward_fn = load_reward_manager(
|
||||
config, tokenizer, num_examine=1, **config.reward_model.get("reward_kwargs", {})
|
||||
)
|
||||
resource_pool_manager = ResourcePoolManager(resource_pool_spec=resource_pool_spec, mapping=mapping)
|
||||
|
||||
from verl.utils.dataset.rl_dataset import collate_fn
|
||||
|
||||
# Use our special dataset
|
||||
train_dataset = AgentDataset(
|
||||
data_files=config.data.train_files,
|
||||
tokenizer=tokenizer,
|
||||
processor=processor,
|
||||
config=config.data,
|
||||
)
|
||||
val_dataset = AgentDataset(
|
||||
data_files=config.data.val_files,
|
||||
tokenizer=tokenizer,
|
||||
processor=processor,
|
||||
config=config.data,
|
||||
)
|
||||
train_sampler = create_rl_sampler(config.data, train_dataset)
|
||||
trainer = AgentLightningTrainer(
|
||||
|
||||
trainer = AgentLightningRayPPOTrainer(
|
||||
config=config,
|
||||
tokenizer=tokenizer,
|
||||
processor=processor,
|
||||
role_worker_mapping=role_worker_mapping,
|
||||
role_worker_mapping=d.role_worker_mapping,
|
||||
resource_pool_manager=resource_pool_manager,
|
||||
ray_worker_group_cls=ray_worker_group_cls,
|
||||
reward_fn=reward_fn,
|
||||
val_reward_fn=val_reward_fn,
|
||||
train_dataset=train_dataset,
|
||||
val_dataset=val_dataset,
|
||||
collate_fn=collate_fn,
|
||||
train_sampler=train_sampler,
|
||||
)
|
||||
trainer.init_workers()
|
||||
|
||||
trainer.fit()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Rollout-level mean policy loss for VERL."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import verl.utils.torch_functional as verl_F
|
||||
from verl.trainer.ppo.core_algos import register_policy_loss
|
||||
|
||||
PER_ROLLOUT_MEAN_LOSS_MODE = "per_rollout_mean"
|
||||
|
||||
|
||||
def normalize_advantages_by_rollout(
|
||||
advantages: torch.Tensor,
|
||||
response_mask: torch.Tensor,
|
||||
rollout_ids: Any,
|
||||
*,
|
||||
num_trained_rows: int,
|
||||
) -> torch.Tensor:
|
||||
"""Normalize each row by its rollout's token count and batch size."""
|
||||
if len(rollout_ids) != advantages.shape[0]:
|
||||
raise ValueError(f"rollout_ids length ({len(rollout_ids)}) must match advantages rows ({advantages.shape[0]})")
|
||||
if num_trained_rows <= 0:
|
||||
raise ValueError("num_trained_rows must be positive")
|
||||
|
||||
row_token_counts = response_mask.sum(dim=-1).to(dtype=advantages.dtype)
|
||||
rollout_token_counts: dict[Any, float] = {}
|
||||
for row_index, rollout_id in enumerate(rollout_ids):
|
||||
rollout_token_counts[rollout_id] = rollout_token_counts.get(rollout_id, 0.0) + float(
|
||||
row_token_counts[row_index].item()
|
||||
)
|
||||
|
||||
row_divisors = torch.tensor(
|
||||
[rollout_token_counts[rollout_id] * num_trained_rows for rollout_id in rollout_ids],
|
||||
dtype=advantages.dtype,
|
||||
device=advantages.device,
|
||||
).clamp_min(1.0)
|
||||
return advantages / row_divisors.unsqueeze(-1)
|
||||
|
||||
|
||||
@register_policy_loss(PER_ROLLOUT_MEAN_LOSS_MODE)
|
||||
def compute_policy_loss_per_rollout_mean(
|
||||
old_log_prob: torch.Tensor,
|
||||
log_prob: torch.Tensor,
|
||||
advantages: torch.Tensor,
|
||||
response_mask: torch.Tensor,
|
||||
loss_agg_mode: str = "token-mean",
|
||||
config: Any | None = None,
|
||||
rollout_is_weights: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""Compute clipped PPO loss from rollout-normalized advantages."""
|
||||
assert config is not None, "per_rollout_mean loss requires the actor config"
|
||||
|
||||
clip_ratio = config.clip_ratio
|
||||
clip_ratio_low = config.clip_ratio_low if config.clip_ratio_low is not None else clip_ratio
|
||||
clip_ratio_high = config.clip_ratio_high if config.clip_ratio_high is not None else clip_ratio
|
||||
clip_ratio_c = config.get("clip_ratio_c", 3.0)
|
||||
assert clip_ratio_c > 1.0, f"clip_ratio_c must be greater than 1.0, got {clip_ratio_c}"
|
||||
|
||||
negative_approx_kl = torch.clamp(log_prob - old_log_prob, min=-20.0, max=20.0)
|
||||
ratio = torch.exp(negative_approx_kl)
|
||||
ppo_kl = verl_F.masked_mean(-negative_approx_kl, response_mask)
|
||||
|
||||
pg_losses1 = -advantages * ratio
|
||||
pg_losses2 = -advantages * torch.clamp(ratio, 1 - clip_ratio_low, 1 + clip_ratio_high)
|
||||
clip_pg_losses1 = torch.maximum(pg_losses1, pg_losses2)
|
||||
pg_clipfrac = verl_F.masked_mean(torch.gt(pg_losses2, pg_losses1).float(), response_mask)
|
||||
|
||||
pg_losses3 = -advantages * clip_ratio_c
|
||||
clip_pg_losses2 = torch.min(pg_losses3, clip_pg_losses1)
|
||||
pg_clipfrac_lower = verl_F.masked_mean(
|
||||
torch.gt(clip_pg_losses1, pg_losses3) * (advantages < 0).float(), response_mask
|
||||
)
|
||||
pg_losses = torch.where(advantages < 0, clip_pg_losses2, clip_pg_losses1)
|
||||
|
||||
if rollout_is_weights is not None:
|
||||
pg_losses = pg_losses * rollout_is_weights
|
||||
|
||||
dp_size = config.global_batch_info.get("dp_size", 1) if config.global_batch_info else 1
|
||||
pg_loss = verl_F.masked_sum(pg_losses, response_mask) * (dp_size or 1)
|
||||
metrics = {
|
||||
"actor/pg_clipfrac": pg_clipfrac.detach().item(),
|
||||
"actor/ppo_kl": ppo_kl.detach().item(),
|
||||
"actor/pg_clipfrac_lower": pg_clipfrac_lower.detach().item(),
|
||||
}
|
||||
return pg_loss, metrics
|
||||
|
||||
|
||||
def register_in_worker() -> None:
|
||||
"""Import hook used by Ray actor processes."""
|
||||
@@ -0,0 +1,558 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Adapters from completed Agent Lightning rollouts to VERL training data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import json
|
||||
import zipfile
|
||||
from typing import Any, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from tensordict import TensorDict
|
||||
from verl import DataProto
|
||||
|
||||
from agentlightning.verl.agl_rollout_manager import CompletedRollout
|
||||
|
||||
_TRACE_MERGE_MISMATCH_WANDB_LIMIT = 100
|
||||
_TRACE_MERGE_MISMATCH_TEXT_LIMIT = 4000
|
||||
_ROLLOUT_TRAJECTORY_WANDB_LIMIT = 24
|
||||
_TRACE_MERGE_MISMATCH_COLUMNS = [
|
||||
"global_steps",
|
||||
"rollout_id",
|
||||
"data_id",
|
||||
"turn_index",
|
||||
"template_mismatch",
|
||||
"retoken_mismatch",
|
||||
"others_mismatch",
|
||||
"prompt_length",
|
||||
"response_length",
|
||||
"previous_trace_length",
|
||||
"current_trace_length",
|
||||
"previous_trace",
|
||||
"current_trace",
|
||||
]
|
||||
_ROLLOUT_TRAJECTORY_COLUMNS = [
|
||||
"global_steps",
|
||||
"trajectory_artifact",
|
||||
"trajectory_artifact_path",
|
||||
"row_count",
|
||||
]
|
||||
|
||||
|
||||
def ids_startswith(full_ids: list[int], prefix_ids: list[int]) -> bool:
|
||||
return full_ids[: len(prefix_ids)] == prefix_ids
|
||||
|
||||
|
||||
def _decode_token_ids(tokenizer: Any | None, ids: list[int]) -> str:
|
||||
if tokenizer is not None:
|
||||
try:
|
||||
text = tokenizer.decode(ids, skip_special_tokens=False)
|
||||
except TypeError:
|
||||
text = tokenizer.decode(ids)
|
||||
except Exception:
|
||||
text = " ".join(str(i) for i in ids)
|
||||
else:
|
||||
text = " ".join(str(i) for i in ids)
|
||||
return text
|
||||
|
||||
|
||||
def _decode_trace_text(tokenizer: Any | None, ids: list[int]) -> str:
|
||||
text = _decode_token_ids(tokenizer, ids)
|
||||
if len(text) > _TRACE_MERGE_MISMATCH_TEXT_LIMIT:
|
||||
truncated = len(text) - _TRACE_MERGE_MISMATCH_TEXT_LIMIT
|
||||
return text[:_TRACE_MERGE_MISMATCH_TEXT_LIMIT] + f"\n...[truncated {truncated} chars]"
|
||||
return text
|
||||
|
||||
|
||||
def _token_ids(value: Any) -> list[int]:
|
||||
if isinstance(value, dict) and isinstance(value.get("token_ids"), list):
|
||||
return value["token_ids"]
|
||||
return []
|
||||
|
||||
|
||||
def _artifact_safe_name(value: Any) -> str:
|
||||
text = str(value)
|
||||
safe = "".join(char if char.isascii() and (char.isalnum() or char in {"-", "_", "."}) else "_" for char in text)
|
||||
return safe or "unknown"
|
||||
|
||||
|
||||
def _build_compact_rollout_trajectory_records(
|
||||
rollouts: list[CompletedRollout],
|
||||
*,
|
||||
tokenizer: Any | None,
|
||||
reward_fillna_value: float,
|
||||
limit: int | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
records: list[dict[str, Any]] = []
|
||||
sorted_rollouts = sorted(rollouts, key=lambda rollout: (rollout.step, rollout.sample_idx_in_step))
|
||||
for rollout in sorted_rollouts:
|
||||
if limit is not None and len(records) >= limit:
|
||||
break
|
||||
if not rollout.triplets:
|
||||
continue
|
||||
last_triplet = rollout.triplets[-1]
|
||||
records.append(
|
||||
{
|
||||
"rollout_id": rollout.rollout_id,
|
||||
"reward": rollout.final_reward if rollout.final_reward is not None else reward_fillna_value,
|
||||
"prompt": _decode_token_ids(tokenizer, _token_ids(last_triplet.prompt)),
|
||||
"response": _decode_token_ids(tokenizer, _token_ids(last_triplet.response)),
|
||||
}
|
||||
)
|
||||
return records
|
||||
|
||||
|
||||
def _build_zipped_jsonl(records: list[dict[str, Any]], jsonl_name: str) -> bytes:
|
||||
jsonl_text = "".join(json.dumps(record, ensure_ascii=False, separators=(",", ":")) + "\n" for record in records)
|
||||
buffer = io.BytesIO()
|
||||
with zipfile.ZipFile(buffer, mode="w", compression=zipfile.ZIP_DEFLATED) as zip_file:
|
||||
zip_file.writestr(jsonl_name, jsonl_text.encode("utf-8"))
|
||||
return buffer.getvalue()
|
||||
|
||||
|
||||
def _upload_trace_merge_mismatches_to_wandb(rows: list[dict[str, Any]], global_steps: int) -> None:
|
||||
try:
|
||||
import wandb
|
||||
|
||||
if wandb.run is None:
|
||||
return
|
||||
table = wandb.Table(columns=cast(list[str | int], _TRACE_MERGE_MISMATCH_COLUMNS))
|
||||
for row in rows:
|
||||
table.add_data(*(row.get(column) for column in _TRACE_MERGE_MISMATCH_COLUMNS))
|
||||
wandb.log({"training/trace_merge_mismatches": table}, step=global_steps)
|
||||
except Exception as exc:
|
||||
print(f"Warning: failed to upload trace merge mismatches to wandb: {exc}")
|
||||
|
||||
|
||||
def _upload_compact_rollout_trajectories_to_wandb(
|
||||
records: list[dict[str, Any]],
|
||||
global_steps: int,
|
||||
*,
|
||||
is_validation: bool = False,
|
||||
) -> None:
|
||||
split = "validation" if is_validation else "train"
|
||||
try:
|
||||
import wandb
|
||||
|
||||
if wandb.run is None:
|
||||
return
|
||||
run = wandb.run
|
||||
artifact_type = f"{split}_trajectories"
|
||||
artifact_path = f"step_{global_steps}/{artifact_type}.jsonl.zip"
|
||||
artifact_name = (
|
||||
f"{split}-trajectories-{_artifact_safe_name(getattr(run, 'id', None) or 'run')}-step-{global_steps}"
|
||||
)
|
||||
artifact = wandb.Artifact(
|
||||
name=artifact_name,
|
||||
type=artifact_type,
|
||||
metadata={"global_steps": global_steps, "row_count": len(records), "format": "jsonl.zip"},
|
||||
)
|
||||
with artifact.new_file(artifact_path, mode="wb") as trajectory_file:
|
||||
trajectory_file.write(_build_zipped_jsonl(records, f"{artifact_type}.jsonl"))
|
||||
run.log_artifact(artifact)
|
||||
|
||||
table = wandb.Table(columns=cast(list[str | int], _ROLLOUT_TRAJECTORY_COLUMNS))
|
||||
table.add_data(global_steps, artifact_name, artifact_path, len(records))
|
||||
table_key = "val/rollout_trajectories" if is_validation else "training/rollout_trajectories"
|
||||
wandb.log({table_key: table}, step=global_steps)
|
||||
except Exception as exc:
|
||||
print(f"Warning: failed to upload {split} trajectories to wandb: {exc}")
|
||||
|
||||
|
||||
def get_left_padded_ids_and_attention_mask(
|
||||
ids: list[int], max_length: int, pad_token_id: int
|
||||
) -> tuple[list[int], list[int]]:
|
||||
seq_len = len(ids)
|
||||
if seq_len >= max_length:
|
||||
return ids[-max_length:], [1] * max_length
|
||||
|
||||
pad_len = max_length - seq_len
|
||||
return [pad_token_id] * pad_len + ids, [0] * pad_len + [1] * seq_len
|
||||
|
||||
|
||||
def get_right_padded_ids_and_attention_mask(
|
||||
ids: list[int], max_length: int, pad_token_id: int
|
||||
) -> tuple[list[int], list[int]]:
|
||||
seq_len = len(ids)
|
||||
if seq_len >= max_length:
|
||||
return ids[:max_length], [1] * max_length
|
||||
|
||||
pad_len = max_length - seq_len
|
||||
return ids + [pad_token_id] * pad_len, [1] * seq_len + [0] * pad_len
|
||||
|
||||
|
||||
class RolloutAdapter:
|
||||
"""Convert completed rollout results into VERL data structures."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
max_prompt_length: int,
|
||||
max_response_length: int,
|
||||
device: torch.device,
|
||||
pad_token_id: int,
|
||||
reward_fillna_value: float = 0.0,
|
||||
trace_aggregator_level: str = "transition",
|
||||
tokenizer: Any | None = None,
|
||||
) -> None:
|
||||
self.max_prompt_length = max_prompt_length
|
||||
self.max_response_length = max_response_length
|
||||
self.device = device
|
||||
self.pad_token_id = pad_token_id
|
||||
self.reward_fillna_value = reward_fillna_value
|
||||
self.trace_aggregator_level = trace_aggregator_level
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def get_train_data_batch(
|
||||
self,
|
||||
completed_rollouts: list[CompletedRollout],
|
||||
*,
|
||||
global_steps: int = 0,
|
||||
) -> tuple[DataProto, dict[str, Any]]:
|
||||
"""Build a VERL training batch from completed rollouts."""
|
||||
level = self.trace_aggregator_level
|
||||
if level not in {"transition", "trajectory"}:
|
||||
raise ValueError(f"Unknown trace_aggregator_level: {level}")
|
||||
|
||||
# Keep rollout randomness within each sample instead of ordering samples by completion time.
|
||||
sorted_rollouts = sorted(completed_rollouts, key=lambda rollout: (rollout.step, rollout.sample_idx_in_step))
|
||||
|
||||
final_rewards: list[float] = []
|
||||
sample_with_reward_count = 0
|
||||
sample_with_trace_count = 0
|
||||
|
||||
input_ids_list: list[list[int]] = []
|
||||
input_attention_mask_list: list[list[int]] = []
|
||||
response_ids_list: list[list[int]] = []
|
||||
response_attention_mask_list: list[list[int]] = []
|
||||
response_mask_list: list[list[int]] = []
|
||||
reward_list: list[float] = []
|
||||
data_id_list: list[str] = []
|
||||
rollout_id_list: list[str] = []
|
||||
turn_index_list: list[int] = []
|
||||
is_drop_list: list[bool] = []
|
||||
response_log_probs_list: list[list[float] | None] = []
|
||||
n_trunc_sample_because_of_response = 0
|
||||
n_skipped_empty_training_rows = 0
|
||||
unmerged_count = 0
|
||||
response_len_per_turn_list: list[int] = []
|
||||
merge_mismatch_rows: list[dict[str, Any]] = []
|
||||
|
||||
def append_training_row(
|
||||
*,
|
||||
rollout_id: str,
|
||||
data_id: str,
|
||||
turn_index: int,
|
||||
prompt_ids: list[int],
|
||||
response_ids: list[int],
|
||||
reward: float,
|
||||
response_mask: list[int] | None = None,
|
||||
response_log_probs: list[float] | None = None,
|
||||
) -> None:
|
||||
nonlocal n_skipped_empty_training_rows, n_trunc_sample_because_of_response
|
||||
if len(prompt_ids) > self.max_prompt_length:
|
||||
prompt_ids = prompt_ids[: self.max_prompt_length]
|
||||
is_drop = True
|
||||
else:
|
||||
is_drop = False
|
||||
|
||||
if len(response_ids) > self.max_response_length:
|
||||
response_ids = response_ids[: self.max_response_length]
|
||||
if response_mask is not None:
|
||||
response_mask = response_mask[: self.max_response_length]
|
||||
if response_log_probs is not None:
|
||||
response_log_probs = response_log_probs[: self.max_response_length]
|
||||
n_trunc_sample_because_of_response += 1
|
||||
|
||||
if response_log_probs is not None and len(response_log_probs) != len(response_ids):
|
||||
response_log_probs = None
|
||||
|
||||
train_token_count = sum(response_mask) if response_mask is not None else len(response_ids)
|
||||
if train_token_count == 0:
|
||||
n_skipped_empty_training_rows += 1
|
||||
return
|
||||
|
||||
one_input_ids, one_input_attention_mask = get_left_padded_ids_and_attention_mask(
|
||||
prompt_ids, self.max_prompt_length, self.pad_token_id
|
||||
)
|
||||
one_response_ids, one_response_attention_mask = get_right_padded_ids_and_attention_mask(
|
||||
response_ids, self.max_response_length, self.pad_token_id
|
||||
)
|
||||
input_ids_list.append(one_input_ids)
|
||||
input_attention_mask_list.append(one_input_attention_mask)
|
||||
response_ids_list.append(one_response_ids)
|
||||
response_attention_mask_list.append(one_response_attention_mask)
|
||||
is_drop_list.append(is_drop)
|
||||
if response_mask is not None:
|
||||
one_response_mask, _ = get_right_padded_ids_and_attention_mask(
|
||||
response_mask, self.max_response_length, 0
|
||||
)
|
||||
response_mask_list.append(one_response_mask)
|
||||
|
||||
response_log_probs_list.append(response_log_probs)
|
||||
|
||||
reward_list.append(reward)
|
||||
data_id_list.append(data_id)
|
||||
rollout_id_list.append(rollout_id)
|
||||
if level == "transition":
|
||||
turn_index_list.append(turn_index)
|
||||
|
||||
for rollout in sorted_rollouts:
|
||||
final_reward = self._fillna_reward(rollout)
|
||||
final_rewards.append(final_reward)
|
||||
if rollout.final_reward is not None:
|
||||
sample_with_reward_count += 1
|
||||
|
||||
if not rollout.triplets:
|
||||
print(f"Warning: No triplets found for training rollout {rollout.rollout_id}, skipping.")
|
||||
continue
|
||||
sample_with_trace_count += 1
|
||||
|
||||
if level == "transition":
|
||||
for turn_index, triplet in enumerate(rollout.triplets):
|
||||
response_ids = triplet.response["token_ids"]
|
||||
log_probs = triplet.response["log_probs"]
|
||||
response_len_per_turn_list.append(len(response_ids))
|
||||
append_training_row(
|
||||
rollout_id=rollout.rollout_id,
|
||||
data_id=rollout.data_id,
|
||||
turn_index=turn_index,
|
||||
prompt_ids=triplet.prompt["token_ids"],
|
||||
response_ids=response_ids,
|
||||
reward=final_reward,
|
||||
response_log_probs=log_probs,
|
||||
)
|
||||
continue
|
||||
else:
|
||||
first_triplet = rollout.triplets[0]
|
||||
group_start_turn_index = 0
|
||||
current_prompt_ids = list(first_triplet.prompt["token_ids"])
|
||||
current_response_ids = list(first_triplet.response["token_ids"])
|
||||
current_context = current_prompt_ids + current_response_ids
|
||||
current_response_mask = [1] * len(current_response_ids)
|
||||
current_response_log_probs: list[float] | None = first_triplet.response["log_probs"]
|
||||
response_len_per_turn_list.append(len(current_response_ids))
|
||||
merged_group_count = 0
|
||||
|
||||
for turn_index, triplet in enumerate(rollout.triplets[1:], start=1):
|
||||
prompt_ids = triplet.prompt["token_ids"]
|
||||
response_ids = triplet.response["token_ids"]
|
||||
log_probs = triplet.response["log_probs"]
|
||||
response_len_per_turn_list.append(len(response_ids))
|
||||
next_context = prompt_ids + response_ids
|
||||
|
||||
if ids_startswith(prompt_ids, current_context):
|
||||
if len(prompt_ids) > len(current_context):
|
||||
observation_ids = prompt_ids[len(current_context) :]
|
||||
current_response_ids += observation_ids
|
||||
current_response_mask += [0] * len(observation_ids)
|
||||
if current_response_log_probs is not None:
|
||||
current_response_log_probs += [0.0] * len(observation_ids)
|
||||
current_response_ids += response_ids
|
||||
current_response_mask += [1] * len(response_ids)
|
||||
if current_response_log_probs is not None:
|
||||
if log_probs is None or len(log_probs) != len(response_ids):
|
||||
current_response_log_probs = None
|
||||
else:
|
||||
current_response_log_probs += list(log_probs)
|
||||
current_context = next_context
|
||||
continue
|
||||
|
||||
if len(merge_mismatch_rows) < _TRACE_MERGE_MISMATCH_WANDB_LIMIT:
|
||||
merge_mismatch_rows.append(
|
||||
{
|
||||
"global_steps": global_steps,
|
||||
"rollout_id": rollout.rollout_id,
|
||||
"data_id": rollout.data_id,
|
||||
"turn_index": turn_index,
|
||||
# Token-prefix failures are classified as other mismatches.
|
||||
"template_mismatch": False,
|
||||
"retoken_mismatch": False,
|
||||
"others_mismatch": True,
|
||||
"prompt_length": len(prompt_ids),
|
||||
"response_length": len(response_ids),
|
||||
"previous_trace_length": len(current_context),
|
||||
"current_trace_length": len(next_context),
|
||||
"previous_trace": _decode_trace_text(self.tokenizer, current_context),
|
||||
"current_trace": _decode_trace_text(self.tokenizer, next_context),
|
||||
}
|
||||
)
|
||||
|
||||
append_training_row(
|
||||
rollout_id=rollout.rollout_id,
|
||||
data_id=rollout.data_id,
|
||||
turn_index=group_start_turn_index,
|
||||
prompt_ids=current_prompt_ids,
|
||||
response_ids=current_response_ids,
|
||||
reward=final_reward,
|
||||
response_mask=current_response_mask,
|
||||
response_log_probs=current_response_log_probs,
|
||||
)
|
||||
merged_group_count += 1
|
||||
|
||||
group_start_turn_index = turn_index
|
||||
current_context = next_context
|
||||
current_prompt_ids = list(prompt_ids)
|
||||
current_response_ids = list(response_ids)
|
||||
current_response_mask = [1] * len(response_ids)
|
||||
current_response_log_probs = log_probs
|
||||
|
||||
append_training_row(
|
||||
rollout_id=rollout.rollout_id,
|
||||
data_id=rollout.data_id,
|
||||
turn_index=group_start_turn_index,
|
||||
prompt_ids=current_prompt_ids,
|
||||
response_ids=current_response_ids,
|
||||
reward=final_reward,
|
||||
response_mask=current_response_mask,
|
||||
response_log_probs=current_response_log_probs,
|
||||
)
|
||||
merged_group_count += 1
|
||||
|
||||
if merged_group_count > 1:
|
||||
unmerged_count += 1
|
||||
|
||||
rollout_trajectory_records = _build_compact_rollout_trajectory_records(
|
||||
sorted_rollouts,
|
||||
tokenizer=self.tokenizer,
|
||||
reward_fillna_value=self.reward_fillna_value,
|
||||
limit=_ROLLOUT_TRAJECTORY_WANDB_LIMIT,
|
||||
)
|
||||
_upload_trace_merge_mismatches_to_wandb(merge_mismatch_rows, global_steps)
|
||||
_upload_compact_rollout_trajectories_to_wandb(rollout_trajectory_records, global_steps)
|
||||
|
||||
n_sample = len(input_ids_list)
|
||||
if n_sample == 0:
|
||||
raise RuntimeError("get_train_data_batch emitted zero training rows.")
|
||||
|
||||
batch_input_ids = torch.LongTensor(input_ids_list).to(self.device)
|
||||
input_attention_mask = torch.LongTensor(input_attention_mask_list).to(self.device)
|
||||
batch_response_ids = torch.LongTensor(response_ids_list).to(self.device)
|
||||
response_attention_mask = torch.LongTensor(response_attention_mask_list).to(self.device)
|
||||
batch_response_mask = torch.LongTensor(response_mask_list).to(self.device) if level == "trajectory" else None
|
||||
|
||||
batch_seq = torch.cat([batch_input_ids, batch_response_ids], dim=-1)
|
||||
attention_mask = torch.cat([input_attention_mask, response_attention_mask], dim=-1)
|
||||
position_ids = torch.clamp(torch.cumsum(attention_mask, dim=-1) - 1, min=0)
|
||||
|
||||
row_has_log_probs_list = [log_probs is not None for log_probs in response_log_probs_list]
|
||||
emit_rollout_log_probs = all(row_has_log_probs_list)
|
||||
if not emit_rollout_log_probs and any(row_has_log_probs_list):
|
||||
print("Warning: Mixed rollout log_probs availability, omitting rollout_log_probs from batch.")
|
||||
|
||||
is_drop_mask = torch.BoolTensor(is_drop_list).to(self.device)
|
||||
scores = torch.tensor(reward_list, dtype=torch.bfloat16).to(self.device)
|
||||
|
||||
token_level_scores = torch.zeros_like(attention_mask, dtype=scores.dtype)
|
||||
token_positions = torch.arange(attention_mask.shape[-1], device=attention_mask.device).unsqueeze(0)
|
||||
eos_mask_idx = torch.argmax(token_positions * attention_mask, dim=-1)
|
||||
token_level_scores[torch.arange(n_sample), eos_mask_idx] = scores
|
||||
token_level_scores = token_level_scores[:, -self.max_response_length :]
|
||||
|
||||
batch_dict = {
|
||||
"prompts": batch_input_ids,
|
||||
"responses": batch_response_ids,
|
||||
"input_ids": batch_seq,
|
||||
"attention_mask": attention_mask,
|
||||
"position_ids": position_ids,
|
||||
"is_drop_mask": is_drop_mask,
|
||||
"token_level_scores": token_level_scores.contiguous(),
|
||||
}
|
||||
if level == "trajectory":
|
||||
assert batch_response_mask is not None
|
||||
batch_dict["response_mask"] = batch_response_mask
|
||||
if emit_rollout_log_probs:
|
||||
padded_log_probs_list = [
|
||||
log_probs + [0.0] * (self.max_response_length - len(log_probs))
|
||||
for log_probs in response_log_probs_list
|
||||
if log_probs is not None
|
||||
]
|
||||
batch_dict["rollout_log_probs"] = torch.tensor(padded_log_probs_list, dtype=torch.float32).to(self.device)
|
||||
|
||||
batch = TensorDict(batch_dict, batch_size=n_sample) # type: ignore[arg-type]
|
||||
data_proto = DataProto(batch=batch)
|
||||
data_proto.non_tensor_batch["data_id_list"] = np.array(data_id_list)
|
||||
data_proto.non_tensor_batch["rollout_id_list"] = np.array(rollout_id_list)
|
||||
if level == "transition":
|
||||
data_proto.non_tensor_batch["turn_index_list"] = np.array(turn_index_list)
|
||||
|
||||
n_response_turns = len(response_len_per_turn_list)
|
||||
data_metrics = {
|
||||
"training/reward": float(np.mean(final_rewards)) if final_rewards else 0.0,
|
||||
"training/n_sample": n_sample,
|
||||
"training/n_rollouts": len(sorted_rollouts),
|
||||
"training/n_rollouts_w_trace": sample_with_trace_count,
|
||||
"training/n_rollouts_w_reward": sample_with_reward_count,
|
||||
"training/n_truncated_sample": n_trunc_sample_because_of_response,
|
||||
"training/n_skipped_empty_rows": n_skipped_empty_training_rows,
|
||||
"training/n_turns": n_response_turns,
|
||||
"response_length/training/avg_by_turn": float(np.mean(response_len_per_turn_list)),
|
||||
"response_length/training/max_by_turn": int(np.max(response_len_per_turn_list)),
|
||||
"response_length/training/min_by_turn": int(np.min(response_len_per_turn_list)),
|
||||
}
|
||||
if level == "trajectory":
|
||||
data_metrics["training/n_unmerged_rollouts"] = unmerged_count
|
||||
data_metrics["training/n_trace_merge_mismatch_rows"] = len(merge_mismatch_rows)
|
||||
|
||||
return data_proto, data_metrics
|
||||
|
||||
def get_test_metrics(self, completed_rollouts: list[CompletedRollout], *, global_steps: int = 0) -> dict[str, Any]:
|
||||
"""Build validation metrics from completed rollouts."""
|
||||
sample_stat_list: list[dict[str, Any]] = []
|
||||
|
||||
for rollout in completed_rollouts:
|
||||
final_reward = self._fillna_reward(rollout)
|
||||
sample_stat: dict[str, Any] = {
|
||||
"reward": final_reward,
|
||||
"has_reward": rollout.final_reward is not None,
|
||||
}
|
||||
if rollout.triplets:
|
||||
response_length_list = [len(triplet.response.get("token_ids") or []) for triplet in rollout.triplets]
|
||||
sample_stat.update(
|
||||
{
|
||||
"total_response_length": np.sum(response_length_list),
|
||||
"mean_response_length": np.mean(response_length_list) if response_length_list else 0,
|
||||
"turn_count": len(rollout.triplets),
|
||||
}
|
||||
)
|
||||
sample_stat_list.append(sample_stat)
|
||||
|
||||
stats_w_trace = [stat for stat in sample_stat_list if "total_response_length" in stat]
|
||||
if not stats_w_trace:
|
||||
raise RuntimeError("get_test_metrics received zero completed rollouts with trace.")
|
||||
|
||||
validation_trajectory_records = _build_compact_rollout_trajectory_records(
|
||||
completed_rollouts,
|
||||
tokenizer=self.tokenizer,
|
||||
reward_fillna_value=self.reward_fillna_value,
|
||||
)
|
||||
_upload_compact_rollout_trajectories_to_wandb(
|
||||
validation_trajectory_records,
|
||||
global_steps,
|
||||
is_validation=True,
|
||||
)
|
||||
|
||||
return {
|
||||
"val/reward": float(np.mean([stat["reward"] for stat in sample_stat_list])),
|
||||
"val/n_rollouts": len(sample_stat_list),
|
||||
"val/n_rollouts_w_trace": len(stats_w_trace),
|
||||
"val/n_rollouts_w_reward": len([stat for stat in sample_stat_list if stat["has_reward"]]),
|
||||
"val/mean_response_length_per_turn": float(
|
||||
np.mean([stat["mean_response_length"] for stat in stats_w_trace])
|
||||
),
|
||||
"val/mean_total_response_length_per_rollout": float(
|
||||
np.mean([stat["total_response_length"] for stat in stats_w_trace])
|
||||
),
|
||||
"val/turn_count": float(np.mean([stat["turn_count"] for stat in stats_w_trace])),
|
||||
}
|
||||
|
||||
def _fillna_reward(self, rollout: CompletedRollout) -> float:
|
||||
if rollout.final_reward is not None:
|
||||
return rollout.final_reward
|
||||
return self.reward_fillna_value
|
||||
|
||||
|
||||
__all__ = ["RolloutAdapter"]
|
||||
@@ -0,0 +1,160 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Rollout-level advantage computation for Agent Lightning training batches."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from verl import DataProto
|
||||
from verl.trainer.ppo.ray_trainer import compute_advantage
|
||||
|
||||
|
||||
def compute_rollout_level_advantage(
|
||||
batch: DataProto,
|
||||
*,
|
||||
adv_estimator: Any,
|
||||
gamma: float,
|
||||
lam: float,
|
||||
num_repeat: int,
|
||||
norm_adv_by_std_in_grpo: bool = True,
|
||||
config: Any | None = None,
|
||||
compute_advantage_fn: Callable[..., DataProto] = compute_advantage,
|
||||
) -> tuple[DataProto, dict[str, int]]:
|
||||
"""Compute advantages once per rollout, then broadcast to rollout triplets."""
|
||||
rollout_ids = _required_non_tensor(batch, "rollout_id_list")
|
||||
if len(rollout_ids) != len(batch):
|
||||
raise RuntimeError(f"rollout_id_list length ({len(rollout_ids)}) must match batch length ({len(batch)})")
|
||||
|
||||
uid_values = batch.non_tensor_batch.get("uid")
|
||||
if uid_values is None:
|
||||
uid_values = batch.non_tensor_batch.get("data_id_list")
|
||||
if uid_values is None:
|
||||
raise RuntimeError("rollout-level advantage requires uid or data_id_list in non_tensor_batch")
|
||||
batch.non_tensor_batch["uid"] = uid_values
|
||||
|
||||
response_mask = _required_tensor(batch, "response_mask")
|
||||
token_level_rewards = _required_tensor(batch, "token_level_rewards")
|
||||
|
||||
rollout_to_indices: dict[Any, list[int]] = {}
|
||||
for row_index, rollout_id in enumerate(rollout_ids):
|
||||
rollout_to_indices.setdefault(rollout_id, []).append(row_index)
|
||||
|
||||
reward_sums = token_level_rewards.sum(dim=-1).detach().float()
|
||||
representative_indices: list[int] = []
|
||||
for rollout_id, row_indices in rollout_to_indices.items():
|
||||
representative_indices.append(row_indices[0])
|
||||
_validate_same_uid(rollout_id, row_indices, uid_values)
|
||||
_validate_same_reward(rollout_id, row_indices, reward_sums)
|
||||
|
||||
rollout_batch = batch[representative_indices]
|
||||
rollout_batch = compute_advantage_fn(
|
||||
rollout_batch,
|
||||
adv_estimator=adv_estimator,
|
||||
gamma=gamma,
|
||||
lam=lam,
|
||||
num_repeat=num_repeat,
|
||||
norm_adv_by_std_in_grpo=norm_adv_by_std_in_grpo,
|
||||
config=config,
|
||||
)
|
||||
|
||||
rollout_scalars = _extract_rollout_scalars(
|
||||
rollout_batch,
|
||||
key="advantages",
|
||||
response_mask=rollout_batch.batch["response_mask"],
|
||||
)
|
||||
batch.batch["advantages"] = _broadcast_rollout_scalars(
|
||||
rollout_ids=rollout_ids,
|
||||
rollout_to_scalar=rollout_scalars,
|
||||
response_mask=response_mask,
|
||||
)
|
||||
|
||||
if "returns" in rollout_batch.batch:
|
||||
return_scalars = _extract_rollout_scalars(
|
||||
rollout_batch,
|
||||
key="returns",
|
||||
response_mask=rollout_batch.batch["response_mask"],
|
||||
)
|
||||
batch.batch["returns"] = _broadcast_rollout_scalars(
|
||||
rollout_ids=rollout_ids,
|
||||
rollout_to_scalar=return_scalars,
|
||||
response_mask=response_mask,
|
||||
)
|
||||
|
||||
metrics = {
|
||||
"training/rollout_level_advantage/n_rows": len(batch),
|
||||
"training/rollout_level_advantage/n_rollouts": len(rollout_to_indices),
|
||||
"training/rollout_level_advantage/n_multi_row_rollouts": sum(
|
||||
1 for row_indices in rollout_to_indices.values() if len(row_indices) > 1
|
||||
),
|
||||
"training/rollout_level_advantage/max_rows_per_rollout": max(
|
||||
(len(row_indices) for row_indices in rollout_to_indices.values()),
|
||||
default=0,
|
||||
),
|
||||
}
|
||||
return batch, metrics
|
||||
|
||||
|
||||
def _required_non_tensor(batch: DataProto, key: str) -> Any:
|
||||
values = batch.non_tensor_batch.get(key)
|
||||
if values is None:
|
||||
raise RuntimeError(f"rollout-level advantage requires {key} in non_tensor_batch")
|
||||
return values
|
||||
|
||||
|
||||
def _required_tensor(batch: DataProto, key: str) -> torch.Tensor:
|
||||
value = batch.batch.get(key)
|
||||
if value is None:
|
||||
raise RuntimeError(f"rollout-level advantage requires {key} in batch")
|
||||
return value
|
||||
|
||||
|
||||
def _validate_same_uid(rollout_id: Any, row_indices: list[int], uid_values: Any) -> None:
|
||||
first_uid = uid_values[row_indices[0]]
|
||||
if any(uid_values[row_index] != first_uid for row_index in row_indices[1:]):
|
||||
raise RuntimeError(f"rollout-level advantage found multiple uid values for rollout_id={rollout_id!r}")
|
||||
|
||||
|
||||
def _validate_same_reward(rollout_id: Any, row_indices: list[int], reward_sums: torch.Tensor) -> None:
|
||||
rollout_rewards = reward_sums[row_indices]
|
||||
if not torch.allclose(rollout_rewards, rollout_rewards[0].expand_as(rollout_rewards)):
|
||||
raise RuntimeError(
|
||||
"rollout-level advantage requires all triplets for the same rollout_id "
|
||||
f"to share the same scalar token_level_rewards sum; got rollout_id={rollout_id!r}"
|
||||
)
|
||||
|
||||
|
||||
def _extract_rollout_scalars(
|
||||
batch: DataProto,
|
||||
*,
|
||||
key: str,
|
||||
response_mask: torch.Tensor,
|
||||
) -> dict[Any, torch.Tensor]:
|
||||
values = _required_tensor(batch, key)
|
||||
rollout_ids = _required_non_tensor(batch, "rollout_id_list")
|
||||
scalars: dict[Any, torch.Tensor] = {}
|
||||
for row_index, rollout_id in enumerate(rollout_ids):
|
||||
masked_values = values[row_index][response_mask[row_index].bool()]
|
||||
if masked_values.numel() == 0:
|
||||
raise RuntimeError(f"rollout-level advantage cannot extract {key} for empty rollout_id={rollout_id!r}")
|
||||
first_value = masked_values[0]
|
||||
if not torch.allclose(masked_values, first_value.expand_as(masked_values)):
|
||||
raise RuntimeError(
|
||||
f"rollout-level advantage requires scalar outcome-style {key}; "
|
||||
f"got non-constant token values for rollout_id={rollout_id!r}"
|
||||
)
|
||||
scalars[rollout_id] = first_value.detach()
|
||||
return scalars
|
||||
|
||||
|
||||
def _broadcast_rollout_scalars(
|
||||
*,
|
||||
rollout_ids: Any,
|
||||
rollout_to_scalar: dict[Any, torch.Tensor],
|
||||
response_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
row_scalars = torch.stack([rollout_to_scalar[rollout_id] for rollout_id in rollout_ids])
|
||||
row_scalars = row_scalars.to(device=response_mask.device)
|
||||
return row_scalars.unsqueeze(-1) * response_mask.to(dtype=row_scalars.dtype)
|
||||
@@ -0,0 +1,41 @@
|
||||
# Installation
|
||||
|
||||
This guide sets up a single-node environment for Agent Lightning v1.0. After completing it, you can run single-machine training jobs.
|
||||
|
||||
Before getting started, install `uv` and NVIDIA CUDA. We support CUDA `12.9` or `13.0`.
|
||||
|
||||
#### Step 1: UV Sync
|
||||
|
||||
From the project root, run:
|
||||
|
||||
```bash
|
||||
cd <this-repo>
|
||||
uv sync
|
||||
```
|
||||
|
||||
This installs the base Python environment into `.venv` under the project root.
|
||||
|
||||
#### Step 2: Install `verl` and FlashAttention
|
||||
|
||||
Agent Lightning uses `verl` as its training backend. The compatible versions of `verl`, `vllm`, and `torch` are tightly coupled, and installing `flash-attn` can also be error-prone. We recommend using `scripts/setup_verl.sh` to install the tested, pinned GPU stack and build `flash-attn` from source.
|
||||
|
||||
Pass the `verl` version and CUDA wheel variant explicitly. The script supports `verl==0.7.1` or `verl==0.8.0`, and CUDA wheel variant `cu129` or `cu130`. We recommend CUDA `13.0` with `verl==0.8.0`:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
bash scripts/setup_verl.sh 0.8.0 cu130
|
||||
|
||||
# or
|
||||
|
||||
bash scripts/setup_verl.sh 0.7.1 cu129
|
||||
```
|
||||
|
||||
For `verl==0.7.1`, the script installs `vllm==0.12.0`. For `verl==0.8.0`, it installs `vllm==0.20.2` first, then installs `verl==0.8.0`. Both paths build `flash-attn==2.8.3` locally against the selected environment. Depending on the number of CPU cores available, the script can take 10-30 minutes to complete.
|
||||
|
||||
#### Step 3: W&B Login
|
||||
|
||||
By default, all tasks upload logs and trajectories to Weights & Biases. Log in to W&B before running a task:
|
||||
|
||||
```bash
|
||||
uv run wandb login
|
||||
```
|
||||
@@ -0,0 +1,55 @@
|
||||
# Quick Start
|
||||
|
||||
This quick start requires only one machine with one A100 GPU. It runs Agent Lightning v1.0 with the **local controller** and provides the shortest path from an installed repository to a real rollout-driven training job.
|
||||
|
||||
## Before you start
|
||||
|
||||
Complete [Installation](00-installation.md), including the `verl` GPU stack.
|
||||
|
||||
> AGL v1.0 itself is lightweight, but policy inference and GRPO updates still require the GPU stack used by `verl` and vLLM.
|
||||
|
||||
## 1. Prepare the example
|
||||
|
||||
Download the Calc-X dataset from [Google Drive](https://drive.google.com/file/d/1FQMyKLLd6hP9dw9rfZn1EZOWNvKaDsqw/view?usp=sharing), then extract it and place these files under `examples/calc_x/data/`:
|
||||
|
||||
```text
|
||||
train.parquet
|
||||
test.parquet
|
||||
test_mini.parquet
|
||||
sample.jsonl
|
||||
```
|
||||
|
||||
Activate the project environment and install the dependencies:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
uv pip install openai httpx sympy \
|
||||
"autogen-agentchat" "autogen-ext[openai]" \
|
||||
"mcp>=1.11.0,<2" mcp-server-calculator
|
||||
```
|
||||
|
||||
## 2. Start one local run
|
||||
|
||||
From the repository root:
|
||||
|
||||
```bash
|
||||
examples/calc_x/run_local.sh
|
||||
```
|
||||
|
||||
The launcher performs four operations:
|
||||
|
||||
1. starts Ray and the `verl`/vLLM model backend;
|
||||
2. starts `agl-server` on port `8181`;
|
||||
3. starts `agl-controller runner_type=local`;
|
||||
4. runs the Calc-X training entrypoint.
|
||||
|
||||
Service logs are written under `/tmp/`.
|
||||
|
||||
Once the task is running, you can view the training results in W&B.
|
||||
|
||||
When you want to stop the run, press `Ctrl+C` once and wait for the script to exit. Do not press `Ctrl+C` repeatedly, as the cleanup process takes some time to stop all resources and processes safely.
|
||||
|
||||
## What's Next
|
||||
|
||||
1. Read [Basics](05-basics.md) to learn the core Agent Lightning >= v1.0 concepts.
|
||||
2. Read the complete [Calc-X example](50-example-calc-x.md), which also covers the Kubernetes controller mode.
|
||||
@@ -0,0 +1,96 @@
|
||||
# Basics
|
||||
|
||||
Agent Lightning v1.0 consists of three main components: the **API Gateway**, the **Rollout Controller**, and the **Customized Trainer**. Together, they connect existing agents to reinforcement learning through an OpenAI-compatible endpoint, without requiring changes to the agent's interaction loop.
|
||||
|
||||
## Overview
|
||||
|
||||
<p align="center">
|
||||
<img src="../images/architecture.jpg" alt="Agent Lightning v1.0 architecture" width="75%">
|
||||
</p>
|
||||
|
||||
The API Gateway stores the core rollouts, model endpoints, and events and provides an OpenAI-compatible model proxy for agent requests. The Rollout Controller launches and manages agent executions as local processes or Kubernetes Jobs. Finally, the Customized Trainer runs model inference and optimization on the GPU side and turns the rollout data collected by the Gateway into policy updates.
|
||||
|
||||
This separation provides several practical advantages:
|
||||
|
||||
- **Zero-code-change agent integration:** an existing agent connects by redirecting its OpenAI-compatible model endpoint to the Gateway.
|
||||
- **Independent resources:** model training and agent execution can run on separate machines or clusters and scale independently.
|
||||
- **Open infrastructure:** agents can run on a self-hosted Kubernetes cluster instead of requiring a commercial sandbox service.
|
||||
|
||||
Below, we briefly introduce each component.
|
||||
|
||||
## API Gateway
|
||||
|
||||
The API Gateway is a lightweight service at the center of Agent Lightning. It stores rollouts, model endpoints, and events. It also provides an OpenAI-compatible proxy for agents to access model inference.
|
||||
|
||||
<p align="center">
|
||||
<img src="../images/agentlightning-schema.jpg" alt="API Gateway objects and rollout state transitions" width="50%">
|
||||
</p>
|
||||
|
||||
### Rollout API
|
||||
|
||||
A **rollout** is one execution of an agent on one input. It has a globally unique ID, an input, user-defined metadata, execution configuration, and a status:
|
||||
|
||||
- `QUEUING`: waiting for the Controller to start the agent;
|
||||
- `RUNNING`: the agent is executing;
|
||||
- `SUCCEEDED`: the execution completed successfully;
|
||||
- `FAILED`: the execution ended with an error or timeout.
|
||||
|
||||
A rollout is not the same as a training example. Algorithms such as GRPO may create several independent rollouts from the same example so that the trainer can compare their rewards.
|
||||
|
||||
The trainer creates rollouts through the Rollout API. The Controller reads queued rollouts and updates their status as execution progresses. Each rollout can also contain append-only events, including:
|
||||
|
||||
- `model_request`, recorded automatically for each model call;
|
||||
- `reward`, normally reported by the agent at the end of execution;
|
||||
- custom events for diagnostics and monitoring.
|
||||
|
||||
Every event is associated with a specific rollout ID and is later exported as training data.
|
||||
|
||||
### OpenAI-compatible proxy
|
||||
|
||||
The Gateway also acts as a reverse proxy. The trainer registers one or more model inference endpoints, and the agent sends its model requests to a rollout-specific Gateway URL. The Gateway forwards each request to the registered model endpoint and records its prompt token IDs, response token IDs, and chosen-token log probabilities as a `model_request` event.
|
||||
|
||||
For example, an OpenAI Chat Completions request for a training rollout is sent to:
|
||||
|
||||
```text
|
||||
POST /proxy/rollout/{rollout_id}/attempt/{attempt_id}/mode/train/openai/v1/chat/completions
|
||||
```
|
||||
|
||||
The corresponding validation path uses `mode/val`. OpenAI-compatible clients can use the path through `/openai/v1` as their base URL and append `/chat/completions` normally.
|
||||
|
||||
Because the rollout ID is part of the proxy URL, every model call is automatically associated with the correct execution. An existing agent only needs to use the provided endpoint; it does not need to implement Agent Lightning's rollout or training logic.
|
||||
|
||||
## Rollout Controller
|
||||
|
||||
The Rollout Controller turns queued rollouts into real agent executions. It continuously reconciles rollout state in the API Gateway with the processes or Jobs it manages, and reports execution progress back to the Gateway.
|
||||
|
||||
<p align="center">
|
||||
<img src="../images/controller-reconciliation.jpg" alt="Controller reconciliation" width="75%">
|
||||
</p>
|
||||
|
||||
The Controller supports two modes:
|
||||
|
||||
- **Local mode:** starts each rollout as a short-lived local subprocess. This mode is convenient for development and debugging when the agent and trainer dependencies can share one machine.
|
||||
- **Kubernetes mode:** creates one Kubernetes Job for each rollout from a user-provided template. This mode isolates agent dependencies and supports concurrent execution on a self-hosted or on-premises cluster.
|
||||
|
||||
The API Gateway remains the source of truth for rollout status. If a process, Kubernetes watch, or network update is interrupted, the Controller retries reconciliation until the execution state converges.
|
||||
|
||||
## Customized Trainer
|
||||
|
||||
The Customized Trainer sits on top of `verl` and connects the training backend to the API Gateway. During each training step, it:
|
||||
|
||||
1. registers the current model inference endpoints;
|
||||
2. creates one or more rollouts for each training input;
|
||||
3. waits for enough rollouts to finish;
|
||||
4. retrieves model requests, rewards, and other events;
|
||||
5. converts the captured calls into `verl` training samples;
|
||||
6. computes advantages and updates the policy.
|
||||
|
||||
The trainer also handles Agent Lightning-specific data processing. It merges consecutive model calls only when their token histories are exactly continuous, computes advantages at the rollout level, and supports rollout-level loss normalization.
|
||||
|
||||
## Configure the components
|
||||
|
||||
The following chapters describe the settings for each component. Start with the trainer to define how rollouts are created and converted into training samples, then configure the server and the Controller that execute them:
|
||||
|
||||
1. [Trainer Configuration](20-trainer-configuration.md)
|
||||
2. [API Gateway Configuration](25-api-gateway-configuration.md)
|
||||
3. [Controller Configuration](30-controller-configuration.md)
|
||||
@@ -0,0 +1,215 @@
|
||||
# Trainer Configuration
|
||||
|
||||
Agent Lightning v1.0 adds its configuration on top of `verl`'s `ppo_trainer` Hydra configuration. The complete default configuration from `agentlightning/verl/config.yaml` is shown below. The following sections explain these settings in detail.
|
||||
|
||||
Complete default configuration added by Agent Lightning:
|
||||
|
||||
```yaml
|
||||
algorithm:
|
||||
enable_rollout_level_advantage: true
|
||||
|
||||
agentlightning:
|
||||
agl_base_url: http://localhost:8080
|
||||
agl_key: ""
|
||||
hooks: null
|
||||
rollout_timeout_seconds: 1800
|
||||
local:
|
||||
agent_class: null
|
||||
env_map: {}
|
||||
k8s:
|
||||
job_template_path: null
|
||||
reward_fillna_value: 0.0
|
||||
max_ppo_update_times: null
|
||||
trace_aggregator:
|
||||
level: trajectory # transition | trajectory
|
||||
trajectory_max_prompt_length: 2048
|
||||
trajectory_max_response_length: 8192
|
||||
async_rollout:
|
||||
enabled: false
|
||||
async_train_batch_size: null
|
||||
|
||||
actor_rollout_ref:
|
||||
actor:
|
||||
policy_loss:
|
||||
loss_mode: per_rollout_mean
|
||||
```
|
||||
|
||||
The configuration above shows the Agent Lightning settings. At runtime, these settings are merged with the original `verl` `ppo_trainer` configuration, whose existing options remain available and take effect as usual.
|
||||
|
||||
## Connect to the API Gateway
|
||||
|
||||
The first group of settings connects the trainer to the Agent Lightning API Gateway:
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `agentlightning.agl_base_url` | `http://localhost:8080` | Gateway URL used by the rollout manager. |
|
||||
| `agentlightning.agl_key` | `""` | Bearer key; must match the API Gateway and Controller. |
|
||||
|
||||
Make sure the machine running the trainer can reach the Gateway at `agentlightning.agl_base_url`. The Hydra `agl_key` value must be identical in the trainer, API Gateway, and Controller configurations.
|
||||
|
||||
## Model and Data
|
||||
|
||||
Model configuration follows the standard `verl` `actor_rollout_ref.model` settings. Set `actor_rollout_ref.model.path` to a Hugging Face model name or local model path:
|
||||
|
||||
```yaml
|
||||
actor_rollout_ref:
|
||||
model:
|
||||
path: Qwen/Qwen2.5-1.5B-Instruct
|
||||
```
|
||||
|
||||
In upstream `verl`, dataset paths are normally configured with `data.train_files` and `data.val_files`. Agent Lightning instead loads the files first and passes the resulting datasets directly to `run_ppo`. This provides additional flexibility: users can pass any dataset as long as it can be represented as a list of JSON objects.
|
||||
|
||||
```python
|
||||
from datasets import Dataset
|
||||
|
||||
from agentlightning.verl.entrypoint import run_ppo
|
||||
|
||||
train_dataset = Dataset.from_parquet("data/train.parquet").to_list()
|
||||
val_dataset = Dataset.from_parquet("data/test.parquet").to_list()
|
||||
|
||||
run_ppo(config, train_dataset=train_dataset, val_dataset=val_dataset)
|
||||
```
|
||||
|
||||
`run_ppo` accepts non-empty in-memory sequences as `train_dataset` and `val_dataset`. Each element is read as a JSON-like object. When the trainer creates a rollout, each element in the list becomes the rollout's `input` field. The Controller can then map fields from `input` into the agent's environment or Kubernetes Job template.
|
||||
|
||||
## Rollout execution
|
||||
|
||||
The Controller has two execution modes: `local` and `k8s`. Configure the matching section below, and the Controller reads that section according to its running mode.
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `agentlightning.local.agent_class` | `null` | Fully qualified Python class imported and started by the Controller in local mode. |
|
||||
| `agentlightning.local.env_map` | `{}` | Maps environment variable names to fields in the rollout `input`. |
|
||||
| `agentlightning.k8s.job_template_path` | `null` | Path to the Jinja Kubernetes Job template used by the Controller in K8s mode. |
|
||||
|
||||
For local execution, set the agent class and map fields from each dataset row into environment variables. For example:
|
||||
|
||||
```yaml
|
||||
agentlightning:
|
||||
local:
|
||||
agent_class: examples.search_r1.agents.search_r1_agent.SearchR1Agent
|
||||
env_map:
|
||||
QUESTION: input.question
|
||||
GOLDEN_ANSWERS: input.golden_answers
|
||||
```
|
||||
|
||||
Here, the Controller imports `SearchR1Agent`, starts one local subprocess for each rollout, and sets `QUESTION` and `GOLDEN_ANSWERS` from that rollout's `input` object.
|
||||
|
||||
In K8s mode, provide a Jinja template that renders to a Kubernetes Job YAML manifest:
|
||||
|
||||
```yaml
|
||||
agentlightning:
|
||||
k8s:
|
||||
job_template_path: examples/calc_x/job-template.yaml
|
||||
```
|
||||
|
||||
The template can use values from the rollout `input`. For example, this fragment replaces the environment-variable values with fields from the current dataset row:
|
||||
|
||||
```yaml
|
||||
env:
|
||||
- name: QUESTION
|
||||
value: {% raw %}{{ input.question | yaml_escape }}{% endraw %}
|
||||
- name: RESULT
|
||||
value: {% raw %}{{ input.result | yaml_escape }}{% endraw %}
|
||||
```
|
||||
|
||||
The trainer reads the Jinja template and includes its text in each rollout. The Controller renders it with that rollout's `input`, then creates one Kubernetes Job per rollout.
|
||||
|
||||
Finally, `agentlightning.rollout_timeout_seconds` sets the maximum execution time for each rollout in both modes. The Controller uses this value and marks a rollout as failed if it does not finish within the configured number of seconds. The default is `1800`.
|
||||
|
||||
## Trace aggregator
|
||||
|
||||
<p align="center">
|
||||
<img src="../images/trajectory-aggregation.jpg" alt="Trajectory aggregation" width="80%">
|
||||
</p>
|
||||
|
||||
The left side of the diagram shows traditional agentic RL, where each rollout corresponds to one training sample. The right side shows Agent Lightning, where one rollout can correspond to multiple training samples. During a rollout, the Gateway collects all raw LLM calls as prompt-response pairs, and the trace aggregator assembles them into training samples using one of the following two modes.
|
||||
|
||||
### Trajectory mode
|
||||
|
||||
`trajectory` is the default and recommended mode:
|
||||
|
||||
```yaml
|
||||
agentlightning:
|
||||
trace_aggregator:
|
||||
level: trajectory
|
||||
trajectory_max_prompt_length: 2048
|
||||
trajectory_max_response_length: 8192
|
||||
```
|
||||
|
||||
The aggregator automatically merges consecutive calls when the next prompt starts with the exact token sequence of the previous prompt and response. Tokens added between calls, such as tool observations, are retained as context but masked from the policy loss. If exact token-prefix continuity is broken, the aggregator starts a new training row instead of merging incompatible calls.
|
||||
|
||||
In this mode:
|
||||
|
||||
- `trajectory_max_prompt_length` limits the initial prompt in each merged training row;
|
||||
- `trajectory_max_response_length` limits all content after the initial prompt. This includes the prompts and responses from later turns, which are merged into the trajectory response sequence.
|
||||
|
||||
We recommend setting `trajectory_max_response_length` relatively high so it can hold multiple turns without truncation. Choose a value that covers the expected combined length of later-turn prompts and responses while fitting the model context window and available GPU memory.
|
||||
|
||||
Training rows whose initial prompt exceeds `trajectory_max_prompt_length` are marked and dropped from the policy-update batch. Content beyond `trajectory_max_response_length`, on the other hand, is truncated to the configured response length.
|
||||
|
||||
The number of dropped and truncated rows is reported to W&B with these metrics:
|
||||
|
||||
- `training/n_sample_dropped/marked` — rows dropped because their prompts exceeded the configured prompt limit;
|
||||
- `training/n_truncated_sample` — rows whose responses were truncated to the configured response limit.
|
||||
|
||||
|
||||
### Transition mode
|
||||
|
||||
In `transition` mode, every model call becomes an independent training row and no calls are merged:
|
||||
|
||||
```yaml
|
||||
agentlightning:
|
||||
trace_aggregator:
|
||||
level: transition
|
||||
|
||||
data:
|
||||
max_prompt_length: 4096
|
||||
max_response_length: 2048
|
||||
```
|
||||
|
||||
Transition mode does not use `trajectory_max_prompt_length` or `trajectory_max_response_length`. It uses the same standard `verl` data limits used for individual vLLM rollout calls:
|
||||
|
||||
- `data.max_prompt_length` limits each call's prompt;
|
||||
- `data.max_response_length` limits each call's response.
|
||||
|
||||
Use transition mode when every request-response call should remain a separate training sample.
|
||||
|
||||
## Algorithm correctness
|
||||
|
||||
The following settings control how rollout data contributes to optimization:
|
||||
|
||||
```yaml
|
||||
algorithm:
|
||||
enable_rollout_level_advantage: true
|
||||
|
||||
actor_rollout_ref:
|
||||
actor:
|
||||
policy_loss:
|
||||
loss_mode: per_rollout_mean
|
||||
|
||||
agentlightning:
|
||||
max_ppo_update_times: 2
|
||||
```
|
||||
|
||||
### Rollout-level advantage
|
||||
|
||||
`algorithm.enable_rollout_level_advantage: true` computes the advantage at the rollout level rather than independently at the training-sample level. This is important because one rollout can produce a variable number of training rows after trace aggregation.
|
||||
|
||||
### Per-rollout mean loss
|
||||
|
||||
`actor_rollout_ref.actor.policy_loss.loss_mode: per_rollout_mean` normalizes the policy loss at the rollout level. It prevents a rollout from receiving more optimization weight only because it produced more training rows.
|
||||
|
||||
For the motivation and detailed formulation of rollout-level advantage and loss normalization, see the [Agent Lightning v1.0 technical report](https://arxiv.org/pdf/2608.17528).
|
||||
|
||||
### Maximum PPO update times
|
||||
|
||||
In extreme cases, trace aggregation may produce too many training samples from one collected batch, which can increase the number of PPO updates and affect training stability. `agentlightning.max_ppo_update_times` limits the maximum number of PPO mini-batch updates performed for one batch.
|
||||
|
||||
The default value is `null`, which applies no explicit update cap. In this case, the trainer uses all complete PPO mini-batches collected for the step; only samples that do not fill a complete mini-batch are dropped for alignment.
|
||||
|
||||
For additional training stability, we recommend setting it to `2`. Samples beyond this limit are dropped before the policy update. The number of samples dropped for mini-batch alignment or this update cap is reported in W&B through `training/n_sample_dropped/same_reward` and `training/n_sample_dropped/random`.
|
||||
|
||||
## Asynchronous training
|
||||
|
||||
Agent Lightning supports collocated asynchronous rollout collection through `agentlightning.async_rollout`. For configuration, behavior, and constraints, see [Asynchronous Training](35-asynchronous-training.md).
|
||||
@@ -0,0 +1,60 @@
|
||||
# API Gateway Configuration
|
||||
|
||||
Start the API Gateway with `agl-server`. It uses Hydra configuration, and the complete default configuration is located at `agentlightning/config/server.yaml`:
|
||||
|
||||
```yaml
|
||||
host: 0.0.0.0
|
||||
port: 8080
|
||||
key: ""
|
||||
default_proxy:
|
||||
model_name: "Qwen/Qwen2.5-7B-Instruct"
|
||||
include_log_probs: true
|
||||
train:
|
||||
temperature: 1
|
||||
val:
|
||||
temperature: 0.7
|
||||
```
|
||||
|
||||
## Override configuration
|
||||
|
||||
Override any setting with a Hydra command-line argument when starting the server. For example:
|
||||
|
||||
```bash
|
||||
agl-server \
|
||||
host=0.0.0.0 \
|
||||
port=8080 \
|
||||
key="$AGL_KEY" \
|
||||
default_proxy.model_name=Qwen/Qwen3-8B
|
||||
```
|
||||
|
||||
## Top-level settings
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `host` | `0.0.0.0` | Uvicorn bind address. |
|
||||
| `port` | `8080` | Uvicorn listen port. |
|
||||
| `key` | `""` | Bearer key for API and proxy routes. Empty disables authentication and logs a warning. |
|
||||
|
||||
Use the same non-empty key in the trainer and Controller.
|
||||
|
||||
## Proxy settings
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `default_proxy.model_name` | `Qwen/Qwen2.5-7B-Instruct` | Registered model name selected for forwarded requests. |
|
||||
| `default_proxy.include_log_probs` | `true` | Ask the train backend for chosen-token log probabilities and token IDs. |
|
||||
| `default_proxy.train.temperature` | `1` | Temperature forced for training rollouts. |
|
||||
| `default_proxy.val.temperature` | `0.7` | Temperature forced for validation rollouts. |
|
||||
|
||||
`default_proxy.model_name` must match the model configured in `verl` at `actor_rollout_ref.model.path`:
|
||||
|
||||
```text
|
||||
server default_proxy.model_name
|
||||
= trainer actor_rollout_ref.model.path
|
||||
```
|
||||
|
||||
A mismatch produces a “model not found” error even if the vLLM endpoint itself is healthy.
|
||||
|
||||
The train and validation temperatures configured here are the values actually used for model requests. Note that `verl` has similar temperature settings, but those values are not used for proxied requests because the proxy replaces them automatically.
|
||||
|
||||
We recommend keeping `default_proxy.include_log_probs: true`. This records rollout log probabilities and allows `verl` to report rollout-correction metrics. Some rollout-correction features also require these log probabilities.
|
||||
@@ -0,0 +1,93 @@
|
||||
# Controller Configuration
|
||||
|
||||
Start the Controller with `agl-controller`. It translates declarative rollouts into real agent executions, uses Hydra configuration, and loads its complete default configuration from `agentlightning/config/controller.yaml`:
|
||||
|
||||
<p align="center">
|
||||
<img src="../images/controller-reconciliation.jpg" alt="Controller reconciliation" width="80%">
|
||||
</p>
|
||||
|
||||
```yaml
|
||||
runner_type: k8s
|
||||
|
||||
agl_server:
|
||||
url: http://localhost:8080
|
||||
agent_url: null
|
||||
key: ""
|
||||
|
||||
k8s_runner:
|
||||
namespace: default
|
||||
ttl_after_finished: 1200
|
||||
max_jobs_per_minute: 100
|
||||
poll_interval: 5
|
||||
|
||||
local_runner:
|
||||
maximum_size: 50
|
||||
poll_interval: 10
|
||||
```
|
||||
|
||||
## Override configuration
|
||||
|
||||
Override any setting with a Hydra command-line argument when starting the Controller. For example:
|
||||
|
||||
```bash
|
||||
agl-controller \
|
||||
runner_type=local \
|
||||
agl_server.url=http://localhost:8080 \
|
||||
agl_server.key="$AGL_KEY" \
|
||||
local_runner.maximum_size=32
|
||||
```
|
||||
## Runner type
|
||||
|
||||
The Controller supports one runner type at a time. Set `runner_type` to either `k8s` or `local`:
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `runner_type` | `k8s` | The single execution backend used by this Controller: `k8s` or `local`. |
|
||||
|
||||
One Controller instance cannot run both modes simultaneously.
|
||||
|
||||
In `k8s` mode, every rollout runs as a Kubernetes Job. The Controller uses the default Kubernetes configuration at `~/.kube/config` on its machine to access the cluster. In `local` mode, every rollout runs as a local subprocess on the Controller machine, with multiple rollouts managed through a local process pool.
|
||||
|
||||
## Connect to the API Gateway
|
||||
|
||||
The Controller configuration contains two API Gateway URLs for two different network paths:
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `agl_server.url` | `http://localhost:8080` | API Gateway URL used by the Controller itself. |
|
||||
| `agl_server.agent_url` | `null` | API Gateway URL used by the Agent. When `null`, it falls back to `agl_server.url`. |
|
||||
| `agl_server.key` | `""` | Bearer key used by the Controller and agents. |
|
||||
|
||||
`agl_server.url` must be reachable from the Controller process. `agl_server.agent_url` must be reachable from the Agent process or pod because it is used to build the Gateway proxy and event URLs injected into that Agent.
|
||||
|
||||
In most cases, `agl_server.agent_url` does not need to be set separately. Leave it as `null`, and Agents automatically use `agl_server.url` to access the API Gateway.
|
||||
|
||||
Set `agl_server.agent_url` only when Agents cannot reach the API Gateway through `agl_server.url`, usually because the Controller and Agents are in different networks. For example, when using the Minikube Docker driver, the Controller runs locally while Agents run inside Minikube, so they do not share the same network. The Controller may use `http://localhost:8080`, while Agents inside Minikube need `http://host.minikube.internal:8080` to access the same API Gateway.
|
||||
|
||||
## K8s runner limits
|
||||
|
||||
The K8s runner provides settings that limit Job creation and clean up completed Jobs:
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `k8s_runner.max_jobs_per_minute` | `100` | Maximum number of Kubernetes Jobs the Controller can create per minute. |
|
||||
| `k8s_runner.ttl_after_finished` | `1200` | Number of seconds a completed Job is retained before Kubernetes removes it automatically. |
|
||||
|
||||
`max_jobs_per_minute` prevents the Controller from creating too many Jobs in a short period. `ttl_after_finished` prevents completed Jobs from accumulating and overloading the Kubernetes API server.
|
||||
|
||||
## Local runner limits
|
||||
|
||||
The local runner limits concurrent processes and periodically synchronizes their state:
|
||||
|
||||
| Key | Default | Description |
|
||||
|---|---:|---|
|
||||
| `local_runner.maximum_size` | `50` | Maximum number of Agent subprocesses managed concurrently on the Controller machine. |
|
||||
| `local_runner.poll_interval` | `10` | Number of seconds between automatic synchronization checks for local process and rollout state. |
|
||||
|
||||
When the process pool reaches `maximum_size`, queued rollouts wait until capacity becomes available.
|
||||
|
||||
## How agents are launched
|
||||
|
||||
In `k8s` mode, the Controller reads the Jinja Job template stored in each rollout. The template originates from `agentlightning.k8s.job_template_path` in the Trainer configuration. The Controller renders the template with the rollout's `input`, applies the Controller settings such as namespace, timeout, and cleanup TTL, and submits the resulting Kubernetes Job. It also injects rollout-specific `AGL_OPENAI_BASE_URL`, `AGL_EVENT_URL`, and `AGL_KEY` values into every container. See [Trainer Configuration](20-trainer-configuration.md#rollout-execution) for the template configuration and Jinja examples, and `agentlightning/controller/k8s_reconciler.py` for the implementation.
|
||||
|
||||
In `local` mode, the Controller reads `agent_class` and `env_map` from the rollout. It imports the Agent class, starts it in a local subprocess, and uses `env_map` to replace environment-variable values with fields from the rollout's `input`. The same rollout-specific Gateway URL, event URL, and key are injected automatically.
|
||||
@@ -0,0 +1,73 @@
|
||||
# Asynchronous Training
|
||||
|
||||
Long-running agents can have very different rollout durations. In synchronous training, one slow rollout can delay the whole update step. Agent Lightning v1.0 supports **collocated asynchronous training**, where rollout generation and model updates share the same GPU pool while unfinished rollout groups carry over to later steps.
|
||||
|
||||

|
||||
|
||||
## Enable asynchronous training
|
||||
|
||||
Enable Agent Lightning asynchronous collection with `agentlightning.async_rollout.enabled`:
|
||||
|
||||
```yaml
|
||||
agentlightning:
|
||||
async_rollout:
|
||||
enabled: true
|
||||
async_train_batch_size: 64
|
||||
```
|
||||
|
||||
You must also set `async_train_batch_size`. It is the number of prompt groups kept active for rollout collection and must be strictly greater than `data.train_batch_size`, which is the number of completed groups consumed by one update:
|
||||
|
||||
```yaml
|
||||
data:
|
||||
train_batch_size: 32
|
||||
|
||||
agentlightning:
|
||||
async_rollout:
|
||||
enabled: true
|
||||
async_train_batch_size: 64
|
||||
```
|
||||
|
||||
A useful starting point is:
|
||||
|
||||
$$B_{async} = 2 B_{train}.$$
|
||||
|
||||
Increase `async_train_batch_size` when rollout durations vary significantly and the resource for running agents has enough capacity. Reduce it when active processes or Kubernetes Jobs consume too many CPU or memory resources.
|
||||
|
||||
## How it works
|
||||
|
||||
The asynchronous collection process is:
|
||||
|
||||
1. The trainer keeps up to `async_train_batch_size` prompt groups active.
|
||||
2. The Controller starts their Agent executions in local processes or Kubernetes Jobs.
|
||||
3. When `data.train_batch_size` complete groups are available, the trainer selects them for the next update instead of waiting for every active group.
|
||||
4. Unfinished groups remain active and carry over to the next collection step.
|
||||
5. Before updating model weights, the Gateway pauses new model requests and waits for requests already in flight to finish.
|
||||
6. The shared GPUs perform the model update, then inference resumes for the next rollout phase.
|
||||
|
||||
Each prompt group remains intact. For example, when `actor_rollout_ref.rollout.n` is `4`, all four sibling rollouts must finish before that group can be used by the optimizer. This preserves GRPO/RLOO group statistics.
|
||||
|
||||
Agents should use a retrying OpenAI or HTTP client. A request arriving while the Gateway is paused receives a retryable response and can continue after inference resumes.
|
||||
|
||||
## Monitoring
|
||||
|
||||
The trainer reports asynchronous collection metrics to W&B:
|
||||
|
||||
| Metric | Interpretation |
|
||||
|---|---|
|
||||
| `training/async/n_prev_carry_over_rollouts` | Rollouts inherited from the previous step. |
|
||||
| `training/async/n_completed_rollouts` | Rollouts consumed by the current step. |
|
||||
| `training/async/n_new_carry_over_rollouts` | Unfinished rollouts carried into the next step. |
|
||||
| `training/async/new_carry_over_age_max_steps` | Oldest carry-over age in optimizer steps. |
|
||||
| `training/async/proxy_inflight_at_pause` | Requests still running when the Gateway pause begins. |
|
||||
| `training/async/proxy_drain_seconds` | Time spent waiting for in-flight requests to finish. |
|
||||
|
||||
## Handle staleness
|
||||
|
||||
Asynchronous rollouts may be generated by an older model version and become stale before they are used for training. To correct this policy mismatch, enable `verl`'s [rollout correction](https://verl.readthedocs.io/en/latest/algo/rollout_corr.html). We recommend token-level importance sampling (TIS) with a clipping threshold of `2`:
|
||||
|
||||
```yaml
|
||||
algorithm:
|
||||
rollout_correction:
|
||||
rollout_is: token
|
||||
rollout_is_threshold: 2
|
||||
```
|
||||
@@ -0,0 +1,73 @@
|
||||
# Calc-X
|
||||
|
||||
| GPU | Model | Controller Mode | Trainer Mode | Code |
|
||||
|---|---|---|---|---|
|
||||
| 1× A100 80GB | `Qwen/Qwen2.5-1.5B-Instruct` | K8s or local | Sync and async | [Source](https://github.com/microsoft/agent-lightning/tree/main/examples/calc_x) |
|
||||
|
||||
Calc-X is a proof-of-concept (POC) example that trains a mathematical reasoning agent on the Calc-X dataset with `verl` and Agent Lightning >=v1.0. It is intentionally lightweight and requires only one GPU. The agent uses AutoGen + MCP calculator tools to solve math problems.
|
||||
|
||||
The example supports two controller modes:
|
||||
|
||||
- **K8s mode:** Minikube provides a minimal Kubernetes environment, and agent rollouts run as Kubernetes Jobs.
|
||||
- **Local mode:** Agent rollouts run directly as local processes without Kubernetes.
|
||||
|
||||
Both synchronous and asynchronous trainer modes are supported.
|
||||
|
||||
## Data Preparation
|
||||
|
||||
Download the Calc-X dataset from [Google Drive](https://drive.google.com/file/d/1FQMyKLLd6hP9dw9rfZn1EZOWNvKaDsqw/view?usp=sharing), then extract it into `examples/calc_x/data/`:
|
||||
|
||||
```bash
|
||||
cd examples/calc_x
|
||||
unzip data/calc-x-data.zip -d data/
|
||||
```
|
||||
|
||||
The expected dataset files are:
|
||||
|
||||
- `data/train.parquet`
|
||||
- `data/test.parquet`
|
||||
- `data/test_mini.parquet`
|
||||
- `data/sample.jsonl`
|
||||
|
||||
## Local Mode
|
||||
|
||||
Make sure you have activated the project environment and installed the following package in Python:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
uv pip install \
|
||||
openai \
|
||||
httpx \
|
||||
sympy \
|
||||
"autogen-agentchat" \
|
||||
"autogen-ext[openai]" \
|
||||
"mcp>=1.11.0,<2" \
|
||||
mcp-server-calculator
|
||||
```
|
||||
|
||||
Then start training:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
bash run_local.sh
|
||||
```
|
||||
|
||||
`run_local.sh` starts `agl-server` and `agl-controller`, and writes their logs under `/tmp/`. The script starts the agent in multi-process mode.
|
||||
When `run_local.sh` exits, it automatically cleans up the server, controller, and agent it started.
|
||||
|
||||
## K8s Mode
|
||||
|
||||
This example uses Minikube to demonstrate the minimal Kubernetes workflow. For production deployments, replace Minikube with a production-grade Kubernetes cluster.
|
||||
|
||||
Make sure you have installed `docker` and `minikube`, then start training by:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
cd examples/calc_x
|
||||
bash run_minikube.sh
|
||||
```
|
||||
|
||||
`run_minikube.sh` starts `agl-server` and `agl-controller`, and writes their logs under `/tmp/`. The script also starts a new local Minikube single-node K8s cluster, and the agent runs in this cluster as Kubernetes Jobs.
|
||||
When `run_minikube.sh` exits, it automatically cleans up the server, controller, and Minikube it started.
|
||||
Minikube needs at least 64 GB of memory; otherwise, it may be killed due to insufficient memory.
|
||||
@@ -0,0 +1,63 @@
|
||||
# GSM8K
|
||||
|
||||
| GPU | Model | Controller Mode | Trainer Mode | Code |
|
||||
|---|---|---|---|---|
|
||||
| 1× A100 80GB | `Qwen/Qwen2.5-1.5B-Instruct` | Local | Sync only | [Source](https://github.com/microsoft/agent-lightning/tree/main/examples/gsm8k) |
|
||||
|
||||
GSM8K trains a grade-school math reasoning agent on the `openai/gsm8k` dataset with `verl` and Agent Lightning >=v1.0.
|
||||
This example runs in local mode and demonstrates support for two API styles:
|
||||
|
||||
1. **Chat Completions API:** the commonly used text-in/text-out API, where the agent sends structured chat messages and receives generated text.
|
||||
2. **Token-in/token-out Completions API:** the agent sends prompt token IDs directly and receives generated token IDs.
|
||||
|
||||
## Data Preparation
|
||||
|
||||
Download the dataset into `~/dataset/gsm8k`:
|
||||
|
||||
```bash
|
||||
hf download openai/gsm8k --repo-type dataset --local-dir ~/dataset/gsm8k
|
||||
```
|
||||
|
||||
The example reads these files by default:
|
||||
|
||||
- `~/dataset/gsm8k/main/train-00000-of-00001.parquet`
|
||||
- `~/dataset/gsm8k/main/test-00000-of-00001.parquet`
|
||||
|
||||
Training uses all samples from `main/train`. Validation uses 100 random samples from `main/test` with seed `42` by default.
|
||||
|
||||
## Training
|
||||
|
||||
Make sure you have activated the project environment and installed the example dependencies:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
uv pip install \
|
||||
datasets \
|
||||
openai \
|
||||
httpx
|
||||
```
|
||||
|
||||
Then start training:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
cd examples/gsm8k
|
||||
bash run_local.sh
|
||||
```
|
||||
|
||||
You can change the validation sample count or seed with:
|
||||
|
||||
```bash
|
||||
bash run_local.sh --val-size 100 --seed 42
|
||||
```
|
||||
|
||||
The local example uses `ChatAgent` with the standard Chat Completions API by default. To demonstrate the token-in/token-out Completions API, use `CompletionAgent` instead:
|
||||
|
||||
```bash
|
||||
bash run_local.sh --api completion
|
||||
```
|
||||
|
||||
In token-in/token-out mode, the agent tokenizes the prompt with the configured model tokenizer, sends prompt token IDs to the OpenAI-compatible Completions endpoint, receives response token IDs, and decodes them locally for answer evaluation.
|
||||
|
||||
`run_local.sh` starts `agl-server`, `agl-controller`, and Ray locally, and writes server/controller logs under `/tmp/`.
|
||||
When the script exits, it cleans up the local server, controller, and Ray process it started.
|
||||
@@ -0,0 +1,49 @@
|
||||
# ScienceWorld
|
||||
|
||||
| GPU | Model | Controller Mode | Trainer Mode | Code |
|
||||
|---|---|---|---|---|
|
||||
| 8× A100 40GB | `Qwen/Qwen2.5-7B-Instruct` | Local | Async only | [Source](https://github.com/microsoft/agent-lightning/tree/main/examples/science_world) |
|
||||
|
||||
ScienceWorld trains an agent with `verl` and Agent Lightning >=v1.0 to solve text-based science tasks from AllenAI's [ScienceWorld](https://github.com/allenai/ScienceWorld).
|
||||
|
||||
This example uses the local controller in asynchronous trainer mode. Each rollout runs as a local process that interacts with a ScienceWorld environment, calls the model through the AGL Gateway, and reports the final reward. It does not require K8s, Docker, or Minikube.
|
||||
|
||||
## Environment Preparation
|
||||
|
||||
Install Java and the example dependencies:
|
||||
|
||||
```bash
|
||||
sudo apt-get install -y default-jre
|
||||
uv pip install scienceworld openai
|
||||
```
|
||||
|
||||
ScienceWorld starts a JVM for each rollout, so Java 1.8 or later is required.
|
||||
|
||||
## Training
|
||||
|
||||
Start local training from the repository root:
|
||||
|
||||
```bash
|
||||
examples/science_world/run_local.sh
|
||||
```
|
||||
|
||||
`run_local.sh` starts `agl-server`, the local `agl-controller`, and the `verl` trainer. The controller launches each rollout as a local process, and the script cleans up the server, controller, and Ray processes when it exits.
|
||||
|
||||
The training dataset is generated automatically from ScienceWorld task names and variation indices. To train on selected tasks or change the number of variations per task:
|
||||
|
||||
```bash
|
||||
examples/science_world/run_local.sh \
|
||||
--task-names find-non-living-thing,find-living-thing \
|
||||
--variations-per-task 50
|
||||
```
|
||||
|
||||
Available runtime settings include:
|
||||
|
||||
| Setting | Default | Description |
|
||||
|---|---|---|
|
||||
| `--task-names` | `all` | Comma-separated task names, or all ScienceWorld tasks |
|
||||
| `--variations-per-task` | `50` | Maximum variations per task |
|
||||
| `--simplification` | `easy` | ScienceWorld simplification preset |
|
||||
| `SW_MAX_STEPS` | `30` | Maximum model turns per rollout |
|
||||
| `SW_ENV_STEP_LIMIT` | `100` | ScienceWorld environment step limit |
|
||||
| `AGL_MAX_TOKENS` | `256` | Maximum tokens per model completion |
|
||||
@@ -0,0 +1,89 @@
|
||||
# Search-R1
|
||||
|
||||
| GPU | Model | Controller Mode | Trainer Mode | Code |
|
||||
|---|---|---|---|---|
|
||||
| 8× A100 40GB | `meta-llama/Llama-3.2-3B-Instruct` | Local | Sync only | [Source](https://github.com/microsoft/agent-lightning/tree/main/examples/search_r1) |
|
||||
|
||||
Search-R1 trains a retrieval-augmented question-answering agent with `verl` and Agent Lightning >=v1.0. During each multi-turn rollout, the agent alternates between model responses and Wikipedia searches before producing a final answer.
|
||||
|
||||
This example is based on [*Search-R1: Training LLMs to Reason and Leverage Search Engines with Reinforcement Learning*](https://arxiv.org/abs/2503.09516) by Jin et al. (2025).
|
||||
|
||||
This example uses the local controller in synchronous trainer mode. Each rollout runs as a local process, calls the policy model through the AGL Gateway, and queries a separate FAISS retrieval service.
|
||||
|
||||
The example supports two API styles:
|
||||
|
||||
1. **Chat Completions API:** the standard text-in/text-out API used by default.
|
||||
2. **Token-in/token-out Completions API:** the agent sends prompt token IDs and receives generated token IDs while preserving the multi-turn token sequence.
|
||||
|
||||
## Data Preparation
|
||||
|
||||
Prepare the Wikipedia corpus, E5 FAISS index, training data, and retriever environment from the repository root:
|
||||
|
||||
```bash
|
||||
examples/search_r1/data_process.sh
|
||||
```
|
||||
|
||||
The script creates a Conda environment named `retriever` and prepares these files:
|
||||
|
||||
- `examples/search_r1/data/wiki-18.jsonl`
|
||||
- `examples/search_r1/data/e5_Flat.index`
|
||||
- `examples/search_r1/data/train.parquet`
|
||||
- `examples/search_r1/data/test.parquet`
|
||||
|
||||
Set `SEARCH_R1_DATA_DIR` before running the script to use a different data directory.
|
||||
|
||||
## Retrieval Service
|
||||
|
||||
Start the retrieval service in a separate terminal and keep it running during training:
|
||||
|
||||
```bash
|
||||
examples/search_r1/retrieval_launch.sh
|
||||
```
|
||||
|
||||
The service listens at `http://127.0.0.1:8000/retrieve` by default. Check that it is ready with:
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:8000/healthz
|
||||
```
|
||||
|
||||
Common retrieval settings include:
|
||||
|
||||
| Setting | Default | Description |
|
||||
|---|---|---|
|
||||
| `SEARCH_R1_DATA_DIR` | `examples/search_r1/data` | Corpus and FAISS index directory |
|
||||
| `SEARCH_R1_RETRIEVAL_PORT` | `8000` | Retrieval service port |
|
||||
| `SEARCH_R1_TOPK` | `3` | Documents returned for each search |
|
||||
| `SEARCH_R1_RETRIEVER_DEVICE` | `auto` | Retriever device: `auto`, `cuda`, `cuda:0`, or `cpu` |
|
||||
|
||||
## Training
|
||||
|
||||
With the retrieval service running, start local training from the repository root:
|
||||
|
||||
```bash
|
||||
examples/search_r1/run.sh
|
||||
```
|
||||
|
||||
`run.sh` starts `agl-server`, the local `agl-controller`, and the `verl` trainer. The script cleans up the server, controller, and Ray processes when it exits.
|
||||
|
||||
The default agent uses the Chat Completions API. To use the token-in/token-out Completions API instead:
|
||||
|
||||
```bash
|
||||
examples/search_r1/run.sh --api-type completion
|
||||
```
|
||||
|
||||
To use different dataset files:
|
||||
|
||||
```bash
|
||||
examples/search_r1/run.sh \
|
||||
--train-file /path/to/train.parquet \
|
||||
--val-file /path/to/test.parquet
|
||||
```
|
||||
|
||||
Agent runtime settings include:
|
||||
|
||||
| Setting | Default | Description |
|
||||
|---|---|---|
|
||||
| `SEARCH_R1_RETRIEVAL_URL` | `http://127.0.0.1:8000/retrieve` | Retrieval endpoint used by rollout agents |
|
||||
| `SEARCH_R1_MAX_TURNS` | `4` | Maximum model/search turns per rollout |
|
||||
| `SEARCH_R1_MAX_TOKENS` | `500` | Maximum generated tokens per model response |
|
||||
| `SEARCH_R1_TEMPERATURE` | `1.0` | Sampling temperature |
|
||||
@@ -0,0 +1,87 @@
|
||||
# LLM-in-Sandbox
|
||||
|
||||
| GPU | Model | Controller Mode | Trainer Mode | Code |
|
||||
|---|---|---|---|---|
|
||||
| 4× A100 80GB | `Qwen/Qwen3-4B-Instruct-2507` | K8s | Sync only | [Source](https://github.com/microsoft/agent-lightning/tree/main/examples/llm-in-sandbox) |
|
||||
|
||||
LLM-in-Sandbox trains a general instruction-following agent with `verl` and Agent Lightning >=v1.0. The agent can manage files, execute code, and use external resources inside an isolated container sandbox.
|
||||
|
||||
This example is based on [*Computer Environments Elicit General Agentic Intelligence in LLMs*](https://arxiv.org/abs/2601.16206) by Cheng et al. (2026).
|
||||
|
||||
This example uses the K8s controller in synchronous trainer mode. Each rollout runs as a Kubernetes Job, while model calls pass through the AGL Gateway to the `verl`-managed vLLM server. The agent dependencies remain isolated from the trainer environment.
|
||||
|
||||
## Environment Preparation
|
||||
|
||||
Use Python 3.12 and install the project environment before running the example. You also need:
|
||||
|
||||
- Docker
|
||||
- Minikube
|
||||
- `kubectl`
|
||||
- Image build support inside Minikube
|
||||
|
||||
The bundled Minikube setup is intended for testing only. For production deployments, replace it with a production-grade Kubernetes cluster.
|
||||
|
||||
## Data Preparation
|
||||
|
||||
The public training and validation data is hosted in the [`daixuancheng/llm-in-sandbox-rl`](https://huggingface.co/datasets/daixuancheng/llm-in-sandbox-rl) dataset on Hugging Face. The upstream [`llm-in-sandbox-rl`](https://github.com/llm-in-sandbox/llm-in-sandbox-rl) repository provides the conversion script used to generate the JSON files expected by this example.
|
||||
|
||||
From the repository root, clone the upstream repository and convert all dataset configurations:
|
||||
|
||||
```bash
|
||||
git clone --depth 1 https://github.com/llm-in-sandbox/llm-in-sandbox-rl.git /tmp/llm-in-sandbox-rl
|
||||
python /tmp/llm-in-sandbox-rl/examples/llm_in_sandbox/convert_llm_sandbox_dataset.py \
|
||||
--all \
|
||||
--output-dir examples/llm-in-sandbox/data
|
||||
```
|
||||
|
||||
The converter downloads the following Hugging Face configurations:
|
||||
|
||||
- Training: `instruct_pretrain` (`train` split, 3,600 samples)
|
||||
- Validation: `math_mini`, `biomed_mini`, and `long_context_mini` (`test` splits)
|
||||
|
||||
The default files used by this example are:
|
||||
|
||||
| Split | Path |
|
||||
|---|---|
|
||||
| Training | `examples/llm-in-sandbox/data/llm_sandbox_instruct_pretrain/train_verl.json` |
|
||||
| Validation | `examples/llm-in-sandbox/data/llm_sandbox_math_mini/test_verl.json` |
|
||||
| Validation | `examples/llm-in-sandbox/data/llm_sandbox_biomed_mini/test_verl.json` |
|
||||
| Validation | `examples/llm-in-sandbox/data/llm_sandbox_long_context_mini/test_verl.json` |
|
||||
|
||||
The command above creates these directories directly; no manual file move is needed. If you generate or download the files separately, place `train_verl.json` and `test_verl.json` in their corresponding directories, or pass those directories to the launcher.
|
||||
|
||||
For validation, select any one or more of `math_mini`, `biomed_mini`, and `long_context_mini`. Separate multiple directories with commas:
|
||||
|
||||
```bash
|
||||
examples/llm-in-sandbox/run.sh \
|
||||
--train-data-dir /path/to/train-data \
|
||||
--val-data-dir /path/to/math-data,/path/to/biomed-data,/path/to/long-context-data
|
||||
```
|
||||
|
||||
## Training
|
||||
|
||||
Start training from the repository root:
|
||||
|
||||
```bash
|
||||
examples/llm-in-sandbox/run.sh
|
||||
```
|
||||
|
||||
The launcher:
|
||||
|
||||
1. creates a local Minikube cluster;
|
||||
2. builds the `llm-in-sandbox-agent:dev` image;
|
||||
3. starts `agl-server` and the K8s `agl-controller`;
|
||||
4. starts the `verl` trainer;
|
||||
5. cleans up the server, controller, and Ray processes when it exits.
|
||||
|
||||
The controller creates one Kubernetes Job for each rollout. Inside the Job, the adapter runs the sandbox agent, routes model calls through the AGL Gateway, evaluates the final answer, and reports the reward.
|
||||
|
||||
Additional `verl` settings can be passed as dotlist overrides:
|
||||
|
||||
```bash
|
||||
examples/llm-in-sandbox/run.sh \
|
||||
trainer.total_epochs=2 \
|
||||
actor_rollout_ref.rollout.n=2
|
||||
```
|
||||
|
||||
Use `Ctrl+C` to stop training and clean up the processes started by the launcher.
|
||||
@@ -0,0 +1,118 @@
|
||||
# Coding Agent
|
||||
|
||||
| GPU | Model | Controller Mode | Trainer Mode | Code |
|
||||
|---|---|---|---|---|
|
||||
| 4× B200 | `Qwen/Qwen3.5-9B` | K8s | Sync and async | [Source](https://github.com/microsoft/agent-lightning/tree/main/examples/swe_smith) |
|
||||
|
||||
The Coding Agent example trains a software-engineering agent on SWE-smith tasks with `verl` and Agent Lightning >=v1.0. Each rollout runs as a Kubernetes Job inside a repository-specific image, edits an isolated checkout, executes tests, and reports the resulting reward to the AGL Gateway.
|
||||
|
||||
This example uses two machines:
|
||||
|
||||
- **Machine A — Kubernetes Controller machine:** connects to the Kubernetes cluster, prepares repository images in the node-accessible Docker runtime, and runs `agl-controller` to create rollout Jobs.
|
||||
- **Machine B — GPU training machine:** provides the GPUs and runs both `agl-server` (the AGL Gateway) and the `verl` trainer with its model backend.
|
||||
|
||||
Machine B's AGL Gateway address must be reachable from Machine A and from the rollout pods in the Kubernetes cluster.
|
||||
|
||||
## Environment Preparation
|
||||
|
||||
On **Machine A (Kubernetes Controller machine)**, activate the project environment and install the dependency used to prepare repository images:
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
uv pip install -r examples/swe_smith/requirements.txt
|
||||
```
|
||||
|
||||
Machine A also requires Docker, `kubectl`, and access to the Kubernetes cluster.
|
||||
|
||||
On **Machine B (GPU training machine)**, install the project and GPU training environment described in the project installation guide. The SWE-smith image-preparation requirements above are not needed on Machine B.
|
||||
|
||||
## Data Preparation
|
||||
|
||||
The provided splits are derived from the original SWE-smith dataset, which contains 59,136 executable software-engineering tasks from 128 Python repositories. We build the training data with the following filtering pipeline:
|
||||
|
||||
1. Remove tasks with an empty problem statement. The original release contains 18,033 such records.
|
||||
2. Remove tasks whose corresponding problem branch is missing from the provided repository image. This affects 1,265 records.
|
||||
3. Remove tasks requiring more than 200 tests, which avoids examples with prohibitively expensive test suites.
|
||||
4. Run Qwen3.5-9B four times on every remaining candidate as a difficulty probe.
|
||||
5. Remove tasks solved in all four probe rollouts because they provide little learning signal.
|
||||
6. Retain tasks with a mixture of successful and failed probe rollouts, yielding approximately 5,000 examples.
|
||||
7. Add a sample of 1,000 tasks that fail all four probes so the training set is not biased toward easier tasks.
|
||||
|
||||
The resulting data contains approximately 6,000 training examples and 400 validation examples. `train_dataset_mixed.jsonl` contains the mixed-difficulty training set, while `val_dataset_filtered.jsonl` contains the filtered validation set.
|
||||
|
||||
Download the pre-split dataset archive from [Google Drive](https://drive.google.com/file/d/1q19DP53l4rldvBR2dkUhbaPI_mHVBVL1/view?usp=drive_link) on **both machines**, then extract it into `examples/swe_smith/`:
|
||||
|
||||
- **Machine A** reads the datasets to determine which repository images must be prepared.
|
||||
- **Machine B** reads the datasets to construct the training and validation inputs.
|
||||
|
||||
The example reads these files by default:
|
||||
|
||||
- `examples/swe_smith/train_dataset_mixed.jsonl`
|
||||
- `examples/swe_smith/val_dataset_filtered.jsonl`
|
||||
|
||||
When using `run.sh`, custom paths can be selected with the `AGL_TRAIN_DATASET_PATH` and `AGL_VAL_DATASET_PATH` environment variables read by the launcher.
|
||||
|
||||
## Repository Image Preparation
|
||||
|
||||
On Machine A, prepare the repository images in the Docker daemon used by the Kubernetes nodes before starting the Controller:
|
||||
|
||||
```bash
|
||||
python examples/swe_smith/pull_images.py \
|
||||
--dataset examples/swe_smith/train_dataset_mixed.jsonl \
|
||||
--dataset examples/swe_smith/val_dataset_filtered.jsonl
|
||||
```
|
||||
|
||||
This command installs the OpenAI client into each required SWE-smith base image and creates the `:openai` tags expected by `job-template-openai.yaml`. Run it again if the datasets introduce new repository images.
|
||||
|
||||
## Training
|
||||
|
||||
The distributed launcher has three roles and must be started in this order:
|
||||
|
||||
```text
|
||||
server → controller → trainer
|
||||
```
|
||||
|
||||
On **Machine B (GPU training machine)**, start the Gateway:
|
||||
|
||||
```bash
|
||||
export AGL_SERVER_PUBLIC_HOST=<address-reachable-from-controller-and-pods>
|
||||
export AGL_KEY=<shared-secret>
|
||||
export AGL_MODEL_NAME=Qwen/Qwen3.5-9B
|
||||
examples/swe_smith/run.sh server
|
||||
```
|
||||
|
||||
On **Machine A (Kubernetes Controller machine)**, start the Controller:
|
||||
|
||||
```bash
|
||||
export AGL_SERVER_PUBLIC_HOST=<gateway-address>
|
||||
export AGL_KEY=<same-shared-secret>
|
||||
export AGL_NAMESPACE=agents
|
||||
examples/swe_smith/run.sh controller
|
||||
```
|
||||
|
||||
After the Gateway and Controller are ready, start the trainer on **Machine B (GPU training machine)**:
|
||||
|
||||
```bash
|
||||
export AGL_KEY=<same-shared-secret>
|
||||
export AGL_MODEL_NAME=Qwen/Qwen3.5-9B
|
||||
examples/swe_smith/run.sh trainer
|
||||
```
|
||||
|
||||
The launcher passes additional arguments to `train_smith_agent.py`, including `verl` dotlist overrides:
|
||||
|
||||
```bash
|
||||
examples/swe_smith/run.sh trainer \
|
||||
trainer.total_training_steps=100 \
|
||||
actor_rollout_ref.rollout.n=4
|
||||
```
|
||||
|
||||
## Preventing Reward Hacking
|
||||
|
||||
A coding agent may obtain the reference fix without solving the task, for example by inspecting Git history, downloading upstream source code with `curl` or `wget`, installing the original package with `pip`, or using Python networking libraries such as `urllib`.
|
||||
|
||||
The SWE agent limits these reward-hacking paths in two ways:
|
||||
|
||||
- **Repository isolation:** before the agent starts, the harness checks out the task branch and moves `.git` outside the visible testbed. Agent commands that invoke Git, access the hidden Git metadata, install packages, download files, or modify the test harness are blocked.
|
||||
- **Network isolation:** we strongly recommend adding a Kubernetes network policy that denies all outbound traffic from agent pods except connections to the AGL Gateway. Without this restriction, an agent may retrieve upstream source code or other external information and obtain reward without solving the task as intended.
|
||||
|
||||
The final reward is computed by running the task-specific `FAIL_TO_PASS` and `PASS_TO_PASS` tests inside the isolated repository environment. These controls are part of the training setup: weakening them can allow the agent to recover reference code and corrupt the reward signal.
|
||||
@@ -0,0 +1,44 @@
|
||||
# Agent Lightning Documentation
|
||||
|
||||
<p align="center">
|
||||
<img src="images/agl-v1.0.svg" alt="Agent Lightning v1.0" width="500">
|
||||
</p>
|
||||
|
||||
Welcome to the Agent Lightning v1.0 documentation. Start with the installation and quick-start guides, then use the configuration guides and examples below to build and train your own agents.
|
||||
|
||||
Agent Lightning v1.0 is a completely redesigned and reimplemented version with the following key features:
|
||||
|
||||
- 🪶 **~3,500 lines of core Python:** Simplicity is the first principle.
|
||||
- 🧩 **Training with real agent harnesses:** Agents interact with the model through the Agent Lightning v1.0 proxy with zero changes while keeping tools, context, control flow, and environments in the loop.
|
||||
- ☸️ **Native Kubernetes support:** Agents run directly as Kubernetes Jobs without relying on external sandbox services.
|
||||
- 💻 **A complete coding-agent training example:** The released pipeline covers data cleaning, reward-hacking prevention, and training scripts.
|
||||
|
||||
For the legacy Agent Lightning releases earlier than v1.0, see the [`v0.x` code branch](https://github.com/microsoft/agent-lightning/tree/v0.x) and the [v0.3.0 documentation](https://microsoft.github.io/agent-lightning/0.3.0/).
|
||||
|
||||
## Getting Started
|
||||
|
||||
| Guide | Description |
|
||||
|---|---|
|
||||
| [Installation](00-installation.md) | Set up the base environment and the tested `verl` GPU stack. |
|
||||
| [Quick Start](01-quick-start.md) | Run a local end-to-end rollout-driven training job. |
|
||||
| [Basics](05-basics.md) | Learn the core components, rollouts, events, and trajectories. |
|
||||
|
||||
## Configuration
|
||||
|
||||
| Guide | Description |
|
||||
|---|---|
|
||||
| [Trainer Configuration](20-trainer-configuration.md) | Configure `verl` integration, rollout collection, and trace aggregation. |
|
||||
| [API Gateway Configuration](25-api-gateway-configuration.md) | Configure the API Gateway and model proxy. |
|
||||
| [Controller Configuration](30-controller-configuration.md) | Configure local and Kubernetes rollout runners. |
|
||||
| [Asynchronous Training](35-asynchronous-training.md) | Configure collocated asynchronous collection and pause/drain behavior. |
|
||||
|
||||
## Examples
|
||||
|
||||
| Example | Description |
|
||||
|---|---|
|
||||
| [Calc-X](50-example-calc-x.md) | Train a math reasoning agent with AutoGen and MCP calculator tools. |
|
||||
| [GSM8K](55-example-gsm8k.md) | Train an agent on grade-school math reasoning tasks. |
|
||||
| [ScienceWorld](60-example-science-world.md) | Train an agent on interactive science tasks in a text environment. |
|
||||
| [Search-R1](65-example-search-r1.md) | Train a multi-turn retrieval and reasoning agent. |
|
||||
| [LLM-in-Sandbox](70-example-llm-in-sandbox.md) | Train a general agent with computer and code execution tools. |
|
||||
| [Coding Agent](75-example-coding-agent.md) | Train a coding agent using repository tests as feedback. |
|
||||
@@ -0,0 +1,4 @@
|
||||
<svg width="16" height="16" viewBox="0 0 16 16" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M8.06935 0.740967C8.46471 0.740967 8.78513 1.06143 8.78513 1.45675C8.78513 1.73028 8.6317 1.96783 8.40619 2.08833V2.6357C8.60265 2.66378 9.18471 2.67955 9.92197 3.01465C10.8483 3.4357 11.121 3.77803 11.4378 4.2357C11.6095 4.48376 11.7205 4.74914 11.7909 5.0357H11.9009C12.273 5.0357 12.5746 5.33732 12.5746 5.70938V6.38307C12.5746 6.75514 12.273 7.05675 11.9009 7.05675H11.8192C11.7591 7.40026 11.6762 7.69328 11.6062 7.85675C11.413 8.30755 11.129 8.48833 10.9746 8.53044C11.0869 8.57254 11.4395 8.67604 11.6483 8.82517C11.943 9.0357 12.2378 9.39174 12.2378 9.75149C12.2378 10.0462 12.1957 10.4673 11.9851 10.6778C11.7074 10.9556 11.2272 11.4778 10.9746 11.6883L6.34303 15.8568L7.35356 13.4989L10.3851 9.54096H8.02724L8.86934 6.59359L8.19658 7.34393L8.19567 7.35149L5.2483 10.762H7.6904L7.10093 12.7831L5.62724 11.8989L4.91146 11.4357C4.53251 11.1831 4.32198 11.0989 4.06935 10.6778C3.91617 10.4225 3.90093 10.0462 3.90093 9.75149C3.90093 9.39174 4.19567 9.0357 4.4904 8.82517C4.69915 8.67604 4.78514 8.61465 4.99567 8.53044C4.82724 8.44623 4.7257 8.30755 4.53251 7.85675C4.46245 7.69328 4.37956 7.40026 4.31951 7.05675H4.23777C3.8657 7.05675 3.56409 6.75514 3.56409 6.38307V5.70938C3.56409 5.33732 3.8657 5.0357 4.23777 5.0357H4.34788C4.41815 4.74914 4.5292 4.48376 4.70093 4.2357C5.01778 3.77803 5.2904 3.4357 6.21672 3.01465C6.95396 2.67955 7.53602 2.66378 7.73251 2.6357V2.08833C7.50704 1.96783 7.35356 1.73028 7.35356 1.45675C7.35356 1.06143 7.67403 0.740967 8.06935 0.740967ZM6.80619 5.0357C6.50389 5.0357 6.25882 5.28077 6.25882 5.58307C6.25882 5.88538 6.50389 6.13044 6.80619 6.13044C7.1085 6.13044 7.35356 5.88538 7.35356 5.58307C7.35356 5.28077 7.1085 5.0357 6.80619 5.0357ZM9.3325 5.0357C9.03018 5.0357 8.78513 5.28077 8.78513 5.58307C8.78513 5.88538 9.03018 6.13044 9.3325 6.13044C9.63481 6.13044 9.87987 5.88538 9.87987 5.58307C9.87987 5.28077 9.63481 5.0357 9.3325 5.0357Z" fill="white"/>
|
||||
<path d="M12.2279 9.63738C12.2342 9.67527 12.2378 9.71342 12.2378 9.75165C12.2378 10.0464 12.1957 10.4674 11.9851 10.678C11.7074 10.9558 11.2272 11.478 10.9746 11.6885L6.34305 15.8569L7.35357 13.499L7.9831 12.677C8.20912 12.6273 8.41543 12.5774 8.57462 12.5306C9.29041 12.3201 10.1325 11.6885 10.7641 11.099C11.2418 10.6531 11.9076 9.97136 12.2279 9.63738ZM9.62725 3.77271C10.3248 3.77271 10.8904 4.33825 10.8904 5.03586V6.80428C10.8904 7.50191 10.3248 8.06744 9.62725 8.06744H8.4483L8.86935 6.59376L8.19659 7.34409L8.19568 7.35165L7.57479 8.06744H6.59568C5.89805 8.06744 5.33252 7.50191 5.33252 6.80428V5.03586C5.33252 4.33825 5.89805 3.77271 6.59568 3.77271H9.62725ZM6.8062 5.03586C6.5039 5.03586 6.25884 5.28093 6.25884 5.58323C6.25884 5.88554 6.5039 6.1306 6.8062 6.1306C7.10851 6.1306 7.35357 5.88554 7.35357 5.58323C7.35357 5.28093 7.10851 5.03586 6.8062 5.03586ZM9.33251 5.03586C9.0302 5.03586 8.78514 5.28093 8.78514 5.58323C8.78514 5.88554 9.0302 6.1306 9.33251 6.1306C9.63483 6.1306 9.87988 5.88554 9.87988 5.58323C9.87988 5.28093 9.63483 5.03586 9.33251 5.03586Z" fill="white"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 3.0 KiB |
@@ -0,0 +1,4 @@
|
||||
<svg width="16" height="16" viewBox="0 0 16 16" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M8.06935 0.740967C8.46471 0.740967 8.78513 1.06143 8.78513 1.45675C8.78513 1.73028 8.6317 1.96783 8.40619 2.08833V2.6357C8.60265 2.66378 9.18471 2.67955 9.92197 3.01465C10.8483 3.4357 11.121 3.77803 11.4378 4.2357C11.6095 4.48376 11.7205 4.74914 11.7909 5.0357H11.9009C12.273 5.0357 12.5746 5.33732 12.5746 5.70938V6.38307C12.5746 6.75514 12.273 7.05675 11.9009 7.05675H11.8192C11.7591 7.40026 11.6762 7.69328 11.6062 7.85675C11.413 8.30755 11.129 8.48833 10.9746 8.53044C11.0869 8.57254 11.4395 8.67604 11.6483 8.82517C11.943 9.0357 12.2378 9.39174 12.2378 9.75149C12.2378 10.0462 12.1957 10.4673 11.9851 10.6778C11.7074 10.9556 11.2272 11.4778 10.9746 11.6883L6.34303 15.8568L7.35356 13.4989L10.3851 9.54096H8.02724L8.86934 6.59359L8.19658 7.34393L8.19567 7.35149L5.2483 10.762H7.6904L7.10093 12.7831L5.62724 11.8989L4.91146 11.4357C4.53251 11.1831 4.32198 11.0989 4.06935 10.6778C3.91617 10.4225 3.90093 10.0462 3.90093 9.75149C3.90093 9.39174 4.19567 9.0357 4.4904 8.82517C4.69915 8.67604 4.78514 8.61465 4.99567 8.53044C4.82724 8.44623 4.7257 8.30755 4.53251 7.85675C4.46245 7.69328 4.37956 7.40026 4.31951 7.05675H4.23777C3.8657 7.05675 3.56409 6.75514 3.56409 6.38307V5.70938C3.56409 5.33732 3.8657 5.0357 4.23777 5.0357H4.34788C4.41815 4.74914 4.5292 4.48376 4.70093 4.2357C5.01778 3.77803 5.2904 3.4357 6.21672 3.01465C6.95396 2.67955 7.53602 2.66378 7.73251 2.6357V2.08833C7.50704 1.96783 7.35356 1.73028 7.35356 1.45675C7.35356 1.06143 7.67403 0.740967 8.06935 0.740967ZM6.80619 5.0357C6.50389 5.0357 6.25882 5.28077 6.25882 5.58307C6.25882 5.88538 6.50389 6.13044 6.80619 6.13044C7.1085 6.13044 7.35356 5.88538 7.35356 5.58307C7.35356 5.28077 7.1085 5.0357 6.80619 5.0357ZM9.3325 5.0357C9.03018 5.0357 8.78513 5.28077 8.78513 5.58307C8.78513 5.88538 9.03018 6.13044 9.3325 6.13044C9.63481 6.13044 9.87987 5.88538 9.87987 5.58307C9.87987 5.28077 9.63481 5.0357 9.3325 5.0357Z" fill="#F69047"/>
|
||||
<path d="M12.2279 9.63738C12.2342 9.67527 12.2378 9.71342 12.2378 9.75165C12.2378 10.0464 12.1957 10.4674 11.9851 10.678C11.7074 10.9558 11.2272 11.478 10.9746 11.6885L6.34305 15.8569L7.35357 13.499L7.9831 12.677C8.20912 12.6273 8.41543 12.5774 8.57462 12.5306C9.29041 12.3201 10.1325 11.6885 10.7641 11.099C11.2418 10.6531 11.9076 9.97136 12.2279 9.63738ZM9.62725 3.77271C10.3248 3.77271 10.8904 4.33825 10.8904 5.03586V6.80428C10.8904 7.50191 10.3248 8.06744 9.62725 8.06744H8.4483L8.86935 6.59376L8.19659 7.34409L8.19568 7.35165L7.57479 8.06744H6.59568C5.89805 8.06744 5.33252 7.50191 5.33252 6.80428V5.03586C5.33252 4.33825 5.89805 3.77271 6.59568 3.77271H9.62725ZM6.8062 5.03586C6.5039 5.03586 6.25884 5.28093 6.25884 5.58323C6.25884 5.88554 6.5039 6.1306 6.8062 6.1306C7.10851 6.1306 7.35357 5.88554 7.35357 5.58323C7.35357 5.28093 7.10851 5.03586 6.8062 5.03586ZM9.33251 5.03586C9.0302 5.03586 8.78514 5.28093 8.78514 5.58323C8.78514 5.88554 9.0302 6.1306 9.33251 6.1306C9.63483 6.1306 9.87988 5.88554 9.87988 5.58323C9.87988 5.28093 9.63483 5.03586 9.33251 5.03586Z" fill="#C45259"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 3.0 KiB |
|
Before Width: | Height: | Size: 4.3 KiB |
|
Before Width: | Height: | Size: 4.5 KiB |
|
Before Width: | Height: | Size: 598 KiB |
|
Before Width: | Height: | Size: 166 KiB |
|
Before Width: | Height: | Size: 1.1 MiB |
|
Before Width: | Height: | Size: 310 KiB |
|
Before Width: | Height: | Size: 235 KiB |
@@ -1,31 +0,0 @@
|
||||
# Server-client Architecture
|
||||
|
||||
Article to be written.
|
||||
|
||||
```mermaid
|
||||
sequenceDiagram
|
||||
participant RL as RL Framework
|
||||
participant TS as Training Server
|
||||
participant AC as Agent Client
|
||||
participant AG as Agent
|
||||
|
||||
AC->>TS: Upload Dataset (1)
|
||||
RL->>TS: Start RL Server (2)
|
||||
TS->>RL: Latest Model (3)
|
||||
|
||||
loop for each batch of tasks
|
||||
loop for each task in the batch
|
||||
AC->>TS: Request Task (4)
|
||||
TS->>AC: Send Task & Model API (5)
|
||||
AC->>AG: Run Agent with Task & Model API (6)
|
||||
loop for each LLM call
|
||||
AC->>AG: Prompt (7)
|
||||
AG->>AC: Response (8)
|
||||
end
|
||||
AG->>AC: Rewarded Trace (9)
|
||||
AC->>TS: Send Rewarded Trace (10)
|
||||
end
|
||||
TS->>RL: Send Batch of Traces (11)
|
||||
RL->>TS: Return Updated Model (12)
|
||||
end
|
||||
```
|
||||
@@ -1,184 +0,0 @@
|
||||
# SQL Agent with Agent Lightning
|
||||
|
||||
> This tutorial is tested with `verl==0.5.0` and `vllm==0.10.0`.
|
||||
|
||||
This example demonstrates how to build and train a self-correcting SQL agent. It leverages [Agent Lightning]({{ config.repo_url }}) and the `verl` framework for Reinforcement Learning (RL) based training, and LangGraph to define the agent's complex, cyclical reasoning workflow. The goal is to fine-tune a Large Language Model (LLM) to accurately convert natural language questions into executable SQL queries.
|
||||
|
||||
## SQL Agent Implementation
|
||||
|
||||
The design of Agent-lightning **allows flexible integration with various agent frameworks**, including AutoGen, CrewAI, OpenAI Agent SDK, LangGraph, and more. It can also work without agent frameworks, allowing you to train an agent built from scratch with Python code. See [our example gallery]({{ config.repo_url }}/tree/{{ config.extra.source_commit }}/examples) for more details.
|
||||
|
||||
The core of the agent is a state machine built with LangGraph, which allows for a robust and transparent workflow. The agent's logic, as visualized below, starts by writing a query, executes it, and then enters a refinement loop where it checks and rewrites the query until it is deemed correct or a turn limit is reached.
|
||||
|
||||
```mermaid
|
||||
---
|
||||
config:
|
||||
flowchart:
|
||||
curve: linear
|
||||
---
|
||||
graph LR;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
write_query(write_query)
|
||||
execute_query(execute_query)
|
||||
check_query(check_query)
|
||||
rewrite_query(rewrite_query)
|
||||
__end__([<p>__end__</p>]):::last
|
||||
__start__ --> write_query;
|
||||
check_query -.-> __end__;
|
||||
check_query -.-> rewrite_query;
|
||||
execute_query --> check_query;
|
||||
rewrite_query --> execute_query;
|
||||
write_query --> execute_query;
|
||||
classDef default fill:#f2f2f2,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#cccccc
|
||||
```
|
||||
|
||||
This workflow is implemented in the `SQLAgent` class within `sql_agent.py`. It consists of the following key steps:
|
||||
|
||||
1. **write_query**: Given a user's question and database schema, the agent makes an initial attempt to write a SQL query.
|
||||
2. **execute_query**: The generated query is run against the target database.
|
||||
3. **check_query**: The agent analyzes the original query and its execution result (or error) to check for mistakes. It uses a specific prompt (`CHECK_QUERY_PROMPT`) to determine if the query is correct.
|
||||
4. **rewrite_query**: If the `check_query` step finds errors, the agent enters this step. It uses the feedback from the previous step to generate a corrected SQL query. The process then loops back to `check_query` for re-evaluation.
|
||||
5. **END**: The loop terminates when `check_query` confirms the query is correct or the maximum number of turns (`max_turns`) is exceeded. One turn corresponds to a complete cycle of `write_query` (if first round), `execute_query`, `check_query`, and potentially `rewrite_query`.
|
||||
|
||||
We aim to train **write_query** and **rewrite_query** step in the setup of this example. The **check_query** step is not trained but will share the same LLM weights as the other steps.
|
||||
|
||||
## Client-Server Training with Agent Lightning
|
||||
|
||||
The training process uses a distributed client-server architecture designed by Agent Lightning to efficiently fine-tune the underlying LLM. This separation allows for scalable data generation across multiple clients while centralizing the computationally intensive model training on a dedicated server with GPUs, and also provides opportunities for customizing algorithms and training strategies (like [prompt optimization]({{ config.repo_url }}/tree/{{ config.extra.source_commit }}/examples/apo)) with minimal code changes.
|
||||
|
||||
* **Training Server (`agentlightning.verl`)**: The server, launched with the first command below, manages the core training loop. It runs an RL algorithm (with `verl` of course) and hosts an OpenAI-compatible LLM endpoint (with `verl`'s async server). The server's sole purpose is to receive interaction data from clients and update the LLM's weights to improve its performance. [This link]({{ config.repo_url }}/tree/{{ config.extra.source_commit }}/agentlightning/verl) points to the implementation of the server, which is built upon `verl`.
|
||||
* **Agent Clients (`sql_agent.py`)**: The clients run the LangGraph agent logic described above. They connect to the server to fetch tasks (natural language questions) and use the server's **OpenAI-compatible endpoint** for all generation steps (`write_query`, `check_query`, `rewrite_query`). After completing a task, the client exports its interaction traces (traced by [AgentOps](https://www.agentops.ai/) and filtered by trace hierarchy), evaluates its correctness to calculate a reward, and sends the entire interaction history (the "trajectory") back to the server for training. To adapt any agent to an "agent client", you do not need to change the agent logic, but only need to invoke the client's `run` method with `agentlightning.trainer`.
|
||||
|
||||

|
||||
|
||||
## Running the Example
|
||||
|
||||
1. Prepare the dataset: download from [here](https://drive.google.com/file/d/1oi9J1jZP9TyM35L85CL3qeGWl2jqlnL6/view) and unzip it to the `data` folder. It's basically a [Spider V1](https://yale-lily.github.io/spider) dataset converted to Parquet format. The dataset contains about 8000 training samples and about 2000 test samples, from which we sampled 500 samples for evaluation.
|
||||
```bash
|
||||
pip install gdown
|
||||
gdown --fuzzy https://drive.google.com/file/d/1oi9J1jZP9TyM35L85CL3qeGWl2jqlnL6/view
|
||||
unzip -q spider-data.zip -d data
|
||||
rm spider-data.zip
|
||||
```
|
||||
|
||||
2. Install the required dependencies:
|
||||
```bash
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
3. Launch the training server:
|
||||
```bash
|
||||
python -m agentlightning.verl \
|
||||
agentlightning.port=9997 \
|
||||
algorithm.adv_estimator=grpo \
|
||||
data.train_files=data/train_spider.parquet \
|
||||
data.val_files=data/test_dev_500.parquet \
|
||||
actor_rollout_ref.rollout.tensor_model_parallel_size=1 \
|
||||
trainer.n_gpus_per_node=1 \
|
||||
data.train_batch_size=32 \
|
||||
actor_rollout_ref.rollout.n=4 \
|
||||
actor_rollout_ref.actor.ppo_mini_batch_size=32 \
|
||||
actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=4 \
|
||||
actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=4 \
|
||||
actor_rollout_ref.rollout.multi_turn.format=hermes \
|
||||
actor_rollout_ref.model.path=meta-llama/Llama-3.2-3B-Instruct \
|
||||
data.max_prompt_length=4096 \
|
||||
data.max_response_length=2048 \
|
||||
data.truncation='error' \
|
||||
trainer.val_before_train=True \
|
||||
actor_rollout_ref.actor.optim.lr=1e-6 \
|
||||
actor_rollout_ref.model.use_remove_padding=True \
|
||||
actor_rollout_ref.actor.use_kl_loss=False \
|
||||
actor_rollout_ref.actor.kl_loss_coef=0.000 \
|
||||
actor_rollout_ref.actor.entropy_coeff=0 \
|
||||
actor_rollout_ref.actor.clip_ratio_low=0.2 \
|
||||
actor_rollout_ref.actor.clip_ratio_high=0.3 \
|
||||
actor_rollout_ref.model.enable_gradient_checkpointing=True \
|
||||
actor_rollout_ref.actor.fsdp_config.param_offload=True \
|
||||
actor_rollout_ref.actor.fsdp_config.optimizer_offload=True \
|
||||
actor_rollout_ref.rollout.name=vllm \
|
||||
actor_rollout_ref.rollout.gpu_memory_utilization=0.8 \
|
||||
actor_rollout_ref.ref.log_prob_micro_batch_size_per_gpu=8 \
|
||||
actor_rollout_ref.ref.fsdp_config.param_offload=True \
|
||||
algorithm.use_kl_in_reward=False \
|
||||
trainer.critic_warmup=0 \
|
||||
trainer.logger=['console','wandb'] \
|
||||
trainer.project_name=AgentLightning \
|
||||
trainer.experiment_name=train_sql_agent \
|
||||
trainer.nnodes=1 \
|
||||
trainer.save_freq=256 \
|
||||
trainer.test_freq=32 \
|
||||
trainer.total_epochs=2
|
||||
```
|
||||
|
||||
4. Launch agent clients that connect with the server:
|
||||
```bash
|
||||
export VERL_API_BASE=http://localhost:9997/ # Same as the server port. This is used for receiving tasks and sending results.
|
||||
python sql_agent.py \
|
||||
--litsqlagent.trained-agents write \ # Will only train the write and rewrite agent.
|
||||
--trainer.n-workers 16 \
|
||||
--litsqlagent.val-temperature 0
|
||||
```
|
||||
|
||||
There is no hard requirement in the launching order of the server and clients. But remember to kill the long-running agent clients after the training is done.
|
||||
|
||||
## Debug the Agent without verl
|
||||
|
||||
You can run the agent client alone without the `verl` server. This is useful for debugging the agent logic and SQL execution.
|
||||
|
||||
1. Copy `.env.example` to `.env` and fill in your OpenAI API key. `VERL_API_BASE` does not really matter here because you are not connecting to the server end.
|
||||
|
||||
2. Run the agent client:
|
||||
```bash
|
||||
dotenv run python sql_agent.py \
|
||||
--litsqlagent.trained-agents write \ # Will only select the trajectories related to write and rewrite.
|
||||
--trainer.n-workers 1 \ # For debug, use single process.
|
||||
--trainer.dev true # Enable the dev debug mode.
|
||||
```
|
||||
|
||||
## Evaluation
|
||||
|
||||
The example is evaluated using Llama-3.2-Instruct models. The models are trained on the Spider dataset for 2 epochs, with evaluation performed on a randomly selected subset of 500 test samples to compute held-out accuracy. The default setup for running agent clients during evaluation is as follows:
|
||||
|
||||
```bash
|
||||
python sql_agent.py \
|
||||
--litsqlagent.trained-agents write \
|
||||
--trainer.n-workers 16 \
|
||||
--trainer.daemon true \
|
||||
--litsqlagent.val-temperature 0 \
|
||||
--litsqlagent.max-turns 3 \
|
||||
--litsqlagent.table-info-truncate 2048 \
|
||||
--litsqlagent.execution-truncate 2048
|
||||
```
|
||||
|
||||
The setup of training server is the same as the command above.
|
||||
|
||||
### W&B Report
|
||||
|
||||
[link](https://api.wandb.ai/links/ultmaster/4cid500g)
|
||||
|
||||
### Performance Metrics
|
||||
|
||||

|
||||
|
||||
| Model | Size | Context | Max Turns | Agents | Acc (Initial) | Acc (Final) | Transitions | Prompt Length | Response Length |
|
||||
|---------------|--------|-----------|-------------|-------------------------------|-----------------|---------------|---------------|-----------------|-------------------|
|
||||
| Llama3.2 | 1B | 2048 | 3 | write|rewrite | 21 | 49.6 | 2.87 → 3.08 | 821.2 | 319.2 → 249.4 |
|
||||
| Llama3.2 | 3B | 2048 | 3 | write|rewrite | 51.8 | 66.4 | 2.20 → 2.72 | 865.6 | 116.2 → 314.3 |
|
||||
|
||||
**Notes:**
|
||||
|
||||
1. **Context Length**: Controlled via `--litsqlagent.table-info-truncate <context-length>` and `--litsqlagent.execution-truncate <context-length>`
|
||||
2. **Max Turns**: Set using `--litsqlagent.max-turns <max-turns>`
|
||||
3. **Agents**: Specified with `--litsqlagent.agents <regex>` (defaults to `write`, which matches both write and rewrite agents)
|
||||
4. **Transitions**: Represents the number of prompt-response pairs traced (collected) during each rollout. Note that this differs from the turn count in the SQL agent workflow, where one turn may encompass 2-3 transitions in the check-rewrite cycle. The number of transitions is also related to which *agents* get involved in the training.
|
||||
5. **Prompt/Response Length**: Average token count per **traced** prompt/transition response.
|
||||
|
||||
### Efficiency Metrics
|
||||
|
||||
| Model | Size | Context | Max Turns | Agents | # GPUs | # Steps | Time (h) | Time/Step (s) | Rollout Time (%) | Update Actor Time (%) |
|
||||
|---------------|--------|-----------|-------------|-------------------------------|----------|-----------|------------|-----------------|--------------------|-------------------------|
|
||||
| Llama3.2 | 1B | 2048 | 3 | write|rewrite | 1 | 436 | 13.06 | 98.9 | 66.7 | 25.2 |
|
||||
| Llama3.2 | 3B | 2048 | 3 | write|rewrite | 2 | 436 | 10.3 | 181.3 | 63.9 | 27.9 |
|
||||
|
After Width: | Height: | Size: 154 KiB |
|
After Width: | Height: | Size: 11 KiB |
|
After Width: | Height: | Size: 302 KiB |
|
After Width: | Height: | Size: 224 KiB |
|
After Width: | Height: | Size: 242 KiB |