diff --git a/.gitattributes b/.gitattributes
index ecb64946bec2fb4fe48f14c41fb11a031ee135f7..038e69d298699083ddf82f7b4cd48e957493d3e8 100644
--- a/.gitattributes
+++ b/.gitattributes
@@ -51,3 +51,95 @@ venv/lib/python3.11/site-packages/PIL/_imagingft.cpython-311-x86_64-linux-gnu.so
venv/lib/python3.11/site-packages/PIL/_imagingmath.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
venv/lib/python3.11/site-packages/PIL/_webp.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
venv/lib/python3.11/site-packages/_brotli.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/ada92cb5d92a588d1b93__mypyc.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/_core.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/attachments/stream.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/codeccontext.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/fifo.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/format.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/frame.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/layout.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/plane.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/resampler.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/audio/stream.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/bitstream.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/buffer.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/bytesource.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/codec/codec.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/codec/context.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/container/core.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/container/input.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/container/output.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/container/pyio.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/container/streams.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/data/stream.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/descriptor.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/dictionary.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/error.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/filter/context.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/filter/filter.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/filter/graph.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/filter/link.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/filter/loudnorm.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/filter/pad.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/format.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/frame.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/logging.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/opaque.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/option.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/packet.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/plane.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/sidedata/motionvectors.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/sidedata/sidedata.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/stream.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/subtitles/codeccontext.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/subtitles/stream.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/subtitles/subtitle.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/utils.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/video/codeccontext.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/video/format.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/video/frame.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/video/plane.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/video/reformatter.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av/video/stream.cpython-311-x86_64-linux-gnu.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libaom-e9efed4a.so.3.2.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libavcodec-fb6c662d.so.61.19.100 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libavfilter-3235a7c8.so.10.4.100 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libavformat-071c54bd.so.61.7.100 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libavutil-2749f1ba.so.59.39.100 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libdav1d-1b53ef2f.so.7.0.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libgmp-a4b719d5.so.10.5.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libgnutls-b9b94016.so.30.36.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libhogweed-9544c08c.so.6.8 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/liblzma-af70179d.so.5.4.4 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libmp3lame-3ecc6556.so.0.0.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libnettle-ecd2e589.so.8.8 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libopencore-amrnb-393dbae2.so.0.0.3 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libopus-59fcdf85.so.0.9.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libspeex-2370356a.so.1.5.2 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libswresample-da7d062e.so.5.3.100 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libswscale-9212cf18.so.8.3.100 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libtwolame-72d74ef7.so.0.0.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libunistring-214e3d6e.so.5.1.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libvorbis-f4a9a6fd.so.0.4.9 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libvorbisenc-0d9d5bdf.so.2.0.12 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libvpx-832f6f52.so.9.0.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libwebp-e16038c7.so.7.1.9 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libx264-2a4c6f6d.so.164 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libx265-d8690e8d.so.199 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libxcb-65da195c.so.1.1.0 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/av.libs/libxml2-cb941fce.so.2.9.13 filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda118.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda121.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda124.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda126.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda128.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda130.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda132.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm64.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm70.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm71.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm714.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm72.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2025.so filter=lfs diff=lfs merge=lfs -text
+venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2026.so filter=lfs diff=lfs merge=lfs -text
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/INSTALLER b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/INSTALLER
new file mode 100644
index 0000000000000000000000000000000000000000..5c69047b2eb8235994febeeae1da4a82365a240a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/INSTALLER
@@ -0,0 +1 @@
+uv
\ No newline at end of file
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/LICENSE b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/LICENSE
new file mode 100644
index 0000000000000000000000000000000000000000..261eeb9e9f8b2b4b0d119366dda99c6fd7d35c64
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/LICENSE
@@ -0,0 +1,201 @@
+ Apache License
+ Version 2.0, January 2004
+ http://www.apache.org/licenses/
+
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
+
+ 1. Definitions.
+
+ "License" shall mean the terms and conditions for use, reproduction,
+ and distribution as defined by Sections 1 through 9 of this document.
+
+ "Licensor" shall mean the copyright owner or entity authorized by
+ the copyright owner that is granting the License.
+
+ "Legal Entity" shall mean the union of the acting entity and all
+ other entities that control, are controlled by, or are under common
+ control with that entity. For the purposes of this definition,
+ "control" means (i) the power, direct or indirect, to cause the
+ direction or management of such entity, whether by contract or
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
+ outstanding shares, or (iii) beneficial ownership of such entity.
+
+ "You" (or "Your") shall mean an individual or Legal Entity
+ exercising permissions granted by this License.
+
+ "Source" form shall mean the preferred form for making modifications,
+ including but not limited to software source code, documentation
+ source, and configuration files.
+
+ "Object" form shall mean any form resulting from mechanical
+ transformation or translation of a Source form, including but
+ not limited to compiled object code, generated documentation,
+ and conversions to other media types.
+
+ "Work" shall mean the work of authorship, whether in Source or
+ Object form, made available under the License, as indicated by a
+ copyright notice that is included in or attached to the work
+ (an example is provided in the Appendix below).
+
+ "Derivative Works" shall mean any work, whether in Source or Object
+ form, that is based on (or derived from) the Work and for which the
+ editorial revisions, annotations, elaborations, or other modifications
+ represent, as a whole, an original work of authorship. For the purposes
+ of this License, Derivative Works shall not include works that remain
+ separable from, or merely link (or bind by name) to the interfaces of,
+ the Work and Derivative Works thereof.
+
+ "Contribution" shall mean any work of authorship, including
+ the original version of the Work and any modifications or additions
+ to that Work or Derivative Works thereof, that is intentionally
+ submitted to Licensor for inclusion in the Work by the copyright owner
+ or by an individual or Legal Entity authorized to submit on behalf of
+ the copyright owner. For the purposes of this definition, "submitted"
+ means any form of electronic, verbal, or written communication sent
+ to the Licensor or its representatives, including but not limited to
+ communication on electronic mailing lists, source code control systems,
+ and issue tracking systems that are managed by, or on behalf of, the
+ Licensor for the purpose of discussing and improving the Work, but
+ excluding communication that is conspicuously marked or otherwise
+ designated in writing by the copyright owner as "Not a Contribution."
+
+ "Contributor" shall mean Licensor and any individual or Legal Entity
+ on behalf of whom a Contribution has been received by Licensor and
+ subsequently incorporated within the Work.
+
+ 2. Grant of Copyright License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ copyright license to reproduce, prepare Derivative Works of,
+ publicly display, publicly perform, sublicense, and distribute the
+ Work and such Derivative Works in Source or Object form.
+
+ 3. Grant of Patent License. Subject to the terms and conditions of
+ this License, each Contributor hereby grants to You a perpetual,
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
+ (except as stated in this section) patent license to make, have made,
+ use, offer to sell, sell, import, and otherwise transfer the Work,
+ where such license applies only to those patent claims licensable
+ by such Contributor that are necessarily infringed by their
+ Contribution(s) alone or by combination of their Contribution(s)
+ with the Work to which such Contribution(s) was submitted. If You
+ institute patent litigation against any entity (including a
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
+ or a Contribution incorporated within the Work constitutes direct
+ or contributory patent infringement, then any patent licenses
+ granted to You under this License for that Work shall terminate
+ as of the date such litigation is filed.
+
+ 4. Redistribution. You may reproduce and distribute copies of the
+ Work or Derivative Works thereof in any medium, with or without
+ modifications, and in Source or Object form, provided that You
+ meet the following conditions:
+
+ (a) You must give any other recipients of the Work or
+ Derivative Works a copy of this License; and
+
+ (b) You must cause any modified files to carry prominent notices
+ stating that You changed the files; and
+
+ (c) You must retain, in the Source form of any Derivative Works
+ that You distribute, all copyright, patent, trademark, and
+ attribution notices from the Source form of the Work,
+ excluding those notices that do not pertain to any part of
+ the Derivative Works; and
+
+ (d) If the Work includes a "NOTICE" text file as part of its
+ distribution, then any Derivative Works that You distribute must
+ include a readable copy of the attribution notices contained
+ within such NOTICE file, excluding those notices that do not
+ pertain to any part of the Derivative Works, in at least one
+ of the following places: within a NOTICE text file distributed
+ as part of the Derivative Works; within the Source form or
+ documentation, if provided along with the Derivative Works; or,
+ within a display generated by the Derivative Works, if and
+ wherever such third-party notices normally appear. The contents
+ of the NOTICE file are for informational purposes only and
+ do not modify the License. You may add Your own attribution
+ notices within Derivative Works that You distribute, alongside
+ or as an addendum to the NOTICE text from the Work, provided
+ that such additional attribution notices cannot be construed
+ as modifying the License.
+
+ You may add Your own copyright statement to Your modifications and
+ may provide additional or different license terms and conditions
+ for use, reproduction, or distribution of Your modifications, or
+ for any such Derivative Works as a whole, provided Your use,
+ reproduction, and distribution of the Work otherwise complies with
+ the conditions stated in this License.
+
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
+ any Contribution intentionally submitted for inclusion in the Work
+ by You to the Licensor shall be under the terms and conditions of
+ this License, without any additional terms or conditions.
+ Notwithstanding the above, nothing herein shall supersede or modify
+ the terms of any separate license agreement you may have executed
+ with Licensor regarding such Contributions.
+
+ 6. Trademarks. This License does not grant permission to use the trade
+ names, trademarks, service marks, or product names of the Licensor,
+ except as required for reasonable and customary use in describing the
+ origin of the Work and reproducing the content of the NOTICE file.
+
+ 7. Disclaimer of Warranty. Unless required by applicable law or
+ agreed to in writing, Licensor provides the Work (and each
+ Contributor provides its Contributions) on an "AS IS" BASIS,
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
+ implied, including, without limitation, any warranties or conditions
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
+ PARTICULAR PURPOSE. You are solely responsible for determining the
+ appropriateness of using or redistributing the Work and assume any
+ risks associated with Your exercise of permissions under this License.
+
+ 8. Limitation of Liability. In no event and under no legal theory,
+ whether in tort (including negligence), contract, or otherwise,
+ unless required by applicable law (such as deliberate and grossly
+ negligent acts) or agreed to in writing, shall any Contributor be
+ liable to You for damages, including any direct, indirect, special,
+ incidental, or consequential damages of any character arising as a
+ result of this License or out of the use or inability to use the
+ Work (including but not limited to damages for loss of goodwill,
+ work stoppage, computer failure or malfunction, or any and all
+ other commercial damages or losses), even if such Contributor
+ has been advised of the possibility of such damages.
+
+ 9. Accepting Warranty or Additional Liability. While redistributing
+ the Work or Derivative Works thereof, You may choose to offer,
+ and charge a fee for, acceptance of support, warranty, indemnity,
+ or other liability obligations and/or rights consistent with this
+ License. However, in accepting such obligations, You may act only
+ on Your own behalf and on Your sole responsibility, not on behalf
+ of any other Contributor, and only if You agree to indemnify,
+ defend, and hold each Contributor harmless for any liability
+ incurred by, or claims asserted against, such Contributor by reason
+ of your accepting any such warranty or additional liability.
+
+ END OF TERMS AND CONDITIONS
+
+ APPENDIX: How to apply the Apache License to your work.
+
+ To apply the Apache License to your work, attach the following
+ boilerplate notice, with the fields enclosed by brackets "[]"
+ replaced with your own identifying information. (Don't include
+ the brackets!) The text should be enclosed in the appropriate
+ comment syntax for the file format. We also recommend that a
+ file or class name and description of purpose be included on the
+ same "printed page" as the copyright notice for easier
+ identification within third-party archives.
+
+ Copyright [yyyy] [name of copyright owner]
+
+ 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.
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/METADATA b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/METADATA
new file mode 100644
index 0000000000000000000000000000000000000000..5b0600e591a917ce30def944158f7796762fb62b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/METADATA
@@ -0,0 +1,380 @@
+Metadata-Version: 2.1
+Name: accelerate
+Version: 1.6.0
+Summary: Accelerate
+Home-page: https://github.com/huggingface/accelerate
+Author: The HuggingFace team
+Author-email: zach.mueller@huggingface.co
+License: Apache
+Keywords: deep learning
+Classifier: Development Status :: 5 - Production/Stable
+Classifier: Intended Audience :: Developers
+Classifier: Intended Audience :: Education
+Classifier: Intended Audience :: Science/Research
+Classifier: License :: OSI Approved :: Apache Software License
+Classifier: Operating System :: OS Independent
+Classifier: Programming Language :: Python :: 3
+Classifier: Programming Language :: Python :: 3.9
+Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
+Requires-Python: >=3.9.0
+Description-Content-Type: text/markdown
+License-File: LICENSE
+Requires-Dist: numpy<3.0.0,>=1.17
+Requires-Dist: packaging>=20.0
+Requires-Dist: psutil
+Requires-Dist: pyyaml
+Requires-Dist: torch>=2.0.0
+Requires-Dist: huggingface-hub>=0.21.0
+Requires-Dist: safetensors>=0.4.3
+Provides-Extra: deepspeed
+Requires-Dist: deepspeed; extra == "deepspeed"
+Provides-Extra: dev
+Requires-Dist: black~=23.1; extra == "dev"
+Requires-Dist: hf-doc-builder>=0.3.0; extra == "dev"
+Requires-Dist: ruff~=0.11.2; extra == "dev"
+Requires-Dist: pytest<=8.0.0,>=7.2.0; extra == "dev"
+Requires-Dist: pytest-xdist; extra == "dev"
+Requires-Dist: pytest-subtests; extra == "dev"
+Requires-Dist: parameterized; extra == "dev"
+Requires-Dist: pytest-order; extra == "dev"
+Requires-Dist: datasets; extra == "dev"
+Requires-Dist: diffusers; extra == "dev"
+Requires-Dist: evaluate; extra == "dev"
+Requires-Dist: torchdata>=0.8.0; extra == "dev"
+Requires-Dist: torchpippy>=0.2.0; extra == "dev"
+Requires-Dist: transformers; extra == "dev"
+Requires-Dist: scipy; extra == "dev"
+Requires-Dist: scikit-learn; extra == "dev"
+Requires-Dist: tqdm; extra == "dev"
+Requires-Dist: bitsandbytes; extra == "dev"
+Requires-Dist: timm; extra == "dev"
+Requires-Dist: rich; extra == "dev"
+Provides-Extra: docs
+Provides-Extra: quality
+Requires-Dist: black~=23.1; extra == "quality"
+Requires-Dist: hf-doc-builder>=0.3.0; extra == "quality"
+Requires-Dist: ruff~=0.11.2; extra == "quality"
+Provides-Extra: rich
+Requires-Dist: rich; extra == "rich"
+Provides-Extra: sagemaker
+Requires-Dist: sagemaker; extra == "sagemaker"
+Provides-Extra: test_dev
+Requires-Dist: datasets; extra == "test-dev"
+Requires-Dist: diffusers; extra == "test-dev"
+Requires-Dist: evaluate; extra == "test-dev"
+Requires-Dist: torchdata>=0.8.0; extra == "test-dev"
+Requires-Dist: torchpippy>=0.2.0; extra == "test-dev"
+Requires-Dist: transformers; extra == "test-dev"
+Requires-Dist: scipy; extra == "test-dev"
+Requires-Dist: scikit-learn; extra == "test-dev"
+Requires-Dist: tqdm; extra == "test-dev"
+Requires-Dist: bitsandbytes; extra == "test-dev"
+Requires-Dist: timm; extra == "test-dev"
+Provides-Extra: test_prod
+Requires-Dist: pytest<=8.0.0,>=7.2.0; extra == "test-prod"
+Requires-Dist: pytest-xdist; extra == "test-prod"
+Requires-Dist: pytest-subtests; extra == "test-prod"
+Requires-Dist: parameterized; extra == "test-prod"
+Requires-Dist: pytest-order; extra == "test-prod"
+Provides-Extra: test_trackers
+Requires-Dist: wandb; extra == "test-trackers"
+Requires-Dist: comet-ml; extra == "test-trackers"
+Requires-Dist: tensorboard; extra == "test-trackers"
+Requires-Dist: dvclive; extra == "test-trackers"
+Requires-Dist: mlflow; extra == "test-trackers"
+Requires-Dist: matplotlib; extra == "test-trackers"
+Provides-Extra: testing
+Requires-Dist: pytest<=8.0.0,>=7.2.0; extra == "testing"
+Requires-Dist: pytest-xdist; extra == "testing"
+Requires-Dist: pytest-subtests; extra == "testing"
+Requires-Dist: parameterized; extra == "testing"
+Requires-Dist: pytest-order; extra == "testing"
+Requires-Dist: datasets; extra == "testing"
+Requires-Dist: diffusers; extra == "testing"
+Requires-Dist: evaluate; extra == "testing"
+Requires-Dist: torchdata>=0.8.0; extra == "testing"
+Requires-Dist: torchpippy>=0.2.0; extra == "testing"
+Requires-Dist: transformers; extra == "testing"
+Requires-Dist: scipy; extra == "testing"
+Requires-Dist: scikit-learn; extra == "testing"
+Requires-Dist: tqdm; extra == "testing"
+Requires-Dist: bitsandbytes; extra == "testing"
+Requires-Dist: timm; extra == "testing"
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
Run your *raw* PyTorch training script on any kind of device
+
+
+
+
+
+
+## Easy to integrate
+
+🤗 Accelerate was created for PyTorch users who like to write the training loop of PyTorch models but are reluctant to write and maintain the boilerplate code needed to use multi-GPUs/TPU/fp16.
+
+🤗 Accelerate abstracts exactly and only the boilerplate code related to multi-GPUs/TPU/fp16 and leaves the rest of your code unchanged.
+
+Here is an example:
+
+```diff
+ import torch
+ import torch.nn.functional as F
+ from datasets import load_dataset
++ from accelerate import Accelerator
+
++ accelerator = Accelerator()
+- device = 'cpu'
++ device = accelerator.device
+
+ model = torch.nn.Transformer().to(device)
+ optimizer = torch.optim.Adam(model.parameters())
+
+ dataset = load_dataset('my_dataset')
+ data = torch.utils.data.DataLoader(dataset, shuffle=True)
+
++ model, optimizer, data = accelerator.prepare(model, optimizer, data)
+
+ model.train()
+ for epoch in range(10):
+ for source, targets in data:
+ source = source.to(device)
+ targets = targets.to(device)
+
+ optimizer.zero_grad()
+
+ output = model(source)
+ loss = F.cross_entropy(output, targets)
+
+- loss.backward()
++ accelerator.backward(loss)
+
+ optimizer.step()
+```
+
+As you can see in this example, by adding 5-lines to any standard PyTorch training script you can now run on any kind of single or distributed node setting (single CPU, single GPU, multi-GPUs and TPUs) as well as with or without mixed precision (fp8, fp16, bf16).
+
+In particular, the same code can then be run without modification on your local machine for debugging or your training environment.
+
+🤗 Accelerate even handles the device placement for you (which requires a few more changes to your code, but is safer in general), so you can even simplify your training loop further:
+
+```diff
+ import torch
+ import torch.nn.functional as F
+ from datasets import load_dataset
++ from accelerate import Accelerator
+
+- device = 'cpu'
++ accelerator = Accelerator()
+
+- model = torch.nn.Transformer().to(device)
++ model = torch.nn.Transformer()
+ optimizer = torch.optim.Adam(model.parameters())
+
+ dataset = load_dataset('my_dataset')
+ data = torch.utils.data.DataLoader(dataset, shuffle=True)
+
++ model, optimizer, data = accelerator.prepare(model, optimizer, data)
+
+ model.train()
+ for epoch in range(10):
+ for source, targets in data:
+- source = source.to(device)
+- targets = targets.to(device)
+
+ optimizer.zero_grad()
+
+ output = model(source)
+ loss = F.cross_entropy(output, targets)
+
+- loss.backward()
++ accelerator.backward(loss)
+
+ optimizer.step()
+```
+
+Want to learn more? Check out the [documentation](https://huggingface.co/docs/accelerate) or have a look at our [examples](https://github.com/huggingface/accelerate/tree/main/examples).
+
+## Launching script
+
+🤗 Accelerate also provides an optional CLI tool that allows you to quickly configure and test your training environment before launching the scripts. No need to remember how to use `torch.distributed.run` or to write a specific launcher for TPU training!
+On your machine(s) just run:
+
+```bash
+accelerate config
+```
+
+and answer the questions asked. This will generate a config file that will be used automatically to properly set the default options when doing
+
+```bash
+accelerate launch my_script.py --args_to_my_script
+```
+
+For instance, here is how you would run the GLUE example on the MRPC task (from the root of the repo):
+
+```bash
+accelerate launch examples/nlp_example.py
+```
+
+This CLI tool is **optional**, and you can still use `python my_script.py` or `python -m torchrun my_script.py` at your convenience.
+
+You can also directly pass in the arguments you would to `torchrun` as arguments to `accelerate launch` if you wish to not run` accelerate config`.
+
+For example, here is how to launch on two GPUs:
+
+```bash
+accelerate launch --multi_gpu --num_processes 2 examples/nlp_example.py
+```
+
+To learn more, check the CLI documentation available [here](https://huggingface.co/docs/accelerate/package_reference/cli).
+
+Or view the configuration zoo [here](https://github.com/huggingface/accelerate/blob/main/examples/config_yaml_templates/)
+
+## Launching multi-CPU run using MPI
+
+🤗 Here is another way to launch multi-CPU run using MPI. You can learn how to install Open MPI on [this page](https://www.open-mpi.org/faq/?category=building#easy-build). You can use Intel MPI or MVAPICH as well.
+Once you have MPI setup on your cluster, just run:
+```bash
+accelerate config
+```
+Answer the questions that are asked, selecting to run using multi-CPU, and answer "yes" when asked if you want accelerate to launch mpirun.
+Then, use `accelerate launch` with your script like:
+```bash
+accelerate launch examples/nlp_example.py
+```
+Alternatively, you can use mpirun directly, without using the CLI like:
+```bash
+mpirun -np 2 python examples/nlp_example.py
+```
+
+## Launching training using DeepSpeed
+
+🤗 Accelerate supports training on single/multiple GPUs using DeepSpeed. To use it, you don't need to change anything in your training code; you can set everything using just `accelerate config`. However, if you desire to tweak your DeepSpeed related args from your Python script, we provide you the `DeepSpeedPlugin`.
+
+```python
+from accelerate import Accelerator, DeepSpeedPlugin
+
+# deepspeed needs to know your gradient accumulation steps beforehand, so don't forget to pass it
+# Remember you still need to do gradient accumulation by yourself, just like you would have done without deepspeed
+deepspeed_plugin = DeepSpeedPlugin(zero_stage=2, gradient_accumulation_steps=2)
+accelerator = Accelerator(mixed_precision='fp16', deepspeed_plugin=deepspeed_plugin)
+
+# How to save your 🤗 Transformer?
+accelerator.wait_for_everyone()
+unwrapped_model = accelerator.unwrap_model(model)
+unwrapped_model.save_pretrained(save_dir, save_function=accelerator.save, state_dict=accelerator.get_state_dict(model))
+```
+
+Note: DeepSpeed support is experimental for now. In case you get into some problem, please open an issue.
+
+## Launching your training from a notebook
+
+🤗 Accelerate also provides a `notebook_launcher` function you can use in a notebook to launch a distributed training. This is especially useful for Colab or Kaggle notebooks with a TPU backend. Just define your training loop in a `training_function` then in your last cell, add:
+
+```python
+from accelerate import notebook_launcher
+
+notebook_launcher(training_function)
+```
+
+An example can be found in [this notebook](https://github.com/huggingface/notebooks/blob/main/examples/accelerate_examples/simple_nlp_example.ipynb). [](https://colab.research.google.com/github/huggingface/notebooks/blob/main/examples/accelerate_examples/simple_nlp_example.ipynb)
+
+## Why should I use 🤗 Accelerate?
+
+You should use 🤗 Accelerate when you want to easily run your training scripts in a distributed environment without having to renounce full control over your training loop. This is not a high-level framework above PyTorch, just a thin wrapper so you don't have to learn a new library. In fact, the whole API of 🤗 Accelerate is in one class, the `Accelerator` object.
+
+## Why shouldn't I use 🤗 Accelerate?
+
+You shouldn't use 🤗 Accelerate if you don't want to write a training loop yourself. There are plenty of high-level libraries above PyTorch that will offer you that, 🤗 Accelerate is not one of them.
+
+## Frameworks using 🤗 Accelerate
+
+If you like the simplicity of 🤗 Accelerate but would prefer a higher-level abstraction around its capabilities, some frameworks and libraries that are built on top of 🤗 Accelerate are listed below:
+
+* [Amphion](https://github.com/open-mmlab/Amphion) is a toolkit for Audio, Music, and Speech Generation. Its purpose is to support reproducible research and help junior researchers and engineers get started in the field of audio, music, and speech generation research and development.
+* [Animus](https://github.com/Scitator/animus) is a minimalistic framework to run machine learning experiments. Animus highlights common "breakpoints" in ML experiments and provides a unified interface for them within [IExperiment](https://github.com/Scitator/animus/blob/main/animus/core.py#L76).
+* [Catalyst](https://github.com/catalyst-team/catalyst#getting-started) is a PyTorch framework for Deep Learning Research and Development. It focuses on reproducibility, rapid experimentation, and codebase reuse so you can create something new rather than write yet another train loop. Catalyst provides a [Runner](https://catalyst-team.github.io/catalyst/api/core.html#runner) to connect all parts of the experiment: hardware backend, data transformations, model training, and inference logic.
+* [fastai](https://github.com/fastai/fastai#installing) is a PyTorch framework for Deep Learning that simplifies training fast and accurate neural nets using modern best practices. fastai provides a [Learner](https://docs.fast.ai/learner.html#Learner) to handle the training, fine-tuning, and inference of deep learning algorithms.
+* [Finetuner](https://github.com/jina-ai/finetuner) is a service that enables models to create higher-quality embeddings for semantic search, visual similarity search, cross-modal text<->image search, recommendation systems, clustering, duplication detection, anomaly detection, or other uses.
+* [InvokeAI](https://github.com/invoke-ai/InvokeAI) is a creative engine for Stable Diffusion models, offering industry-leading WebUI, terminal usage support, and serves as the foundation for many commercial products.
+* [Kornia](https://kornia.readthedocs.io/en/latest/get-started/introduction.html) is a differentiable library that allows classical computer vision to be integrated into deep learning models. Kornia provides a [Trainer](https://kornia.readthedocs.io/en/latest/x.html#kornia.x.Trainer) with the specific purpose to train and fine-tune the supported deep learning algorithms within the library.
+* [Open Assistant](https://projects.laion.ai/Open-Assistant/) is a chat-based assistant that understands tasks, can interact with their party systems, and retrieve information dynamically to do so.
+* [pytorch-accelerated](https://github.com/Chris-hughes10/pytorch-accelerated) is a lightweight training library, with a streamlined feature set centered around a general-purpose [Trainer](https://pytorch-accelerated.readthedocs.io/en/latest/trainer.html), that places a huge emphasis on simplicity and transparency; enabling users to understand exactly what is going on under the hood, but without having to write and maintain the boilerplate themselves!
+* [Stable Diffusion web UI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) is an open-source browser-based easy-to-use interface based on the Gradio library for Stable Diffusion.
+* [torchkeras](https://github.com/lyhue1991/torchkeras) is a simple tool for training pytorch model just in a keras style, a dynamic and beautiful plot is provided in notebook to monitor your loss or metric.
+* [transformers](https://github.com/huggingface/transformers) as a tool for helping train state-of-the-art machine learning models in PyTorch, Tensorflow, and JAX. (Accelerate is the backend for the PyTorch side).
+
+
+## Installation
+
+This repository is tested on Python 3.8+ and PyTorch 1.10.0+
+
+You should install 🤗 Accelerate in a [virtual environment](https://docs.python.org/3/library/venv.html). If you're unfamiliar with Python virtual environments, check out the [user guide](https://packaging.python.org/guides/installing-using-pip-and-virtual-environments/).
+
+First, create a virtual environment with the version of Python you're going to use and activate it.
+
+Then, you will need to install PyTorch: refer to the [official installation page](https://pytorch.org/get-started/locally/#start-locally) regarding the specific install command for your platform. Then 🤗 Accelerate can be installed using pip as follows:
+
+```bash
+pip install accelerate
+```
+
+## Supported integrations
+
+- CPU only
+- multi-CPU on one node (machine)
+- multi-CPU on several nodes (machines)
+- single GPU
+- multi-GPU on one node (machine)
+- multi-GPU on several nodes (machines)
+- TPU
+- FP16/BFloat16 mixed precision
+- FP8 mixed precision with [Transformer Engine](https://github.com/NVIDIA/TransformerEngine) or [MS-AMP](https://github.com/Azure/MS-AMP/)
+- DeepSpeed support (Experimental)
+- PyTorch Fully Sharded Data Parallel (FSDP) support (Experimental)
+- Megatron-LM support (Experimental)
+
+## Citing 🤗 Accelerate
+
+If you use 🤗 Accelerate in your publication, please cite it by using the following BibTeX entry.
+
+```bibtex
+@Misc{accelerate,
+ title = {Accelerate: Training and inference at scale made simple, efficient and adaptable.},
+ author = {Sylvain Gugger and Lysandre Debut and Thomas Wolf and Philipp Schmid and Zachary Mueller and Sourab Mangrulkar and Marc Sun and Benjamin Bossan},
+ howpublished = {\url{https://github.com/huggingface/accelerate}},
+ year = {2022}
+}
+```
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/RECORD b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/RECORD
new file mode 100644
index 0000000000000000000000000000000000000000..6a75fe83db6c788837e887b88d2a2b14679cabf1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/RECORD
@@ -0,0 +1,95 @@
+../../../bin/accelerate,sha256=yGR8aKf_gKSS8x9pCBUGuD7rCpYLVZLhIyVaQrvwtWE,327
+../../../bin/accelerate-config,sha256=dUFEPDzU0dxtO1aioRHSqdcaCFuHDWurFs1a1A0pd5s,319
+../../../bin/accelerate-estimate-memory,sha256=1x5il15x5leo8DeYfFmUgzfQRVUVSUgcRQn1rioSBKk,321
+../../../bin/accelerate-launch,sha256=_Er34YoJFY7jHgz-unzt4xu-p-8x13fnFQtvuwW1YgY,319
+../../../bin/accelerate-merge-weights,sha256=H8qIBwDIWnizGG4GmuXRZ0K3iFjN9kl9orgQOh3zzlU,318
+accelerate-1.6.0.dist-info/INSTALLER,sha256=5hhM4Q4mYTT9z6QB6PGpUAW81PGNFrYrdXMj4oM_6ak,2
+accelerate-1.6.0.dist-info/LICENSE,sha256=xx0jnfkXJvxRnG63LTGOxlggYnIysveWIZ6H3PNdCrQ,11357
+accelerate-1.6.0.dist-info/METADATA,sha256=zT5ADQHZZeLT4qEiGMNSG4cT7hCnQplwyshDyeDyZNo,19421
+accelerate-1.6.0.dist-info/RECORD,,
+accelerate-1.6.0.dist-info/REQUESTED,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+accelerate-1.6.0.dist-info/WHEEL,sha256=GV9aMThwP_4oNCtvEC2ec3qUYutgWeAzklro_0m4WJQ,91
+accelerate-1.6.0.dist-info/entry_points.txt,sha256=Vpy8gUGfZ-1VnM2229fb8CpJNLBdMH_wtJ9PQ7b_2tQ,296
+accelerate-1.6.0.dist-info/top_level.txt,sha256=esVfdxTidsjQ90zsN_rPpjLFJ4ijRlx4mnLrG09hlt4,11
+accelerate/__init__.py,sha256=r3I-pArsQK9ZrH3XgnjeCoXo4l-DEFOWQhjj3BguTZc,1504
+accelerate/accelerator.py,sha256=G952noNHGPrl-poK6qAj1OY32kGjmN5S13v8zy7H63E,173175
+accelerate/big_modeling.py,sha256=IMiAtiuZQpwSyk2jQsoYC2uWzfRUSpCg7FiThSvjfKw,29702
+accelerate/checkpointing.py,sha256=BaDOrpQzRI2U1BvN2vK4lepTRNNqbxGd4QPa1zOShoc,13612
+accelerate/commands/__init__.py,sha256=m1PPTDT4ziIAvM0-FDSgIMIZ69Konn126s6LwuzH6v8,606
+accelerate/commands/accelerate_cli.py,sha256=SkwFad6Z1ZsGjtm7TiXFq8je-akshp_0WxX_6rGSBw8,1972
+accelerate/commands/config/__init__.py,sha256=iJK8dgj3pc5Vdr1E7UuGoFu-BlybyXLxYDoTg9gXngE,1645
+accelerate/commands/config/cluster.py,sha256=w0L3zTyZp4sjDpCrM3NxOjxZ0kyPJZmzi06pFZmbM2c,37472
+accelerate/commands/config/config.py,sha256=FuRlQvOjgATEtyqOSsGD-KEtOCvACOHjs2C-krrtldk,3035
+accelerate/commands/config/config_args.py,sha256=xn6M8iJnlFycosDlbM0BE86r9RxfdDwHtIlk-UUq7UM,10082
+accelerate/commands/config/config_utils.py,sha256=mdvZE9fpllfD8S4Blhqk3nLqQ5m14WJ0jQ1yh768H10,3177
+accelerate/commands/config/default.py,sha256=sPgQVt_0zk68KlupQFqt8B6JUoPMFPxXmXr7xFM-EN8,6212
+accelerate/commands/config/sagemaker.py,sha256=GjHE2-h4tRr1P_PFtMF3miiAtJlzkbHbMb6kFXqn8eo,10341
+accelerate/commands/config/update.py,sha256=NXW1J7GkUHpg71QlIXsmMB_0z8S8IZo2FWax5POwrhc,2395
+accelerate/commands/env.py,sha256=-B3FPX4S705A-P_tyLKm_JzGpz-TeKqFNPdNWDAdGIM,4156
+accelerate/commands/estimate.py,sha256=Qduq4xudVyIede37BMEe1rNhXf-rfW-MHV2KtwxdfEA,12585
+accelerate/commands/launch.py,sha256=7DI42Uw4kf_peOpY5TUA1V2yz7cuSO3cYnLgiI5G1Vs,47496
+accelerate/commands/menu/__init__.py,sha256=uqSlBM0TFHBwzdv3p3SXfpAk1lZFp4h1a7mbBdscPHs,645
+accelerate/commands/menu/cursor.py,sha256=-lmpJVAzvNc0c3EOtSuLoKB59zqylVCbYyWLPnrOmvQ,2028
+accelerate/commands/menu/helpers.py,sha256=KrSB5fJjH4MUEUAQJ6bYaN16AYcnl9UalDrPD3DYeeg,1483
+accelerate/commands/menu/input.py,sha256=T8Mdd-Y_OURgqfDV9qZh4Wf6hmT22AneNtJzj4JA1Rk,2512
+accelerate/commands/menu/keymap.py,sha256=eXj-suyYs1m5dEHoUKN4mKAMLc8DWHnwhP6G6JSU0jQ,4086
+accelerate/commands/menu/selection_menu.py,sha256=bxy-DHaKKC6SCToOlMBv5_z0MdUzylEg6Sio9OuV3GM,4921
+accelerate/commands/merge.py,sha256=quDKckN3vKn9nsGjdwfoojnfTMFdKRRUkY1DYuuNNmc,2388
+accelerate/commands/test.py,sha256=YrPYEaAACOGZ6btn2MV6NbMSEdBUcMWADLbQWaZSHtk,2149
+accelerate/commands/to_fsdp2.py,sha256=gfbhoUT4qFB3LVDMNmckElgLG0yWm8aj_aofszeiJmM,5991
+accelerate/commands/tpu.py,sha256=KyxDP7IuveidZrbW4rx2s8Ku3o_ptI6tzwr_R7ck0os,5548
+accelerate/commands/utils.py,sha256=aT8xUCe2pCkFII7yZxcfaohEjgBAzMUM7WiD4UuWSOY,4150
+accelerate/data_loader.py,sha256=yArisKhfuIJzDD7vuOgZAqEJNUC8tgl2L8ay92rgtfY,64551
+accelerate/hooks.py,sha256=lYtYSIqEQnZOImgj2UMTngQPkcQDEHS2klwak1oHD6w,32248
+accelerate/inference.py,sha256=NLANdzXm5PwmDWbPYkFmoRoQSLLvuhfvIG33xfpapT0,7668
+accelerate/launchers.py,sha256=QIqUVkDc-oTmWf00L8kas7u2RBEwOYoRi8M2Our0DAs,13721
+accelerate/local_sgd.py,sha256=aCj_yqXK_FhhZRWEpzXIkgXBERH6fC3HyrC3nsOj1uA,4160
+accelerate/logging.py,sha256=4XcgY_BV7Qn_enh2tZ-8fNtuaE_3n-LsYJbgwhRx_PI,5042
+accelerate/memory_utils.py,sha256=3R5LoeHl6GgTZ-IMPrDZMdaEehWarGdPqODushb-6pg,862
+accelerate/optimizer.py,sha256=QfgCkQ5dA-fLSi_Z7CBPRCObXA1rL9zxHg4tyKCEg2A,8113
+accelerate/scheduler.py,sha256=des_4M_Tt1W8gCYZZbLla0GHBEgJY3Wx2EGBQPTzeiY,4238
+accelerate/state.py,sha256=YYpuPqXeNjz5_Y71h0zmCu13cBuDmQ8lw6fAmoSWUFk,55457
+accelerate/test_utils/__init__.py,sha256=8xikmLMAM6_6CwVF6tsdsv4XzgWkHAk2tZBdV9DxIH8,1749
+accelerate/test_utils/examples.py,sha256=IN4n2lxA95hexE2rojsyyjhpXLbXnbmjTzd8UTws5_4,7257
+accelerate/test_utils/scripts/__init__.py,sha256=m1PPTDT4ziIAvM0-FDSgIMIZ69Konn126s6LwuzH6v8,606
+accelerate/test_utils/scripts/external_deps/__init__.py,sha256=m1PPTDT4ziIAvM0-FDSgIMIZ69Konn126s6LwuzH6v8,606
+accelerate/test_utils/scripts/external_deps/test_checkpointing.py,sha256=XHaNRmnrARd1izXFjWGi5UjYGas-4vqayW51jAHBPCA,10699
+accelerate/test_utils/scripts/external_deps/test_ds_multiple_model.py,sha256=Cg4-h0B4UcOQ5CxXjIdrsPVR5fFsWCv24DqZGjXEwW8,13790
+accelerate/test_utils/scripts/external_deps/test_metrics.py,sha256=Ev2XKaiwmznoxKujskAAuISGChW646MOiyf0CXEPb9Y,12168
+accelerate/test_utils/scripts/external_deps/test_peak_memory_usage.py,sha256=9Yn9Rc7d-yWr1fU0RagASPG5l8vrKeHVYbuYABbA-fU,12498
+accelerate/test_utils/scripts/external_deps/test_performance.py,sha256=Di6LT19bCBLlWmCBSu_jjdqR2EqngXpvUOGDBx8GfZE,10432
+accelerate/test_utils/scripts/external_deps/test_pippy.py,sha256=ocZntbmAduln2ma4LeEA9o-S8hla3YXCJ_A8hEcWHgs,4762
+accelerate/test_utils/scripts/external_deps/test_zero3_integration.py,sha256=P9alBOHZ9Lfqs5LoRP7bCbXl-tnsNrBkvJZGseibBeA,1665
+accelerate/test_utils/scripts/test_cli.py,sha256=qfk1aYFtdvYFCYPkl05602SNGvk08QTv0xZVVcFVtzM,833
+accelerate/test_utils/scripts/test_ddp_comm_hook.py,sha256=k_-2MBjLKNdMGIcneTbuGd84K05Wp1GEQX6DUVF9UBw,3566
+accelerate/test_utils/scripts/test_distributed_data_loop.py,sha256=RUWTwd7DIpr2fl7JtKOsvTjMiJioTxO8FdSr2Lw_5uI,15137
+accelerate/test_utils/scripts/test_merge_weights.py,sha256=dssMnAoZt291vNLbPhPOTQUooh0leg_0erQh0uZH6aU,6125
+accelerate/test_utils/scripts/test_notebook.py,sha256=qfIy3IvH74-kGn8nadBn_k7qrviqvsxy5ijsnUhuY6o,3894
+accelerate/test_utils/scripts/test_ops.py,sha256=Bcs-h8EMJwULTfbizlFN5qkv3JraWEpoSZWMn-HswiI,6265
+accelerate/test_utils/scripts/test_script.py,sha256=8-53hIVQXD28HQT4h2Ijy6yGCHfTWDAf1-HOi4UtDng,34219
+accelerate/test_utils/scripts/test_sync.py,sha256=PDe8sYZLCL2LKjj_L9b-Bh2BjAjeii9EZ8sZNfuYx5s,18817
+accelerate/test_utils/testing.py,sha256=x9RK70VgAMyHlo5xut7P85j-9kdAnlfQe_4jwSPpMv4,27807
+accelerate/test_utils/training.py,sha256=jO5YEIr34jAcnJ_9WNp_x3zuHzSam_I6IgMvmcGm7yI,6456
+accelerate/tracking.py,sha256=ucpsoYAT3pVXgOfwDdXf6qTugY2-tk-EINvZtfmRitM,42756
+accelerate/utils/__init__.py,sha256=wjpXyvFxS-ed3Stwm_IHIlBmsmP7KRyAljQ_Qss-OWw,7802
+accelerate/utils/ao.py,sha256=koMiji7AG1kJMRMkJnwSnpuycfx4lPY3CNnpNx2ZqzM,4736
+accelerate/utils/bnb.py,sha256=KCbg6LUt4eXvPHVnKh7rSVcPwDnzxY_Ii7yYmK5bNGw,20737
+accelerate/utils/constants.py,sha256=hc24V0pgxWdBQwS6SXxDKwuIni2pCnzdfvMOX1XI9Os,3264
+accelerate/utils/dataclasses.py,sha256=E7CnCbfskpzxzSorst95Via_XE39t0NP_UGYgJUris0,131486
+accelerate/utils/deepspeed.py,sha256=QYIXv5LwHXw7wBFFo-7a0t86MbwNAfieJkkBaLGA6wI,14064
+accelerate/utils/environment.py,sha256=h0zacbBkAp9szltTf5-aTr5NcbVsQp7wl6DFWp8XNuI,15257
+accelerate/utils/fsdp_utils.py,sha256=Q2tc9EakwBjuYlyXvQrBLV97r6cdReRft6KeS1P_Vb4,28938
+accelerate/utils/imports.py,sha256=YI1ebPJAuxarclENTfzvDPPGf6jeEKnVQ42taFPuqh0,16759
+accelerate/utils/launch.py,sha256=nN4ykAtnEL3oITLTejABltdpS3OivcE2COmX-BnWuY4,31195
+accelerate/utils/megatron_lm.py,sha256=FnIF-niZjvdMk9ymafZWEPjDho_Q_P98C69qc9g5r_E,58059
+accelerate/utils/memory.py,sha256=lDHqW7Ue_CPmw_DWgNxX_B3HY71_srAFdgR10XiVRSM,6960
+accelerate/utils/modeling.py,sha256=_xSTiH7zSsffZULSTJuzcDK6IaWImEMOcbq1xqeI7GY,92319
+accelerate/utils/offload.py,sha256=VFaL8oSJzqZ_47VuUQ69xZi9bF2heRSFoOSnnOxbGXc,7825
+accelerate/utils/operations.py,sha256=VWPYvtrO4UGX5JmisanXzLLUbhAeL8kQk0yYc66bQ2M,31055
+accelerate/utils/other.py,sha256=iiLZcKEAlK2Sj_wt03gAEGKrk7_NZFwbmy9cgEppRPw,13231
+accelerate/utils/random.py,sha256=Xv_ZJm9eaC2Q7rgZy9OpOunKuTingMiDQCH00qhNVxE,6220
+accelerate/utils/rich.py,sha256=8JZX_uGMQX-BufdXxJpdne7BWd1KyLHSgbiGxrDMYr8,847
+accelerate/utils/torch_xla.py,sha256=Pq1tuqN0X_pWDVza6YgjfO45uoJdoRVRForLeLQzFus,1908
+accelerate/utils/tqdm.py,sha256=k8e9JnieTEQHCCNBaiBys7hPxWlEbyRASdIma-qy_X8,1657
+accelerate/utils/transformer_engine.py,sha256=498Y3z2BkbybYLtBiuF_TJgt8Iii943s4wgRAV8FDC4,6372
+accelerate/utils/versions.py,sha256=UgmcbjBm--6CIx1ZamSAMjAK_B_2l48LbeaNygqej8M,2149
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/REQUESTED b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/REQUESTED
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/WHEEL b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/WHEEL
new file mode 100644
index 0000000000000000000000000000000000000000..dcfdc6e359074689c0bdb567634b4f84add7849c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/WHEEL
@@ -0,0 +1,5 @@
+Wheel-Version: 1.0
+Generator: setuptools (75.1.0)
+Root-Is-Purelib: true
+Tag: py3-none-any
+
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/entry_points.txt b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/entry_points.txt
new file mode 100644
index 0000000000000000000000000000000000000000..8b9bf6b798b250a47a3febdf0e32c88507fbf86d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/entry_points.txt
@@ -0,0 +1,6 @@
+[console_scripts]
+accelerate = accelerate.commands.accelerate_cli:main
+accelerate-config = accelerate.commands.config:main
+accelerate-estimate-memory = accelerate.commands.estimate:main
+accelerate-launch = accelerate.commands.launch:main
+accelerate-merge-weights = accelerate.commands.merge:main
diff --git a/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/top_level.txt b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/top_level.txt
new file mode 100644
index 0000000000000000000000000000000000000000..a9368375be0e0e13fdad0eea4b92541bd9e1f594
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate-1.6.0.dist-info/top_level.txt
@@ -0,0 +1 @@
+accelerate
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_ddp_comm_hook.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_ddp_comm_hook.py
new file mode 100644
index 0000000000000000000000000000000000000000..0db5844e026d1c035670e518a8f81d33136ea665
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_ddp_comm_hook.py
@@ -0,0 +1,85 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 torch
+
+from accelerate import Accelerator, DDPCommunicationHookType, DistributedDataParallelKwargs, PartialState
+from accelerate.utils import is_hpu_available
+
+
+class MockModel(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ torch.manual_seed(0)
+ self.p = torch.nn.Parameter(torch.randn(40, 20))
+
+ def forward(self, x, rank):
+ return self.p * (x ** (1 + rank))
+
+
+def _run_and_get_grads(model, rank):
+ torch.manual_seed(2024)
+ input = torch.randn(40, 20)
+ output = model(input, rank)
+ output.mean().backward()
+ param = next(model.parameters())
+ return param.grad
+
+
+def test_ddp_comm_hook(comm_hook, comm_wrapper, comm_state_option):
+ ddp_kwargs = DistributedDataParallelKwargs(
+ comm_hook=comm_hook,
+ comm_wrapper=comm_wrapper,
+ comm_state_option=comm_state_option,
+ )
+ accelerator = Accelerator(kwargs_handlers=[ddp_kwargs])
+
+ model = accelerator.prepare(MockModel())
+ hook_grads = _run_and_get_grads(model, accelerator.local_process_index)
+
+ reference_model = torch.nn.parallel.DistributedDataParallel(
+ MockModel().to(accelerator.device),
+ device_ids=[accelerator.local_process_index],
+ output_device=accelerator.local_process_index,
+ )
+ reference_grads = _run_and_get_grads(reference_model, accelerator.local_process_index)
+
+ torch.testing.assert_close(hook_grads, reference_grads, rtol=1e-2, atol=1e-2)
+
+
+def main():
+ for comm_hook, comm_wrapper, comm_state_option in [
+ (DDPCommunicationHookType.NO, DDPCommunicationHookType.NO, {}),
+ (DDPCommunicationHookType.FP16, DDPCommunicationHookType.NO, {}),
+ (DDPCommunicationHookType.BF16, DDPCommunicationHookType.NO, {}),
+ (DDPCommunicationHookType.POWER_SGD, DDPCommunicationHookType.NO, {}),
+ (DDPCommunicationHookType.POWER_SGD, DDPCommunicationHookType.FP16, {}),
+ (DDPCommunicationHookType.POWER_SGD, DDPCommunicationHookType.BF16, {}),
+ (DDPCommunicationHookType.POWER_SGD, DDPCommunicationHookType.NO, {"matrix_approximation_rank": 2}),
+ (DDPCommunicationHookType.BATCHED_POWER_SGD, DDPCommunicationHookType.NO, {}),
+ (DDPCommunicationHookType.BATCHED_POWER_SGD, DDPCommunicationHookType.FP16, {}),
+ (DDPCommunicationHookType.BATCHED_POWER_SGD, DDPCommunicationHookType.BF16, {}),
+ ]:
+ if is_hpu_available():
+ HPU_UNSUPPORTED_COMM_HOOKS = {DDPCommunicationHookType.FP16, DDPCommunicationHookType.BF16}
+ if comm_hook in HPU_UNSUPPORTED_COMM_HOOKS or comm_wrapper in HPU_UNSUPPORTED_COMM_HOOKS:
+ print(f"Skipping test DDP comm hook: {comm_hook}, comm wrapper: {comm_wrapper} on HPU")
+ continue
+
+ print(f"Test DDP comm hook: {comm_hook}, comm wrapper: {comm_wrapper}")
+ test_ddp_comm_hook(comm_hook, comm_wrapper, comm_state_option)
+ PartialState().destroy_process_group()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_distributed_data_loop.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_distributed_data_loop.py
new file mode 100644
index 0000000000000000000000000000000000000000..08cbbeb844bcc3189c99d328580c251aeafc2052
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_distributed_data_loop.py
@@ -0,0 +1,410 @@
+#!/usr/bin/env python
+
+# Copyright 2021 The HuggingFace Team. All rights reserved.
+#
+# 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 pickle
+import tempfile
+import warnings
+from unittest.mock import Mock
+
+import torch
+from torch.utils.data import (
+ BatchSampler,
+ DataLoader,
+ Dataset,
+ IterableDataset,
+ RandomSampler,
+ TensorDataset,
+ default_collate,
+)
+
+from accelerate.accelerator import Accelerator, DataLoaderConfiguration
+from accelerate.utils.dataclasses import DistributedType
+
+
+NUM_ELEMENTS = 22
+NUM_WORKERS = 4
+BATCH_SIZE = 4
+
+
+class DummyDataset(Dataset):
+ def __len__(self):
+ return NUM_ELEMENTS
+
+ def __getitem__(self, index):
+ squeeze = False
+
+ if isinstance(index, int):
+ index = [index]
+ squeeze = True
+ elif isinstance(index, slice):
+ index = list(range(*index.indices(self.size)))
+ else:
+ index = list(index)
+
+ batch = [{"index": i, "label": i % 2, "random_augmentation": torch.rand(1).item()} for i in index]
+
+ if squeeze:
+ batch = batch[0]
+
+ return batch
+
+
+class DummyIterableDataset(IterableDataset):
+ def __init__(self, data):
+ self.data = data
+
+ def __iter__(self):
+ yield from self.data
+
+
+def create_accelerator(even_batches=True):
+ dataloader_config = DataLoaderConfiguration(even_batches=even_batches)
+ accelerator = Accelerator(dataloader_config=dataloader_config)
+ assert accelerator.num_processes == 2, "this script expects that two GPUs are available"
+ return accelerator
+
+
+def create_dataloader(
+ accelerator: Accelerator, dataset_size: int, batch_size: int, iterable: bool = False, shuffle: bool = False
+):
+ """
+ Create a simple DataLoader to use during the test cases
+ """
+ values = torch.as_tensor(range(dataset_size))
+ if shuffle:
+ values = values[torch.randperm(values.size(0))]
+ if iterable:
+ dataset = DummyIterableDataset(values)
+ else:
+ dataset = TensorDataset(torch.as_tensor(range(dataset_size)))
+
+ dl = DataLoader(dataset, batch_size=batch_size)
+ dl = accelerator.prepare(dl)
+
+ return dl
+
+
+def verify_dataloader_batch_sizes(
+ accelerator: Accelerator,
+ dataset_size: int,
+ batch_size: int,
+ process_0_expected_batch_sizes: list[int],
+ process_1_expected_batch_sizes: list[int],
+):
+ """
+ A helper function for verifying the batch sizes coming from a prepared dataloader in each process
+ """
+ dl = create_dataloader(accelerator=accelerator, dataset_size=dataset_size, batch_size=batch_size)
+
+ batch_sizes = [len(batch[0]) for batch in dl]
+
+ if accelerator.process_index == 0:
+ assert batch_sizes == process_0_expected_batch_sizes
+ elif accelerator.process_index == 1:
+ assert batch_sizes == process_1_expected_batch_sizes
+
+
+def test_default_ensures_even_batch_sizes():
+ accelerator = create_accelerator()
+
+ # without padding, we would expect a different number of batches
+ verify_dataloader_batch_sizes(
+ accelerator,
+ dataset_size=3,
+ batch_size=1,
+ process_0_expected_batch_sizes=[1, 1],
+ process_1_expected_batch_sizes=[1, 1],
+ )
+
+ # without padding, we would expect the same number of batches, but different sizes
+ verify_dataloader_batch_sizes(
+ accelerator,
+ dataset_size=7,
+ batch_size=2,
+ process_0_expected_batch_sizes=[2, 2],
+ process_1_expected_batch_sizes=[2, 2],
+ )
+
+
+def test_can_disable_even_batches():
+ accelerator = create_accelerator(even_batches=False)
+
+ verify_dataloader_batch_sizes(
+ accelerator,
+ dataset_size=3,
+ batch_size=1,
+ process_0_expected_batch_sizes=[1, 1],
+ process_1_expected_batch_sizes=[1],
+ )
+
+ verify_dataloader_batch_sizes(
+ accelerator,
+ dataset_size=7,
+ batch_size=2,
+ process_0_expected_batch_sizes=[2, 2],
+ process_1_expected_batch_sizes=[2, 1],
+ )
+
+
+def test_can_join_uneven_inputs():
+ accelerator = create_accelerator(even_batches=False)
+
+ model = torch.nn.Linear(1, 1)
+ ddp_model = accelerator.prepare(model)
+
+ dl = create_dataloader(accelerator, dataset_size=3, batch_size=1)
+
+ batch_idxs = []
+ with accelerator.join_uneven_inputs([ddp_model]):
+ for batch_idx, batch in enumerate(dl):
+ output = ddp_model(batch[0].float())
+ loss = output.sum()
+ loss.backward()
+ batch_idxs.append(batch_idx)
+
+ accelerator.wait_for_everyone()
+
+ if accelerator.process_index == 0:
+ assert batch_idxs == [0, 1]
+ elif accelerator.process_index == 1:
+ assert batch_idxs == [0]
+
+
+def test_join_raises_warning_for_non_ddp_distributed(accelerator):
+ with warnings.catch_warnings(record=True) as w:
+ with accelerator.join_uneven_inputs([Mock()]):
+ pass
+
+ assert issubclass(w[-1].category, UserWarning)
+ assert "only supported for multi-GPU" in str(w[-1].message)
+
+
+def test_join_can_override_even_batches():
+ default_even_batches = True
+ overridden_even_batches = False
+ accelerator = create_accelerator(even_batches=default_even_batches)
+ model = torch.nn.Linear(1, 1)
+ ddp_model = accelerator.prepare(model)
+ train_dl = create_dataloader(accelerator, dataset_size=3, batch_size=1)
+ valid_dl = create_dataloader(accelerator, dataset_size=3, batch_size=1)
+
+ with accelerator.join_uneven_inputs([ddp_model], even_batches=overridden_even_batches):
+ train_dl_overridden_value = train_dl.batch_sampler.even_batches
+ valid_dl_overridden_value = valid_dl.batch_sampler.even_batches
+
+ assert train_dl_overridden_value == overridden_even_batches
+ assert valid_dl_overridden_value == overridden_even_batches
+ assert train_dl.batch_sampler.even_batches == default_even_batches
+ assert valid_dl.batch_sampler.even_batches == default_even_batches
+
+
+def test_join_can_override_for_mixed_type_dataloaders():
+ default_even_batches = True
+ overridden_even_batches = False
+ accelerator = create_accelerator(even_batches=default_even_batches)
+ model = torch.nn.Linear(1, 1)
+ ddp_model = accelerator.prepare(model)
+ create_dataloader(accelerator, dataset_size=3, batch_size=1, iterable=True)
+ batch_dl = create_dataloader(accelerator, dataset_size=3, batch_size=1)
+
+ with warnings.catch_warnings():
+ warnings.filterwarnings("ignore")
+ try:
+ with accelerator.join_uneven_inputs([ddp_model], even_batches=overridden_even_batches):
+ batch_dl_overridden_value = batch_dl.batch_sampler.even_batches
+ except AttributeError:
+ # ensure attribute error is not raised when processing iterable dl
+ raise AssertionError
+
+ assert batch_dl_overridden_value == overridden_even_batches
+ assert batch_dl.batch_sampler.even_batches == default_even_batches
+
+
+def test_join_raises_warning_for_iterable_when_overriding_even_batches():
+ accelerator = create_accelerator()
+ model = torch.nn.Linear(1, 1)
+ ddp_model = accelerator.prepare(model)
+ create_dataloader(accelerator, dataset_size=3, batch_size=1, iterable=True)
+
+ with warnings.catch_warnings(record=True) as w:
+ with accelerator.join_uneven_inputs([ddp_model], even_batches=False):
+ pass
+
+ assert issubclass(w[-1].category, UserWarning)
+ assert "only supported for map-style datasets" in str(w[-1].message)
+
+
+def test_pickle_accelerator():
+ accelerator = create_accelerator()
+ data_loader = create_dataloader(accelerator, dataset_size=32, batch_size=4)
+ _ = accelerator.prepare(data_loader)
+ pickled_accelerator = pickle.dumps(accelerator)
+ unpickled_accelerator = pickle.loads(pickled_accelerator)
+ # TODO: Maybe this should be implemented as __eq__ for AcceleratorState?
+ assert accelerator.state.__dict__ == unpickled_accelerator.state.__dict__
+
+
+def test_data_loader(data_loader, accelerator):
+ # Prepare the DataLoader
+ data_loader = accelerator.prepare(data_loader)
+
+ all_examples = []
+ for i, batch in enumerate(data_loader):
+ index, _ = accelerator.gather_for_metrics((batch["index"], batch["label"]))
+ all_examples.extend(index.detach().cpu().numpy().tolist())
+
+ # Sort the examples
+ sorted_all_examples = sorted(all_examples)
+
+ # Check if all elements are present in the sorted list of iterated samples
+ assert len(set(sorted_all_examples)) == NUM_ELEMENTS, (
+ "Not all the dataset elements have been iterated in an epoch due to duplication of samples across processes."
+ )
+
+
+def test_stateful_dataloader(accelerator):
+ """
+ Tests that a stateful dataloader can be iterated over, saved after a few batches using `load_state_dict`, and then
+ resumed from the saved state.
+
+ The result should be the same as the rest of the data that iterated over after saving.
+ """
+ old_dataloader_config = accelerator.dataloader_config
+ try:
+ accelerator.dataloader_config = DataLoaderConfiguration(use_stateful_dataloader=True)
+ prepared_dl = create_dataloader(
+ accelerator, dataset_size=32 * accelerator.num_processes, batch_size=4, iterable=True, shuffle=True
+ )
+ untrained_batches = []
+ # Calculate what step that will be
+ total_batches = 32 * accelerator.num_processes // (4 * accelerator.num_processes)
+ last_batch_num = total_batches - 1
+ for step, batch in enumerate(prepared_dl):
+ # Step just before
+ if step == last_batch_num - 1:
+ state_dict = prepared_dl.state_dict()
+ if step >= last_batch_num:
+ # Otherwise grab the "unseen" batches
+ untrained_batches.append(batch)
+ not_skipped_batches = accelerator.gather(untrained_batches)
+ prepared_dl.load_state_dict(state_dict)
+ resumed_batches = []
+ for batch in prepared_dl:
+ resumed_batches.append(batch)
+ resumed_batches = accelerator.gather(resumed_batches)
+ for b1, b2 in zip(not_skipped_batches, resumed_batches):
+ for v1, v2 in zip(b1, b2):
+ assert torch.equal(v1, v2), f"Batch {b1} and {b2} are not equal"
+ finally:
+ accelerator.dataloader_config = old_dataloader_config
+
+
+def test_stateful_dataloader_save_state(accelerator):
+ """
+ Tests that a stateful dataloader can be iterated over, saved after a few batches using `Accelerator.save_state`,
+ and then resumed from the saved state.
+
+ The result should be the same as the rest of the data that iterated over after saving.
+ """
+ old_dataloader_config = accelerator.dataloader_config
+ try:
+ with tempfile.TemporaryDirectory() as tmpdir:
+ accelerator.dataloader_config = DataLoaderConfiguration(use_stateful_dataloader=True)
+ prepared_dl = create_dataloader(
+ accelerator, dataset_size=32 * accelerator.num_processes, batch_size=4, iterable=True, shuffle=True
+ )
+ untrained_batches = []
+ # Calculate what step that will be
+ total_batches = 32 * accelerator.num_processes // (4 * accelerator.num_processes)
+ last_batch_num = total_batches - 1
+ for step, batch in enumerate(prepared_dl):
+ # Step just before
+ if step == last_batch_num - 1:
+ accelerator.save_state(tmpdir)
+ if step >= last_batch_num:
+ # Otherwise grab the "unseen" batches
+ untrained_batches.append(batch)
+ not_skipped_batches = accelerator.gather(untrained_batches)
+ accelerator.load_state(tmpdir)
+ resumed_batches = []
+ for batch in prepared_dl:
+ resumed_batches.append(batch)
+ resumed_batches = accelerator.gather(resumed_batches)
+ for b1, b2 in zip(not_skipped_batches, resumed_batches):
+ for v1, v2 in zip(b1, b2):
+ assert torch.equal(v1, v2), f"Batch {b1} and {b2} are not equal"
+ finally:
+ accelerator.dataloader_config = old_dataloader_config
+
+
+def main():
+ accelerator = create_accelerator()
+ torch.manual_seed(accelerator.process_index)
+
+ accelerator.print("Test that even_batches variable ensures uniform batches across processes")
+ test_default_ensures_even_batch_sizes()
+
+ accelerator.print("Run tests with even_batches disabled")
+ test_can_disable_even_batches()
+
+ accelerator.print("Test joining uneven inputs")
+ test_can_join_uneven_inputs()
+
+ accelerator.print("Test overriding even_batches when joining uneven inputs")
+ test_join_can_override_even_batches()
+
+ accelerator.print("Test overriding even_batches for mixed dataloader types")
+ test_join_can_override_for_mixed_type_dataloaders()
+
+ accelerator.print("Test overriding even_batches raises a warning for iterable dataloaders")
+ test_join_raises_warning_for_iterable_when_overriding_even_batches()
+
+ accelerator.print("Test join with non DDP distributed raises warning")
+ original_state = accelerator.state.distributed_type
+ accelerator.state.distributed_type = DistributedType.FSDP
+ test_join_raises_warning_for_non_ddp_distributed(accelerator)
+ accelerator.state.distributed_type = original_state
+
+ accelerator.print("Test pickling an accelerator")
+ test_pickle_accelerator()
+
+ dataset = DummyDataset()
+
+ accelerator.print("Test DataLoader with shuffle=False")
+ loader = DataLoader(dataset, shuffle=False, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS)
+ test_data_loader(loader, accelerator)
+
+ accelerator.print("Test DataLoader with shuffle=True")
+ loader = DataLoader(dataset, shuffle=True, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS)
+ test_data_loader(loader, accelerator)
+
+ accelerator.print("Test DataLoader with batch_sampler")
+ sampler = BatchSampler(RandomSampler(dataset), batch_size=BATCH_SIZE, drop_last=False)
+ loader = DataLoader(dataset, batch_sampler=sampler, num_workers=NUM_WORKERS)
+ test_data_loader(loader, accelerator)
+
+ accelerator.print("Test DataLoader with sampler as an instance of `BatchSampler`")
+ sampler = BatchSampler(RandomSampler(dataset), batch_size=BATCH_SIZE, drop_last=False)
+ loader = DataLoader(dataset, sampler=sampler, batch_size=None, collate_fn=default_collate, num_workers=NUM_WORKERS)
+ test_data_loader(loader, accelerator)
+ test_stateful_dataloader(accelerator)
+ test_stateful_dataloader_save_state(accelerator)
+
+ accelerator.end_training()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_merge_weights.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_merge_weights.py
new file mode 100644
index 0000000000000000000000000000000000000000..8671cf99ecc8aa9149ab5de98deb73e1083cecab
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_merge_weights.py
@@ -0,0 +1,162 @@
+# Copyright 2024 The HuggingFace Team. All rights reserved.
+#
+# 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 gc
+import logging
+import shutil
+from pathlib import Path
+
+import torch
+from safetensors.torch import load_file
+from torch.distributed.fsdp.fully_sharded_data_parallel import ShardingStrategy, StateDictType
+from torch.utils.data import DataLoader
+
+from accelerate import Accelerator, FullyShardedDataParallelPlugin
+from accelerate.commands.merge import merge_command, merge_command_parser
+from accelerate.state import AcceleratorState
+from accelerate.test_utils import torch_device
+from accelerate.test_utils.training import RegressionDataset
+from accelerate.utils import merge_fsdp_weights, patch_environment, save_fsdp_model
+
+
+logging.basicConfig(level=logging.INFO)
+
+parser = merge_command_parser()
+
+
+class TinyModel(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.linear1 = torch.nn.Linear(16, 16)
+ self.activation = torch.nn.ReLU()
+ self.linear2 = torch.nn.Linear(16, 16)
+ self.softmax = torch.nn.Softmax()
+
+ def forward(self, x):
+ return self.linear2(self.activation(self.linear1(x)))
+
+
+def setup():
+ if AcceleratorState._shared_state != {}:
+ AcceleratorState()._reset_state()
+ plugin = FullyShardedDataParallelPlugin(
+ sharding_strategy=ShardingStrategy.FULL_SHARD, state_dict_type=StateDictType.SHARDED_STATE_DICT
+ )
+ model = TinyModel()
+ with patch_environment(fsdp_auto_wrap_policy="SIZE_BASED_WRAP"):
+ plugin.set_auto_wrap_policy(model)
+ accelerator = Accelerator(fsdp_plugin=plugin)
+ model = accelerator.prepare(model)
+ return model, plugin, accelerator
+
+
+def mock_training(accelerator, model):
+ train_set = RegressionDataset(length=128, seed=42)
+ train_dl = DataLoader(train_set, batch_size=16, shuffle=False)
+ optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
+
+ train_dl, model, optimizer = accelerator.prepare(train_dl, model, optimizer)
+ for _ in range(3):
+ for batch in train_dl:
+ model.zero_grad()
+ output = model(batch["x"])
+ loss = torch.nn.functional.mse_loss(output, batch["y"])
+ accelerator.backward(loss)
+ optimizer.step()
+ return model
+
+
+def check_weights(operation, state_1, state_2):
+ for weight_1, weight_2 in zip(state_1.values(), state_2.values()):
+ if str(weight_1.device) != torch_device:
+ weight_1 = weight_1.to(torch_device)
+ if str(weight_2.device) != torch_device:
+ weight_2 = weight_2.to(torch_device)
+ if operation == "same":
+ assert torch.allclose(weight_1, weight_2)
+ else:
+ assert not torch.allclose(weight_1, weight_2)
+
+
+def check_safetensors_weights(path, model):
+ safe_state_dict = load_file(path / "model.safetensors")
+ safe_loaded_model = TinyModel()
+ check_weights("diff", model.state_dict(), safe_loaded_model.state_dict())
+ safe_loaded_model.load_state_dict(safe_state_dict)
+ check_weights("same", model.state_dict(), safe_loaded_model.state_dict())
+
+
+def check_pytorch_weights(path, model):
+ nonsafe_state_dict = torch.load(path / "pytorch_model.bin")
+ nonsafe_loaded_model = TinyModel()
+ check_weights("diff", model.state_dict(), nonsafe_loaded_model.state_dict())
+ nonsafe_loaded_model.load_state_dict(nonsafe_state_dict)
+ check_weights("same", model.state_dict(), nonsafe_loaded_model.state_dict())
+
+
+def test_merge_weights_safetensors(model, path):
+ # Should now be saved at `path/merged.safetensors`
+ merge_fsdp_weights(path / "pytorch_model_fsdp_0", path, safe_serialization=True)
+ check_safetensors_weights(path, model)
+
+
+def test_merge_weights_command_safetensors(model, path):
+ args = parser.parse_args([str(path / "pytorch_model_fsdp_0"), str(path)])
+ merge_command(args)
+ check_safetensors_weights(path, model)
+
+
+def test_merge_weights_pytorch(model, path):
+ # Should now be saved at `path/merged.bin`
+ merge_fsdp_weights(path / "pytorch_model_fsdp_0", path, safe_serialization=False)
+ check_pytorch_weights(path, model)
+
+
+def test_merge_weights_command_pytorch(model, path):
+ args = parser.parse_args([str(path / "pytorch_model_fsdp_0"), str(path), "--unsafe_serialization"])
+ merge_command(args)
+ check_pytorch_weights(path, model)
+
+
+if __name__ == "__main__":
+ # Note this test requires at least two accelerators!
+ model, plugin, accelerator = setup()
+ if accelerator.num_processes > 1:
+ try:
+ # Initial setup for things
+ out_path = Path("test_merge_weights_fsdp_weights")
+ if not out_path.exists():
+ out_path.mkdir(parents=True, exist_ok=True)
+
+ # Train briefly once weights aren't the baseline
+ model = mock_training(accelerator, model)
+ accelerator.wait_for_everyone()
+
+ gc.collect() # Needed for some lingering refs after training
+ save_fsdp_model(plugin, accelerator, model, out_path)
+ accelerator.wait_for_everyone()
+
+ # Finally we can test
+ test_merge_weights_safetensors(model, out_path)
+ test_merge_weights_command_safetensors(model, out_path)
+ test_merge_weights_pytorch(model, out_path)
+ test_merge_weights_command_pytorch(model, out_path)
+ except Exception:
+ raise
+ finally:
+ # Cleanup in case of any failures
+ if accelerator.is_main_process:
+ shutil.rmtree(out_path)
+ accelerator.wait_for_everyone()
+ accelerator.end_training()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_notebook.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_notebook.py
new file mode 100644
index 0000000000000000000000000000000000000000..267c11b50b22250e781f94e3643b8895cc6aeb02
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_notebook.py
@@ -0,0 +1,118 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+"""
+Test file to ensure that in general certain situational setups for notebooks work.
+"""
+
+import os
+import time
+from multiprocessing import Queue
+
+from pytest import mark, raises
+from torch.distributed.elastic.multiprocessing.errors import ChildFailedError
+
+from accelerate import PartialState, notebook_launcher
+from accelerate.test_utils import require_bnb
+from accelerate.utils import is_bnb_available
+
+
+def basic_function():
+ # Just prints the PartialState
+ print(f"PartialState:\n{PartialState()}")
+
+
+def tough_nut_function(queue: Queue):
+ if queue.empty():
+ return
+ trial = queue.get()
+ if trial > 0:
+ queue.put(trial - 1)
+ raise RuntimeError("The nut hasn't cracked yet! Try again.")
+
+ print(f"PartialState:\n{PartialState()}")
+
+
+def bipolar_sleep_function(sleep_sec: int):
+ state = PartialState()
+ if state.process_index % 2 == 0:
+ raise RuntimeError("I'm an even process. I don't like to sleep.")
+ else:
+ time.sleep(sleep_sec)
+
+
+NUM_PROCESSES = int(os.environ.get("ACCELERATE_NUM_PROCESSES", 1))
+
+
+def test_can_initialize():
+ notebook_launcher(basic_function, (), num_processes=NUM_PROCESSES)
+
+
+@mark.skipif(NUM_PROCESSES < 2, reason="Need at least 2 processes to test static rendezvous backends")
+def test_static_rdzv_backend():
+ notebook_launcher(basic_function, (), num_processes=NUM_PROCESSES, rdzv_backend="static")
+
+
+@mark.skipif(NUM_PROCESSES < 2, reason="Need at least 2 processes to test c10d rendezvous backends")
+def test_c10d_rdzv_backend():
+ notebook_launcher(basic_function, (), num_processes=NUM_PROCESSES, rdzv_backend="c10d")
+
+
+@mark.skipif(NUM_PROCESSES < 2, reason="Need at least 2 processes to test fault tolerance")
+def test_fault_tolerant(max_restarts: int = 3):
+ queue = Queue()
+ queue.put(max_restarts)
+ notebook_launcher(tough_nut_function, (queue,), num_processes=NUM_PROCESSES, max_restarts=max_restarts)
+
+
+@mark.skipif(NUM_PROCESSES < 2, reason="Need at least 2 processes to test monitoring")
+def test_monitoring(monitor_interval: float = 0.01, sleep_sec: int = 100):
+ start_time = time.time()
+ with raises(ChildFailedError, match="I'm an even process. I don't like to sleep."):
+ notebook_launcher(
+ bipolar_sleep_function,
+ (sleep_sec,),
+ num_processes=NUM_PROCESSES,
+ monitor_interval=monitor_interval,
+ )
+ assert time.time() - start_time < sleep_sec, "Monitoring did not stop the process in time."
+
+
+@require_bnb
+def test_problematic_imports():
+ with raises(RuntimeError, match="Please keep these imports"):
+ import bitsandbytes as bnb # noqa: F401
+
+ notebook_launcher(basic_function, (), num_processes=NUM_PROCESSES)
+
+
+def main():
+ print("Test basic notebook can be ran")
+ test_can_initialize()
+ print("Test static rendezvous backend")
+ test_static_rdzv_backend()
+ print("Test c10d rendezvous backend")
+ test_c10d_rdzv_backend()
+ print("Test fault tolerant")
+ test_fault_tolerant()
+ print("Test monitoring")
+ test_monitoring()
+ if is_bnb_available():
+ print("Test problematic imports (bnb)")
+ test_problematic_imports()
+ if NUM_PROCESSES > 1:
+ PartialState().destroy_process_group()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_ops.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..f8f535d7b25a7bda527901787261591364545c09
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_ops.py
@@ -0,0 +1,181 @@
+#!/usr/bin/env python
+
+# Copyright 2023 The HuggingFace Team. All rights reserved.
+#
+# 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 torch
+
+from accelerate import PartialState
+from accelerate.test_utils.testing import assert_exception
+from accelerate.utils.dataclasses import DistributedType
+from accelerate.utils.operations import (
+ DistributedOperationException,
+ broadcast,
+ copy_tensor_to_devices,
+ gather,
+ gather_object,
+ pad_across_processes,
+ reduce,
+)
+
+
+def create_tensor(state):
+ return (torch.arange(state.num_processes) + 1.0 + (state.num_processes * state.process_index)).to(state.device)
+
+
+def test_gather(state):
+ tensor = create_tensor(state)
+ gathered_tensor = gather(tensor)
+ assert gathered_tensor.tolist() == list(range(1, state.num_processes**2 + 1))
+
+
+def test_gather_object(state):
+ # Gather objects in TorchXLA is not supported.
+ if state.distributed_type == DistributedType.XLA:
+ return
+ obj = [state.process_index]
+ gathered_obj = gather_object(obj)
+ assert len(gathered_obj) == state.num_processes, f"{gathered_obj}, {len(gathered_obj)} != {state.num_processes}"
+ assert gathered_obj == list(range(state.num_processes)), f"{gathered_obj} != {list(range(state.num_processes))}"
+
+
+def test_gather_non_contigous(state):
+ # Skip this test because the 'is_contiguous' function of XLA tensor always returns True.
+ if state.distributed_type == DistributedType.XLA:
+ return
+
+ # Create a non-contiguous tensor (enforce non-contiguity after device memory allocation)
+ tensor = torch.arange(12, device=state.device).view(4, 3).t()
+ assert not tensor.is_contiguous()
+ # Shouldn't error out
+ _ = gather(tensor)
+
+
+def test_broadcast(state):
+ tensor = create_tensor(state)
+ broadcasted_tensor = broadcast(tensor)
+ assert broadcasted_tensor.shape == torch.Size([state.num_processes])
+ assert broadcasted_tensor.tolist() == list(range(1, state.num_processes + 1))
+
+
+def test_pad_across_processes(state):
+ # We need to pad the tensor with one more element if we are the main process
+ # to ensure that we can pad
+ if state.is_main_process:
+ tensor = torch.arange(state.num_processes + 1).to(state.device)
+ else:
+ tensor = torch.arange(state.num_processes).to(state.device)
+ padded_tensor = pad_across_processes(tensor)
+ assert padded_tensor.shape == torch.Size([state.num_processes + 1])
+ if not state.is_main_process:
+ assert padded_tensor.tolist() == list(range(0, state.num_processes)) + [0]
+
+
+def test_reduce_sum(state):
+ # For now runs on only two processes
+ if state.num_processes != 2:
+ return
+ tensor = create_tensor(state)
+ reduced_tensor = reduce(tensor, "sum")
+ truth_tensor = torch.tensor([4.0, 6]).to(state.device)
+ assert torch.allclose(reduced_tensor, truth_tensor), f"{reduced_tensor} != {truth_tensor}"
+
+
+def test_reduce_mean(state):
+ # For now runs on only two processes
+ if state.num_processes != 2:
+ return
+ tensor = create_tensor(state)
+ reduced_tensor = reduce(tensor, "mean")
+ truth_tensor = torch.tensor([2.0, 3]).to(state.device)
+ assert torch.allclose(reduced_tensor, truth_tensor), f"{reduced_tensor} != {truth_tensor}"
+
+
+def test_op_checker(state):
+ # Must be in a distributed state, and gathering is currently not supported in TorchXLA.
+ if state.distributed_type in [DistributedType.NO, DistributedType.XLA]:
+ return
+ state.debug = True
+ # `pad_across_processes`
+ if state.process_index == 0:
+ data = {"tensor": torch.tensor([[0.0, 1, 2, 3, 4]]).to(state.device)}
+ else:
+ data = {"tensor": torch.tensor([[[0.0, 1, 2, 3, 4, 5]]]).to(state.device)}
+
+ with assert_exception(DistributedOperationException):
+ pad_across_processes(data, dim=0)
+
+ # `reduce`
+ if state.process_index == 0:
+ data = {"tensor": torch.tensor([[0.0, 1, 2, 3, 4]]).to(state.device)}
+ else:
+ data = {"tensor": torch.tensor([[[0.0, 1, 2, 3, 4], [5, 6, 7, 8, 9]]]).to(state.device)}
+
+ with assert_exception(DistributedOperationException):
+ reduce(data)
+
+ # `broadcast`
+ if state.process_index == 0:
+ data = {"tensor": torch.tensor([[0.0, 1, 2, 3, 4]]).to(state.device)}
+ else:
+ data = {"tensor": torch.tensor([[[0.0, 1, 2, 3, 4], [5, 6, 7, 8, 9]]]).to(state.device)}
+
+ with assert_exception(DistributedOperationException):
+ broadcast(data)
+
+ state.debug = False
+
+
+def test_copy_tensor_to_devices(state):
+ if state.distributed_type not in [DistributedType.MULTI_GPU, DistributedType.XLA]:
+ return
+ if state.is_main_process:
+ tensor = torch.tensor([1, 2, 3], dtype=torch.int).to(state.device)
+ else:
+ tensor = None
+ tensor = copy_tensor_to_devices(tensor)
+ assert torch.allclose(tensor, torch.tensor([1, 2, 3], dtype=torch.int, device=state.device))
+
+
+def _mp_fn(index):
+ # For xla_spawn (TPUs)
+ main()
+
+
+def main():
+ state = PartialState()
+ state.print(f"State: {state}")
+ state.print("testing gather")
+ test_gather(state)
+ state.print("testing gather_object")
+ test_gather_object(state)
+ state.print("testing gather non-contigous")
+ test_gather_non_contigous(state)
+ state.print("testing broadcast")
+ test_broadcast(state)
+ state.print("testing pad_across_processes")
+ test_pad_across_processes(state)
+ state.print("testing reduce_sum")
+ test_reduce_sum(state)
+ state.print("testing reduce_mean")
+ test_reduce_mean(state)
+ state.print("testing op_checker")
+ test_op_checker(state)
+ state.print("testing sending tensors across devices")
+ test_copy_tensor_to_devices(state)
+ state.destroy_process_group()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_script.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_script.py
new file mode 100644
index 0000000000000000000000000000000000000000..6912ba2fa850d6a14636dedcfd8e75f75b501d58
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_script.py
@@ -0,0 +1,901 @@
+#!/usr/bin/env python
+
+# Copyright 2021 The HuggingFace Team. All rights reserved.
+#
+# 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 contextlib
+import io
+import math
+import time
+from copy import deepcopy
+from pathlib import Path
+
+import numpy as np
+import torch
+from torch.utils.data import DataLoader, Dataset
+
+from accelerate import Accelerator
+from accelerate.data_loader import SeedableRandomSampler, prepare_data_loader
+from accelerate.state import AcceleratorState
+from accelerate.test_utils import RegressionDataset, are_the_same_tensors
+from accelerate.utils import (
+ DataLoaderConfiguration,
+ DistributedType,
+ gather,
+ is_bf16_available,
+ is_datasets_available,
+ is_fp16_available,
+ is_hpu_available,
+ is_ipex_available,
+ is_pytest_available,
+ is_xpu_available,
+ set_seed,
+ synchronize_rng_states,
+)
+
+
+# TODO: remove RegressionModel4XPU once ccl support empty buffer in broadcasting.
+if is_xpu_available():
+ from accelerate.test_utils import RegressionModel4XPU as RegressionModel
+else:
+ from accelerate.test_utils import RegressionModel
+
+if is_hpu_available():
+ ATOL = 1e-3
+ RTOL = 1e-3
+else:
+ ATOL = 1e-6
+ RTOL = 1e-6
+
+
+def generate_baseline_dataloader(train_set, generator, batch_size, use_seedable_sampler=False):
+ "Creates a dataloader that can also use the `SeedableRandomSampler`"
+ if use_seedable_sampler:
+ # The SeedableRandomSampler is needed during distributed setups
+ # for full reproducibility across processes with the `DataLoader`
+ sampler = SeedableRandomSampler(
+ generator=generator,
+ data_source=train_set,
+ num_samples=len(train_set),
+ )
+ return DataLoader(train_set, batch_size=batch_size, sampler=sampler)
+ else:
+ return DataLoader(train_set, batch_size=batch_size, shuffle=True, generator=generator)
+
+
+def print_main(state):
+ print(f"Printing from the main process {state.process_index}")
+
+
+def print_local_main(state):
+ print(f"Printing from the local main process {state.local_process_index}")
+
+
+def print_last(state):
+ print(f"Printing from the last process {state.process_index}")
+
+
+def print_on(state, process_idx):
+ print(f"Printing from process {process_idx}: {state.process_index}")
+
+
+def process_execution_check():
+ accelerator = Accelerator()
+ num_processes = accelerator.num_processes
+ # Test main_process_first context manager
+ path = Path("check_main_process_first.txt")
+ with accelerator.main_process_first():
+ if accelerator.is_main_process:
+ time.sleep(0.1) # ensure main process takes longest
+ with open(path, "a+") as f:
+ f.write("Currently in the main process\n")
+ else:
+ with open(path, "a+") as f:
+ f.write("Now on another process\n")
+ accelerator.wait_for_everyone()
+
+ if accelerator.is_main_process:
+ with open(path) as f:
+ text = "".join(f.readlines())
+ try:
+ assert text.startswith("Currently in the main process\n"), "Main process was not first"
+ if num_processes > 1:
+ assert text.endswith("Now on another process\n"), "Main process was not first"
+ assert text.count("Now on another process\n") == accelerator.num_processes - 1, (
+ f"Only wrote to file {text.count('Now on another process') + 1} times, not {accelerator.num_processes}"
+ )
+ except AssertionError:
+ path.unlink()
+ raise
+
+ if accelerator.is_main_process and path.exists():
+ path.unlink()
+ accelerator.wait_for_everyone()
+ # Test the decorators
+ f = io.StringIO()
+ with contextlib.redirect_stdout(f):
+ accelerator.on_main_process(print_main)(accelerator.state)
+ result = f.getvalue().rstrip()
+ if accelerator.is_main_process:
+ assert result == "Printing from the main process 0", f"{result} != Printing from the main process 0"
+ else:
+ assert f.getvalue().rstrip() == "", f'{result} != ""'
+ f.truncate(0)
+ f.seek(0)
+
+ with contextlib.redirect_stdout(f):
+ accelerator.on_local_main_process(print_local_main)(accelerator.state)
+ if accelerator.is_local_main_process:
+ assert f.getvalue().rstrip() == "Printing from the local main process 0"
+ else:
+ assert f.getvalue().rstrip() == ""
+ f.truncate(0)
+ f.seek(0)
+
+ with contextlib.redirect_stdout(f):
+ accelerator.on_last_process(print_last)(accelerator.state)
+ if accelerator.is_last_process:
+ assert f.getvalue().rstrip() == f"Printing from the last process {accelerator.state.num_processes - 1}"
+ else:
+ assert f.getvalue().rstrip() == ""
+ f.truncate(0)
+ f.seek(0)
+
+ for process_idx in range(num_processes):
+ with contextlib.redirect_stdout(f):
+ accelerator.on_process(print_on, process_index=process_idx)(accelerator.state, process_idx)
+ if accelerator.process_index == process_idx:
+ assert f.getvalue().rstrip() == f"Printing from process {process_idx}: {accelerator.process_index}"
+ else:
+ assert f.getvalue().rstrip() == ""
+ f.truncate(0)
+ f.seek(0)
+
+
+def init_state_check():
+ # Test we can instantiate this twice in a row.
+ state = AcceleratorState()
+ if state.local_process_index == 0:
+ print("Testing, testing. 1, 2, 3.")
+ print(state)
+
+
+def rng_sync_check():
+ state = AcceleratorState()
+ synchronize_rng_states(["torch"])
+ assert are_the_same_tensors(torch.get_rng_state()), "RNG states improperly synchronized on CPU."
+ if state.distributed_type == DistributedType.MULTI_GPU:
+ synchronize_rng_states(["cuda"])
+ assert are_the_same_tensors(torch.cuda.get_rng_state()), "RNG states improperly synchronized on GPU."
+ elif state.distributed_type == DistributedType.MULTI_XPU:
+ synchronize_rng_states(["xpu"])
+ assert are_the_same_tensors(torch.xpu.get_rng_state()), "RNG states improperly synchronized on XPU."
+ generator = torch.Generator()
+ synchronize_rng_states(["generator"], generator=generator)
+ assert are_the_same_tensors(generator.get_state()), "RNG states improperly synchronized in generator."
+
+ if state.local_process_index == 0:
+ print("All rng are properly synched.")
+
+
+def dl_preparation_check():
+ state = AcceleratorState()
+ length = 32 * state.num_processes
+
+ dl = DataLoader(range(length), batch_size=8)
+ dl = prepare_data_loader(dl, state.device, state.num_processes, state.process_index, put_on_device=True)
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result)
+
+ assert torch.equal(result.cpu(), torch.arange(0, length).long()), "Wrong non-shuffled dataloader result."
+
+ dl = DataLoader(range(length), batch_size=8)
+ dl = prepare_data_loader(
+ dl,
+ state.device,
+ state.num_processes,
+ state.process_index,
+ put_on_device=True,
+ split_batches=True,
+ )
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result)
+ assert torch.equal(result.cpu(), torch.arange(0, length).long()), "Wrong non-shuffled dataloader result."
+
+ if state.process_index == 0:
+ print("Non-shuffled dataloader passing.")
+
+ dl = DataLoader(range(length), batch_size=8, shuffle=True)
+ dl = prepare_data_loader(dl, state.device, state.num_processes, state.process_index, put_on_device=True)
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result).tolist()
+ result.sort()
+ assert result == list(range(length)), "Wrong shuffled dataloader result."
+
+ dl = DataLoader(range(length), batch_size=8, shuffle=True)
+ dl = prepare_data_loader(
+ dl,
+ state.device,
+ state.num_processes,
+ state.process_index,
+ put_on_device=True,
+ split_batches=True,
+ )
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result).tolist()
+ result.sort()
+ assert result == list(range(length)), "Wrong shuffled dataloader result."
+
+ if state.local_process_index == 0:
+ print("Shuffled dataloader passing.")
+
+
+def central_dl_preparation_check():
+ state = AcceleratorState()
+ length = 32 * state.num_processes
+
+ dl = DataLoader(range(length), batch_size=8)
+ dl = prepare_data_loader(
+ dl, state.device, state.num_processes, state.process_index, put_on_device=True, dispatch_batches=True
+ )
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result)
+ assert torch.equal(result.cpu(), torch.arange(0, length).long()), "Wrong non-shuffled dataloader result."
+
+ dl = DataLoader(range(length), batch_size=8)
+ dl = prepare_data_loader(
+ dl,
+ state.device,
+ state.num_processes,
+ state.process_index,
+ put_on_device=True,
+ split_batches=True,
+ dispatch_batches=True,
+ )
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result)
+ assert torch.equal(result.cpu(), torch.arange(0, length).long()), "Wrong non-shuffled dataloader result."
+
+ if state.process_index == 0:
+ print("Non-shuffled central dataloader passing.")
+
+ dl = DataLoader(range(length), batch_size=8, shuffle=True)
+ dl = prepare_data_loader(
+ dl, state.device, state.num_processes, state.process_index, put_on_device=True, dispatch_batches=True
+ )
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result).tolist()
+ result.sort()
+ assert result == list(range(length)), "Wrong shuffled dataloader result."
+
+ dl = DataLoader(range(length), batch_size=8, shuffle=True)
+ dl = prepare_data_loader(
+ dl,
+ state.device,
+ state.num_processes,
+ state.process_index,
+ put_on_device=True,
+ split_batches=True,
+ dispatch_batches=True,
+ )
+ result = []
+ for batch in dl:
+ result.append(gather(batch))
+ result = torch.cat(result).tolist()
+ result.sort()
+ assert result == list(range(length)), "Wrong shuffled dataloader result."
+
+ if state.local_process_index == 0:
+ print("Shuffled central dataloader passing.")
+
+
+def custom_sampler_check():
+ state = AcceleratorState()
+
+ class CustomDataset(Dataset):
+ def __init__(self, data):
+ self.data = data
+
+ def __len__(self):
+ return len(self.data)
+
+ def __getitem__(self, index):
+ return self.data[index]
+
+ class CustomBatchSampler:
+ def __init__(self, dataset_length: int, batch_size: int, shuffle: bool = True):
+ self.batch_size = batch_size
+ self.data_index = np.arange(dataset_length)
+ self.shuffle = shuffle
+
+ def __iter__(self):
+ num_batches = len(self)
+ if self.shuffle:
+ index = np.random.permutation(self.data_index)
+ else:
+ index = self.data_index
+ output = np.array_split(index, num_batches)
+ yield from output
+
+ def __len__(self):
+ return math.ceil(len(self.data_index) / self.batch_size)
+
+ dataset = CustomDataset(range(32 * state.num_processes))
+ sampler = CustomBatchSampler(len(dataset), batch_size=8)
+ dl = DataLoader(dataset, batch_sampler=sampler)
+ dl = prepare_data_loader(dl, state.device, state.num_processes, state.process_index)
+ # We need just ensure that `dl.batch_sampler` (or `dl.batch_sampler.batch_sampler` is indeed the old batch sampler
+ if hasattr(dl.batch_sampler, "batch_sampler"):
+ assert isinstance(dl.batch_sampler.batch_sampler, CustomBatchSampler), (
+ "Custom sampler was changed after calling `prepare_data_loader`"
+ )
+ else:
+ assert isinstance(dl.batch_sampler, CustomBatchSampler), (
+ "Custom sampler was changed after calling `prepare_data_loader`"
+ )
+
+
+def check_seedable_sampler():
+ # Set seed
+ set_seed(42)
+ train_set = RegressionDataset(length=10, seed=42)
+ train_dl = DataLoader(train_set, batch_size=2, shuffle=True)
+
+ config = DataLoaderConfiguration(use_seedable_sampler=True)
+ accelerator = Accelerator(dataloader_config=config)
+ train_dl = accelerator.prepare(train_dl)
+ original_items = []
+ for _ in range(3):
+ for batch in train_dl:
+ original_items.append(batch["x"])
+ original_items = torch.cat(original_items)
+
+ # Set seed again and the epoch
+ set_seed(42)
+ train_dl.set_epoch(0)
+ new_items = []
+ for _ in range(3):
+ for batch in train_dl:
+ new_items.append(batch["x"])
+ new_items = torch.cat(new_items)
+ assert torch.allclose(original_items, new_items), "Did not obtain the same items with the same seed and epoch."
+
+
+def check_seedable_sampler_in_batch_sampler_shard():
+ set_seed(42)
+
+ config = DataLoaderConfiguration(use_seedable_sampler=True)
+ accelerator = Accelerator(dataloader_config=config)
+ assert accelerator.num_processes > 1, "This test requires more than one process."
+
+ dataloader = DataLoader(list(range(10)), batch_size=1, shuffle=True)
+ prepared_data_loader = prepare_data_loader(
+ dataloader=dataloader,
+ use_seedable_sampler=True,
+ )
+
+ target_sampler = prepared_data_loader.batch_sampler.batch_sampler.sampler
+ assert isinstance(target_sampler, SeedableRandomSampler), (
+ "Sampler in BatchSamplerShard is not SeedableRandomSampler."
+ )
+
+
+def check_seedable_sampler_with_data_seed():
+ # Set seed
+ set_seed(42)
+ data_seed = 42
+ train_set = RegressionDataset(length=10, seed=42)
+ train_dl = DataLoader(train_set, batch_size=2, shuffle=True)
+
+ config = DataLoaderConfiguration(use_seedable_sampler=True, data_seed=data_seed)
+ accelerator = Accelerator(dataloader_config=config)
+ prepared_dl = accelerator.prepare(train_dl)
+ original_items = []
+ for _ in range(3):
+ for batch in prepared_dl:
+ original_items.append(batch["x"])
+ original_items = torch.cat(original_items)
+
+ # Set new data seed
+ config.data_seed = 43
+ accelerator = Accelerator(dataloader_config=config)
+ prepared_dl = accelerator.prepare(train_dl)
+ new_items = []
+ for _ in range(3):
+ for batch in prepared_dl:
+ new_items.append(batch["x"])
+ new_items = torch.cat(new_items)
+ assert not torch.allclose(original_items, new_items), "Obtained the same items with different data seed."
+
+
+def mock_training(length, batch_size, generator, use_seedable_sampler=False):
+ set_seed(42)
+ generator.manual_seed(42)
+ train_set = RegressionDataset(length=length, seed=42)
+
+ train_dl = generate_baseline_dataloader(train_set, generator, batch_size, use_seedable_sampler)
+ model = RegressionModel()
+ optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
+ for epoch in range(3):
+ for batch in train_dl:
+ model.zero_grad()
+ output = model(batch["x"])
+ loss = torch.nn.functional.mse_loss(output, batch["y"])
+ loss.backward()
+ optimizer.step()
+ return train_set, model
+
+
+def training_check(use_seedable_sampler=False):
+ state = AcceleratorState()
+ generator = torch.Generator()
+ batch_size = 8
+ length = batch_size * 4 * state.num_processes
+
+ train_set, old_model = mock_training(length, batch_size * state.num_processes, generator, use_seedable_sampler)
+ assert are_the_same_tensors(old_model.a), "Did not obtain the same model on both processes."
+ assert are_the_same_tensors(old_model.b), "Did not obtain the same model on both processes."
+
+ accelerator = Accelerator()
+ train_dl = generate_baseline_dataloader(train_set, generator, batch_size, use_seedable_sampler)
+ model = RegressionModel()
+ optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
+
+ train_dl, model, optimizer = accelerator.prepare(train_dl, model, optimizer)
+ set_seed(42)
+ generator.manual_seed(42)
+ for _ in range(3):
+ for batch in train_dl:
+ model.zero_grad()
+ output = model(batch["x"])
+ loss = torch.nn.functional.mse_loss(output, batch["y"])
+ accelerator.backward(loss)
+ optimizer.step()
+
+ model = accelerator.unwrap_model(model).cpu()
+ torch.testing.assert_close(
+ old_model.a,
+ model.a,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+ torch.testing.assert_close(
+ old_model.b,
+ model.b,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+
+ accelerator.print("Training yielded the same results on one CPU or distributed setup with no batch split.")
+
+ dataloader_config = DataLoaderConfiguration(split_batches=True, use_seedable_sampler=use_seedable_sampler)
+ accelerator = Accelerator(dataloader_config=dataloader_config)
+ train_dl = generate_baseline_dataloader(
+ train_set, generator, batch_size * state.num_processes, use_seedable_sampler
+ )
+ model = RegressionModel()
+ optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
+
+ train_dl, model, optimizer = accelerator.prepare(train_dl, model, optimizer)
+ set_seed(42)
+ generator.manual_seed(42)
+ for _ in range(3):
+ for batch in train_dl:
+ model.zero_grad()
+ output = model(batch["x"])
+ loss = torch.nn.functional.mse_loss(output, batch["y"])
+ accelerator.backward(loss)
+ optimizer.step()
+
+ model = accelerator.unwrap_model(model).cpu()
+ torch.testing.assert_close(
+ old_model.a,
+ model.a,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+ torch.testing.assert_close(
+ old_model.b,
+ model.b,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+
+ accelerator.print("Training yielded the same results on one CPU or distributed setup with batch split.")
+
+ # FP32 wrapper check
+ if torch.cuda.is_available():
+ # Mostly a test that model.forward will have autocast when running unwrap_model(model, keep_fp32_wrapper=True)
+ print("Keep fp32 wrapper check.")
+ AcceleratorState._reset_state()
+ accelerator = Accelerator(mixed_precision="fp16")
+
+ model = torch.nn.Linear(2, 4)
+ model = accelerator.prepare(model)
+ model_with_fp32_wrapper = accelerator.unwrap_model(model, keep_fp32_wrapper=True)
+
+ # Run forward with fp16 as input.
+ # When the model is with mixed precision wrapper, no error will be raised.
+ input_tensor = torch.Tensor([1, 2]).to(dtype=torch.float16, device=accelerator.device)
+ output = model_with_fp32_wrapper(input_tensor)
+
+ # BF16 support
+ if is_bf16_available():
+ # Mostly a test that BF16 doesn't crash as the operation inside the model is not converted to BF16
+ print("BF16 training check.")
+ AcceleratorState._reset_state()
+ dataloader_config = DataLoaderConfiguration(use_seedable_sampler=use_seedable_sampler)
+ accelerator = Accelerator(mixed_precision="bf16", dataloader_config=dataloader_config)
+ train_dl = generate_baseline_dataloader(train_set, generator, batch_size, use_seedable_sampler)
+ model = RegressionModel()
+ optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
+
+ train_dl, model, optimizer = accelerator.prepare(train_dl, model, optimizer)
+ set_seed(42)
+ generator.manual_seed(42)
+ for _ in range(3):
+ for batch in train_dl:
+ model.zero_grad()
+ output = model(batch["x"])
+ loss = torch.nn.functional.mse_loss(output, batch["y"])
+ accelerator.backward(loss)
+ optimizer.step()
+
+ model = accelerator.unwrap_model(model).cpu()
+ torch.testing.assert_close(
+ old_model.a,
+ model.a,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+ torch.testing.assert_close(
+ old_model.b,
+ model.b,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+
+ # FP16 support (HPU fp16 model seems to be off by 10% from the CPU, which is a lot of numerical error)
+ if is_fp16_available() and not is_hpu_available():
+ # Mostly a test that FP16 doesn't crash as the operation inside the model is not converted to FP16
+ print("FP16 training check.")
+ AcceleratorState._reset_state()
+ dataloader_config = DataLoaderConfiguration(use_seedable_sampler=use_seedable_sampler)
+ accelerator = Accelerator(mixed_precision="fp16", dataloader_config=dataloader_config)
+ train_dl = generate_baseline_dataloader(train_set, generator, batch_size, use_seedable_sampler)
+ model = RegressionModel()
+ optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
+
+ train_dl, model, optimizer = accelerator.prepare(train_dl, model, optimizer)
+ set_seed(42)
+ generator.manual_seed(42)
+ for _ in range(3):
+ for batch in train_dl:
+ model.zero_grad()
+ output = model(batch["x"])
+ loss = torch.nn.functional.mse_loss(output, batch["y"])
+ accelerator.backward(loss)
+ optimizer.step()
+
+ model = accelerator.unwrap_model(model).cpu()
+ torch.testing.assert_close(
+ old_model.a,
+ model.a,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+ torch.testing.assert_close(
+ old_model.b,
+ model.b,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+
+ # IPEX support is only for CPU
+ if is_ipex_available():
+ print("ipex BF16 training check.")
+ AcceleratorState._reset_state()
+ dataloader_config = DataLoaderConfiguration(use_seedable_sampler=use_seedable_sampler)
+ accelerator = Accelerator(mixed_precision="bf16", cpu=True, dataloader_config=dataloader_config)
+ train_dl = generate_baseline_dataloader(train_set, generator, batch_size, use_seedable_sampler)
+ model = RegressionModel()
+ optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
+
+ train_dl, model, optimizer = accelerator.prepare(train_dl, model, optimizer)
+ set_seed(42)
+ generator.manual_seed(42)
+ for _ in range(3):
+ for batch in train_dl:
+ model.zero_grad()
+ output = model(batch["x"])
+ loss = torch.nn.functional.mse_loss(output, batch["y"])
+ accelerator.backward(loss)
+ optimizer.step()
+
+ model = accelerator.unwrap_model(model).cpu()
+ torch.testing.assert_close(
+ old_model.a,
+ model.a,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+ torch.testing.assert_close(
+ old_model.b,
+ model.b,
+ atol=ATOL,
+ rtol=RTOL,
+ msg=lambda msg: f"Did not obtain the same model on CPU or distributed training.\n{msg}",
+ )
+
+
+def test_split_between_processes_dataset(datasets_Dataset):
+ state = AcceleratorState()
+ data = datasets_Dataset.from_list([dict(k=v) for v in range(2 * state.num_processes)])
+ with state.split_between_processes(data, apply_padding=False) as results:
+ assert len(results) == 2, (
+ f"Each process did not have two items. Process index: {state.process_index}; Length: {len(results)}"
+ )
+
+ data = datasets_Dataset.from_list([dict(k=v) for v in range(2 * state.num_processes - 1)])
+ with state.split_between_processes(data, apply_padding=False) as results:
+ if state.is_last_process:
+ assert len(results) == 1, (
+ f"Last process did not receive a single item. Process index: {state.process_index}; Length: {len(results)}"
+ )
+ else:
+ assert len(results) == 2, (
+ f"One of the intermediate processes did not receive two items. Process index: {state.process_index}; Length: {len(results)}"
+ )
+
+ data = datasets_Dataset.from_list([dict(k=v) for v in range(2 * state.num_processes - 1)])
+ with state.split_between_processes(data, apply_padding=True) as results:
+ if state.num_processes == 1:
+ assert len(results) == 1, (
+ f"Single process did not receive a single item. Process index: {state.process_index}; Length: {len(results)}"
+ )
+ else:
+ assert len(results) == 2, (
+ f"Each process did not have two items. Process index: {state.process_index}; Length: {len(results)}"
+ )
+
+ state.wait_for_everyone()
+
+
+def test_split_between_processes_list():
+ state = AcceleratorState()
+ data = list(range(0, 2 * state.num_processes))
+ with state.split_between_processes(data) as results:
+ assert len(results) == 2, (
+ f"Each process did not have two items. Process index: {state.process_index}; Length: {len(results)}"
+ )
+
+ data = list(range(0, (3 * state.num_processes) - 1))
+ with state.split_between_processes(data, apply_padding=True) as results:
+ if state.is_last_process:
+ # Test that the last process gets the extra item(s)
+ num_samples_per_device = math.ceil(len(data) / state.num_processes)
+ assert len(results) == num_samples_per_device, (
+ f"Last process did not get the extra item(s). Process index: {state.process_index}; Length: {len(results)}"
+ )
+ state.wait_for_everyone()
+
+
+def test_split_between_processes_nested_dict():
+ state = AcceleratorState()
+ a = [1, 2, 3, 4, 5, 6, 7, 8]
+ b = ["a", "b", "c", "d", "e", "f", "g", "h"]
+ c = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8])
+ if state.num_processes in (1, 2, 4):
+ data = {"a": a, "b": b, "c": c}
+ data_copy = deepcopy(data)
+ with state.split_between_processes(data) as results:
+ if state.process_index == 0:
+ assert results["a"] == data_copy["a"][: 8 // state.num_processes]
+ elif state.num_processes == 2:
+ assert results["a"] == data_copy["a"][4:]
+ elif state.process_index == 3:
+ # We return a list each time
+ assert results["a"] == data_copy["a"][-2:], f"Expected: {data_copy['a'][-2]}, Actual: {results['a']}"
+ if state.process_index == 0:
+ assert results["b"] == data_copy["b"][: 8 // state.num_processes]
+ elif state.num_processes == 2:
+ assert results["b"] == data_copy["b"][4:]
+ elif state.process_index == 3:
+ assert results["b"] == data_copy["b"][-2:]
+ if state.process_index == 0:
+ assert torch.allclose(results["c"], data_copy["c"][: 8 // state.num_processes]), (
+ f"Did not obtain expected values on process 0, expected `{data['c'][: 8 // state.num_processes]}`, received: {results['c']}"
+ )
+ elif state.num_processes == 2:
+ assert torch.allclose(results["c"], data_copy["c"][4:]), (
+ f"Did not obtain expected values on process 2, expected `{data['c'][4:]}`, received: {results['c']}"
+ )
+ elif state.process_index == 3:
+ assert torch.allclose(results["c"], data_copy["c"][-2:]), (
+ f"Did not obtain expected values on process 4, expected `{data['c'][-2:]}`, received: {results['c']}"
+ )
+
+ state.wait_for_everyone()
+
+
+def test_split_between_processes_tensor():
+ state = AcceleratorState()
+ if state.num_processes > 1:
+ data = torch.tensor([[0, 1, 2, 3], [4, 5, 6, 7]]).to(state.device)
+ with state.split_between_processes(data) as results:
+ if state.process_index == 0:
+ expected = torch.tensor([[0, 1, 2, 3]]).to(state.device)
+ else:
+ expected = torch.tensor([[4, 5, 6, 7]]).to(state.device)
+ torch.testing.assert_close(results, expected)
+ state.wait_for_everyone()
+
+
+def test_split_between_processes_evenly():
+ state = AcceleratorState()
+ if state.num_processes in (1, 2, 4, 8):
+ data = list(range(17))
+ num_samples_per_process = len(data) // state.num_processes
+ num_extras = len(data) % state.num_processes
+ with state.split_between_processes(data) as results:
+ if state.process_index < num_extras:
+ assert len(results) == num_samples_per_process + 1, (
+ f"Each Process should have even elements. Expected: {num_samples_per_process + 1}, Actual: {len(results)}"
+ )
+ else:
+ assert len(results) == num_samples_per_process, (
+ f"Each Process should have even elements. Expected: {num_samples_per_process}, Actual: {len(results)}"
+ )
+ state.wait_for_everyone()
+
+
+def test_trigger():
+ accelerator = Accelerator()
+ # should start with being false
+ assert accelerator.check_trigger() is False
+
+ # set a breakpoint on the main process
+ if accelerator.is_main_process:
+ accelerator.set_trigger()
+
+ # check it's been activated across all processes
+ # calls `all_reduce` and triggers a sync
+ assert accelerator.check_trigger() is True
+
+ # check it's been reset after the sync
+ assert accelerator.check_trigger() is False
+
+
+def test_reinstantiated_state():
+ import pytest
+
+ AcceleratorState._reset_state()
+ simple_model = torch.nn.Linear(1, 1)
+ # First define an accelerator
+ accelerator = Accelerator()
+ # Then call `reset_state`, breaking the state existing in the accelerator
+ AcceleratorState._reset_state()
+ # Now try and prepare a simple model, should raise the custom error early
+ with pytest.raises(AttributeError) as cm:
+ accelerator.prepare(simple_model)
+ assert "`AcceleratorState` object has no attribute" in str(cm.value.args[0])
+ assert "This happens if `AcceleratorState._reset_state()`" in str(cm.value.args[0])
+
+
+def main():
+ accelerator = Accelerator()
+ state = accelerator.state
+ if state.local_process_index == 0:
+ print("**Initialization**")
+ init_state_check()
+ state.wait_for_everyone()
+
+ if state.distributed_type == DistributedType.MULTI_GPU:
+ num_processes_per_node = torch.cuda.device_count()
+ else:
+ num_processes_per_node = state.num_processes
+
+ # We only run this test on non-multinode
+ if num_processes_per_node == state.num_processes:
+ if state.process_index == 0:
+ print("\n**Test process execution**")
+ process_execution_check()
+
+ if state.process_index == 0:
+ print("\n**Test split between processes as a list**")
+ test_split_between_processes_list()
+
+ if state.process_index == 0:
+ print("\n**Test split between processes as a dict**")
+ test_split_between_processes_nested_dict()
+
+ if state.process_index == 0:
+ print("\n**Test split between processes as a tensor**")
+ test_split_between_processes_tensor()
+
+ if state.process_index == 0:
+ print("\n**Test split between processes evenly**")
+ test_split_between_processes_evenly()
+
+ if state.process_index == 0:
+ print("\n**Test split between processes as a datasets.Dataset**")
+ if is_datasets_available():
+ from datasets import Dataset as datasets_Dataset
+
+ test_split_between_processes_dataset(datasets_Dataset)
+ else:
+ print("Skipped because Hugging Face datasets is not available")
+
+ if state.local_process_index == 0:
+ print("\n**Test random number generator synchronization**")
+ rng_sync_check()
+
+ if state.local_process_index == 0:
+ print("\n**DataLoader integration test**")
+ dl_preparation_check()
+ if state.distributed_type != DistributedType.XLA:
+ central_dl_preparation_check()
+ custom_sampler_check()
+ check_seedable_sampler()
+ check_seedable_sampler_with_data_seed()
+
+ if state.num_processes > 1:
+ check_seedable_sampler_in_batch_sampler_shard()
+
+ # Trainings are not exactly the same in DeepSpeed and CPU mode
+ if state.distributed_type == DistributedType.DEEPSPEED:
+ return
+
+ if state.local_process_index == 0:
+ print("\n**Training integration test**")
+ training_check(use_seedable_sampler=False)
+ training_check(use_seedable_sampler=True)
+
+ if state.local_process_index == 0:
+ print("\n**Breakpoint trigger test**")
+ test_trigger()
+
+ if is_pytest_available():
+ if state.local_process_index == 0:
+ print("\n**Test reinstantiated state**")
+ test_reinstantiated_state()
+
+ state.destroy_process_group()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_sync.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_sync.py
new file mode 100644
index 0000000000000000000000000000000000000000..44e1ecc1d59c5691f284282fb8cd2259c8a70658
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/scripts/test_sync.py
@@ -0,0 +1,410 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+from copy import deepcopy
+
+import torch
+import torch.nn.functional as F
+from torch.optim import AdamW
+from torch.optim.lr_scheduler import LambdaLR
+from torch.utils.data import DataLoader
+
+from accelerate.accelerator import Accelerator, DataLoaderConfiguration, GradientAccumulationPlugin
+from accelerate.state import GradientState
+from accelerate.test_utils import RegressionDataset, RegressionModel
+from accelerate.utils import DistributedType, set_seed
+
+
+def check_model_parameters(model_a, model_b, did_step, iteration, **kwargs):
+ for param, grad_param in zip(model_a.parameters(), model_b.parameters()):
+ if not param.requires_grad:
+ continue
+ if not did_step:
+ # Grads should not be in sync
+ assert torch.allclose(param.grad, grad_param.grad, **kwargs) is False, (
+ f"Gradients in sync when they should not be at iteration {iteration}:\nmodel_a grad ({param.grad}) == model_b grad ({grad_param.grad})"
+ )
+ else:
+ # Grads should be in sync
+ assert torch.allclose(param.grad, grad_param.grad, **kwargs) is True, (
+ f"Gradients not in sync when they should be at iteration {iteration}:\nmodel_a grad ({param.grad}) != model_b grad ({grad_param.grad})"
+ )
+
+
+def step_model(model, input, target, accelerator, do_backward=True):
+ model.train()
+ output = model(input)
+ loss = F.mse_loss(output, target.to(output.device))
+ if not do_backward:
+ loss /= accelerator.gradient_accumulation_steps
+ loss.backward()
+ else:
+ accelerator.backward(loss)
+
+
+def get_training_setup(accelerator, sched=False):
+ "Returns everything needed to perform basic training"
+ set_seed(42)
+ model = RegressionModel()
+ ddp_model = deepcopy(model)
+ dset = RegressionDataset(length=80)
+ dataloader = DataLoader(dset, batch_size=16)
+ model.to(accelerator.device)
+ if sched:
+ opt = AdamW(params=model.parameters(), lr=1e-3)
+ ddp_opt = AdamW(params=ddp_model.parameters(), lr=1e-3)
+ sched = LambdaLR(opt, lr_lambda=lambda epoch: epoch**0.65)
+ ddp_sched = LambdaLR(ddp_opt, lr_lambda=lambda epoch: epoch**0.65)
+ # Make a copy of `model`
+ if sched:
+ ddp_model, ddp_opt, ddp_sched, dataloader = accelerator.prepare(ddp_model, ddp_opt, ddp_sched, dataloader)
+ else:
+ ddp_model, dataloader = accelerator.prepare(ddp_model, dataloader)
+ if sched:
+ return (model, opt, sched, dataloader, ddp_model, ddp_opt, ddp_sched)
+ return model, ddp_model, dataloader
+
+
+def test_noop_sync(accelerator):
+ # Test when on a single CPU or GPU that the context manager does nothing
+ model, ddp_model, dataloader = get_training_setup(accelerator)
+ # Use a single batch
+ ddp_input, ddp_target = next(iter(dataloader)).values()
+ for iteration in range(3):
+ # Gather the distributed inputs and targs for the base model
+ input, target = accelerator.gather((ddp_input, ddp_target))
+ input, target = input.to(accelerator.device), target.to(accelerator.device)
+ # Perform our initial ground truth step in non "DDP"
+ step_model(model, input, target, accelerator)
+ # Do "gradient accumulation" (noop)
+ if iteration % 2 == 0:
+ # Accumulate grads locally
+ with accelerator.no_sync(ddp_model):
+ step_model(ddp_model, ddp_input, ddp_target, accelerator)
+ else:
+ # Sync grads
+ step_model(ddp_model, ddp_input, ddp_target, accelerator)
+
+ # Since `no_sync` is a noop, `ddp_model` and `model` grads should always be in sync
+ check_model_parameters(model, ddp_model, True, iteration)
+ for param, ddp_param in zip(model.parameters(), ddp_model.parameters()):
+ if not param.requires_grad:
+ continue
+ assert torch.allclose(param.grad, ddp_param.grad), (
+ f"Gradients not in sync when they should be:\nModel grad ({param.grad}) != DDP grad ({ddp_param.grad})"
+ )
+
+ # Shuffle ddp_input on each iteration
+ torch.manual_seed(1337 + iteration)
+ ddp_input = ddp_input[torch.randperm(len(ddp_input))]
+
+
+def test_distributed_sync(accelerator):
+ # Test on distributed setup that context manager behaves properly
+ model, ddp_model, dataloader = get_training_setup(accelerator)
+ # Use a single batch
+ ddp_input, ddp_target = next(iter(dataloader)).values()
+ for iteration in range(3):
+ # Gather the distributed inputs and targs for the base model
+ input, target = accelerator.gather((ddp_input, ddp_target))
+ input, target = input.to(accelerator.device), target.to(accelerator.device)
+ # Perform our initial ground truth step in non "DDP"
+ step_model(model, input, target, accelerator)
+ # Do "gradient accumulation" (noop)
+ if iteration % 2 == 0:
+ # Accumulate grads locally
+ with accelerator.no_sync(ddp_model):
+ step_model(ddp_model, ddp_input, ddp_target, accelerator)
+ else:
+ # Sync grads
+ step_model(ddp_model, ddp_input, ddp_target, accelerator)
+
+ # DDP model and model should only be in sync when not (iteration % 2 == 0)
+ for param, ddp_param in zip(model.parameters(), ddp_model.parameters()):
+ if not param.requires_grad:
+ continue
+ if iteration % 2 == 0:
+ # Grads should not be in sync
+ assert torch.allclose(param.grad, ddp_param.grad) is False, (
+ f"Gradients in sync when they should not be:\nModel grad ({param.grad}) == DDP grad ({ddp_param.grad})"
+ )
+ else:
+ # Grads should be in sync
+ assert torch.allclose(param.grad, ddp_param.grad) is True, (
+ f"Gradients not in sync when they should be:\nModel grad ({param.grad}) != DDP grad ({ddp_param.grad})"
+ )
+
+ # Shuffle ddp_input on each iteration
+ torch.manual_seed(1337 + iteration)
+ ddp_input = ddp_input[torch.randperm(len(ddp_input))]
+
+
+def test_distributed_sync_multiple_fwd(accelerator):
+ # Test on distributed setup that context manager behaves properly when used with multiple forwards followed by multiple backwards
+ model, ddp_model, dataloader = get_training_setup(accelerator)
+ # Do multiple forwards
+ losses = []
+ num_iterations = 3
+ for iteration in range(num_iterations):
+ ddp_input, ddp_target = next(iter(dataloader)).values()
+
+ # Gather the distributed inputs and targs for the base model
+ input, target = accelerator.gather((ddp_input, ddp_target))
+ input, target = input.to(accelerator.device), target.to(accelerator.device)
+
+ # Perform our initial ground truth step in non "DDP"
+ step_model(model, input, target, accelerator)
+
+ # Accumulate grads locally
+ with accelerator.no_sync(ddp_model):
+ ddp_output = ddp_model(ddp_input)
+ loss = F.mse_loss(ddp_output, ddp_target.to(ddp_output.device))
+ losses.append(loss)
+
+ # Do multiple backwards and sync only at the last backward
+ for iteration in range(num_iterations):
+ loss = losses[iteration]
+
+ if iteration < num_iterations - 1:
+ # Accumulate grads locally
+ accelerator.backward(loss)
+
+ # DDP model and model should only be in sync after last backward
+ for param, ddp_param in zip(model.parameters(), ddp_model.parameters()):
+ if not param.requires_grad:
+ continue
+ # Grads should not be in sync
+ assert torch.allclose(param.grad, ddp_param.grad) is False, (
+ f"Gradients in sync when they should not be:\nModel grad ({param.grad}) == DDP grad ({ddp_param.grad})"
+ )
+
+ else:
+ # Sync grads if last backward
+ with accelerator.trigger_sync_in_backward(ddp_model):
+ accelerator.backward(loss)
+
+ # DDP model and model should only be in sync after last backward
+ for param, ddp_param in zip(model.parameters(), ddp_model.parameters()):
+ if not param.requires_grad:
+ continue
+ # Grads should be in sync
+ assert torch.allclose(param.grad, ddp_param.grad) is True, (
+ f"Gradients not in sync when they should be:\nModel grad ({param.grad}) != DDP grad ({ddp_param.grad})"
+ )
+
+
+def test_gradient_accumulation(split_batches=False, dispatch_batches=False, sync_each_batch=False):
+ gradient_accumulation_plugin = GradientAccumulationPlugin(num_steps=2, sync_each_batch=sync_each_batch)
+ dataloader_config = DataLoaderConfiguration(split_batches=split_batches, dispatch_batches=dispatch_batches)
+ accelerator = Accelerator(
+ dataloader_config=dataloader_config,
+ gradient_accumulation_plugin=gradient_accumulation_plugin,
+ )
+ # Test that context manager behaves properly
+ model, ddp_model, dataloader = get_training_setup(accelerator)
+ for iteration, batch in enumerate(dataloader):
+ ddp_input, ddp_target = batch.values()
+ # Gather the distributed inputs and targs for the base model
+ input, target = accelerator.gather((ddp_input, ddp_target))
+ input, target = input.to(accelerator.device), target.to(accelerator.device)
+ # Perform our initial ground truth step in non "DDP"
+ step_model(model, input, target, accelerator, False)
+ # Do "gradient accumulation" (noop)
+ with accelerator.accumulate(ddp_model):
+ step_model(ddp_model, ddp_input, ddp_target, accelerator)
+
+ # DDP model and model should only be in sync when not (iteration % 2 == 0)
+ for param, ddp_param in zip(model.parameters(), ddp_model.parameters()):
+ if not param.requires_grad:
+ continue
+ if ((iteration + 1) % 2 == 0) or (iteration == len(dataloader) - 1) or sync_each_batch:
+ # Grads should be in sync
+ assert torch.allclose(param.grad, ddp_param.grad) is True, (
+ f"Gradients not in sync when they should be at iteration {iteration}:\nModel grad ({param.grad}) != DDP grad ({ddp_param.grad})"
+ )
+ else:
+ # Grads should not be in sync
+ assert torch.allclose(param.grad, ddp_param.grad) is False, (
+ f"Gradients in sync when they should not be at iteration {iteration}:\nModel grad ({param.grad}) == DDP grad ({ddp_param.grad})"
+ )
+
+ # Shuffle ddp_input on each iteration
+ torch.manual_seed(1337 + iteration)
+ ddp_input = ddp_input[torch.randperm(len(ddp_input))]
+ GradientState._reset_state()
+
+
+def test_gradient_accumulation_with_opt_and_scheduler(
+ split_batches=False, dispatch_batches=False, sync_each_batch=False
+):
+ gradient_accumulation_plugin = GradientAccumulationPlugin(num_steps=2, sync_each_batch=sync_each_batch)
+ dataloader_config = DataLoaderConfiguration(split_batches=split_batches, dispatch_batches=dispatch_batches)
+ accelerator = Accelerator(
+ dataloader_config=dataloader_config,
+ gradient_accumulation_plugin=gradient_accumulation_plugin,
+ )
+ # Test that context manager behaves properly
+ model, opt, sched, dataloader, ddp_model, ddp_opt, ddp_sched = get_training_setup(accelerator, True)
+ for iteration, batch in enumerate(dataloader):
+ ddp_input, ddp_target = batch.values()
+ # Gather the distributed inputs and targs for the base model
+ input, target = accelerator.gather((ddp_input, ddp_target))
+ input, target = input.to(accelerator.device), target.to(accelerator.device)
+ # Perform our initial ground truth step in non "DDP"
+ model.train()
+ ddp_model.train()
+ step_model(model, input, target, accelerator, False)
+ opt.step()
+
+ if ((iteration + 1) % 2 == 0) or ((iteration + 1) == len(dataloader)):
+ if split_batches:
+ sched.step()
+ else:
+ for _ in range(accelerator.num_processes):
+ sched.step()
+
+ # Perform gradient accumulation under wrapper
+ with accelerator.accumulate(ddp_model):
+ step_model(ddp_model, ddp_input, ddp_target, accelerator)
+ ddp_opt.step()
+ ddp_sched.step()
+
+ # Learning rates should be the same
+ assert opt.param_groups[0]["lr"] == ddp_opt.param_groups[0]["lr"], (
+ f"Learning rates found in each optimizer did not align\nopt: {opt.param_groups[0]['lr']}\nDDP opt: {ddp_opt.param_groups[0]['lr']}\n"
+ )
+ did_step = (((iteration + 1) % 2) == 0) or ((iteration + 1) == len(dataloader))
+ if accelerator.num_processes > 1:
+ check_model_parameters(
+ model,
+ ddp_model,
+ did_step or sync_each_batch, # syncs at each grad_accum interval of if sync_each_batch==True
+ iteration,
+ rtol=1e-3, # needs a relative tolerance due to roundoff errors
+ )
+
+ if did_step:
+ opt.zero_grad() # flush gradients every accum step
+ ddp_opt.zero_grad()
+
+ # Shuffle ddp_input on each iteration
+ torch.manual_seed(1337 + iteration)
+ GradientState._reset_state()
+
+
+def test_dataloader_break():
+ accelerator = Accelerator()
+ first_dset = RegressionDataset(length=80)
+ first_dataloader = DataLoader(first_dset, batch_size=16)
+ second_dset = RegressionDataset(length=96)
+ second_dataloader = DataLoader(second_dset, batch_size=16)
+ first_dataloader, second_dataloader = accelerator.prepare(first_dataloader, second_dataloader)
+
+ assert accelerator.gradient_state.active_dataloader is None
+ for iteration, _ in enumerate(first_dataloader):
+ assert id(accelerator.gradient_state.active_dataloader) == id(first_dataloader)
+ if iteration < len(first_dataloader) - 1:
+ assert not accelerator.gradient_state.end_of_dataloader
+ if iteration == 1:
+ for batch_num, _ in enumerate(second_dataloader):
+ assert id(accelerator.gradient_state.active_dataloader) == id(second_dataloader)
+ if batch_num < len(second_dataloader) - 1:
+ assert not accelerator.gradient_state.end_of_dataloader
+ else:
+ assert accelerator.gradient_state.end_of_dataloader
+ else:
+ assert accelerator.gradient_state.end_of_dataloader
+ assert accelerator.gradient_state.active_dataloader is None
+
+
+def main():
+ accelerator = Accelerator()
+ state = accelerator.state
+ if state.local_process_index == 0:
+ print("**Test `accumulate` gradient accumulation with dataloader break**")
+ if state.distributed_type != DistributedType.XLA:
+ test_dataloader_break()
+ if state.distributed_type == DistributedType.NO:
+ if state.local_process_index == 0:
+ print("**Test NOOP `no_sync` context manager**")
+ test_noop_sync(accelerator)
+ if state.distributed_type in (
+ DistributedType.MULTI_GPU,
+ DistributedType.MULTI_NPU,
+ DistributedType.MULTI_MLU,
+ DistributedType.MULTI_SDAA,
+ DistributedType.MULTI_MUSA,
+ DistributedType.MULTI_CPU,
+ DistributedType.MULTI_HPU,
+ ):
+ if state.local_process_index == 0:
+ print("**Test Distributed `no_sync` context manager**")
+ test_distributed_sync(accelerator)
+ if state.local_process_index == 0:
+ print("**Test Distributed `no_sync` context manager with multiple forwards**")
+ test_distributed_sync_multiple_fwd(accelerator)
+ if state.distributed_type in (
+ DistributedType.MULTI_GPU,
+ DistributedType.MULTI_NPU,
+ DistributedType.MULTI_MLU,
+ DistributedType.MULTI_SDAA,
+ DistributedType.MULTI_MUSA,
+ DistributedType.MULTI_HPU,
+ ):
+ for split_batch in [True, False]:
+ for dispatch_batches in [True, False]:
+ for sync_each_batch in [True, False]:
+ if state.local_process_index == 0:
+ print(
+ "**Test `accumulate` gradient accumulation, ",
+ f"`split_batches={split_batch}` and `dispatch_batches={dispatch_batches}` and `sync_each_batch={sync_each_batch}`**",
+ )
+ test_gradient_accumulation(split_batch, dispatch_batches, sync_each_batch)
+
+ # Currently will break on torch 2.0 +, need to investigate why
+ if state.local_process_index == 0:
+ print(
+ "**Test `accumulate` gradient accumulation with optimizer and scheduler, ",
+ "`split_batches=False`, `dispatch_batches=False`, `sync_each_batch=False`**",
+ )
+ test_gradient_accumulation_with_opt_and_scheduler()
+ if state.distributed_type in (
+ DistributedType.MULTI_GPU,
+ DistributedType.MULTI_NPU,
+ DistributedType.MULTI_MLU,
+ DistributedType.MULTI_SDAA,
+ DistributedType.MULTI_MUSA,
+ DistributedType.MULTI_HPU,
+ ):
+ for split_batch in [True, False]:
+ for dispatch_batches in [True, False]:
+ for sync_each_batch in [True, False]:
+ if not split_batch and not dispatch_batches and not sync_each_batch:
+ continue
+ if state.local_process_index == 0:
+ print(
+ "**Test `accumulate` gradient accumulation with optimizer and scheduler, ",
+ f"`split_batches={split_batch}` and `dispatch_batches={dispatch_batches}` and `sync_each_batch={sync_each_batch}`**",
+ )
+ test_gradient_accumulation_with_opt_and_scheduler(split_batch, dispatch_batches, sync_each_batch)
+ state.destroy_process_group()
+
+
+def _mp_fn(index):
+ # For xla_spawn (TPUs)
+ main()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/testing.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/testing.py
new file mode 100644
index 0000000000000000000000000000000000000000..79a981b8daaf7b702b2624969d11f270f2e0764e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/testing.py
@@ -0,0 +1,814 @@
+# Copyright 2021 The HuggingFace Team. All rights reserved.
+#
+# 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
+import inspect
+import io
+import os
+import shutil
+import subprocess
+import sys
+import tempfile
+import unittest
+from contextlib import contextmanager
+from functools import partial
+from pathlib import Path
+from typing import Union
+from unittest import mock
+
+import torch
+
+import accelerate
+
+from ..state import AcceleratorState
+from ..utils import (
+ check_cuda_fp8_capability,
+ gather,
+ is_bnb_available,
+ is_clearml_available,
+ is_comet_ml_available,
+ is_cuda_available,
+ is_datasets_available,
+ is_deepspeed_available,
+ is_dvclive_available,
+ is_fp8_available,
+ is_fp16_available,
+ is_habana_gaudi1,
+ is_hpu_available,
+ is_import_timer_available,
+ is_matplotlib_available,
+ is_mlflow_available,
+ is_mlu_available,
+ is_mps_available,
+ is_musa_available,
+ is_npu_available,
+ is_pandas_available,
+ is_pippy_available,
+ is_pytest_available,
+ is_schedulefree_available,
+ is_sdaa_available,
+ is_tensorboard_available,
+ is_timm_available,
+ is_torch_version,
+ is_torch_xla_available,
+ is_torchao_available,
+ is_torchdata_stateful_dataloader_available,
+ is_torchvision_available,
+ is_transformer_engine_available,
+ is_transformers_available,
+ is_triton_available,
+ is_wandb_available,
+ is_xpu_available,
+ str_to_bool,
+)
+
+
+def get_backend():
+ if is_torch_xla_available():
+ return "xla", torch.cuda.device_count(), torch.cuda.memory_allocated
+ elif is_cuda_available():
+ return "cuda", torch.cuda.device_count(), torch.cuda.memory_allocated
+ elif is_mps_available(min_version="2.0"):
+ return "mps", 1, torch.mps.current_allocated_memory
+ elif is_mps_available():
+ return "mps", 1, lambda: 0
+ elif is_mlu_available():
+ return "mlu", torch.mlu.device_count(), torch.mlu.memory_allocated
+ elif is_sdaa_available():
+ return "sdaa", torch.sdaa.device_count(), torch.sdaa.memory_allocated
+ elif is_musa_available():
+ return "musa", torch.musa.device_count(), torch.musa.memory_allocated
+ elif is_npu_available():
+ return "npu", torch.npu.device_count(), torch.npu.memory_allocated
+ elif is_xpu_available():
+ return "xpu", torch.xpu.device_count(), torch.xpu.memory_allocated
+ elif is_hpu_available():
+ return "hpu", torch.hpu.device_count(), torch.hpu.memory_allocated
+ else:
+ return "cpu", 1, lambda: 0
+
+
+torch_device, device_count, memory_allocated_func = get_backend()
+
+
+def get_launch_command(**kwargs) -> list:
+ """
+ Wraps around `kwargs` to help simplify launching from `subprocess`.
+
+ Example:
+ ```python
+ # returns ['accelerate', 'launch', '--num_processes=2', '--device_count=2']
+ get_launch_command(num_processes=2, device_count=2)
+ ```
+ """
+ command = ["accelerate", "launch"]
+ for k, v in kwargs.items():
+ if isinstance(v, bool) and v:
+ command.append(f"--{k}")
+ elif v is not None:
+ command.append(f"--{k}={v}")
+ return command
+
+
+DEFAULT_LAUNCH_COMMAND = get_launch_command(num_processes=device_count, monitor_interval=0.1)
+
+
+def parse_flag_from_env(key, default=False):
+ try:
+ value = os.environ[key]
+ except KeyError:
+ # KEY isn't set, default to `default`.
+ _value = default
+ else:
+ # KEY is set, convert it to True or False.
+ try:
+ _value = str_to_bool(value)
+ except ValueError:
+ # More values are supported, but let's keep the message simple.
+ raise ValueError(f"If set, {key} must be yes or no.")
+ return _value
+
+
+_run_slow_tests = parse_flag_from_env("RUN_SLOW", default=False)
+
+
+def skip(test_case):
+ "Decorator that skips a test unconditionally"
+ return unittest.skip("Test was skipped")(test_case)
+
+
+def slow(test_case):
+ """
+ Decorator marking a test as slow. Slow tests are skipped by default. Set the RUN_SLOW environment variable to a
+ truthy value to run them.
+ """
+ return unittest.skipUnless(_run_slow_tests, "test is slow")(test_case)
+
+
+def require_cpu(test_case):
+ """
+ Decorator marking a test that must be only ran on the CPU. These tests are skipped when a GPU is available.
+ """
+ return unittest.skipUnless(torch_device == "cpu", "test requires only a CPU")(test_case)
+
+
+def require_non_cpu(test_case):
+ """
+ Decorator marking a test that requires a hardware accelerator backend. These tests are skipped when there are no
+ hardware accelerator available.
+ """
+ return unittest.skipUnless(torch_device != "cpu", "test requires a GPU")(test_case)
+
+
+def require_cuda(test_case):
+ """
+ Decorator marking a test that requires CUDA. These tests are skipped when there are no GPU available or when
+ TorchXLA is available.
+ """
+ return unittest.skipUnless(is_cuda_available() and not is_torch_xla_available(), "test requires a GPU")(test_case)
+
+
+def require_cuda_or_hpu(test_case):
+ """
+ Decorator marking a test that requires CUDA or HPU. These tests are skipped when there are no GPU available or when
+ TorchXLA is available.
+ """
+ return unittest.skipUnless(
+ (is_cuda_available() and not is_torch_xla_available()) or is_hpu_available(), "test requires a GPU or HPU"
+ )(test_case)
+
+
+def require_xpu(test_case):
+ """
+ Decorator marking a test that requires XPU. These tests are skipped when there are no XPU available.
+ """
+ return unittest.skipUnless(is_xpu_available(), "test requires a XPU")(test_case)
+
+
+def require_cuda_or_xpu(test_case):
+ """
+ Decorator marking a test that requires CUDA or XPU. These tests are skipped when there are no GPU available or when
+ TorchXLA is available.
+ """
+ cuda_condition = is_cuda_available() and not is_torch_xla_available()
+ xpu_condition = is_xpu_available()
+ return unittest.skipUnless(cuda_condition or xpu_condition, "test requires a CUDA GPU or XPU")(test_case)
+
+
+def require_non_xpu(test_case):
+ """
+ Decorator marking a test that should be skipped for XPU.
+ """
+ return unittest.skipUnless(torch_device != "xpu", "test requires a non-XPU")(test_case)
+
+
+def require_non_hpu(test_case):
+ """
+ Decorator marking a test that should be skipped for HPU.
+ """
+ return unittest.skipUnless(torch_device != "hpu", "test requires a non-HPU")(test_case)
+
+
+def require_fp16(test_case):
+ """
+ Decorator marking a test that requires FP16. These tests are skipped when FP16 is not supported.
+ """
+
+ return unittest.skipUnless(is_fp16_available(), "test requires FP16 support")(test_case)
+
+
+def require_fp8(test_case):
+ """
+ Decorator marking a test that requires FP8. These tests are skipped when FP8 is not supported.
+ """
+
+ # is_fp8_available only checks for libraries
+ # ideally it should check for device capability as well
+ fp8_is_available = is_fp8_available()
+
+ if torch.cuda.is_available() and not check_cuda_fp8_capability():
+ fp8_is_available = False
+
+ if is_hpu_available() and is_habana_gaudi1():
+ fp8_is_available = False
+
+ return unittest.skipUnless(fp8_is_available, "test requires FP8 support")(test_case)
+
+
+def require_mlu(test_case):
+ """
+ Decorator marking a test that requires MLU. These tests are skipped when there are no MLU available.
+ """
+ return unittest.skipUnless(is_mlu_available(), "test require a MLU")(test_case)
+
+
+def require_sdaa(test_case):
+ """
+ Decorator marking a test that requires SDAA. These tests are skipped when there are no SDAA available.
+ """
+ return unittest.skipUnless(is_sdaa_available(), "test require a SDAA")(test_case)
+
+
+def require_musa(test_case):
+ """
+ Decorator marking a test that requires MUSA. These tests are skipped when there are no MUSA available.
+ """
+ return unittest.skipUnless(is_musa_available(), "test require a MUSA")(test_case)
+
+
+def require_npu(test_case):
+ """
+ Decorator marking a test that requires NPU. These tests are skipped when there are no NPU available.
+ """
+ return unittest.skipUnless(is_npu_available(), "test require a NPU")(test_case)
+
+
+def require_mps(test_case):
+ """
+ Decorator marking a test that requires MPS backend. These tests are skipped when torch doesn't support `mps`
+ backend.
+ """
+ return unittest.skipUnless(is_mps_available(), "test requires a `mps` backend support in `torch`")(test_case)
+
+
+def require_huggingface_suite(test_case):
+ """
+ Decorator marking a test that requires transformers and datasets. These tests are skipped when they are not.
+ """
+ return unittest.skipUnless(
+ is_transformers_available() and is_datasets_available(),
+ "test requires the Hugging Face suite",
+ )(test_case)
+
+
+def require_transformers(test_case):
+ """
+ Decorator marking a test that requires transformers. These tests are skipped when they are not.
+ """
+ return unittest.skipUnless(is_transformers_available(), "test requires the transformers library")(test_case)
+
+
+def require_timm(test_case):
+ """
+ Decorator marking a test that requires timm. These tests are skipped when they are not.
+ """
+ return unittest.skipUnless(is_timm_available(), "test requires the timm library")(test_case)
+
+
+def require_torchvision(test_case):
+ """
+ Decorator marking a test that requires torchvision. These tests are skipped when they are not.
+ """
+ return unittest.skipUnless(is_torchvision_available(), "test requires the torchvision library")(test_case)
+
+
+def require_triton(test_case):
+ """
+ Decorator marking a test that requires triton. These tests are skipped when they are not.
+ """
+ return unittest.skipUnless(is_triton_available(), "test requires the triton library")(test_case)
+
+
+def require_schedulefree(test_case):
+ """
+ Decorator marking a test that requires schedulefree. These tests are skipped when they are not.
+ """
+ return unittest.skipUnless(is_schedulefree_available(), "test requires the schedulefree library")(test_case)
+
+
+def require_bnb(test_case):
+ """
+ Decorator marking a test that requires bitsandbytes. These tests are skipped when they are not.
+ """
+ return unittest.skipUnless(is_bnb_available(), "test requires the bitsandbytes library")(test_case)
+
+
+def require_tpu(test_case):
+ """
+ Decorator marking a test that requires TPUs. These tests are skipped when there are no TPUs available.
+ """
+ return unittest.skipUnless(is_torch_xla_available(check_is_tpu=True), "test requires TPU")(test_case)
+
+
+def require_non_torch_xla(test_case):
+ """
+ Decorator marking a test as requiring an environment without TorchXLA. These tests are skipped when TorchXLA is
+ available.
+ """
+ return unittest.skipUnless(not is_torch_xla_available(), "test requires an env without TorchXLA")(test_case)
+
+
+def require_single_device(test_case):
+ """
+ Decorator marking a test that requires a single device. These tests are skipped when there is no hardware
+ accelerator available or number of devices is more than one.
+ """
+ return unittest.skipUnless(
+ torch_device != "cpu" and device_count == 1, "test requires a single device accelerator"
+ )(test_case)
+
+
+def require_single_gpu(test_case):
+ """
+ Decorator marking a test that requires CUDA on a single GPU. These tests are skipped when there are no GPU
+ available or number of GPUs is more than one.
+ """
+ return unittest.skipUnless(torch.cuda.device_count() == 1, "test requires a GPU")(test_case)
+
+
+def require_single_xpu(test_case):
+ """
+ Decorator marking a test that requires CUDA on a single XPU. These tests are skipped when there are no XPU
+ available or number of xPUs is more than one.
+ """
+ return unittest.skipUnless(torch.xpu.device_count() == 1, "test requires a XPU")(test_case)
+
+
+def require_multi_device(test_case):
+ """
+ Decorator marking a test that requires a multi-device setup. These tests are skipped on a machine without multiple
+ devices.
+ """
+ return unittest.skipUnless(device_count > 1, "test requires multiple hardware accelerators")(test_case)
+
+
+def require_multi_gpu(test_case):
+ """
+ Decorator marking a test that requires a multi-GPU setup. These tests are skipped on a machine without multiple
+ GPUs.
+ """
+ return unittest.skipUnless(torch.cuda.device_count() > 1, "test requires multiple GPUs")(test_case)
+
+
+def require_multi_xpu(test_case):
+ """
+ Decorator marking a test that requires a multi-XPU setup. These tests are skipped on a machine without multiple
+ XPUs.
+ """
+ return unittest.skipUnless(torch.xpu.device_count() > 1, "test requires multiple XPUs")(test_case)
+
+
+def require_multi_gpu_or_xpu(test_case):
+ """
+ Decorator marking a test that requires a multi-GPU setup. These tests are skipped on a machine without multiple
+ GPUs or XPUs.
+ """
+ return unittest.skipUnless(
+ (is_cuda_available() or is_xpu_available()) and device_count > 1, "test requires multiple GPUs or XPUs"
+ )(test_case)
+
+
+def require_deepspeed(test_case):
+ """
+ Decorator marking a test that requires DeepSpeed installed. These tests are skipped when DeepSpeed isn't installed
+ """
+ return unittest.skipUnless(is_deepspeed_available(), "test requires DeepSpeed")(test_case)
+
+
+def require_tp(test_case):
+ """
+ Decorator marking a test that requires TP installed. These tests are skipped when TP isn't installed
+ """
+ return unittest.skipUnless(is_torch_version(">=", "2.3.0"), "test requires torch version >= 2.3.0")(test_case)
+
+
+def require_torch_min_version(test_case=None, version=None):
+ """
+ Decorator marking that a test requires a particular torch version to be tested. These tests are skipped when an
+ installed torch version is less than the required one.
+ """
+ if test_case is None:
+ return partial(require_torch_min_version, version=version)
+ return unittest.skipUnless(is_torch_version(">=", version), f"test requires torch version >= {version}")(test_case)
+
+
+def require_tensorboard(test_case):
+ """
+ Decorator marking a test that requires tensorboard installed. These tests are skipped when tensorboard isn't
+ installed
+ """
+ return unittest.skipUnless(is_tensorboard_available(), "test requires Tensorboard")(test_case)
+
+
+def require_wandb(test_case):
+ """
+ Decorator marking a test that requires wandb installed. These tests are skipped when wandb isn't installed
+ """
+ return unittest.skipUnless(is_wandb_available(), "test requires wandb")(test_case)
+
+
+def require_comet_ml(test_case):
+ """
+ Decorator marking a test that requires comet_ml installed. These tests are skipped when comet_ml isn't installed
+ """
+ return unittest.skipUnless(is_comet_ml_available(), "test requires comet_ml")(test_case)
+
+
+def require_clearml(test_case):
+ """
+ Decorator marking a test that requires clearml installed. These tests are skipped when clearml isn't installed
+ """
+ return unittest.skipUnless(is_clearml_available(), "test requires clearml")(test_case)
+
+
+def require_dvclive(test_case):
+ """
+ Decorator marking a test that requires dvclive installed. These tests are skipped when dvclive isn't installed
+ """
+ return unittest.skipUnless(is_dvclive_available(), "test requires dvclive")(test_case)
+
+
+def require_pandas(test_case):
+ """
+ Decorator marking a test that requires pandas installed. These tests are skipped when pandas isn't installed
+ """
+ return unittest.skipUnless(is_pandas_available(), "test requires pandas")(test_case)
+
+
+def require_mlflow(test_case):
+ """
+ Decorator marking a test that requires mlflow installed. These tests are skipped when mlflow isn't installed
+ """
+ return unittest.skipUnless(is_mlflow_available(), "test requires mlflow")(test_case)
+
+
+def require_pippy(test_case):
+ """
+ Decorator marking a test that requires pippy installed. These tests are skipped when pippy isn't installed It is
+ also checked if the test is running on a Gaudi1 device which doesn't support pippy.
+ """
+ return unittest.skipUnless(is_pippy_available() and not is_habana_gaudi1(), "test requires pippy")(test_case)
+
+
+def require_import_timer(test_case):
+ """
+ Decorator marking a test that requires tuna interpreter installed. These tests are skipped when tuna isn't
+ installed
+ """
+ return unittest.skipUnless(is_import_timer_available(), "test requires tuna interpreter")(test_case)
+
+
+def require_transformer_engine(test_case):
+ """
+ Decorator marking a test that requires transformers engine installed. These tests are skipped when transformers
+ engine isn't installed
+ """
+ return unittest.skipUnless(is_transformer_engine_available(), "test requires transformers engine")(test_case)
+
+
+def require_torchao(test_case):
+ """
+ Decorator marking a test that requires torchao installed. These tests are skipped when torchao isn't installed
+ """
+ return unittest.skipUnless(is_torchao_available(), "test requires torchao")(test_case)
+
+
+def require_matplotlib(test_case):
+ """
+ Decorator marking a test that requires matplotlib installed. These tests are skipped when matplotlib isn't
+ installed
+ """
+ return unittest.skipUnless(is_matplotlib_available(), "test requires matplotlib")(test_case)
+
+
+_atleast_one_tracker_available = (
+ any([is_wandb_available(), is_tensorboard_available()]) and not is_comet_ml_available()
+)
+
+
+def require_trackers(test_case):
+ """
+ Decorator marking that a test requires at least one tracking library installed. These tests are skipped when none
+ are installed
+ """
+ return unittest.skipUnless(
+ _atleast_one_tracker_available,
+ "test requires at least one tracker to be available and for `comet_ml` to not be installed",
+ )(test_case)
+
+
+def require_torchdata_stateful_dataloader(test_case):
+ """
+ Decorator marking a test that requires torchdata.stateful_dataloader.
+
+ These tests are skipped when torchdata with stateful_dataloader module isn't installed.
+
+ """
+ return unittest.skipUnless(
+ is_torchdata_stateful_dataloader_available(), "test requires torchdata.stateful_dataloader"
+ )(test_case)
+
+
+def run_first(test_case):
+ """
+ Decorator marking a test with order(1). When pytest-order plugin is installed, tests marked with this decorator are
+ garanteed to run first.
+
+ This is especially useful in some test settings like on a Gaudi instance where a Gaudi device can only be used by a
+ single process at a time. So we make sure all tests that run in a subprocess are launched first, to avoid device
+ allocation conflicts.
+
+ If pytest is not installed, test will be returned as is.
+ """
+
+ if is_pytest_available():
+ import pytest
+
+ return pytest.mark.order(1)(test_case)
+ return test_case
+
+
+class TempDirTestCase(unittest.TestCase):
+ """
+ A TestCase class that keeps a single `tempfile.TemporaryDirectory` open for the duration of the class, wipes its
+ data at the start of a test, and then destroyes it at the end of the TestCase.
+
+ Useful for when a class or API requires a single constant folder throughout it's use, such as Weights and Biases
+
+ The temporary directory location will be stored in `self.tmpdir`
+ """
+
+ clear_on_setup = True
+
+ @classmethod
+ def setUpClass(cls):
+ "Creates a `tempfile.TemporaryDirectory` and stores it in `cls.tmpdir`"
+ cls.tmpdir = Path(tempfile.mkdtemp())
+
+ @classmethod
+ def tearDownClass(cls):
+ "Remove `cls.tmpdir` after test suite has finished"
+ if os.path.exists(cls.tmpdir):
+ shutil.rmtree(cls.tmpdir)
+
+ def setUp(self):
+ "Destroy all contents in `self.tmpdir`, but not `self.tmpdir`"
+ if self.clear_on_setup:
+ for path in self.tmpdir.glob("**/*"):
+ if path.is_file():
+ path.unlink()
+ elif path.is_dir():
+ shutil.rmtree(path)
+
+
+class AccelerateTestCase(unittest.TestCase):
+ """
+ A TestCase class that will reset the accelerator state at the end of every test. Every test that checks or utilizes
+ the `AcceleratorState` class should inherit from this to avoid silent failures due to state being shared between
+ tests.
+ """
+
+ def tearDown(self):
+ super().tearDown()
+ # Reset the state of the AcceleratorState singleton.
+ AcceleratorState._reset_state(True)
+
+
+class MockingTestCase(unittest.TestCase):
+ """
+ A TestCase class designed to dynamically add various mockers that should be used in every test, mimicking the
+ behavior of a class-wide mock when defining one normally will not do.
+
+ Useful when a mock requires specific information available only initialized after `TestCase.setUpClass`, such as
+ setting an environment variable with that information.
+
+ The `add_mocks` function should be ran at the end of a `TestCase`'s `setUp` function, after a call to
+ `super().setUp()` such as:
+ ```python
+ def setUp(self):
+ super().setUp()
+ mocks = mock.patch.dict(os.environ, {"SOME_ENV_VAR", "SOME_VALUE"})
+ self.add_mocks(mocks)
+ ```
+ """
+
+ def add_mocks(self, mocks: Union[mock.Mock, list[mock.Mock]]):
+ """
+ Add custom mocks for tests that should be repeated on each test. Should be called during
+ `MockingTestCase.setUp`, after `super().setUp()`.
+
+ Args:
+ mocks (`mock.Mock` or list of `mock.Mock`):
+ Mocks that should be added to the `TestCase` after `TestCase.setUpClass` has been run
+ """
+ self.mocks = mocks if isinstance(mocks, (tuple, list)) else [mocks]
+ for m in self.mocks:
+ m.start()
+ self.addCleanup(m.stop)
+
+
+def are_the_same_tensors(tensor):
+ state = AcceleratorState()
+ tensor = tensor[None].clone().to(state.device)
+ tensors = gather(tensor).cpu()
+ tensor = tensor[0].cpu()
+ for i in range(tensors.shape[0]):
+ if not torch.equal(tensors[i], tensor):
+ return False
+ return True
+
+
+class _RunOutput:
+ def __init__(self, returncode, stdout, stderr):
+ self.returncode = returncode
+ self.stdout = stdout
+ self.stderr = stderr
+
+
+async def _read_stream(stream, callback):
+ while True:
+ line = await stream.readline()
+ if line:
+ callback(line)
+ else:
+ break
+
+
+async def _stream_subprocess(cmd, env=None, stdin=None, timeout=None, quiet=False, echo=False) -> _RunOutput:
+ if echo:
+ print("\nRunning: ", " ".join(cmd))
+
+ p = await asyncio.create_subprocess_exec(
+ cmd[0],
+ *cmd[1:],
+ stdin=stdin,
+ stdout=asyncio.subprocess.PIPE,
+ stderr=asyncio.subprocess.PIPE,
+ env=env,
+ )
+
+ # note: there is a warning for a possible deadlock when using `wait` with huge amounts of data in the pipe
+ # https://docs.python.org/3/library/asyncio-subprocess.html#asyncio.asyncio.subprocess.Process.wait
+ #
+ # If it starts hanging, will need to switch to the following code. The problem is that no data
+ # will be seen until it's done and if it hangs for example there will be no debug info.
+ # out, err = await p.communicate()
+ # return _RunOutput(p.returncode, out, err)
+
+ out = []
+ err = []
+
+ def tee(line, sink, pipe, label=""):
+ line = line.decode("utf-8").rstrip()
+ sink.append(line)
+ if not quiet:
+ print(label, line, file=pipe)
+
+ # XXX: the timeout doesn't seem to make any difference here
+ await asyncio.wait(
+ [
+ asyncio.create_task(_read_stream(p.stdout, lambda l: tee(l, out, sys.stdout, label="stdout:"))),
+ asyncio.create_task(_read_stream(p.stderr, lambda l: tee(l, err, sys.stderr, label="stderr:"))),
+ ],
+ timeout=timeout,
+ )
+ return _RunOutput(await p.wait(), out, err)
+
+
+def execute_subprocess_async(cmd: list, env=None, stdin=None, timeout=180, quiet=False, echo=True) -> _RunOutput:
+ # Cast every path in `cmd` to a string
+ for i, c in enumerate(cmd):
+ if isinstance(c, Path):
+ cmd[i] = str(c)
+ loop = asyncio.get_event_loop()
+ result = loop.run_until_complete(
+ _stream_subprocess(cmd, env=env, stdin=stdin, timeout=timeout, quiet=quiet, echo=echo)
+ )
+
+ cmd_str = " ".join(cmd)
+ if result.returncode > 0:
+ stderr = "\n".join(result.stderr)
+ raise RuntimeError(
+ f"'{cmd_str}' failed with returncode {result.returncode}\n\n"
+ f"The combined stderr from workers follows:\n{stderr}"
+ )
+
+ return result
+
+
+class SubprocessCallException(Exception):
+ pass
+
+
+def run_command(command: list[str], return_stdout=False, env=None):
+ """
+ Runs `command` with `subprocess.check_output` and will potentially return the `stdout`. Will also properly capture
+ if an error occured while running `command`
+ """
+ # Cast every path in `command` to a string
+ for i, c in enumerate(command):
+ if isinstance(c, Path):
+ command[i] = str(c)
+ if env is None:
+ env = os.environ.copy()
+ try:
+ output = subprocess.check_output(command, stderr=subprocess.STDOUT, env=env)
+ if return_stdout:
+ if hasattr(output, "decode"):
+ output = output.decode("utf-8")
+ return output
+ except subprocess.CalledProcessError as e:
+ raise SubprocessCallException(
+ f"Command `{' '.join(command)}` failed with the following error:\n\n{e.output.decode()}"
+ ) from e
+
+
+def path_in_accelerate_package(*components: str) -> Path:
+ """
+ Get a path within the `accelerate` package's directory.
+
+ Args:
+ *components: Components of the path to join after the package directory.
+
+ Returns:
+ `Path`: The path to the requested file or directory.
+ """
+
+ accelerate_package_dir = Path(inspect.getfile(accelerate)).parent
+ return accelerate_package_dir.joinpath(*components)
+
+
+@contextmanager
+def assert_exception(exception_class: Exception, msg: str = None) -> bool:
+ """
+ Context manager to assert that the right `Exception` class was raised.
+
+ If `msg` is provided, will check that the message is contained in the raised exception.
+ """
+ was_ran = False
+ try:
+ yield
+ was_ran = True
+ except Exception as e:
+ assert isinstance(e, exception_class), f"Expected exception of type {exception_class} but got {type(e)}"
+ if msg is not None:
+ assert msg in str(e), f"Expected message '{msg}' to be in exception but got '{str(e)}'"
+ if was_ran:
+ raise AssertionError(f"Expected exception of type {exception_class} but ran without issue.")
+
+
+def capture_call_output(func, *args, **kwargs):
+ """
+ Takes in a `func` with `args` and `kwargs` and returns the captured stdout as a string
+ """
+ captured_output = io.StringIO()
+ original_stdout = sys.stdout
+ try:
+ sys.stdout = captured_output
+ func(*args, **kwargs)
+ except Exception as e:
+ raise e
+ finally:
+ sys.stdout = original_stdout
+ return captured_output.getvalue()
diff --git a/venv/lib/python3.11/site-packages/accelerate/test_utils/training.py b/venv/lib/python3.11/site-packages/accelerate/test_utils/training.py
new file mode 100644
index 0000000000000000000000000000000000000000..e71896c1f98bf47093d772cc77d7f27ba30c7ee8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/test_utils/training.py
@@ -0,0 +1,162 @@
+# Copyright 2021 The HuggingFace Team. All rights reserved.
+#
+# 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 numpy as np
+import torch
+from torch.utils.data import DataLoader
+
+from accelerate.utils.dataclasses import DistributedType
+
+
+class RegressionDataset:
+ def __init__(self, a=2, b=3, length=64, seed=None):
+ rng = np.random.default_rng(seed)
+ self.length = length
+ self.x = rng.normal(size=(length,)).astype(np.float32)
+ self.y = a * self.x + b + rng.normal(scale=0.1, size=(length,)).astype(np.float32)
+
+ def __len__(self):
+ return self.length
+
+ def __getitem__(self, i):
+ return {"x": self.x[i], "y": self.y[i]}
+
+
+class RegressionModel4XPU(torch.nn.Module):
+ def __init__(self, a=0, b=0, double_output=False):
+ super().__init__()
+ self.a = torch.nn.Parameter(torch.tensor([2, 3]).float())
+ self.b = torch.nn.Parameter(torch.tensor([2, 3]).float())
+ self.first_batch = True
+
+ def forward(self, x=None):
+ if self.first_batch:
+ print(f"Model dtype: {self.a.dtype}, {self.b.dtype}. Input dtype: {x.dtype}")
+ self.first_batch = False
+ return x * self.a[0] + self.b[0]
+
+
+class RegressionModel(torch.nn.Module):
+ def __init__(self, a=0, b=0, double_output=False):
+ super().__init__()
+ self.a = torch.nn.Parameter(torch.tensor(a).float())
+ self.b = torch.nn.Parameter(torch.tensor(b).float())
+ self.first_batch = True
+
+ def forward(self, x=None):
+ if self.first_batch:
+ print(f"Model dtype: {self.a.dtype}, {self.b.dtype}. Input dtype: {x.dtype}")
+ self.first_batch = False
+ return x * self.a + self.b
+
+
+def mocked_dataloaders(accelerator, batch_size: int = 16):
+ from datasets import load_dataset
+ from transformers import AutoTokenizer
+
+ tokenizer = AutoTokenizer.from_pretrained("bert-base-cased")
+ data_files = {"train": "tests/test_samples/MRPC/train.csv", "validation": "tests/test_samples/MRPC/dev.csv"}
+ datasets = load_dataset("csv", data_files=data_files)
+ label_list = datasets["train"].unique("label")
+
+ label_to_id = {v: i for i, v in enumerate(label_list)}
+
+ def tokenize_function(examples):
+ # max_length=None => use the model max length (it's actually the default)
+ outputs = tokenizer(
+ examples["sentence1"], examples["sentence2"], truncation=True, max_length=None, padding="max_length"
+ )
+ if "label" in examples:
+ outputs["labels"] = [label_to_id[l] for l in examples["label"]]
+ return outputs
+
+ # Apply the method we just defined to all the examples in all the splits of the dataset
+ tokenized_datasets = datasets.map(
+ tokenize_function,
+ batched=True,
+ remove_columns=["sentence1", "sentence2", "label"],
+ )
+
+ def collate_fn(examples):
+ # On TPU it's best to pad everything to the same length or training will be very slow.
+ if accelerator.distributed_type == DistributedType.XLA:
+ return tokenizer.pad(examples, padding="max_length", max_length=128, return_tensors="pt")
+ return tokenizer.pad(examples, padding="longest", return_tensors="pt")
+
+ # Instantiate dataloaders.
+ train_dataloader = DataLoader(tokenized_datasets["train"], shuffle=True, collate_fn=collate_fn, batch_size=2)
+ eval_dataloader = DataLoader(tokenized_datasets["validation"], shuffle=False, collate_fn=collate_fn, batch_size=1)
+
+ return train_dataloader, eval_dataloader
+
+
+def mocked_dataloaders_for_autoregressive_models(accelerator, batch_size: int = 16):
+ from datasets import load_dataset
+ from transformers import AutoTokenizer
+
+ tokenizer = AutoTokenizer.from_pretrained("HuggingFaceTB/SmolLM-360M")
+ tokenizer.pad_token = tokenizer.eos_token
+
+ data_files = {"train": "tests/test_samples/MRPC/train.csv", "validation": "tests/test_samples/MRPC/dev.csv"}
+ datasets = load_dataset("csv", data_files=data_files)
+
+ def tokenize_function(examples):
+ # max_length=None => use the model max length (it's actually the default)
+ outputs = tokenizer(examples["sentence1"], truncation=True, max_length=None, return_attention_mask=False)
+ return outputs
+
+ # Apply the method we just defined to all the examples in all the splits of the dataset
+ # starting with the main process first:
+ with accelerator.main_process_first():
+ tokenized_datasets = datasets.map(
+ tokenize_function,
+ batched=True,
+ remove_columns=["sentence1", "sentence2", "label"],
+ )
+
+ def collate_fn(examples):
+ # On TPU it's best to pad everything to the same length or training will be very slow.
+ max_length = (
+ 128
+ if accelerator.distributed_type == DistributedType.XLA
+ else max([len(e["input_ids"]) for e in examples])
+ )
+ # When using mixed precision we want round multiples of 8/16
+ if accelerator.mixed_precision == "fp8":
+ pad_to_multiple_of = 16
+ elif accelerator.mixed_precision != "no":
+ pad_to_multiple_of = 8
+ else:
+ pad_to_multiple_of = None
+
+ batch = tokenizer.pad(
+ examples,
+ padding="max_length",
+ max_length=max_length + 1,
+ pad_to_multiple_of=pad_to_multiple_of,
+ return_tensors="pt",
+ )
+
+ batch["labels"] = batch["input_ids"][:, 1:]
+ batch["input_ids"] = batch["input_ids"][:, :-1]
+
+ batch["labels"] = torch.where(batch["labels"] == tokenizer.pad_token_id, -100, batch["labels"])
+
+ return batch
+
+ # Instantiate dataloaders.
+ train_dataloader = DataLoader(tokenized_datasets["train"], shuffle=False, collate_fn=collate_fn, batch_size=2)
+ eval_dataloader = DataLoader(tokenized_datasets["validation"], shuffle=False, collate_fn=collate_fn, batch_size=1)
+
+ return train_dataloader, eval_dataloader
diff --git a/venv/lib/python3.11/site-packages/accelerate/tracking.py b/venv/lib/python3.11/site-packages/accelerate/tracking.py
new file mode 100644
index 0000000000000000000000000000000000000000..765f7adbf763b94ed900bfe22b7cd3f31cab0eb9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/tracking.py
@@ -0,0 +1,1089 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+# Expectation:
+# Provide a project dir name, then each type of logger gets stored in project/{`logging_dir`}
+
+import json
+import os
+import time
+from functools import wraps
+from typing import Any, Optional, Union
+
+import yaml
+from packaging import version
+
+from .logging import get_logger
+from .state import PartialState
+from .utils import (
+ LoggerType,
+ compare_versions,
+ is_aim_available,
+ is_clearml_available,
+ is_comet_ml_available,
+ is_dvclive_available,
+ is_mlflow_available,
+ is_tensorboard_available,
+ is_wandb_available,
+ listify,
+)
+
+
+_available_trackers = []
+
+if is_tensorboard_available():
+ _available_trackers.append(LoggerType.TENSORBOARD)
+
+if is_wandb_available():
+ _available_trackers.append(LoggerType.WANDB)
+
+if is_comet_ml_available():
+ _available_trackers.append(LoggerType.COMETML)
+
+if is_aim_available():
+ _available_trackers.append(LoggerType.AIM)
+
+if is_mlflow_available():
+ _available_trackers.append(LoggerType.MLFLOW)
+
+if is_clearml_available():
+ _available_trackers.append(LoggerType.CLEARML)
+
+if is_dvclive_available():
+ _available_trackers.append(LoggerType.DVCLIVE)
+
+logger = get_logger(__name__)
+
+
+def on_main_process(function):
+ """
+ Decorator to selectively run the decorated function on the main process only based on the `main_process_only`
+ attribute in a class.
+
+ Checks at function execution rather than initialization time, not triggering the initialization of the
+ `PartialState`.
+ """
+
+ @wraps(function)
+ def execute_on_main_process(self, *args, **kwargs):
+ if getattr(self, "main_process_only", False):
+ return PartialState().on_main_process(function)(self, *args, **kwargs)
+ else:
+ return function(self, *args, **kwargs)
+
+ return execute_on_main_process
+
+
+def get_available_trackers():
+ "Returns a list of all supported available trackers in the system"
+ return _available_trackers
+
+
+class GeneralTracker:
+ """
+ A base Tracker class to be used for all logging integration implementations.
+
+ Each function should take in `**kwargs` that will automatically be passed in from a base dictionary provided to
+ [`Accelerator`].
+
+ Should implement `name`, `requires_logging_directory`, and `tracker` properties such that:
+
+ `name` (`str`): String representation of the tracker class name, such as "TensorBoard" `requires_logging_directory`
+ (`bool`): Whether the logger requires a directory to store their logs. `tracker` (`object`): Should return internal
+ tracking mechanism used by a tracker class (such as the `run` for wandb)
+
+ Implementations can also include a `main_process_only` (`bool`) attribute to toggle if relevent logging, init, and
+ other functions should occur on the main process or across all processes (by default will use `True`)
+ """
+
+ main_process_only = True
+
+ def __init__(self, _blank=False):
+ if not _blank:
+ err = ""
+ if not hasattr(self, "name"):
+ err += "`name`"
+ if not hasattr(self, "requires_logging_directory"):
+ if len(err) > 0:
+ err += ", "
+ err += "`requires_logging_directory`"
+
+ # as tracker is a @property that relies on post-init
+ if "tracker" not in dir(self):
+ if len(err) > 0:
+ err += ", "
+ err += "`tracker`"
+ if len(err) > 0:
+ raise NotImplementedError(
+ f"The implementation for this tracker class is missing the following "
+ f"required attributes. Please define them in the class definition: "
+ f"{err}"
+ )
+
+ def store_init_configuration(self, values: dict):
+ """
+ Logs `values` as hyperparameters for the run. Implementations should use the experiment configuration
+ functionality of a tracking API.
+
+ Args:
+ values (Dictionary `str` to `bool`, `str`, `float` or `int`):
+ Values to be stored as initial hyperparameters as key-value pairs. The values need to have type `bool`,
+ `str`, `float`, `int`, or `None`.
+ """
+ pass
+
+ def log(self, values: dict, step: Optional[int], **kwargs):
+ """
+ Logs `values` to the current run. Base `log` implementations of a tracking API should go in here, along with
+ special behavior for the `step parameter.
+
+ Args:
+ values (Dictionary `str` to `str`, `float`, or `int`):
+ Values to be logged as key-value pairs. The values need to have type `str`, `float`, or `int`.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ """
+ pass
+
+ def finish(self):
+ """
+ Should run any finalizing functions within the tracking API. If the API should not have one, just don't
+ overwrite that method.
+ """
+ pass
+
+
+class TensorBoardTracker(GeneralTracker):
+ """
+ A `Tracker` class that supports `tensorboard`. Should be initialized at the start of your script.
+
+ Args:
+ run_name (`str`):
+ The name of the experiment run
+ logging_dir (`str`, `os.PathLike`):
+ Location for TensorBoard logs to be stored.
+ **kwargs (additional keyword arguments, *optional*):
+ Additional key word arguments passed along to the `tensorboard.SummaryWriter.__init__` method.
+ """
+
+ name = "tensorboard"
+ requires_logging_directory = True
+
+ @on_main_process
+ def __init__(self, run_name: str, logging_dir: Union[str, os.PathLike], **kwargs):
+ try:
+ from torch.utils import tensorboard
+ except ModuleNotFoundError:
+ import tensorboardX as tensorboard
+ super().__init__()
+ self.run_name = run_name
+ self.logging_dir = os.path.join(logging_dir, run_name)
+ self.writer = tensorboard.SummaryWriter(self.logging_dir, **kwargs)
+ logger.debug(f"Initialized TensorBoard project {self.run_name} logging to {self.logging_dir}")
+ logger.debug(
+ "Make sure to log any initial configurations with `self.store_init_configuration` before training!"
+ )
+
+ @property
+ def tracker(self):
+ return self.writer
+
+ @on_main_process
+ def store_init_configuration(self, values: dict):
+ """
+ Logs `values` as hyperparameters for the run. Should be run at the beginning of your experiment. Stores the
+ hyperparameters in a yaml file for future use.
+
+ Args:
+ values (Dictionary `str` to `bool`, `str`, `float` or `int`):
+ Values to be stored as initial hyperparameters as key-value pairs. The values need to have type `bool`,
+ `str`, `float`, `int`, or `None`.
+ """
+ self.writer.add_hparams(values, metric_dict={})
+ self.writer.flush()
+ project_run_name = time.time()
+ dir_name = os.path.join(self.logging_dir, str(project_run_name))
+ os.makedirs(dir_name, exist_ok=True)
+ with open(os.path.join(dir_name, "hparams.yml"), "w") as outfile:
+ try:
+ yaml.dump(values, outfile)
+ except yaml.representer.RepresenterError:
+ logger.error("Serialization to store hyperparameters failed")
+ raise
+ logger.debug("Stored initial configuration hyperparameters to TensorBoard and hparams yaml file")
+
+ @on_main_process
+ def log(self, values: dict, step: Optional[int] = None, **kwargs):
+ """
+ Logs `values` to the current run.
+
+ Args:
+ values (Dictionary `str` to `str`, `float`, `int` or `dict` of `str` to `float`/`int`):
+ Values to be logged as key-value pairs. The values need to have type `str`, `float`, `int` or `dict` of
+ `str` to `float`/`int`.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to either `SummaryWriter.add_scaler`,
+ `SummaryWriter.add_text`, or `SummaryWriter.add_scalers` method based on the contents of `values`.
+ """
+ values = listify(values)
+ for k, v in values.items():
+ if isinstance(v, (int, float)):
+ self.writer.add_scalar(k, v, global_step=step, **kwargs)
+ elif isinstance(v, str):
+ self.writer.add_text(k, v, global_step=step, **kwargs)
+ elif isinstance(v, dict):
+ self.writer.add_scalars(k, v, global_step=step, **kwargs)
+ self.writer.flush()
+ logger.debug("Successfully logged to TensorBoard")
+
+ @on_main_process
+ def log_images(self, values: dict, step: Optional[int], **kwargs):
+ """
+ Logs `images` to the current run.
+
+ Args:
+ values (Dictionary `str` to `List` of `np.ndarray` or `PIL.Image`):
+ Values to be logged as key-value pairs. The values need to have type `List` of `np.ndarray` or
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to the `SummaryWriter.add_image` method.
+ """
+ for k, v in values.items():
+ self.writer.add_images(k, v, global_step=step, **kwargs)
+ logger.debug("Successfully logged images to TensorBoard")
+
+ @on_main_process
+ def finish(self):
+ """
+ Closes `TensorBoard` writer
+ """
+ self.writer.close()
+ logger.debug("TensorBoard writer closed")
+
+
+class WandBTracker(GeneralTracker):
+ """
+ A `Tracker` class that supports `wandb`. Should be initialized at the start of your script.
+
+ Args:
+ run_name (`str`):
+ The name of the experiment run.
+ **kwargs (additional keyword arguments, *optional*):
+ Additional key word arguments passed along to the `wandb.init` method.
+ """
+
+ name = "wandb"
+ requires_logging_directory = False
+ main_process_only = False
+
+ @on_main_process
+ def __init__(self, run_name: str, **kwargs):
+ super().__init__()
+ self.run_name = run_name
+
+ import wandb
+
+ self.run = wandb.init(project=self.run_name, **kwargs)
+ logger.debug(f"Initialized WandB project {self.run_name}")
+ logger.debug(
+ "Make sure to log any initial configurations with `self.store_init_configuration` before training!"
+ )
+
+ @property
+ def tracker(self):
+ return self.run
+
+ @on_main_process
+ def store_init_configuration(self, values: dict):
+ """
+ Logs `values` as hyperparameters for the run. Should be run at the beginning of your experiment.
+
+ Args:
+ values (Dictionary `str` to `bool`, `str`, `float` or `int`):
+ Values to be stored as initial hyperparameters as key-value pairs. The values need to have type `bool`,
+ `str`, `float`, `int`, or `None`.
+ """
+ import wandb
+
+ wandb.config.update(values, allow_val_change=True)
+ logger.debug("Stored initial configuration hyperparameters to WandB")
+
+ @on_main_process
+ def log(self, values: dict, step: Optional[int] = None, **kwargs):
+ """
+ Logs `values` to the current run.
+
+ Args:
+ values (Dictionary `str` to `str`, `float`, `int` or `dict` of `str` to `float`/`int`):
+ Values to be logged as key-value pairs. The values need to have type `str`, `float`, `int` or `dict` of
+ `str` to `float`/`int`.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to the `wandb.log` method.
+ """
+ self.run.log(values, step=step, **kwargs)
+ logger.debug("Successfully logged to WandB")
+
+ @on_main_process
+ def log_images(self, values: dict, step: Optional[int] = None, **kwargs):
+ """
+ Logs `images` to the current run.
+
+ Args:
+ values (Dictionary `str` to `List` of `np.ndarray` or `PIL.Image`):
+ Values to be logged as key-value pairs. The values need to have type `List` of `np.ndarray` or
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to the `wandb.log` method.
+ """
+ import wandb
+
+ for k, v in values.items():
+ self.log({k: [wandb.Image(image) for image in v]}, step=step, **kwargs)
+ logger.debug("Successfully logged images to WandB")
+
+ @on_main_process
+ def log_table(
+ self,
+ table_name: str,
+ columns: list[str] = None,
+ data: list[list[Any]] = None,
+ dataframe: Any = None,
+ step: Optional[int] = None,
+ **kwargs,
+ ):
+ """
+ Log a Table containing any object type (text, image, audio, video, molecule, html, etc). Can be defined either
+ with `columns` and `data` or with `dataframe`.
+
+ Args:
+ table_name (`str`):
+ The name to give to the logged table on the wandb workspace
+ columns (list of `str`, *optional*):
+ The name of the columns on the table
+ data (List of List of Any data type, *optional*):
+ The data to be logged in the table
+ dataframe (Any data type, *optional*):
+ The data to be logged in the table
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ """
+ import wandb
+
+ values = {table_name: wandb.Table(columns=columns, data=data, dataframe=dataframe)}
+ self.log(values, step=step, **kwargs)
+
+ @on_main_process
+ def finish(self):
+ """
+ Closes `wandb` writer
+ """
+ self.run.finish()
+ logger.debug("WandB run closed")
+
+
+class CometMLTracker(GeneralTracker):
+ """
+ A `Tracker` class that supports `comet_ml`. Should be initialized at the start of your script.
+
+ API keys must be stored in a Comet config file.
+
+ Note:
+ For `comet_ml` versions < 3.41.0, additional keyword arguments are passed to `comet_ml.Experiment` instead:
+ https://www.comet.com/docs/v2/api-and-sdk/python-sdk/reference/Experiment/#comet_ml.Experiment.__init__
+
+ Args:
+ run_name (`str`):
+ The name of the experiment run.
+ **kwargs (additional keyword arguments, *optional*):
+ Additional key word arguments passed along to the `comet_ml.start` method:
+ https://www.comet.com/docs/v2/api-and-sdk/python-sdk/reference/start/
+ """
+
+ name = "comet_ml"
+ requires_logging_directory = False
+
+ @on_main_process
+ def __init__(self, run_name: str, **kwargs):
+ super().__init__()
+ self.run_name = run_name
+
+ import comet_ml
+
+ comet_version = version.parse(comet_ml.__version__)
+ if compare_versions(comet_version, ">=", "3.41.0"):
+ self.writer = comet_ml.start(project_name=run_name, **kwargs)
+ else:
+ logger.info("Update `comet_ml` (>=3.41.0) for experiment reuse and offline support.")
+ self.writer = comet_ml.Experiment(project_name=run_name, **kwargs)
+
+ logger.debug(f"Initialized CometML project {self.run_name}")
+ logger.debug(
+ "Make sure to log any initial configurations with `self.store_init_configuration` before training!"
+ )
+
+ @property
+ def tracker(self):
+ return self.writer
+
+ @on_main_process
+ def store_init_configuration(self, values: dict):
+ """
+ Logs `values` as hyperparameters for the run. Should be run at the beginning of your experiment.
+
+ Args:
+ values (Dictionary `str` to `bool`, `str`, `float` or `int`):
+ Values to be stored as initial hyperparameters as key-value pairs. The values need to have type `bool`,
+ `str`, `float`, `int`, or `None`.
+ """
+ self.writer.log_parameters(values)
+ logger.debug("Stored initial configuration hyperparameters to Comet")
+
+ @on_main_process
+ def log(self, values: dict, step: Optional[int] = None, **kwargs):
+ """
+ Logs `values` to the current run.
+
+ Args:
+ values (Dictionary `str` to `str`, `float`, `int` or `dict` of `str` to `float`/`int`):
+ Values to be logged as key-value pairs. The values need to have type `str`, `float`, `int` or `dict` of
+ `str` to `float`/`int`.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to either `Experiment.log_metric`, `Experiment.log_other`,
+ or `Experiment.log_metrics` method based on the contents of `values`.
+ """
+ if step is not None:
+ self.writer.set_step(step)
+ for k, v in values.items():
+ if isinstance(v, (int, float)):
+ self.writer.log_metric(k, v, step=step, **kwargs)
+ elif isinstance(v, str):
+ self.writer.log_other(k, v, **kwargs)
+ elif isinstance(v, dict):
+ self.writer.log_metrics(v, step=step, **kwargs)
+ logger.debug("Successfully logged to Comet")
+
+ @on_main_process
+ def finish(self):
+ """
+ Flush `comet-ml` writer
+ """
+ self.writer.end()
+ logger.debug("Comet run flushed")
+
+
+class AimTracker(GeneralTracker):
+ """
+ A `Tracker` class that supports `aim`. Should be initialized at the start of your script.
+
+ Args:
+ run_name (`str`):
+ The name of the experiment run.
+ **kwargs (additional keyword arguments, *optional*):
+ Additional key word arguments passed along to the `Run.__init__` method.
+ """
+
+ name = "aim"
+ requires_logging_directory = True
+
+ @on_main_process
+ def __init__(self, run_name: str, logging_dir: Optional[Union[str, os.PathLike]] = ".", **kwargs):
+ self.run_name = run_name
+
+ from aim import Run
+
+ self.writer = Run(repo=logging_dir, **kwargs)
+ self.writer.name = self.run_name
+ logger.debug(f"Initialized Aim project {self.run_name}")
+ logger.debug(
+ "Make sure to log any initial configurations with `self.store_init_configuration` before training!"
+ )
+
+ @property
+ def tracker(self):
+ return self.writer
+
+ @on_main_process
+ def store_init_configuration(self, values: dict):
+ """
+ Logs `values` as hyperparameters for the run. Should be run at the beginning of your experiment.
+
+ Args:
+ values (`dict`):
+ Values to be stored as initial hyperparameters as key-value pairs.
+ """
+ self.writer["hparams"] = values
+
+ @on_main_process
+ def log(self, values: dict, step: Optional[int], **kwargs):
+ """
+ Logs `values` to the current run.
+
+ Args:
+ values (`dict`):
+ Values to be logged as key-value pairs.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to the `Run.track` method.
+ """
+ # Note: replace this with the dictionary support when merged
+ for key, value in values.items():
+ self.writer.track(value, name=key, step=step, **kwargs)
+
+ @on_main_process
+ def log_images(self, values: dict, step: Optional[int] = None, kwargs: Optional[dict[str, dict]] = None):
+ """
+ Logs `images` to the current run.
+
+ Args:
+ values (`Dict[str, Union[np.ndarray, PIL.Image, Tuple[np.ndarray, str], Tuple[PIL.Image, str]]]`):
+ Values to be logged as key-value pairs. The values need to have type `np.ndarray` or PIL.Image. If a
+ tuple is provided, the first element should be the image and the second element should be the caption.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs (`Dict[str, dict]`):
+ Additional key word arguments passed along to the `Run.Image` and `Run.track` method specified by the
+ keys `aim_image` and `track`, respectively.
+ """
+ import aim
+
+ aim_image_kw = {}
+ track_kw = {}
+
+ if kwargs is not None:
+ aim_image_kw = kwargs.get("aim_image", {})
+ track_kw = kwargs.get("track", {})
+
+ for key, value in values.items():
+ if isinstance(value, tuple):
+ img, caption = value
+ else:
+ img, caption = value, ""
+ aim_image = aim.Image(img, caption=caption, **aim_image_kw)
+ self.writer.track(aim_image, name=key, step=step, **track_kw)
+
+ @on_main_process
+ def finish(self):
+ """
+ Closes `aim` writer
+ """
+ self.writer.close()
+
+
+class MLflowTracker(GeneralTracker):
+ """
+ A `Tracker` class that supports `mlflow`. Should be initialized at the start of your script.
+
+ Args:
+ experiment_name (`str`, *optional*):
+ Name of the experiment. Environment variable MLFLOW_EXPERIMENT_NAME has priority over this argument.
+ logging_dir (`str` or `os.PathLike`, defaults to `"."`):
+ Location for mlflow logs to be stored.
+ run_id (`str`, *optional*):
+ If specified, get the run with the specified UUID and log parameters and metrics under that run. The run’s
+ end time is unset and its status is set to running, but the run’s other attributes (source_version,
+ source_type, etc.) are not changed. Environment variable MLFLOW_RUN_ID has priority over this argument.
+ tags (`Dict[str, str]`, *optional*):
+ An optional `dict` of `str` keys and values, or a `str` dump from a `dict`, to set as tags on the run. If a
+ run is being resumed, these tags are set on the resumed run. If a new run is being created, these tags are
+ set on the new run. Environment variable MLFLOW_TAGS has priority over this argument.
+ nested_run (`bool`, *optional*, defaults to `False`):
+ Controls whether run is nested in parent run. True creates a nested run. Environment variable
+ MLFLOW_NESTED_RUN has priority over this argument.
+ run_name (`str`, *optional*):
+ Name of new run (stored as a mlflow.runName tag). Used only when `run_id` is unspecified.
+ description (`str`, *optional*):
+ An optional string that populates the description box of the run. If a run is being resumed, the
+ description is set on the resumed run. If a new run is being created, the description is set on the new
+ run.
+ """
+
+ name = "mlflow"
+ requires_logging_directory = False
+
+ @on_main_process
+ def __init__(
+ self,
+ experiment_name: str = None,
+ logging_dir: Optional[Union[str, os.PathLike]] = None,
+ run_id: Optional[str] = None,
+ tags: Optional[Union[dict[str, Any], str]] = None,
+ nested_run: Optional[bool] = False,
+ run_name: Optional[str] = None,
+ description: Optional[str] = None,
+ ):
+ experiment_name = os.environ.get("MLFLOW_EXPERIMENT_NAME", experiment_name)
+ run_id = os.environ.get("MLFLOW_RUN_ID", run_id)
+ tags = os.environ.get("MLFLOW_TAGS", tags)
+ if isinstance(tags, str):
+ tags = json.loads(tags)
+
+ nested_run = os.environ.get("MLFLOW_NESTED_RUN", nested_run)
+
+ import mlflow
+
+ exps = mlflow.search_experiments(filter_string=f"name = '{experiment_name}'")
+ if len(exps) > 0:
+ if len(exps) > 1:
+ logger.warning("Multiple experiments with the same name found. Using first one.")
+ experiment_id = exps[0].experiment_id
+ else:
+ experiment_id = mlflow.create_experiment(
+ name=experiment_name,
+ artifact_location=logging_dir,
+ tags=tags,
+ )
+
+ self.active_run = mlflow.start_run(
+ run_id=run_id,
+ experiment_id=experiment_id,
+ run_name=run_name,
+ nested=nested_run,
+ tags=tags,
+ description=description,
+ )
+
+ logger.debug(f"Initialized mlflow experiment {experiment_name}")
+ logger.debug(
+ "Make sure to log any initial configurations with `self.store_init_configuration` before training!"
+ )
+
+ @property
+ def tracker(self):
+ return self.active_run
+
+ @on_main_process
+ def store_init_configuration(self, values: dict):
+ """
+ Logs `values` as hyperparameters for the run. Should be run at the beginning of your experiment.
+
+ Args:
+ values (`dict`):
+ Values to be stored as initial hyperparameters as key-value pairs.
+ """
+ import mlflow
+
+ for name, value in list(values.items()):
+ # internally, all values are converted to str in MLflow
+ if len(str(value)) > mlflow.utils.validation.MAX_PARAM_VAL_LENGTH:
+ logger.warning_once(
+ f'Accelerate is attempting to log a value of "{value}" for key "{name}" as a parameter. MLflow\'s'
+ f" log_param() only accepts values no longer than {mlflow.utils.validation.MAX_PARAM_VAL_LENGTH} characters so we dropped this attribute."
+ )
+ del values[name]
+
+ values_list = list(values.items())
+
+ # MLflow cannot log more than 100 values in one go, so we have to split it
+ for i in range(0, len(values_list), mlflow.utils.validation.MAX_PARAMS_TAGS_PER_BATCH):
+ mlflow.log_params(dict(values_list[i : i + mlflow.utils.validation.MAX_PARAMS_TAGS_PER_BATCH]))
+
+ logger.debug("Stored initial configuration hyperparameters to MLflow")
+
+ @on_main_process
+ def log(self, values: dict, step: Optional[int]):
+ """
+ Logs `values` to the current run.
+
+ Args:
+ values (`dict`):
+ Values to be logged as key-value pairs.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ """
+ metrics = {}
+ for k, v in values.items():
+ if isinstance(v, (int, float)):
+ metrics[k] = v
+ else:
+ logger.warning_once(
+ f'MLflowTracker is attempting to log a value of "{v}" of type {type(v)} for key "{k}" as a metric. '
+ "MLflow's log_metric() only accepts float and int types so we dropped this attribute."
+ )
+ import mlflow
+
+ mlflow.log_metrics(metrics, step=step)
+ logger.debug("Successfully logged to mlflow")
+
+ @on_main_process
+ def log_figure(self, figure: Any, artifact_file: str, **save_kwargs):
+ """
+ Logs an figure to the current run.
+
+ Args:
+ figure (Any):
+ The figure to be logged.
+ artifact_file (`str`, *optional*):
+ The run-relative artifact file path in posixpath format to which the image is saved.
+ If not provided, the image is saved to a default location.
+ **kwargs:
+ Additional keyword arguments passed to the underlying mlflow.log_image function.
+ """
+ import mlflow
+
+ mlflow.log_figure(figure=figure, artifact_file=artifact_file, **save_kwargs)
+ logger.debug("Successfully logged image to mlflow")
+
+ @on_main_process
+ def log_artifacts(self, local_dir: str, artifact_path: Optional[str] = None):
+ """
+ Logs an artifacts (all content of a dir) to the current run.
+
+ local_dir (`str`):
+ Path to the directory to be logged as an artifact.
+ artifact_path (`str`, *optional*):
+ Directory within the run's artifact directory where the artifact will be logged. If omitted, the
+ artifact will be logged to the root of the run's artifact directory. The run step. If included, the
+ artifact will be affiliated with this step.
+ """
+ import mlflow
+
+ mlflow.log_artifacts(local_dir=local_dir, artifact_path=artifact_path)
+ logger.debug("Successfully logged artofact to mlflow")
+
+ @on_main_process
+ def log_artifact(self, local_path: str, artifact_path: Optional[str] = None):
+ """
+ Logs an artifact (file) to the current run.
+
+ local_path (`str`):
+ Path to the file to be logged as an artifact.
+ artifact_path (`str`, *optional*):
+ Directory within the run's artifact directory where the artifact will be logged. If omitted, the
+ artifact will be logged to the root of the run's artifact directory. The run step. If included, the
+ artifact will be affiliated with this step.
+ """
+ import mlflow
+
+ mlflow.log_artifact(local_path=local_path, artifact_path=artifact_path)
+ logger.debug("Successfully logged artofact to mlflow")
+
+ @on_main_process
+ def finish(self):
+ """
+ End the active MLflow run.
+ """
+ import mlflow
+
+ mlflow.end_run()
+
+
+class ClearMLTracker(GeneralTracker):
+ """
+ A `Tracker` class that supports `clearml`. Should be initialized at the start of your script.
+
+ Args:
+ run_name (`str`, *optional*):
+ Name of the experiment. Environment variables `CLEARML_PROJECT` and `CLEARML_TASK` have priority over this
+ argument.
+ **kwargs (additional keyword arguments, *optional*):
+ Kwargs passed along to the `Task.__init__` method.
+ """
+
+ name = "clearml"
+ requires_logging_directory = False
+
+ @on_main_process
+ def __init__(self, run_name: str = None, **kwargs):
+ from clearml import Task
+
+ current_task = Task.current_task()
+ self._initialized_externally = False
+ if current_task:
+ self._initialized_externally = True
+ self.task = current_task
+ return
+
+ kwargs.setdefault("project_name", os.environ.get("CLEARML_PROJECT", run_name))
+ kwargs.setdefault("task_name", os.environ.get("CLEARML_TASK", run_name))
+ self.task = Task.init(**kwargs)
+
+ @property
+ def tracker(self):
+ return self.task
+
+ @on_main_process
+ def store_init_configuration(self, values: dict):
+ """
+ Connect configuration dictionary to the Task object. Should be run at the beginning of your experiment.
+
+ Args:
+ values (`dict`):
+ Values to be stored as initial hyperparameters as key-value pairs.
+ """
+ return self.task.connect_configuration(values)
+
+ @on_main_process
+ def log(self, values: dict[str, Union[int, float]], step: Optional[int] = None, **kwargs):
+ """
+ Logs `values` dictionary to the current run. The dictionary keys must be strings. The dictionary values must be
+ ints or floats
+
+ Args:
+ values (`Dict[str, Union[int, float]]`):
+ Values to be logged as key-value pairs. If the key starts with 'eval_'/'test_'/'train_', the value will
+ be reported under the 'eval'/'test'/'train' series and the respective prefix will be removed.
+ Otherwise, the value will be reported under the 'train' series, and no prefix will be removed.
+ step (`int`, *optional*):
+ If specified, the values will be reported as scalars, with the iteration number equal to `step`.
+ Otherwise they will be reported as single values.
+ kwargs:
+ Additional key word arguments passed along to the `clearml.Logger.report_single_value` or
+ `clearml.Logger.report_scalar` methods.
+ """
+ clearml_logger = self.task.get_logger()
+ for k, v in values.items():
+ if not isinstance(v, (int, float)):
+ logger.warning_once(
+ "Accelerator is attempting to log a value of "
+ f'"{v}" of type {type(v)} for key "{k}" as a scalar. '
+ "This invocation of ClearML logger's report_scalar() "
+ "is incorrect so we dropped this attribute."
+ )
+ continue
+ if step is None:
+ clearml_logger.report_single_value(name=k, value=v, **kwargs)
+ continue
+ title, series = ClearMLTracker._get_title_series(k)
+ clearml_logger.report_scalar(title=title, series=series, value=v, iteration=step, **kwargs)
+
+ @on_main_process
+ def log_images(self, values: dict, step: Optional[int] = None, **kwargs):
+ """
+ Logs `images` to the current run.
+
+ Args:
+ values (`Dict[str, List[Union[np.ndarray, PIL.Image]]`):
+ Values to be logged as key-value pairs. The values need to have type `List` of `np.ndarray` or
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to the `clearml.Logger.report_image` method.
+ """
+ clearml_logger = self.task.get_logger()
+ for k, v in values.items():
+ title, series = ClearMLTracker._get_title_series(k)
+ clearml_logger.report_image(title=title, series=series, iteration=step, image=v, **kwargs)
+
+ @on_main_process
+ def log_table(
+ self,
+ table_name: str,
+ columns: list[str] = None,
+ data: list[list[Any]] = None,
+ dataframe: Any = None,
+ step: Optional[int] = None,
+ **kwargs,
+ ):
+ """
+ Log a Table to the task. Can be defined eitherwith `columns` and `data` or with `dataframe`.
+
+ Args:
+ table_name (`str`):
+ The name of the table
+ columns (list of `str`, *optional*):
+ The name of the columns on the table
+ data (List of List of Any data type, *optional*):
+ The data to be logged in the table. If `columns` is not specified, then the first entry in data will be
+ the name of the columns of the table
+ dataframe (Any data type, *optional*):
+ The data to be logged in the table
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to the `clearml.Logger.report_table` method.
+ """
+ to_report = dataframe
+ if dataframe is None:
+ if data is None:
+ raise ValueError(
+ "`ClearMLTracker.log_table` requires that `data` to be supplied if `dataframe` is `None`"
+ )
+ to_report = [columns] + data if columns else data
+ title, series = ClearMLTracker._get_title_series(table_name)
+ self.task.get_logger().report_table(title=title, series=series, table_plot=to_report, iteration=step, **kwargs)
+
+ @on_main_process
+ def finish(self):
+ """
+ Close the ClearML task. If the task was initialized externally (e.g. by manually calling `Task.init`), this
+ function is a noop
+ """
+ if self.task and not self._initialized_externally:
+ self.task.close()
+
+ @staticmethod
+ def _get_title_series(name):
+ for prefix in ["eval", "test", "train"]:
+ if name.startswith(prefix + "_"):
+ return name[len(prefix) + 1 :], prefix
+ return name, "train"
+
+
+class DVCLiveTracker(GeneralTracker):
+ """
+ A `Tracker` class that supports `dvclive`. Should be initialized at the start of your script.
+
+ Args:
+ run_name (`str`, *optional*):
+ Ignored for dvclive. See `kwargs` instead.
+ kwargs:
+ Additional key word arguments passed along to [`dvclive.Live()`](https://dvc.org/doc/dvclive/live).
+
+ Example:
+
+ ```py
+ from accelerate import Accelerator
+
+ accelerator = Accelerator(log_with="dvclive")
+ accelerator.init_trackers(project_name="my_project", init_kwargs={"dvclive": {"dir": "my_directory"}})
+ ```
+ """
+
+ name = "dvclive"
+ requires_logging_directory = False
+
+ @on_main_process
+ def __init__(self, run_name: Optional[str] = None, live: Optional[Any] = None, **kwargs):
+ from dvclive import Live
+
+ super().__init__()
+ self.live = live if live is not None else Live(**kwargs)
+
+ @property
+ def tracker(self):
+ return self.live
+
+ @on_main_process
+ def store_init_configuration(self, values: dict):
+ """
+ Logs `values` as hyperparameters for the run. Should be run at the beginning of your experiment. Stores the
+ hyperparameters in a yaml file for future use.
+
+ Args:
+ values (Dictionary `str` to `bool`, `str`, `float`, `int`, or a List or Dict of those types):
+ Values to be stored as initial hyperparameters as key-value pairs. The values need to have type `bool`,
+ `str`, `float`, or `int`.
+ """
+ self.live.log_params(values)
+
+ @on_main_process
+ def log(self, values: dict, step: Optional[int] = None, **kwargs):
+ """
+ Logs `values` to the current run.
+
+ Args:
+ values (Dictionary `str` to `str`, `float`, or `int`):
+ Values to be logged as key-value pairs. The values need to have type `str`, `float`, or `int`.
+ step (`int`, *optional*):
+ The run step. If included, the log will be affiliated with this step.
+ kwargs:
+ Additional key word arguments passed along to `dvclive.Live.log_metric()`.
+ """
+ from dvclive.plots import Metric
+
+ if step is not None:
+ self.live.step = step
+ for k, v in values.items():
+ if Metric.could_log(v):
+ self.live.log_metric(k, v, **kwargs)
+ else:
+ logger.warning_once(
+ "Accelerator attempted to log a value of "
+ f'"{v}" of type {type(v)} for key "{k}" as a scalar. '
+ "This invocation of DVCLive's Live.log_metric() "
+ "is incorrect so we dropped this attribute."
+ )
+ self.live.next_step()
+
+ @on_main_process
+ def finish(self):
+ """
+ Closes `dvclive.Live()`.
+ """
+ self.live.end()
+
+
+LOGGER_TYPE_TO_CLASS = {
+ "aim": AimTracker,
+ "comet_ml": CometMLTracker,
+ "mlflow": MLflowTracker,
+ "tensorboard": TensorBoardTracker,
+ "wandb": WandBTracker,
+ "clearml": ClearMLTracker,
+ "dvclive": DVCLiveTracker,
+}
+
+
+def filter_trackers(
+ log_with: list[Union[str, LoggerType, GeneralTracker]],
+ logging_dir: Union[str, os.PathLike] = None,
+):
+ """
+ Takes in a list of potential tracker types and checks that:
+ - The tracker wanted is available in that environment
+ - Filters out repeats of tracker types
+ - If `all` is in `log_with`, will return all trackers in the environment
+ - If a tracker requires a `logging_dir`, ensures that `logging_dir` is not `None`
+
+ Args:
+ log_with (list of `str`, [`~utils.LoggerType`] or [`~tracking.GeneralTracker`], *optional*):
+ A list of loggers to be setup for experiment tracking. Should be one or several of:
+
+ - `"all"`
+ - `"tensorboard"`
+ - `"wandb"`
+ - `"comet_ml"`
+ - `"mlflow"`
+ - `"dvclive"`
+ If `"all"` is selected, will pick up all available trackers in the environment and initialize them. Can
+ also accept implementations of `GeneralTracker` for custom trackers, and can be combined with `"all"`.
+ logging_dir (`str`, `os.PathLike`, *optional*):
+ A path to a directory for storing logs of locally-compatible loggers.
+ """
+ loggers = []
+ if log_with is not None:
+ if not isinstance(log_with, (list, tuple)):
+ log_with = [log_with]
+ if "all" in log_with or LoggerType.ALL in log_with:
+ loggers = [o for o in log_with if issubclass(type(o), GeneralTracker)] + get_available_trackers()
+ else:
+ for log_type in log_with:
+ if log_type not in LoggerType and not issubclass(type(log_type), GeneralTracker):
+ raise ValueError(f"Unsupported logging capability: {log_type}. Choose between {LoggerType.list()}")
+ if issubclass(type(log_type), GeneralTracker):
+ loggers.append(log_type)
+ else:
+ log_type = LoggerType(log_type)
+ if log_type not in loggers:
+ if log_type in get_available_trackers():
+ tracker_init = LOGGER_TYPE_TO_CLASS[str(log_type)]
+ if tracker_init.requires_logging_directory:
+ if logging_dir is None:
+ raise ValueError(
+ f"Logging with `{log_type}` requires a `logging_dir` to be passed in."
+ )
+ loggers.append(log_type)
+ else:
+ logger.debug(f"Tried adding logger {log_type}, but package is unavailable in the system.")
+
+ return loggers
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/__init__.py b/venv/lib/python3.11/site-packages/accelerate/utils/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..1c626cba9bcf4619e572d4deae388ace600fefec
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/__init__.py
@@ -0,0 +1,291 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+from .ao import convert_model_to_fp8_ao, filter_first_and_last_linear_layers, has_ao_layers
+from .constants import (
+ MITA_PROFILING_AVAILABLE_PYTORCH_VERSION,
+ MODEL_NAME,
+ OPTIMIZER_NAME,
+ PROFILE_PATTERN_NAME,
+ RNG_STATE_NAME,
+ SAFE_MODEL_NAME,
+ SAFE_WEIGHTS_INDEX_NAME,
+ SAFE_WEIGHTS_NAME,
+ SAFE_WEIGHTS_PATTERN_NAME,
+ SAMPLER_NAME,
+ SCALER_NAME,
+ SCHEDULER_NAME,
+ TORCH_DISTRIBUTED_OPERATION_TYPES,
+ TORCH_LAUNCH_PARAMS,
+ WEIGHTS_INDEX_NAME,
+ WEIGHTS_NAME,
+ WEIGHTS_PATTERN_NAME,
+ XPU_PROFILING_AVAILABLE_PYTORCH_VERSION,
+)
+from .dataclasses import (
+ AORecipeKwargs,
+ AutocastKwargs,
+ BnbQuantizationConfig,
+ ComputeEnvironment,
+ CustomDtype,
+ DataLoaderConfiguration,
+ DDPCommunicationHookType,
+ DeepSpeedPlugin,
+ DistributedDataParallelKwargs,
+ DistributedType,
+ DynamoBackend,
+ FP8RecipeKwargs,
+ FullyShardedDataParallelPlugin,
+ GradientAccumulationPlugin,
+ GradScalerKwargs,
+ InitProcessGroupKwargs,
+ KwargsHandler,
+ LoggerType,
+ MegatronLMPlugin,
+ MSAMPRecipeKwargs,
+ PrecisionType,
+ ProfileKwargs,
+ ProjectConfiguration,
+ RNGType,
+ SageMakerDistributedType,
+ TensorInformation,
+ TERecipeKwargs,
+ TorchDynamoPlugin,
+ TorchTensorParallelPlugin,
+ add_model_config_to_megatron_parser,
+)
+from .environment import (
+ are_libraries_initialized,
+ check_cuda_fp8_capability,
+ check_cuda_p2p_ib_support,
+ clear_environment,
+ convert_dict_to_env_variables,
+ get_cpu_distributed_information,
+ get_gpu_info,
+ get_int_from_env,
+ parse_choice_from_env,
+ parse_flag_from_env,
+ patch_environment,
+ purge_accelerate_environment,
+ set_numa_affinity,
+ str_to_bool,
+)
+from .imports import (
+ deepspeed_required,
+ get_ccl_version,
+ is_4bit_bnb_available,
+ is_8bit_bnb_available,
+ is_aim_available,
+ is_bf16_available,
+ is_bitsandbytes_multi_backend_available,
+ is_bnb_available,
+ is_boto3_available,
+ is_ccl_available,
+ is_clearml_available,
+ is_comet_ml_available,
+ is_cuda_available,
+ is_datasets_available,
+ is_deepspeed_available,
+ is_dvclive_available,
+ is_fp8_available,
+ is_fp16_available,
+ is_habana_gaudi1,
+ is_hpu_available,
+ is_import_timer_available,
+ is_ipex_available,
+ is_lomo_available,
+ is_matplotlib_available,
+ is_megatron_lm_available,
+ is_mlflow_available,
+ is_mlu_available,
+ is_mps_available,
+ is_msamp_available,
+ is_musa_available,
+ is_npu_available,
+ is_pandas_available,
+ is_peft_available,
+ is_pippy_available,
+ is_pynvml_available,
+ is_pytest_available,
+ is_rich_available,
+ is_sagemaker_available,
+ is_schedulefree_available,
+ is_sdaa_available,
+ is_tensorboard_available,
+ is_timm_available,
+ is_torch_xla_available,
+ is_torchao_available,
+ is_torchdata_available,
+ is_torchdata_stateful_dataloader_available,
+ is_torchvision_available,
+ is_transformer_engine_available,
+ is_transformers_available,
+ is_triton_available,
+ is_wandb_available,
+ is_weights_only_available,
+ is_xccl_available,
+ is_xpu_available,
+ torchao_required,
+)
+from .modeling import (
+ align_module_device,
+ calculate_maximum_sizes,
+ check_device_map,
+ check_tied_parameters_in_config,
+ check_tied_parameters_on_same_device,
+ compute_module_sizes,
+ convert_file_size_to_int,
+ dtype_byte_size,
+ find_tied_parameters,
+ get_balanced_memory,
+ get_grad_scaler,
+ get_max_layer_size,
+ get_max_memory,
+ get_mixed_precision_context_manager,
+ has_offloaded_params,
+ id_tensor_storage,
+ infer_auto_device_map,
+ is_peft_model,
+ load_checkpoint_in_model,
+ load_offloaded_weights,
+ load_state_dict,
+ named_module_tensors,
+ retie_parameters,
+ set_module_tensor_to_device,
+)
+from .offload import (
+ OffloadedWeightsLoader,
+ PrefixedDataset,
+ extract_submodules_state_dict,
+ load_offloaded_weight,
+ offload_state_dict,
+ offload_weight,
+ save_offload_index,
+)
+from .operations import (
+ CannotPadNestedTensorWarning,
+ GatheredParameters,
+ broadcast,
+ broadcast_object_list,
+ concatenate,
+ convert_outputs_to_fp32,
+ convert_to_fp32,
+ copy_tensor_to_devices,
+ find_batch_size,
+ find_device,
+ gather,
+ gather_object,
+ get_data_structure,
+ honor_type,
+ ignorant_find_batch_size,
+ initialize_tensors,
+ is_namedtuple,
+ is_tensor_information,
+ is_torch_tensor,
+ listify,
+ pad_across_processes,
+ pad_input_tensors,
+ recursively_apply,
+ reduce,
+ send_to_device,
+ slice_tensors,
+)
+from .versions import compare_versions, is_torch_version
+
+
+if is_deepspeed_available():
+ from .deepspeed import (
+ DeepSpeedEngineWrapper,
+ DeepSpeedOptimizerWrapper,
+ DeepSpeedSchedulerWrapper,
+ DummyOptim,
+ DummyScheduler,
+ HfDeepSpeedConfig,
+ get_active_deepspeed_plugin,
+ map_pytorch_optim_to_deepspeed,
+ )
+
+from .bnb import has_4bit_bnb_layers, load_and_quantize_model
+from .fsdp_utils import (
+ disable_fsdp_ram_efficient_loading,
+ enable_fsdp_ram_efficient_loading,
+ ensure_weights_retied,
+ fsdp2_load_full_state_dict,
+ fsdp2_prepare_model,
+ fsdp2_switch_optimizer_parameters,
+ get_fsdp2_grad_scaler,
+ load_fsdp_model,
+ load_fsdp_optimizer,
+ merge_fsdp_weights,
+ save_fsdp_model,
+ save_fsdp_optimizer,
+)
+from .launch import (
+ PrepareForLaunch,
+ _filter_args,
+ prepare_deepspeed_cmd_env,
+ prepare_multi_gpu_env,
+ prepare_sagemager_args_inputs,
+ prepare_simple_launcher_cmd_env,
+ prepare_tpu,
+)
+
+# For docs
+from .megatron_lm import (
+ AbstractTrainStep,
+ BertTrainStep,
+ GPTTrainStep,
+ MegatronLMDummyDataLoader,
+ MegatronLMDummyScheduler,
+ T5TrainStep,
+ avg_losses_across_data_parallel_group,
+)
+
+
+if is_megatron_lm_available():
+ from .megatron_lm import (
+ MegatronEngine,
+ MegatronLMOptimizerWrapper,
+ MegatronLMSchedulerWrapper,
+ gather_across_data_parallel_groups,
+ )
+ from .megatron_lm import initialize as megatron_lm_initialize
+ from .megatron_lm import prepare_data_loader as megatron_lm_prepare_data_loader
+ from .megatron_lm import prepare_model_optimizer_scheduler as megatron_lm_prepare_model_optimizer_scheduler
+ from .megatron_lm import prepare_optimizer as megatron_lm_prepare_optimizer
+ from .megatron_lm import prepare_scheduler as megatron_lm_prepare_scheduler
+from .memory import find_executable_batch_size, release_memory
+from .other import (
+ check_os_kernel,
+ clean_state_dict_for_safetensors,
+ convert_bytes,
+ extract_model_from_parallel,
+ get_module_children_bottom_up,
+ get_pretty_name,
+ is_port_in_use,
+ load,
+ merge_dicts,
+ recursive_getattr,
+ save,
+ wait_for_everyone,
+ write_basic_config,
+)
+from .random import set_seed, synchronize_rng_state, synchronize_rng_states
+from .torch_xla import install_xla
+from .tqdm import tqdm
+from .transformer_engine import (
+ apply_fp8_autowrap,
+ contextual_fp8_autocast,
+ convert_model,
+ has_transformer_engine_layers,
+)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/ao.py b/venv/lib/python3.11/site-packages/accelerate/utils/ao.py
new file mode 100644
index 0000000000000000000000000000000000000000..53e9c85fe70df91c7ea3d0396aa75705744f2223
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/ao.py
@@ -0,0 +1,139 @@
+# Copyright 2025 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+"""
+Needed utilities for torchao FP8 training.
+"""
+
+from functools import partial
+from typing import Callable, Optional
+
+import torch
+
+from .imports import is_torchao_available, torchao_required
+
+
+if is_torchao_available():
+ from torchao.float8.float8_linear import Float8LinearConfig
+
+
+def find_first_last_linear_layers(model: torch.nn.Module):
+ """
+ Finds the first and last linear layer names in a model.
+
+ This is needed during FP8 to avoid issues with instability by keeping the first and last layers unquantized.
+
+ Ref: https://x.com/xariusrke/status/1826669142604141052
+ """
+ first_linear, last_linear = None, None
+ for name, module in model.named_modules():
+ if isinstance(module, torch.nn.Linear):
+ if first_linear is None:
+ first_linear = name
+ last_linear = name
+ return first_linear, last_linear
+
+
+def filter_linear_layers(module, fqn: str, layers_to_filter: list[str]) -> bool:
+ """
+ A function which will check if `module` is:
+ - a `torch.nn.Linear` layer
+ - has in_features and out_features divisible by 16
+ - is not part of `layers_to_filter`
+
+ Args:
+ module (`torch.nn.Module`):
+ The module to check.
+ fqn (`str`):
+ The fully qualified name of the layer.
+ layers_to_filter (`List[str]`):
+ The list of layers to filter.
+ """
+ if isinstance(module, torch.nn.Linear):
+ if module.in_features % 16 != 0 or module.out_features % 16 != 0:
+ return False
+ if fqn in layers_to_filter:
+ return False
+ return True
+
+
+def filter_first_and_last_linear_layers(module, fqn: str) -> bool:
+ """
+ A filter function which will filter out all linear layers except the first and last.
+
+
+
+ For stability reasons, we skip the first and last linear layers Otherwise can lead to the model not training or
+ converging properly
+
+
+
+ Args:
+ module (`torch.nn.Module`):
+ The module to check.
+ fqn (`str`):
+ The fully qualified name of the layer.
+ """
+ first_linear, last_linear = find_first_last_linear_layers(module)
+ return filter_linear_layers(module, fqn, layers_to_filter=[first_linear, last_linear])
+
+
+@torchao_required
+def has_ao_layers(model: torch.nn.Module):
+ from torchao.float8.float8_linear import Float8Linear
+
+ for name, module in model.named_modules():
+ if isinstance(module, Float8Linear):
+ return True
+ return False
+
+
+@torchao_required
+def convert_model_to_fp8_ao(
+ model: torch.nn.Module,
+ config: Optional["Float8LinearConfig"] = None,
+ module_filter_func: Optional[Callable] = filter_first_and_last_linear_layers,
+):
+ """
+ Converts all `nn.Linear` layers in the model (except the first and last) to torchao's `Float8Linear` layer inplace.
+
+ Args:
+ model (`torch.nn.Module`):
+ The model to convert.
+ config (`torchao.float8.Float8LinearConfig`, *optional*):
+ The configuration for the FP8 training. Recommended to utilize
+ `torchao.float8.recipe_name_to_linear_config` to generate this. In general, the default config should be
+ sufficient (what is passed when set to `None`).
+ module_filter_func (`Callable`, *optional*, defaults to `filter_linear_layers`):
+ Optional function that must take in a module and layer name, and returns a boolean indicating whether the
+ module should be converted to FP8. Defaults to `filter_linear_layers`. See it for an example.
+
+ Example:
+
+ ```python
+ from accelerate.utils.ao import convert_model_to_fp8_ao
+
+ model = MyModel()
+ model.to("cuda")
+ convert_to_float8_training(model)
+
+ model.train()
+ ```
+ """
+ from torchao.float8 import convert_to_float8_training
+
+ first_linear, last_linear = find_first_last_linear_layers(model)
+ if module_filter_func is None:
+ module_filter_func = partial(filter_linear_layers, layers_to_filter=[first_linear, last_linear])
+ convert_to_float8_training(model, module_filter_fn=module_filter_func, config=config)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/bnb.py b/venv/lib/python3.11/site-packages/accelerate/utils/bnb.py
new file mode 100644
index 0000000000000000000000000000000000000000..f78820aa61a9f556ac3435d17c4648c8de32efd2
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/bnb.py
@@ -0,0 +1,470 @@
+# Copyright 2023 The HuggingFace Team. All rights reserved.
+#
+# 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 logging
+import os
+from copy import deepcopy
+from typing import Optional, Union
+
+import torch
+import torch.nn as nn
+
+from accelerate.utils.imports import (
+ is_4bit_bnb_available,
+ is_8bit_bnb_available,
+)
+
+from ..big_modeling import dispatch_model, init_empty_weights
+from .dataclasses import BnbQuantizationConfig
+from .modeling import (
+ find_tied_parameters,
+ get_balanced_memory,
+ infer_auto_device_map,
+ load_checkpoint_in_model,
+ offload_weight,
+ set_module_tensor_to_device,
+)
+
+
+logger = logging.getLogger(__name__)
+
+
+def load_and_quantize_model(
+ model: torch.nn.Module,
+ bnb_quantization_config: BnbQuantizationConfig,
+ weights_location: Union[str, os.PathLike] = None,
+ device_map: Optional[dict[str, Union[int, str, torch.device]]] = None,
+ no_split_module_classes: Optional[list[str]] = None,
+ max_memory: Optional[dict[Union[int, str], Union[int, str]]] = None,
+ offload_folder: Optional[Union[str, os.PathLike]] = None,
+ offload_state_dict: bool = False,
+):
+ """
+ This function will quantize the input model with the associated config passed in `bnb_quantization_config`. If the
+ model is in the meta device, we will load and dispatch the weights according to the `device_map` passed. If the
+ model is already loaded, we will quantize the model and put the model on the GPU,
+
+ Args:
+ model (`torch.nn.Module`):
+ Input model. The model can be already loaded or on the meta device
+ bnb_quantization_config (`BnbQuantizationConfig`):
+ The bitsandbytes quantization parameters
+ weights_location (`str` or `os.PathLike`):
+ The folder weights_location to load. It can be:
+ - a path to a file containing a whole model state dict
+ - a path to a `.json` file containing the index to a sharded checkpoint
+ - a path to a folder containing a unique `.index.json` file and the shards of a checkpoint.
+ - a path to a folder containing a unique pytorch_model.bin file.
+ device_map (`Dict[str, Union[int, str, torch.device]]`, *optional*):
+ A map that specifies where each submodule should go. It doesn't need to be refined to each parameter/buffer
+ name, once a given module name is inside, every submodule of it will be sent to the same device.
+ no_split_module_classes (`List[str]`, *optional*):
+ A list of layer class names that should never be split across device (for instance any layer that has a
+ residual connection).
+ max_memory (`Dict`, *optional*):
+ A dictionary device identifier to maximum memory. Will default to the maximum memory available if unset.
+ offload_folder (`str` or `os.PathLike`, *optional*):
+ If the `device_map` contains any value `"disk"`, the folder where we will offload weights.
+ offload_state_dict (`bool`, *optional*, defaults to `False`):
+ If `True`, will temporarily offload the CPU state dict on the hard drive to avoid getting out of CPU RAM if
+ the weight of the CPU state dict + the biggest shard does not fit.
+
+ Returns:
+ `torch.nn.Module`: The quantized model
+ """
+
+ load_in_4bit = bnb_quantization_config.load_in_4bit
+ load_in_8bit = bnb_quantization_config.load_in_8bit
+
+ if load_in_8bit and not is_8bit_bnb_available():
+ raise ImportError(
+ "You have a version of `bitsandbytes` that is not compatible with 8bit quantization,"
+ " make sure you have the latest version of `bitsandbytes` installed."
+ )
+ if load_in_4bit and not is_4bit_bnb_available():
+ raise ValueError(
+ "You have a version of `bitsandbytes` that is not compatible with 4bit quantization,"
+ "make sure you have the latest version of `bitsandbytes` installed."
+ )
+
+ modules_on_cpu = []
+ # custom device map
+ if isinstance(device_map, dict) and len(device_map.keys()) > 1:
+ modules_on_cpu = [key for key, value in device_map.items() if value in ["disk", "cpu"]]
+
+ # We keep some modules such as the lm_head in their original dtype for numerical stability reasons
+ if bnb_quantization_config.skip_modules is None:
+ bnb_quantization_config.skip_modules = get_keys_to_not_convert(model)
+
+ # add cpu modules to skip modules only for 4-bit modules
+ if load_in_4bit:
+ bnb_quantization_config.skip_modules.extend(modules_on_cpu)
+ modules_to_not_convert = bnb_quantization_config.skip_modules
+
+ # We add the modules we want to keep in full precision
+ if bnb_quantization_config.keep_in_fp32_modules is None:
+ bnb_quantization_config.keep_in_fp32_modules = []
+ keep_in_fp32_modules = bnb_quantization_config.keep_in_fp32_modules
+ modules_to_not_convert.extend(keep_in_fp32_modules)
+
+ # compatibility with peft
+ model.is_loaded_in_4bit = load_in_4bit
+ model.is_loaded_in_8bit = load_in_8bit
+
+ model_device = get_parameter_device(model)
+ if model_device.type != "meta":
+ # quantization of an already loaded model
+ logger.warning(
+ "It is not recommended to quantize a loaded model. "
+ "The model should be instantiated under the `init_empty_weights` context manager."
+ )
+ model = replace_with_bnb_layers(model, bnb_quantization_config, modules_to_not_convert=modules_to_not_convert)
+ # convert param to the right dtype
+ dtype = bnb_quantization_config.torch_dtype
+ for name, param in model.state_dict().items():
+ if any(module_to_keep_in_fp32 in name for module_to_keep_in_fp32 in keep_in_fp32_modules):
+ param.to(torch.float32)
+ if param.dtype != torch.float32:
+ name = name.replace(".weight", "").replace(".bias", "")
+ param = getattr(model, name, None)
+ if param is not None:
+ param.to(torch.float32)
+ elif torch.is_floating_point(param):
+ param.to(dtype)
+ if model_device.type == "cuda":
+ # move everything to cpu in the first place because we can't do quantization if the weights are already on cuda
+ model.cuda(torch.cuda.current_device())
+ torch.cuda.empty_cache()
+ elif torch.cuda.is_available():
+ model.to(torch.cuda.current_device())
+ elif torch.xpu.is_available():
+ model.to(torch.xpu.current_device())
+ else:
+ raise RuntimeError("No GPU found. A GPU is needed for quantization.")
+ logger.info(
+ f"The model device type is {model_device.type}. However, gpu is needed for quantization."
+ "We move the model to gpu."
+ )
+ return model
+
+ elif weights_location is None:
+ raise RuntimeError(
+ f"`weights_location` needs to be the folder path containing the weights of the model, but we found {weights_location} "
+ )
+
+ else:
+ with init_empty_weights():
+ model = replace_with_bnb_layers(
+ model, bnb_quantization_config, modules_to_not_convert=modules_to_not_convert
+ )
+ device_map = get_quantized_model_device_map(
+ model,
+ bnb_quantization_config,
+ device_map,
+ max_memory=max_memory,
+ no_split_module_classes=no_split_module_classes,
+ )
+ if offload_state_dict is None and device_map is not None and "disk" in device_map.values():
+ offload_state_dict = True
+
+ offload = any(x in list(device_map.values()) for x in ["cpu", "disk"])
+
+ load_checkpoint_in_model(
+ model,
+ weights_location,
+ device_map,
+ dtype=bnb_quantization_config.torch_dtype,
+ offload_folder=offload_folder,
+ offload_state_dict=offload_state_dict,
+ keep_in_fp32_modules=bnb_quantization_config.keep_in_fp32_modules,
+ offload_8bit_bnb=load_in_8bit and offload,
+ )
+ return dispatch_model(model, device_map=device_map, offload_dir=offload_folder)
+
+
+def get_quantized_model_device_map(
+ model, bnb_quantization_config, device_map=None, max_memory=None, no_split_module_classes=None
+):
+ if device_map is None:
+ if torch.cuda.is_available():
+ device_map = {"": torch.cuda.current_device()}
+ elif torch.xpu.is_available():
+ device_map = {"": torch.xpu.current_device()}
+ else:
+ raise RuntimeError("No GPU found. A GPU is needed for quantization.")
+ logger.info("The device_map was not initialized.Setting device_map to `{'':torch.cuda.current_device()}`.")
+
+ if isinstance(device_map, str):
+ if device_map not in ["auto", "balanced", "balanced_low_0", "sequential"]:
+ raise ValueError(
+ "If passing a string for `device_map`, please choose 'auto', 'balanced', 'balanced_low_0' or "
+ "'sequential'."
+ )
+
+ special_dtypes = {}
+ special_dtypes.update(
+ {
+ name: bnb_quantization_config.torch_dtype
+ for name, _ in model.named_parameters()
+ if any(m in name for m in bnb_quantization_config.skip_modules)
+ }
+ )
+ special_dtypes.update(
+ {
+ name: torch.float32
+ for name, _ in model.named_parameters()
+ if any(m in name for m in bnb_quantization_config.keep_in_fp32_modules)
+ }
+ )
+
+ kwargs = {}
+ kwargs["special_dtypes"] = special_dtypes
+ kwargs["no_split_module_classes"] = no_split_module_classes
+ kwargs["dtype"] = bnb_quantization_config.target_dtype
+
+ # get max_memory for each device.
+ if device_map != "sequential":
+ max_memory = get_balanced_memory(
+ model,
+ low_zero=(device_map == "balanced_low_0"),
+ max_memory=max_memory,
+ **kwargs,
+ )
+
+ kwargs["max_memory"] = max_memory
+ device_map = infer_auto_device_map(model, **kwargs)
+
+ if isinstance(device_map, dict):
+ # check if don't have any quantized module on the cpu
+ modules_not_to_convert = bnb_quantization_config.skip_modules + bnb_quantization_config.keep_in_fp32_modules
+
+ device_map_without_some_modules = {
+ key: device_map[key] for key in device_map.keys() if key not in modules_not_to_convert
+ }
+ for device in ["cpu", "disk"]:
+ if device in device_map_without_some_modules.values():
+ if bnb_quantization_config.load_in_4bit:
+ raise ValueError(
+ """
+ Some modules are dispatched on the CPU or the disk. Make sure you have enough GPU RAM to fit
+ the quantized model. If you want to dispatch the model on the CPU or the disk while keeping
+ these modules in `torch_dtype`, you need to pass a custom `device_map` to
+ `load_and_quantize_model`. Check
+ https://huggingface.co/docs/accelerate/main/en/usage_guides/quantization#offload-modules-to-cpu-and-disk
+ for more details.
+ """
+ )
+ else:
+ logger.info(
+ "Some modules are are offloaded to the CPU or the disk. Note that these modules will be converted to 8-bit"
+ )
+ del device_map_without_some_modules
+ return device_map
+
+
+def replace_with_bnb_layers(model, bnb_quantization_config, modules_to_not_convert=None, current_key_name=None):
+ """
+ A helper function to replace all `torch.nn.Linear` modules by `bnb.nn.Linear8bit` modules or by `bnb.nn.Linear4bit`
+ modules from the `bitsandbytes`library. The function will be run recursively and replace `torch.nn.Linear` modules.
+
+ Parameters:
+ model (`torch.nn.Module`):
+ Input model or `torch.nn.Module` as the function is run recursively.
+ modules_to_not_convert (`List[str]`):
+ Names of the modules to not quantize convert. In practice we keep the `lm_head` in full precision for
+ numerical stability reasons.
+ current_key_name (`List[str]`, *optional*):
+ An array to track the current key of the recursion. This is used to check whether the current key (part of
+ it) is not in the list of modules to not convert.
+ """
+
+ if modules_to_not_convert is None:
+ modules_to_not_convert = []
+
+ model, has_been_replaced = _replace_with_bnb_layers(
+ model, bnb_quantization_config, modules_to_not_convert, current_key_name
+ )
+ if not has_been_replaced:
+ logger.warning(
+ "You are loading your model in 8bit or 4bit but no linear modules were found in your model."
+ " this can happen for some architectures such as gpt2 that uses Conv1D instead of Linear layers."
+ " Please double check your model architecture, or submit an issue on github if you think this is"
+ " a bug."
+ )
+ return model
+
+
+def _replace_with_bnb_layers(
+ model,
+ bnb_quantization_config,
+ modules_to_not_convert=None,
+ current_key_name=None,
+):
+ """
+ Private method that wraps the recursion for module replacement.
+
+ Returns the converted model and a boolean that indicates if the conversion has been successfull or not.
+ """
+ # bitsandbytes will initialize CUDA on import, so it needs to be imported lazily
+ import bitsandbytes as bnb
+
+ has_been_replaced = False
+ for name, module in model.named_children():
+ if current_key_name is None:
+ current_key_name = []
+ current_key_name.append(name)
+ if isinstance(module, nn.Linear) and name not in modules_to_not_convert:
+ # Check if the current key is not in the `modules_to_not_convert`
+ current_key_name_str = ".".join(current_key_name)
+ proceed = True
+ for key in modules_to_not_convert:
+ if (
+ (key in current_key_name_str) and (key + "." in current_key_name_str)
+ ) or key == current_key_name_str:
+ proceed = False
+ break
+ if proceed:
+ # Load bnb module with empty weight and replace ``nn.Linear` module
+ if bnb_quantization_config.load_in_8bit:
+ bnb_module = bnb.nn.Linear8bitLt(
+ module.in_features,
+ module.out_features,
+ module.bias is not None,
+ has_fp16_weights=False,
+ threshold=bnb_quantization_config.llm_int8_threshold,
+ )
+ elif bnb_quantization_config.load_in_4bit:
+ bnb_module = bnb.nn.Linear4bit(
+ module.in_features,
+ module.out_features,
+ module.bias is not None,
+ bnb_quantization_config.bnb_4bit_compute_dtype,
+ compress_statistics=bnb_quantization_config.bnb_4bit_use_double_quant,
+ quant_type=bnb_quantization_config.bnb_4bit_quant_type,
+ )
+ else:
+ raise ValueError("load_in_8bit and load_in_4bit can't be both False")
+ bnb_module.weight.data = module.weight.data
+ if module.bias is not None:
+ bnb_module.bias.data = module.bias.data
+ bnb_module.requires_grad_(False)
+ setattr(model, name, bnb_module)
+ has_been_replaced = True
+ if len(list(module.children())) > 0:
+ _, _has_been_replaced = _replace_with_bnb_layers(
+ module, bnb_quantization_config, modules_to_not_convert, current_key_name
+ )
+ has_been_replaced = has_been_replaced | _has_been_replaced
+ # Remove the last key for recursion
+ current_key_name.pop(-1)
+ return model, has_been_replaced
+
+
+def get_keys_to_not_convert(model):
+ r"""
+ An utility function to get the key of the module to keep in full precision if any For example for CausalLM modules
+ we may want to keep the lm_head in full precision for numerical stability reasons. For other architectures, we want
+ to keep the tied weights of the model. The function will return a list of the keys of the modules to not convert in
+ int8.
+
+ Parameters:
+ model (`torch.nn.Module`):
+ Input model
+ """
+ # Create a copy of the model
+ with init_empty_weights():
+ tied_model = deepcopy(model) # this has 0 cost since it is done inside `init_empty_weights` context manager`
+
+ tied_params = find_tied_parameters(tied_model)
+ # For compatibility with Accelerate < 0.18
+ if isinstance(tied_params, dict):
+ tied_keys = sum(list(tied_params.values()), []) + list(tied_params.keys())
+ else:
+ tied_keys = sum(tied_params, [])
+ has_tied_params = len(tied_keys) > 0
+
+ # Check if it is a base model
+ is_base_model = False
+ if hasattr(model, "base_model_prefix"):
+ is_base_model = not hasattr(model, model.base_model_prefix)
+
+ # Ignore this for base models (BertModel, GPT2Model, etc.)
+ if (not has_tied_params) and is_base_model:
+ return []
+
+ # otherwise they have an attached head
+ list_modules = list(model.named_children())
+ list_last_module = [list_modules[-1][0]]
+
+ # add last module together with tied weights
+ intersection = set(list_last_module) - set(tied_keys)
+ list_untouched = list(set(tied_keys)) + list(intersection)
+
+ # remove ".weight" from the keys
+ names_to_remove = [".weight", ".bias"]
+ filtered_module_names = []
+ for name in list_untouched:
+ for name_to_remove in names_to_remove:
+ if name_to_remove in name:
+ name = name.replace(name_to_remove, "")
+ filtered_module_names.append(name)
+
+ return filtered_module_names
+
+
+def has_4bit_bnb_layers(model):
+ """Check if we have `bnb.nn.Linear4bit` or `bnb.nn.Linear8bitLt` layers inside our model"""
+ # bitsandbytes will initialize CUDA on import, so it needs to be imported lazily
+ import bitsandbytes as bnb
+
+ for m in model.modules():
+ if isinstance(m, bnb.nn.Linear4bit):
+ return True
+ return False
+
+
+def get_parameter_device(parameter: nn.Module):
+ return next(parameter.parameters()).device
+
+
+def quantize_and_offload_8bit(model, param, param_name, new_dtype, offload_folder, offload_index, fp16_statistics):
+ # if it is not quantized, we quantize and offload the quantized weights and the SCB stats
+ if fp16_statistics is None:
+ set_module_tensor_to_device(model, param_name, 0, dtype=new_dtype, value=param)
+ tensor_name = param_name
+ module = model
+ if "." in tensor_name:
+ splits = tensor_name.split(".")
+ for split in splits[:-1]:
+ new_module = getattr(module, split)
+ if new_module is None:
+ raise ValueError(f"{module} has no attribute {split}.")
+ module = new_module
+ tensor_name = splits[-1]
+ # offload weights
+ module._parameters[tensor_name].requires_grad = False
+ offload_weight(module._parameters[tensor_name], param_name, offload_folder, index=offload_index)
+ if hasattr(module._parameters[tensor_name], "SCB"):
+ offload_weight(
+ module._parameters[tensor_name].SCB,
+ param_name.replace("weight", "SCB"),
+ offload_folder,
+ index=offload_index,
+ )
+ else:
+ offload_weight(param, param_name, offload_folder, index=offload_index)
+ offload_weight(fp16_statistics, param_name.replace("weight", "SCB"), offload_folder, index=offload_index)
+
+ set_module_tensor_to_device(model, param_name, "meta", dtype=new_dtype, value=torch.empty(*param.size()))
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/constants.py b/venv/lib/python3.11/site-packages/accelerate/utils/constants.py
new file mode 100644
index 0000000000000000000000000000000000000000..9b9ded4880ea35d5cceeae41d663950329c603ad
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/constants.py
@@ -0,0 +1,92 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 operator as op
+
+
+SCALER_NAME = "scaler.pt"
+MODEL_NAME = "pytorch_model"
+SAFE_MODEL_NAME = "model"
+RNG_STATE_NAME = "random_states"
+OPTIMIZER_NAME = "optimizer"
+SCHEDULER_NAME = "scheduler"
+SAMPLER_NAME = "sampler"
+PROFILE_PATTERN_NAME = "profile_{suffix}.json"
+WEIGHTS_NAME = f"{MODEL_NAME}.bin"
+WEIGHTS_PATTERN_NAME = "pytorch_model{suffix}.bin"
+WEIGHTS_INDEX_NAME = f"{WEIGHTS_NAME}.index.json"
+SAFE_WEIGHTS_NAME = f"{SAFE_MODEL_NAME}.safetensors"
+SAFE_WEIGHTS_PATTERN_NAME = "model{suffix}.safetensors"
+SAFE_WEIGHTS_INDEX_NAME = f"{SAFE_WEIGHTS_NAME}.index.json"
+SAGEMAKER_PYTORCH_VERSION = "1.10.2"
+SAGEMAKER_PYTHON_VERSION = "py38"
+SAGEMAKER_TRANSFORMERS_VERSION = "4.17.0"
+SAGEMAKER_PARALLEL_EC2_INSTANCES = ["ml.p3.16xlarge", "ml.p3dn.24xlarge", "ml.p4dn.24xlarge"]
+FSDP_SHARDING_STRATEGY = ["FULL_SHARD", "SHARD_GRAD_OP", "NO_SHARD", "HYBRID_SHARD", "HYBRID_SHARD_ZERO2"]
+FSDP_AUTO_WRAP_POLICY = ["TRANSFORMER_BASED_WRAP", "SIZE_BASED_WRAP", "NO_WRAP"]
+FSDP_BACKWARD_PREFETCH = ["BACKWARD_PRE", "BACKWARD_POST", "NO_PREFETCH"]
+FSDP_STATE_DICT_TYPE = ["FULL_STATE_DICT", "LOCAL_STATE_DICT", "SHARDED_STATE_DICT"]
+FSDP2_STATE_DICT_TYPE = ["SHARDED_STATE_DICT"]
+FSDP_PYTORCH_VERSION = (
+ "2.1.0.a0+32f93b1" # Technically should be 2.1.0, but MS-AMP uses this specific prerelease in their Docker image.
+)
+FSDP2_PYTORCH_VERSION = "2.5.1"
+FSDP_MODEL_NAME = "pytorch_model_fsdp"
+DEEPSPEED_MULTINODE_LAUNCHERS = ["pdsh", "standard", "openmpi", "mvapich", "mpich", "nossh", "slurm"]
+TORCH_DYNAMO_MODES = ["default", "reduce-overhead", "max-autotune"]
+ELASTIC_LOG_LINE_PREFIX_TEMPLATE_PYTORCH_VERSION = "2.2.0"
+XPU_PROFILING_AVAILABLE_PYTORCH_VERSION = "2.4.0"
+MITA_PROFILING_AVAILABLE_PYTORCH_VERSION = "2.1.0"
+BETA_TP_AVAILABLE_PYTORCH_VERSION = "2.3.0"
+BETA_TP_AVAILABLE_TRANSFORMERS_VERSION = "4.47.0"
+
+STR_OPERATION_TO_FUNC = {">": op.gt, ">=": op.ge, "==": op.eq, "!=": op.ne, "<=": op.le, "<": op.lt}
+
+# These are the args for `torch.distributed.launch` for pytorch < 1.9
+TORCH_LAUNCH_PARAMS = [
+ "nnodes",
+ "nproc_per_node",
+ "rdzv_backend",
+ "rdzv_endpoint",
+ "rdzv_id",
+ "rdzv_conf",
+ "standalone",
+ "max_restarts",
+ "monitor_interval",
+ "start_method",
+ "role",
+ "module",
+ "m",
+ "no_python",
+ "run_path",
+ "log_dir",
+ "r",
+ "redirects",
+ "t",
+ "tee",
+ "node_rank",
+ "master_addr",
+ "master_port",
+]
+
+CUDA_DISTRIBUTED_TYPES = ["DEEPSPEED", "MULTI_GPU", "FSDP", "MEGATRON_LM", "TP"]
+TORCH_DISTRIBUTED_OPERATION_TYPES = CUDA_DISTRIBUTED_TYPES + [
+ "MULTI_NPU",
+ "MULTI_MLU",
+ "MULTI_SDAA",
+ "MULTI_MUSA",
+ "MULTI_XPU",
+ "MULTI_CPU",
+ "MULTI_HPU",
+]
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/dataclasses.py b/venv/lib/python3.11/site-packages/accelerate/utils/dataclasses.py
new file mode 100644
index 0000000000000000000000000000000000000000..23446b899db3d7aa4d469a381b484c4721aade7b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/dataclasses.py
@@ -0,0 +1,2775 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+"""
+General namespace and dataclass related classes
+"""
+
+import argparse
+import copy
+import enum
+import functools
+import logging
+import os
+import warnings
+from collections.abc import Iterable
+from contextlib import contextmanager
+from dataclasses import dataclass, field
+from datetime import timedelta
+from typing import TYPE_CHECKING, Any, Callable, Literal, Optional, Union, get_args
+
+import torch
+
+from .constants import (
+ BETA_TP_AVAILABLE_PYTORCH_VERSION,
+ FSDP_AUTO_WRAP_POLICY,
+ FSDP_BACKWARD_PREFETCH,
+ FSDP_SHARDING_STRATEGY,
+ MITA_PROFILING_AVAILABLE_PYTORCH_VERSION,
+ XPU_PROFILING_AVAILABLE_PYTORCH_VERSION,
+)
+from .environment import parse_flag_from_env, str_to_bool
+from .imports import (
+ is_cuda_available,
+ is_hpu_available,
+ is_mlu_available,
+ is_msamp_available,
+ is_musa_available,
+ is_npu_available,
+ is_transformer_engine_available,
+ is_xpu_available,
+)
+from .versions import compare_versions, is_torch_version
+
+
+if TYPE_CHECKING:
+ # Mock imports for type checking
+ from torchao.float8 import Float8LinearConfig
+
+logger = logging.getLogger(__name__)
+
+
+class KwargsHandler:
+ """
+ Internal mixin that implements a `to_kwargs()` method for a dataclass.
+ """
+
+ def to_dict(self):
+ return copy.deepcopy(self.__dict__)
+
+ def to_kwargs(self):
+ """
+ Returns a dictionary containing the attributes with values different from the default of this class.
+ """
+ # import clear_environment here to avoid circular import problem
+ from .environment import clear_environment
+
+ with clear_environment():
+ default_dict = self.__class__().to_dict()
+ this_dict = self.to_dict()
+ return {k: v for k, v in this_dict.items() if default_dict[k] != v}
+
+
+class EnumWithContains(enum.EnumMeta):
+ "A metaclass that adds the ability to check if `self` contains an item with the `in` operator"
+
+ def __contains__(cls, item):
+ try:
+ cls(item)
+ except ValueError:
+ return False
+ return True
+
+
+class BaseEnum(enum.Enum, metaclass=EnumWithContains):
+ "An enum class that can get the value of an item with `str(Enum.key)`"
+
+ def __str__(self):
+ return self.value
+
+ @classmethod
+ def list(cls):
+ "Method to list all the possible items in `cls`"
+ return list(map(str, cls))
+
+
+@dataclass
+class AutocastKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize how `torch.autocast` behaves. Please refer to the
+ documentation of this [context manager](https://pytorch.org/docs/stable/amp.html#torch.autocast) for more
+ information on each argument.
+
+ Example:
+
+ ```python
+ from accelerate import Accelerator
+ from accelerate.utils import AutocastKwargs
+
+ kwargs = AutocastKwargs(cache_enabled=True)
+ accelerator = Accelerator(kwargs_handlers=[kwargs])
+ ```
+ """
+
+ enabled: bool = True
+ cache_enabled: bool = None
+
+
+class DDPCommunicationHookType(BaseEnum):
+ """
+ Represents a type of communication hook used in DDP.
+
+ Values:
+
+ - **NO** -- no communication hook
+ - **FP16** -- DDP communication hook to compress the gradients in FP16
+ - **BF16** -- DDP communication hook to compress the gradients in BF16
+ - **POWER_SGD** -- DDP communication hook to use PowerSGD
+ - **BATCHED_POWER_SGD** -- DDP communication hook to use batched PowerSGD
+ """
+
+ NO = "no"
+ FP16 = "fp16"
+ BF16 = "bf16"
+ POWER_SGD = "power_sgd"
+ BATCHED_POWER_SGD = "batched_power_sgd"
+
+
+@dataclass
+class DistributedDataParallelKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize how your model is wrapped in a
+ `torch.nn.parallel.DistributedDataParallel`. Please refer to the documentation of this
+ [wrapper](https://pytorch.org/docs/stable/generated/torch.nn.parallel.DistributedDataParallel.html) for more
+ information on each argument.
+
+
+
+ `gradient_as_bucket_view` is only available in PyTorch 1.7.0 and later versions.
+
+ `static_graph` is only available in PyTorch 1.11.0 and later versions.
+
+
+
+ Example:
+
+ ```python
+ from accelerate import Accelerator
+ from accelerate.utils import DistributedDataParallelKwargs
+
+ kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
+ accelerator = Accelerator(kwargs_handlers=[kwargs])
+ ```
+ """
+
+ dim: int = 0
+ broadcast_buffers: bool = True
+ bucket_cap_mb: int = 25
+ find_unused_parameters: bool = False
+ check_reduction: bool = False
+ gradient_as_bucket_view: bool = False
+ static_graph: bool = False
+
+ comm_hook: DDPCommunicationHookType = DDPCommunicationHookType.NO
+ comm_wrapper: Literal[
+ DDPCommunicationHookType.NO, DDPCommunicationHookType.FP16, DDPCommunicationHookType.BF16
+ ] = DDPCommunicationHookType.NO
+ comm_state_option: dict = field(default_factory=dict)
+
+ def to_dict(self, ignore_keys=("comm_hook", "comm_wrapper", "comm_state_option")):
+ return {k: v for k, v in super().to_dict().items() if k not in ignore_keys}
+
+ def register_comm_hook(self, model):
+ from torch.distributed.algorithms.ddp_comm_hooks import default_hooks, powerSGD_hook
+
+ hook_map: dict[DDPCommunicationHookType, Callable] = {
+ DDPCommunicationHookType.FP16: default_hooks.fp16_compress_hook,
+ DDPCommunicationHookType.BF16: default_hooks.bf16_compress_hook,
+ DDPCommunicationHookType.POWER_SGD: powerSGD_hook.powerSGD_hook,
+ DDPCommunicationHookType.BATCHED_POWER_SGD: powerSGD_hook.batched_powerSGD_hook,
+ }
+
+ wrapper_map: dict[DDPCommunicationHookType, Callable] = {
+ DDPCommunicationHookType.FP16: default_hooks.fp16_compress_wrapper,
+ DDPCommunicationHookType.BF16: default_hooks.bf16_compress_wrapper,
+ }
+
+ hook: Optional[Callable] = hook_map.get(self.comm_hook)
+ wrapper: Optional[Callable] = wrapper_map.get(self.comm_wrapper)
+
+ if hook and wrapper:
+ hook = wrapper(hook)
+
+ if hook:
+ state = (
+ powerSGD_hook.PowerSGDState(None, **self.comm_state_option)
+ if self.comm_hook in (DDPCommunicationHookType.POWER_SGD, DDPCommunicationHookType.BATCHED_POWER_SGD)
+ else None
+ )
+ model.register_comm_hook(
+ state=state,
+ hook=hook,
+ )
+
+
+@dataclass
+class GradScalerKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize the behavior of mixed precision, specifically how the
+ `torch.cuda.amp.GradScaler` used is created. Please refer to the documentation of this
+ [scaler](https://pytorch.org/docs/stable/amp.html?highlight=gradscaler) for more information on each argument.
+
+
+
+ `GradScaler` is only available in PyTorch 1.5.0 and later versions.
+
+
+
+ Example:
+
+ ```python
+ from accelerate import Accelerator
+ from accelerate.utils import GradScalerKwargs
+
+ kwargs = GradScalerKwargs(backoff_factor=0.25)
+ accelerator = Accelerator(kwargs_handlers=[kwargs])
+ ```
+ """
+
+ init_scale: float = 65536.0
+ growth_factor: float = 2.0
+ backoff_factor: float = 0.5
+ growth_interval: int = 2000
+ enabled: bool = True
+
+
+@dataclass
+class InitProcessGroupKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize the initialization of the distributed processes. Please refer
+ to the documentation of this
+ [method](https://pytorch.org/docs/stable/distributed.html#torch.distributed.init_process_group) for more
+ information on each argument.
+
+ Note: If `timeout` is set to `None`, the default will be based upon how `backend` is set.
+
+ ```python
+ from datetime import timedelta
+ from accelerate import Accelerator
+ from accelerate.utils import InitProcessGroupKwargs
+
+ kwargs = InitProcessGroupKwargs(timeout=timedelta(seconds=800))
+ accelerator = Accelerator(kwargs_handlers=[kwargs])
+ ```
+ """
+
+ backend: Optional[str] = "nccl"
+ init_method: Optional[str] = None
+ timeout: Optional[timedelta] = None
+
+ def __post_init__(self):
+ if self.timeout is None:
+ seconds = 1800 if self.backend != "nccl" else 600
+ self.timeout = timedelta(seconds=seconds)
+
+
+# Literals
+Backend = Literal["MSAMP", "TE"]
+OptLevel = Literal["O1", "O2"]
+FP8Format = Literal["E4M3", "HYBRID"]
+AmaxComputeAlgorithm = Literal["max", "most_recent"]
+
+
+# FP8 training recipe kwargs
+@dataclass
+class AORecipeKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize the initialization of the recipe for FP8 mixed precision
+ training with `torchao` FP8.
+
+ Args:
+ config (`torchao.float8.Float8LinearConfig`, *optional*, default to `None`):
+ The configuration for the FP8 training. In general, the default config should be sufficient.
+ module_filter_func (`Callable`, *optional*, default to `None`):
+ Optional function that must take in a module and layer name, and returns a boolean indicating whether the
+ module should be converted to FP8. Defaults to `accelerate.utils.ao.filter_linear_layers`. See it for an
+ example.
+ """
+
+ config: Optional["Float8LinearConfig"] = None
+ module_filter_func: Optional[Callable] = None
+
+
+@dataclass
+class TERecipeKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize the initialization of the recipe for FP8 mixed precision
+ training with `transformer-engine`.
+
+
+
+ For more information on the args, please refer to the API
+ [documentation](https://docs.nvidia.com/deeplearning/transformer-engine/user-guide/api/common.html).
+
+
+
+ ```python
+ from accelerate import Accelerator
+ from accelerate.utils import TERecipeKwargs
+
+ kwargs = TERecipeKwargs(fp8_format="HYBRID")
+ accelerator = Accelerator(mixed_precision="fp8", kwargs_handlers=[kwargs])
+ ```
+
+ Args:
+ use_autocast_during_eval (`bool`, *optional*, default to `False`):
+ Whether to use FP8 autocast during eval mode. Generally better metrics are found when this is `False`.
+ margin (`int`, *optional*, default to 0):
+ The margin to use for the gradient scaling.
+ interval (`int`, *optional*, default to 1):
+ The interval to use for how often the scaling factor is recomputed.
+ fp8_format (`str`, *optional*, default to "HYBRID"):
+ The format to use for the FP8 recipe. Must be one of `HYBRID` or `E4M3`. (Generally `HYBRID` for training,
+ `E4M3` for evaluation)
+ amax_history_len (`int`, *optional*, default to 1024):
+ The length of the history to use for the scaling factor computation
+ amax_compute_algo (`str`, *optional*, default to "most_recent"):
+ The algorithm to use for the scaling factor computation. Must be one of `max` or `most_recent`.
+ override_linear_precision (`tuple` of three `bool`, *optional*, default to `(False, False, False)`):
+ Whether or not to execute `fprop`, `dgrad`, and `wgrad` GEMMS in higher precision.
+ """
+
+ use_autocast_during_eval: bool = None
+ margin: int = None
+ interval: int = None
+ fp8_format: FP8Format = None
+ amax_history_len: int = None
+ amax_compute_algo: AmaxComputeAlgorithm = None
+ override_linear_precision: tuple[bool, bool, bool] = None
+
+ def __post_init__(self):
+ env_prefix = "ACCELERATE_FP8_"
+ if not is_transformer_engine_available():
+ raise ImportError("TransformerEngine is not available. Please install it or use a different backend.")
+ if self.use_autocast_during_eval is None:
+ self.use_autocast_during_eval = parse_flag_from_env(env_prefix + "USE_AUTOCAST_DURING_EVAL")
+ if self.margin is None:
+ self.margin = int(os.environ.get(env_prefix + "MARGIN", 0))
+ if self.interval is None:
+ self.interval = int(os.environ.get(env_prefix + "INTERVAL", 1))
+ if self.fp8_format is None:
+ self.fp8_format = os.environ.get(env_prefix + "FORMAT", "HYBRID")
+ self.fp8_format = self.fp8_format.upper()
+ if self.fp8_format not in get_args(FP8Format):
+ raise ValueError(f"`fp8_format` must be one of {' or '.join(get_args(FP8Format))}.")
+ if self.amax_compute_algo is None:
+ self.amax_compute_algo = os.environ.get(env_prefix + "AMAX_COMPUTE_ALGO", "most_recent")
+ self.amax_compute_algo = self.amax_compute_algo.lower()
+ if self.amax_compute_algo not in get_args(AmaxComputeAlgorithm):
+ raise ValueError(f"`amax_compute_algo` must be one of {' or '.join(get_args(AmaxComputeAlgorithm))}")
+ if self.amax_history_len is None:
+ self.amax_history_len = int(os.environ.get(env_prefix + "AMAX_HISTORY_LEN", 1024))
+ if self.override_linear_precision is None:
+ fprop = parse_flag_from_env(env_prefix + "OVERRIDE_FPROP")
+ dgrad = parse_flag_from_env(env_prefix + "OVERRIDE_DGRAD")
+ wgrad = parse_flag_from_env(env_prefix + "OVERRIDE_WGRAD")
+ self.override_linear_precision = (fprop, dgrad, wgrad)
+
+
+@dataclass
+class MSAMPRecipeKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize the initialization of the recipe for FP8 mixed precision
+ training with `ms-amp`.
+ """
+
+ opt_level: OptLevel = None
+
+ def __post_init__(self):
+ env_prefix = "ACCELERATE_FP8_"
+ if self.opt_level is None:
+ self.opt_level = os.environ.get(env_prefix + "OPT_LEVEL", "O2")
+ if self.opt_level not in get_args(OptLevel):
+ raise ValueError(f"`opt_level` must be one of {' or '.join(get_args(OptLevel))}")
+
+
+@dataclass
+class FP8RecipeKwargs(TERecipeKwargs, MSAMPRecipeKwargs):
+ """
+ Deprecated. Please use one of the proper FP8 recipe kwargs classes such as `TERecipeKwargs` or `MSAMPRecipeKwargs`
+ instead.
+ """
+
+ backend: Backend = None
+
+ def __post_init__(self):
+ env_prefix = "ACCELERATE_FP8_"
+ warnings.warn(
+ "FP8RecipeKwargs is deprecated and will be removed in Accelerate v2.0.0. "
+ "Please use one of the proper FP8 recipe kwargs classes such as TERecipeKwargs or MSAMPRecipeKwargs instead.",
+ FutureWarning,
+ )
+ default_backend = "msamp" if is_msamp_available() else "te"
+ if self.backend is None:
+ self.backend = os.environ.get(env_prefix + "BACKEND", default_backend)
+ self.backend = self.backend.upper()
+ if self.backend not in get_args(Backend):
+ raise ValueError("`backend` must be 'MSAMP' or 'TE' (TransformerEngine) to use `FP8RecipeKwargs`.")
+ super().__post_init__()
+
+
+# Literal
+ProfilerActivity = Literal["cpu", "xpu", "mtia", "cuda", "hpu"]
+
+
+@dataclass
+class ProfileKwargs(KwargsHandler):
+ """
+ Use this object in your [`Accelerator`] to customize the initialization of the profiler. Please refer to the
+ documentation of this [context manager](https://pytorch.org/docs/stable/profiler.html#torch.profiler.profile) for
+ more information on each argument.
+
+
+
+ `torch.profiler` is only available in PyTorch 1.8.1 and later versions.
+
+
+
+ Example:
+
+ ```python
+ from accelerate import Accelerator
+ from accelerate.utils import ProfileKwargs
+
+ kwargs = ProfileKwargs(activities=["cpu", "cuda"])
+ accelerator = Accelerator(kwargs_handlers=[kwargs])
+ ```
+
+ Args:
+ activities (`List[str]`, *optional*, default to `None`):
+ The list of activity groups to use in profiling. Must be one of `"cpu"`, `"xpu"`, `"mtia"`, "hpu" or
+ `"cuda"`.
+ schedule_option (`Dict[str, int]`, *optional*, default to `None`):
+ The schedule option to use for the profiler. Available keys are `wait`, `warmup`, `active`, `repeat` and
+ `skip_first`. The profiler will skip the first `skip_first` steps, then wait for `wait` steps, then do the
+ warmup for the next `warmup` steps, then do the active recording for the next `active` steps and then
+ repeat the cycle starting with `wait` steps. The optional number of cycles is specified with the `repeat`
+ parameter, the zero value means that the cycles will continue until the profiling is finished.
+ on_trace_ready (`Callable`, *optional*, default to `None`):
+ Callable that is called at each step when schedule returns `ProfilerAction.RECORD_AND_SAVE` during the
+ profiling.
+ record_shapes (`bool`, *optional*, default to `False`):
+ Save information about operator’s input shapes.
+ profile_memory (`bool`, *optional*, default to `False`):
+ Track tensor memory allocation/deallocation
+ with_stack (`bool`, *optional*, default to `False`):
+ Record source information (file and line number) for the ops.
+ with_flops (`bool`, *optional*, default to `False`):
+ Use formula to estimate the FLOPS of specific operators
+ with_modules (`bool`, *optional*, default to `False`):
+ Record module hierarchy (including function names) corresponding to the callstack of the op.
+ output_trace_dir (`str`, *optional*, default to `None`):
+ Exports the collected trace in Chrome JSON format. Chrome use 'chrome://tracing' view json file. Defaults
+ to None, which means profiling does not store json files.
+ """
+
+ activities: Optional[list[ProfilerActivity]] = None
+ schedule_option: Optional[dict[str, int]] = None
+ on_trace_ready: Optional[Callable] = None
+ record_shapes: bool = False
+ profile_memory: bool = False
+ with_stack: bool = False
+ with_flops: bool = False
+ with_modules: bool = False
+ output_trace_dir: Optional[str] = None
+
+ def _get_profiler_activity(self, activity: ProfilerActivity) -> torch.profiler.ProfilerActivity:
+ """Get the profiler activity from the string.
+
+ Args:
+ activity (str): The profiler activity name.
+
+ Returns:
+ torch.profiler.ProfilerActivity: The profiler activity.
+ """
+
+ profiler_activity_map: dict[str, torch.profiler.ProfilerActivity] = {
+ "cpu": torch.profiler.ProfilerActivity.CPU,
+ "cuda": torch.profiler.ProfilerActivity.CUDA,
+ }
+
+ if is_hpu_available():
+ profiler_activity_map["hpu"] = torch.profiler.ProfilerActivity.HPU
+
+ if is_torch_version(">=", XPU_PROFILING_AVAILABLE_PYTORCH_VERSION):
+ if torch.xpu.is_available():
+ profiler_activity_map["xpu"] = torch.profiler.ProfilerActivity.XPU
+
+ if is_torch_version(">=", MITA_PROFILING_AVAILABLE_PYTORCH_VERSION):
+ if torch.mtia.is_available():
+ profiler_activity_map["mtia"] = torch.profiler.ProfilerActivity.MTIA
+
+ if activity not in profiler_activity_map:
+ raise ValueError(f"Invalid profiler activity: {activity}. Must be one of {list(profiler_activity_map)}.")
+ return profiler_activity_map[activity]
+
+ def build(self) -> torch.profiler.profile:
+ """
+ Build a profiler object with the current configuration.
+
+ Returns:
+ torch.profiler.profile: The profiler object.
+ """
+ activities: Optional[list[ProfilerActivity]] = None
+ if self.activities is not None:
+ activities = [self._get_profiler_activity(activity) for activity in self.activities]
+ schedule: Optional[torch.profiler.schedule] = None
+ if self.schedule_option is not None:
+ schedule = torch.profiler.schedule(**self.schedule_option)
+
+ return torch.profiler.profile(
+ activities=activities,
+ schedule=schedule,
+ on_trace_ready=self.on_trace_ready,
+ record_shapes=self.record_shapes,
+ profile_memory=self.profile_memory,
+ with_stack=self.with_stack,
+ with_flops=self.with_flops,
+ with_modules=self.with_modules,
+ )
+
+
+class DistributedType(str, enum.Enum):
+ """
+ Represents a type of distributed environment.
+
+ Values:
+
+ - **NO** -- Not a distributed environment, just a single process.
+ - **MULTI_CPU** -- Distributed on multiple CPU nodes.
+ - **MULTI_GPU** -- Distributed on multiple GPUs.
+ - **MULTI_MLU** -- Distributed on multiple MLUs.
+ - **MULTI_SDAA** -- Distributed on multiple SDAAs.
+ - **MULTI_MUSA** -- Distributed on multiple MUSAs.
+ - **MULTI_NPU** -- Distributed on multiple NPUs.
+ - **MULTI_XPU** -- Distributed on multiple XPUs.
+ - **MULTI_HPU** -- Distributed on multiple HPUs.
+ - **DEEPSPEED** -- Using DeepSpeed.
+ - **XLA** -- Using TorchXLA.
+ """
+
+ # Subclassing str as well as Enum allows the `DistributedType` to be JSON-serializable out of the box.
+ NO = "NO"
+ MULTI_CPU = "MULTI_CPU"
+ MULTI_GPU = "MULTI_GPU"
+ MULTI_NPU = "MULTI_NPU"
+ MULTI_MLU = "MULTI_MLU"
+ MULTI_SDAA = "MULTI_SDAA"
+ MULTI_MUSA = "MULTI_MUSA"
+ MULTI_XPU = "MULTI_XPU"
+ DEEPSPEED = "DEEPSPEED"
+ FSDP = "FSDP"
+ TP = "TP"
+ XLA = "XLA"
+ MEGATRON_LM = "MEGATRON_LM"
+ MULTI_HPU = "MULTI_HPU"
+
+
+class SageMakerDistributedType(str, enum.Enum):
+ """
+ Represents a type of distributed environment.
+
+ Values:
+
+ - **NO** -- Not a distributed environment, just a single process.
+ - **DATA_PARALLEL** -- using sagemaker distributed data parallelism.
+ - **MODEL_PARALLEL** -- using sagemaker distributed model parallelism.
+ """
+
+ # Subclassing str as well as Enum allows the `SageMakerDistributedType` to be JSON-serializable out of the box.
+ NO = "NO"
+ DATA_PARALLEL = "DATA_PARALLEL"
+ MODEL_PARALLEL = "MODEL_PARALLEL"
+
+
+class FP8BackendType(str, enum.Enum):
+ """
+ Represents the backend used for FP8.
+
+ Values:
+
+ - **TE** -- using TransformerEngine.
+ - **MSAMP** -- using msamp.
+ """
+
+ # Subclassing str as well as Enum allows the `FP8BackendType` to be JSON-serializable out of the box.
+ TE = "TE"
+ MSAMP = "MSAMP"
+
+
+class ComputeEnvironment(str, enum.Enum):
+ """
+ Represents a type of the compute environment.
+
+ Values:
+
+ - **LOCAL_MACHINE** -- private/custom cluster hardware.
+ - **AMAZON_SAGEMAKER** -- Amazon SageMaker as compute environment.
+ """
+
+ # Subclassing str as well as Enum allows the `ComputeEnvironment` to be JSON-serializable out of the box.
+ LOCAL_MACHINE = "LOCAL_MACHINE"
+ AMAZON_SAGEMAKER = "AMAZON_SAGEMAKER"
+
+
+class DynamoBackend(str, BaseEnum):
+ """
+ Represents a dynamo backend (see https://pytorch.org/docs/stable/torch.compiler.html).
+
+ Values:
+
+ - **NO** -- Do not use torch dynamo.
+ - **EAGER** -- Uses PyTorch to run the extracted GraphModule. This is quite useful in debugging TorchDynamo
+ issues.
+ - **AOT_EAGER** -- Uses AotAutograd with no compiler, i.e, just using PyTorch eager for the AotAutograd's
+ extracted forward and backward graphs. This is useful for debugging, and unlikely to give speedups.
+ - **INDUCTOR** -- Uses TorchInductor backend with AotAutograd and cudagraphs by leveraging codegened Triton
+ kernels. [Read
+ more](https://dev-discuss.pytorch.org/t/torchinductor-a-pytorch-native-compiler-with-define-by-run-ir-and-symbolic-shapes/747)
+ - **AOT_TS_NVFUSER** -- nvFuser with AotAutograd/TorchScript. [Read
+ more](https://dev-discuss.pytorch.org/t/tracing-with-primitives-update-1-nvfuser-and-its-primitives/593)
+ - **NVPRIMS_NVFUSER** -- nvFuser with PrimTorch. [Read
+ more](https://dev-discuss.pytorch.org/t/tracing-with-primitives-update-1-nvfuser-and-its-primitives/593)
+ - **CUDAGRAPHS** -- cudagraphs with AotAutograd. [Read more](https://github.com/pytorch/torchdynamo/pull/757)
+ - **OFI** -- Uses Torchscript optimize_for_inference. Inference only. [Read
+ more](https://pytorch.org/docs/stable/generated/torch.jit.optimize_for_inference.html)
+ - **FX2TRT** -- Uses Nvidia TensorRT for inference optimizations. Inference only. [Read
+ more](https://github.com/pytorch/TensorRT/blob/master/docsrc/tutorials/getting_started_with_fx_path.rst)
+ - **ONNXRT** -- Uses ONNXRT for inference on CPU/GPU. Inference only. [Read more](https://onnxruntime.ai/)
+ - **TENSORRT** -- Uses ONNXRT to run TensorRT for inference optimizations. [Read
+ more](https://github.com/onnx/onnx-tensorrt)
+ - **AOT_TORCHXLA_TRACE_ONCE** -- Uses Pytorch/XLA with TorchDynamo optimization, for training. [Read
+ more](https://github.com/pytorch/xla/blob/r2.0/docs/dynamo.md)
+ - **TORCHXLA_TRACE_ONCE** -- Uses Pytorch/XLA with TorchDynamo optimization, for inference. [Read
+ more](https://github.com/pytorch/xla/blob/r2.0/docs/dynamo.md)
+ - **IPEX** -- Uses IPEX for inference on CPU. Inference only. [Read
+ more](https://github.com/intel/intel-extension-for-pytorch).
+ - **TVM** -- Uses Apach TVM for inference optimizations. [Read more](https://tvm.apache.org/)
+ - **HPU_BACKEND** -- Uses HPU backend for inference optimizations.
+
+ """
+
+ # Subclassing str as well as Enum allows the `SageMakerDistributedType` to be JSON-serializable out of the box.
+ NO = "NO"
+ EAGER = "EAGER"
+ AOT_EAGER = "AOT_EAGER"
+ INDUCTOR = "INDUCTOR"
+ AOT_TS_NVFUSER = "AOT_TS_NVFUSER"
+ NVPRIMS_NVFUSER = "NVPRIMS_NVFUSER"
+ CUDAGRAPHS = "CUDAGRAPHS"
+ OFI = "OFI"
+ FX2TRT = "FX2TRT"
+ ONNXRT = "ONNXRT"
+ TENSORRT = "TENSORRT"
+ AOT_TORCHXLA_TRACE_ONCE = "AOT_TORCHXLA_TRACE_ONCE"
+ TORCHXLA_TRACE_ONCE = "TORCHXLA_TRACE_ONCE"
+ IPEX = "IPEX"
+ TVM = "TVM"
+ HPU_BACKEND = "HPU_BACKEND"
+
+
+class LoggerType(BaseEnum):
+ """Represents a type of supported experiment tracker
+
+ Values:
+
+ - **ALL** -- all available trackers in the environment that are supported
+ - **TENSORBOARD** -- TensorBoard as an experiment tracker
+ - **WANDB** -- wandb as an experiment tracker
+ - **COMETML** -- comet_ml as an experiment tracker
+ - **DVCLIVE** -- dvclive as an experiment tracker
+ """
+
+ ALL = "all"
+ AIM = "aim"
+ TENSORBOARD = "tensorboard"
+ WANDB = "wandb"
+ COMETML = "comet_ml"
+ MLFLOW = "mlflow"
+ CLEARML = "clearml"
+ DVCLIVE = "dvclive"
+
+
+class PrecisionType(str, BaseEnum):
+ """Represents a type of precision used on floating point values
+
+ Values:
+
+ - **NO** -- using full precision (FP32)
+ - **FP16** -- using half precision
+ - **BF16** -- using brain floating point precision
+ """
+
+ NO = "no"
+ FP8 = "fp8"
+ FP16 = "fp16"
+ BF16 = "bf16"
+
+
+class RNGType(BaseEnum):
+ TORCH = "torch"
+ CUDA = "cuda"
+ MLU = "mlu"
+ SDAA = "sdaa"
+ MUSA = "musa"
+ NPU = "npu"
+ XLA = "xla"
+ XPU = "xpu"
+ HPU = "hpu"
+ GENERATOR = "generator"
+
+
+class CustomDtype(enum.Enum):
+ r"""
+ An enum that contains multiple custom dtypes that can be used for `infer_auto_device_map`.
+ """
+
+ FP8 = "fp8"
+ INT4 = "int4"
+ INT2 = "int2"
+
+
+# data classes
+
+
+@dataclass
+class TensorInformation:
+ shape: torch.Size
+ dtype: torch.dtype
+
+
+@dataclass
+class DataLoaderConfiguration:
+ """
+ Configuration for dataloader-related items when calling `accelerator.prepare`.
+
+ Args:
+ split_batches (`bool`, defaults to `False`):
+ Whether or not the accelerator should split the batches yielded by the dataloaders across the devices. If
+ `True`, the actual batch size used will be the same on any kind of distributed processes, but it must be a
+ round multiple of `num_processes` you are using. If `False`, actual batch size used will be the one set in
+ your script multiplied by the number of processes.
+ dispatch_batches (`bool`, defaults to `None`):
+ If set to `True`, the dataloader prepared by the Accelerator is only iterated through on the main process
+ and then the batches are split and broadcast to each process. Will default to `True` for `DataLoader` whose
+ underlying dataset is an `IterableDataset`, `False` otherwise.
+ even_batches (`bool`, defaults to `True`):
+ If set to `True`, in cases where the total batch size across all processes does not exactly divide the
+ dataset, samples at the start of the dataset will be duplicated so the batch can be divided equally among
+ all workers.
+ use_seedable_sampler (`bool`, defaults to `False`):
+ Whether or not use a fully seedable random sampler ([`data_loader.SeedableRandomSampler`]). Ensures
+ training results are fully reproducable using a different sampling technique. While seed-to-seed results
+ may differ, on average the differences are neglible when using multiple different seeds to compare. Should
+ also be ran with [`~utils.set_seed`] for the best results.
+ data_seed (`int`, defaults to `None`):
+ The seed to use for the underlying generator when using `use_seedable_sampler`. If `None`, the generator
+ will use the current default seed from torch.
+ non_blocking (`bool`, defaults to `False`):
+ If set to `True`, the dataloader prepared by the Accelerator will utilize non-blocking host-to-device
+ transfers, allowing for better overlap between dataloader communication and computation. Recommended that
+ the prepared dataloader has `pin_memory` set to `True` to work properly.
+ use_stateful_dataloader (`bool`, defaults to `False`):
+ If set to `True`, the dataloader prepared by the Accelerator will be backed by
+ [torchdata.StatefulDataLoader](https://github.com/pytorch/data/tree/main/torchdata/stateful_dataloader).
+ This requires `torchdata` version 0.8.0 or higher that supports StatefulDataLoader to be installed.
+ """
+
+ split_batches: bool = field(
+ default=False,
+ metadata={
+ "help": "Whether or not the accelerator should split the batches yielded by the dataloaders across the devices. If"
+ " `True` the actual batch size used will be the same on any kind of distributed processes, but it must be a"
+ " round multiple of the `num_processes` you are using. If `False`, actual batch size used will be the one set"
+ " in your script multiplied by the number of processes."
+ },
+ )
+ dispatch_batches: bool = field(
+ default=None,
+ metadata={
+ "help": "If set to `True`, the dataloader prepared by the Accelerator is only iterated through on the main process"
+ " and then the batches are split and broadcast to each process. Will default to `True` for `DataLoader` whose"
+ " underlying dataset is an `IterableDataset`, `False` otherwise."
+ },
+ )
+ even_batches: bool = field(
+ default=True,
+ metadata={
+ "help": "If set to `True`, in cases where the total batch size across all processes does not exactly divide the"
+ " dataset, samples at the start of the dataset will be duplicated so the batch can be divided equally among"
+ " all workers."
+ },
+ )
+ use_seedable_sampler: bool = field(
+ default=False,
+ metadata={
+ "help": "Whether or not use a fully seedable random sampler ([`data_loader.SeedableRandomSampler`])."
+ "Ensures training results are fully reproducable using a different sampling technique. "
+ "While seed-to-seed results may differ, on average the differences are neglible when using"
+ "multiple different seeds to compare. Should also be ran with [`~utils.set_seed`] for the best results."
+ },
+ )
+ data_seed: int = field(
+ default=None,
+ metadata={
+ "help": "The seed to use for the underlying generator when using `use_seedable_sampler`. If `None`, the generator"
+ " will use the current default seed from torch."
+ },
+ )
+ non_blocking: bool = field(
+ default=False,
+ metadata={
+ "help": "If set to `True`, the dataloader prepared by the Accelerator will utilize non-blocking host-to-device"
+ " transfers, allowing for better overlap between dataloader communication and computation. Recommended that the"
+ " prepared dataloader has `pin_memory` set to `True` to work properly."
+ },
+ )
+ use_stateful_dataloader: bool = field(
+ default=False,
+ metadata={
+ "help": "If set to `True`, the dataloader prepared by the Accelerator will be backed by "
+ "[torchdata.StatefulDataLoader](https://github.com/pytorch/data/tree/main/torchdata/stateful_dataloader). This requires `torchdata` version 0.8.0 or higher that supports StatefulDataLoader to be installed."
+ },
+ )
+
+
+@dataclass
+class ProjectConfiguration:
+ """
+ Configuration for the Accelerator object based on inner-project needs.
+
+ Args:
+ project_dir (`str`, defaults to `None`):
+ A path to a directory for storing data.
+ logging_dir (`str`, defaults to `None`):
+ A path to a directory for storing logs of locally-compatible loggers. If None, defaults to `project_dir`.
+ automatic_checkpoint_naming (`bool`, defaults to `False`):
+ Whether saved states should be automatically iteratively named.
+ total_limit (`int`, defaults to `None`):
+ The maximum number of total saved states to keep.
+ iteration (`int`, defaults to `0`):
+ The current save iteration.
+ save_on_each_node (`bool`, defaults to `False`):
+ When doing multi-node distributed training, whether to save models and checkpoints on each node, or only on
+ the main one.
+ """
+
+ project_dir: str = field(default=None, metadata={"help": "A path to a directory for storing data."})
+ logging_dir: str = field(
+ default=None,
+ metadata={
+ "help": "A path to a directory for storing logs of locally-compatible loggers. If None, defaults to `project_dir`."
+ },
+ )
+ automatic_checkpoint_naming: bool = field(
+ default=False,
+ metadata={"help": "Whether saved states should be automatically iteratively named."},
+ )
+
+ total_limit: int = field(
+ default=None,
+ metadata={"help": "The maximum number of total saved states to keep."},
+ )
+
+ iteration: int = field(
+ default=0,
+ metadata={"help": "The current save iteration."},
+ )
+
+ save_on_each_node: bool = field(
+ default=False,
+ metadata={
+ "help": (
+ "When doing multi-node distributed training, whether to save models and checkpoints on each node, or"
+ " only on the main one"
+ )
+ },
+ )
+
+ def set_directories(self, project_dir: str = None):
+ "Sets `self.project_dir` and `self.logging_dir` to the appropriate values."
+ self.project_dir = project_dir
+ if self.logging_dir is None:
+ self.logging_dir = project_dir
+
+ def __post_init__(self):
+ self.set_directories(self.project_dir)
+
+
+@dataclass
+class GradientAccumulationPlugin(KwargsHandler):
+ """
+ A plugin to configure gradient accumulation behavior. You can only pass one of `gradient_accumulation_plugin` or
+ `gradient_accumulation_steps` to [`Accelerator`]. Passing both raises an error.
+
+ Parameters:
+ num_steps (`int`):
+ The number of steps to accumulate gradients for.
+ adjust_scheduler (`bool`, *optional*, defaults to `True`):
+ Whether to adjust the scheduler steps to account for the number of steps being accumulated. Should be
+ `True` if the used scheduler was not adjusted for gradient accumulation.
+ sync_with_dataloader (`bool`, *optional*, defaults to `True`):
+ Whether to synchronize setting the gradients when at the end of the dataloader.
+ sync_each_batch (`bool`, *optional*):
+ Whether to synchronize setting the gradients at each data batch. Seting to `True` may reduce memory
+ requirements when using gradient accumulation with distributed training, at expense of speed.
+
+ Example:
+
+ ```python
+ from accelerate.utils import GradientAccumulationPlugin
+
+ gradient_accumulation_plugin = GradientAccumulationPlugin(num_steps=2)
+ accelerator = Accelerator(gradient_accumulation_plugin=gradient_accumulation_plugin)
+ ```
+ """
+
+ num_steps: int = field(default=None, metadata={"help": "The number of steps to accumulate gradients for."})
+ adjust_scheduler: bool = field(
+ default=True,
+ metadata={
+ "help": "Whether to adjust the scheduler steps to account for the number of steps being accumulated. Should be `True` if the used scheduler was not adjusted for gradient accumulation."
+ },
+ )
+ sync_with_dataloader: bool = field(
+ default=True,
+ metadata={
+ "help": "Whether to synchronize setting the gradients when at the end of the dataloader. Should only be set to `False` if you know what you're doing."
+ },
+ )
+ sync_each_batch: bool = field(
+ default=False,
+ metadata={
+ "help": "Whether to synchronize setting the gradients at each data batch. Setting to `True` may reduce memory requirements when using gradient accumulation with distributed training, at expense of speed."
+ },
+ )
+
+
+@dataclass
+class TorchDynamoPlugin(KwargsHandler):
+ """
+ This plugin is used to compile a model with PyTorch 2.0
+
+ Args:
+ backend (`DynamoBackend`, defaults to `None`):
+ A valid Dynamo backend. See https://pytorch.org/docs/stable/torch.compiler.html for more details.
+ mode (`str`, defaults to `None`):
+ Possible options are 'default', 'reduce-overhead' or 'max-autotune'.
+ fullgraph (`bool`, defaults to `None`):
+ Whether it is ok to break model into several subgraphs.
+ dynamic (`bool`, defaults to `None`):
+ Whether to use dynamic shape for tracing.
+ options (`Any`, defaults to `None`):
+ A dictionary of options to pass to the backend.
+ disable (`bool`, defaults to `False`):
+ Turn torch.compile() into a no-op for testing
+ """
+
+ backend: DynamoBackend = field(
+ default=None,
+ metadata={"help": f"Possible options are {[b.value.lower() for b in DynamoBackend]}"},
+ )
+ mode: str = field(
+ default=None, metadata={"help": "Possible options are 'default', 'reduce-overhead' or 'max-autotune'"}
+ )
+ fullgraph: bool = field(default=None, metadata={"help": "Whether it is ok to break model into several subgraphs"})
+ dynamic: bool = field(default=None, metadata={"help": "Whether to use dynamic shape for tracing"})
+ options: Any = field(default=None, metadata={"help": "A dictionary of options to pass to the backend."})
+ disable: bool = field(default=False, metadata={"help": "Turn torch.compile() into a no-op for testing"})
+
+ def __post_init__(self):
+ prefix = "ACCELERATE_DYNAMO_"
+ if self.backend is None:
+ self.backend = os.environ.get(prefix + "BACKEND", "no")
+ self.backend = DynamoBackend(self.backend.upper())
+ if self.mode is None:
+ self.mode = os.environ.get(prefix + "MODE", "default")
+ if self.fullgraph is None:
+ self.fullgraph = str_to_bool(os.environ.get(prefix + "USE_FULLGRAPH", "False")) == 1
+ if self.dynamic is None:
+ self.dynamic = str_to_bool(os.environ.get(prefix + "USE_DYNAMIC", "False")) == 1
+
+ def to_dict(self):
+ dynamo_config = copy.deepcopy(self.__dict__)
+ dynamo_config["backend"] = dynamo_config["backend"].value.lower()
+ return dynamo_config
+
+
+@dataclass
+class DeepSpeedPlugin:
+ """
+ This plugin is used to integrate DeepSpeed.
+
+ Args:
+ hf_ds_config (`Any`, defaults to `None`):
+ Path to DeepSpeed config file or dict or an object of class `accelerate.utils.deepspeed.HfDeepSpeedConfig`.
+ gradient_accumulation_steps (`int`, defaults to `None`):
+ Number of steps to accumulate gradients before updating optimizer states. If not set, will use the value
+ from the `Accelerator` directly.
+ gradient_clipping (`float`, defaults to `None`):
+ Enable gradient clipping with value.
+ zero_stage (`int`, defaults to `None`):
+ Possible options are 0, 1, 2, 3. Default will be taken from environment variable.
+ is_train_batch_min (`bool`, defaults to `True`):
+ If both train & eval dataloaders are specified, this will decide the `train_batch_size`.
+ offload_optimizer_device (`str`, defaults to `None`):
+ Possible options are none|cpu|nvme. Only applicable with ZeRO Stages 2 and 3.
+ offload_param_device (`str`, defaults to `None`):
+ Possible options are none|cpu|nvme. Only applicable with ZeRO Stage 3.
+ offload_optimizer_nvme_path (`str`, defaults to `None`):
+ Possible options are /nvme|/local_nvme. Only applicable with ZeRO Stage 3.
+ offload_param_nvme_path (`str`, defaults to `None`):
+ Possible options are /nvme|/local_nvme. Only applicable with ZeRO Stage 3.
+ zero3_init_flag (`bool`, defaults to `None`):
+ Flag to indicate whether to save 16-bit model. Only applicable with ZeRO Stage-3.
+ zero3_save_16bit_model (`bool`, defaults to `None`):
+ Flag to indicate whether to save 16-bit model. Only applicable with ZeRO Stage-3.
+ transformer_moe_cls_names (`str`, defaults to `None`):
+ Comma-separated list of Transformers MoE layer class names (case-sensitive). For example,
+ `MixtralSparseMoeBlock`, `Qwen2MoeSparseMoeBlock`, `JetMoEAttention`, `JetMoEBlock`, etc.
+ enable_msamp (`bool`, defaults to `None`):
+ Flag to indicate whether to enable MS-AMP backend for FP8 training.
+ msasmp_opt_level (`Optional[Literal["O1", "O2"]]`, defaults to `None`):
+ Optimization level for MS-AMP (defaults to 'O1'). Only applicable if `enable_msamp` is True. Should be one
+ of ['O1' or 'O2'].
+ """
+
+ hf_ds_config: Any = field(
+ default=None,
+ metadata={
+ "help": "path to DeepSpeed config file or dict or an object of class `accelerate.utils.deepspeed.HfDeepSpeedConfig`."
+ },
+ )
+ gradient_accumulation_steps: int = field(
+ default=None,
+ metadata={
+ "help": "Number of steps to accumulate gradients before updating optimizer states. If not set, will use the value from the `Accelerator` directly."
+ },
+ )
+ gradient_clipping: float = field(default=None, metadata={"help": "Enable gradient clipping with value"})
+ zero_stage: int = field(
+ default=None,
+ metadata={"help": "Possible options are 0,1,2,3; Default will be taken from environment variable"},
+ )
+ is_train_batch_min: bool = field(
+ default=True,
+ metadata={"help": "If both train & eval dataloaders are specified, this will decide the train_batch_size"},
+ )
+ offload_optimizer_device: str = field(
+ default=None,
+ metadata={"help": "Possible options are none|cpu|nvme. Only applicable with ZeRO Stages 2 and 3."},
+ )
+ offload_param_device: str = field(
+ default=None,
+ metadata={"help": "Possible options are none|cpu|nvme. Only applicable with ZeRO Stage 3."},
+ )
+ offload_optimizer_nvme_path: str = field(
+ default=None,
+ metadata={"help": "Possible options are /nvme|/local_nvme. Only applicable with ZeRO Stage 3."},
+ )
+ offload_param_nvme_path: str = field(
+ default=None,
+ metadata={"help": "Possible options are /nvme|/local_nvme. Only applicable with ZeRO Stage 3."},
+ )
+ zero3_init_flag: bool = field(
+ default=None,
+ metadata={
+ "help": "Flag to indicate whether to enable `deepspeed.zero.Init` for constructing massive models."
+ "Only applicable with ZeRO Stage-3."
+ },
+ )
+ zero3_save_16bit_model: bool = field(
+ default=None,
+ metadata={"help": "Flag to indicate whether to save 16-bit model. Only applicable with ZeRO Stage-3."},
+ )
+ transformer_moe_cls_names: str = field(
+ default=None,
+ metadata={
+ "help": "comma-separated list of transformers MoE layer class names (case-sensitive), e.g : "
+ " `MixtralSparseMoeBlock`, `Qwen2MoeSparseMoeBlock`, `JetMoEAttention,JetMoEBlock` ..."
+ },
+ )
+ enable_msamp: bool = field(
+ default=None,
+ metadata={"help": "Flag to indicate whether to enable MS-AMP backend for FP8 training."},
+ )
+ msamp_opt_level: Optional[Literal["O1", "O2"]] = field(
+ default=None,
+ metadata={
+ "help": "Optimization level for MS-AMP (defaults to 'O1'). Only applicable if `enable_msamp` is True. Should be one of ['O1' or 'O2']."
+ },
+ )
+
+ def __post_init__(self):
+ from .deepspeed import HfDeepSpeedConfig
+
+ if self.gradient_accumulation_steps is None:
+ gas = os.environ.get("ACCELERATE_GRADIENT_ACCUMULATION_STEPS", "auto")
+ self.gradient_accumulation_steps = int(gas) if gas.isdigit() else gas
+
+ if self.gradient_clipping is None:
+ gradient_clipping = os.environ.get("ACCELERATE_GRADIENT_CLIPPING", "auto")
+ self.gradient_clipping = gradient_clipping if gradient_clipping == "auto" else float(gradient_clipping)
+
+ if self.zero_stage is None:
+ self.zero_stage = int(os.environ.get("ACCELERATE_DEEPSPEED_ZERO_STAGE", 2))
+
+ if self.offload_optimizer_device is None:
+ self.offload_optimizer_device = os.environ.get("ACCELERATE_DEEPSPEED_OFFLOAD_OPTIMIZER_DEVICE", "none")
+
+ if self.offload_param_device is None:
+ self.offload_param_device = os.environ.get("ACCELERATE_DEEPSPEED_OFFLOAD_PARAM_DEVICE", "none")
+
+ if self.offload_optimizer_nvme_path is None:
+ self.offload_optimizer_nvme_path = os.environ.get(
+ "ACCELERATE_DEEPSPEED_OFFLOAD_OPTIMIZER_NVME_PATH", "none"
+ )
+
+ if self.offload_param_nvme_path is None:
+ self.offload_param_nvme_path = os.environ.get("ACCELERATE_DEEPSPEED_OFFLOAD_PARAM_NVME_PATH", "none")
+
+ if self.zero3_save_16bit_model is None:
+ self.zero3_save_16bit_model = (
+ os.environ.get("ACCELERATE_DEEPSPEED_ZERO3_SAVE_16BIT_MODEL", "false") == "true"
+ )
+ if self.enable_msamp is None:
+ self.enable_msamp = os.environ.get("ACCELERATE_FP8_BACKEND", None) == "MSAMP"
+
+ if self.msamp_opt_level is None:
+ self.msamp_opt_level = os.environ.get("ACCELERATE_FP8_OPT_LEVEL", "O1")
+
+ if self.hf_ds_config is None:
+ self.hf_ds_config = os.environ.get("ACCELERATE_DEEPSPEED_CONFIG_FILE", "none")
+ if (
+ isinstance(self.hf_ds_config, dict)
+ or (isinstance(self.hf_ds_config, str) and self.hf_ds_config != "none")
+ or isinstance(self.hf_ds_config, HfDeepSpeedConfig)
+ ):
+ if not isinstance(self.hf_ds_config, HfDeepSpeedConfig):
+ self.hf_ds_config = HfDeepSpeedConfig(self.hf_ds_config)
+ if "gradient_accumulation_steps" not in self.hf_ds_config.config:
+ self.hf_ds_config.config["gradient_accumulation_steps"] = 1
+ if "zero_optimization" not in self.hf_ds_config.config:
+ raise ValueError("Please specify the ZeRO optimization config in the DeepSpeed config.")
+
+ self._deepspeed_config_checks()
+ plugin_to_config_mapping = {
+ "gradient_accumulation_steps": "gradient_accumulation_steps",
+ "gradient_clipping": "gradient_clipping",
+ "zero_stage": "zero_optimization.stage",
+ "offload_optimizer_device": "zero_optimization.offload_optimizer.device",
+ "offload_param_device": "zero_optimization.offload_param.device",
+ "offload_param_nvme_path": "zero_optimization.offload_param.nvme_path",
+ "offload_optimizer_nvme_path": "zero_optimization.offload_optimizer.nvme_path",
+ "zero3_save_16bit_model": "zero_optimization.stage3_gather_16bit_weights_on_model_save",
+ }
+ kwargs = {v: getattr(self, k) for k, v in plugin_to_config_mapping.items() if getattr(self, k) is not None}
+ for key in kwargs.keys():
+ self.fill_match(key, **kwargs, must_match=False)
+ self.hf_ds_config.set_stage_and_offload()
+
+ # filling the missing values in the class attributes from the DeepSpeed config
+ # when using the DeepSpeed config file.
+ for key, value in plugin_to_config_mapping.items():
+ config_value = self.hf_ds_config.get_value(value)
+ if config_value is not None and config_value != "auto":
+ setattr(self, key, config_value)
+ else:
+ config = {
+ "train_batch_size": "auto",
+ "train_micro_batch_size_per_gpu": "auto",
+ "gradient_accumulation_steps": self.gradient_accumulation_steps,
+ "zero_optimization": {
+ "stage": self.zero_stage,
+ "offload_optimizer": {
+ "device": self.offload_optimizer_device,
+ "nvme_path": self.offload_optimizer_nvme_path
+ if self.offload_optimizer_device == "nvme"
+ else None,
+ },
+ "offload_param": {
+ "device": self.offload_param_device,
+ "nvme_path": self.offload_param_nvme_path if self.offload_param_device == "nvme" else None,
+ },
+ "stage3_gather_16bit_weights_on_model_save": self.zero3_save_16bit_model,
+ },
+ }
+ if self.gradient_clipping:
+ config["gradient_clipping"] = self.gradient_clipping
+ self.hf_ds_config = HfDeepSpeedConfig(config)
+
+ self.deepspeed_config = self.hf_ds_config.config
+ self.deepspeed_config["steps_per_print"] = float("inf") # this will stop deepspeed from logging @ stdout
+ if self.zero3_init_flag is None:
+ self.zero3_init_flag = (
+ str_to_bool(os.environ.get("ACCELERATE_DEEPSPEED_ZERO3_INIT", str(self.hf_ds_config.is_zero3()))) == 1
+ )
+ if self.zero3_init_flag and not self.hf_ds_config.is_zero3():
+ warnings.warn("DeepSpeed Zero3 Init flag is only applicable for ZeRO Stage 3. Setting it to False.")
+ self.zero3_init_flag = False
+ # NOTE: Set to False by default, will be set to `True` automatically if it's the first plugin passed
+ # to the `Accelerator`'s `deepspeed_plugin` param, *or* `AcceleratorState().enable_deepspeed_plugin(plugin_key)` is manually called
+ self._set_selected(False)
+
+ # Ignore if it's already set
+ if self.enable_msamp and "msamp" not in self.deepspeed_config:
+ if self.zero_stage == 3:
+ raise NotImplementedError(
+ "MS-AMP is not supported for ZeRO Stage 3. Please use ZeRO Stage 0, 1, or 2 instead."
+ )
+ if self.msamp_opt_level not in ["O1", "O2"]:
+ raise ValueError("Invalid optimization level for MS-AMP. Please use one of ['O1' or'O2'].")
+ self.deepspeed_config["msamp"] = {"enabled": True, "opt_level": self.msamp_opt_level}
+
+ def fill_match(self, ds_key_long, mismatches=None, must_match=True, **kwargs):
+ mismatches = [] if mismatches is None else mismatches
+ config, ds_key = self.hf_ds_config.find_config_node(ds_key_long)
+ if config is None:
+ return
+
+ if config.get(ds_key) == "auto":
+ if ds_key_long in kwargs:
+ config[ds_key] = kwargs[ds_key_long]
+ return
+ else:
+ raise ValueError(
+ f"`{ds_key_long}` not found in kwargs. "
+ f"Please specify `{ds_key_long}` without `auto` (set to correct value) in the DeepSpeed config file or "
+ "pass it in kwargs."
+ )
+
+ if not must_match:
+ return
+
+ ds_val = config.get(ds_key)
+ if ds_val is not None and ds_key_long in kwargs:
+ if ds_val != kwargs[ds_key_long]:
+ mismatches.append(f"- ds {ds_key_long}={ds_val} vs arg {ds_key_long}={kwargs[ds_key_long]}")
+
+ def is_auto(self, ds_key_long):
+ val = self.hf_ds_config.get_value(ds_key_long)
+ if val is None:
+ return False
+ else:
+ return val == "auto"
+
+ def get_value(self, ds_key_long, default=None):
+ return self.hf_ds_config.get_value(ds_key_long, default)
+
+ def deepspeed_config_process(self, prefix="", mismatches=None, config=None, must_match=True, **kwargs):
+ """Process the DeepSpeed config with the values from the kwargs."""
+ mismatches = [] if mismatches is None else mismatches
+ if config is None:
+ config = self.deepspeed_config
+ for key, value in config.items():
+ if isinstance(value, dict):
+ self.deepspeed_config_process(
+ prefix=prefix + key + ".", mismatches=mismatches, config=value, must_match=must_match, **kwargs
+ )
+ else:
+ self.fill_match(prefix + key, mismatches, must_match=must_match, **kwargs)
+ if len(mismatches) > 0 and prefix == "":
+ mismatches_msg = "\n".join(mismatches)
+ raise ValueError(
+ "Please correct the following DeepSpeed config values that mismatch kwargs "
+ f" values:\n{mismatches_msg}\nThe easiest method is to set these DeepSpeed config values to 'auto'."
+ )
+
+ def set_mixed_precision(self, mixed_precision):
+ ds_config = self.deepspeed_config
+ kwargs = {
+ "fp16.enabled": mixed_precision == "fp16",
+ # When training in fp8, we still rely on bf16 autocast for the core mixed precision
+ "bf16.enabled": mixed_precision in ("bf16", "fp8"),
+ }
+ if mixed_precision == "fp16":
+ if "fp16" not in ds_config:
+ ds_config["fp16"] = {"enabled": True, "auto_cast": True}
+ elif mixed_precision in ("bf16", "fp8"):
+ if "bf16" not in ds_config:
+ ds_config["bf16"] = {"enabled": True}
+
+ if mixed_precision == "fp8" and self.enable_msamp:
+ if "msamp" not in ds_config:
+ ds_config["msamp"] = {"enabled": True, "opt_level": self.msamp_opt_level}
+
+ if mixed_precision != "no":
+ diff_dtype = "bf16" if mixed_precision == "fp16" else "fp16"
+ if str(ds_config.get(diff_dtype, {}).get("enabled", "False")).lower() == "true":
+ raise ValueError(
+ f"`--mixed_precision` arg cannot be set to `{mixed_precision}` when `{diff_dtype}` is set in the DeepSpeed config file."
+ )
+ for dtype in ["fp16", "bf16"]:
+ if dtype not in ds_config:
+ ds_config[dtype] = {"enabled": False}
+ self.fill_match("fp16.enabled", must_match=False, **kwargs)
+ self.fill_match("bf16.enabled", must_match=False, **kwargs)
+
+ def set_deepspeed_weakref(self):
+ from .imports import is_transformers_available
+
+ ds_config = copy.deepcopy(self.deepspeed_config)
+ if self.zero3_init_flag:
+ if not is_transformers_available():
+ raise Exception(
+ "When `zero3_init_flag` is set, it requires Transformers to be installed. "
+ "Please run `pip install transformers`."
+ )
+ if "gradient_accumulation_steps" not in ds_config or ds_config["gradient_accumulation_steps"] == "auto":
+ ds_config["gradient_accumulation_steps"] = 1
+ if "train_micro_batch_size_per_gpu" not in ds_config or ds_config["train_micro_batch_size_per_gpu"] == "auto":
+ ds_config["train_micro_batch_size_per_gpu"] = 1
+ if ds_config.get("train_batch_size", None) == "auto":
+ del ds_config["train_batch_size"]
+
+ if compare_versions("transformers", "<", "4.46"):
+ from transformers.deepspeed import HfDeepSpeedConfig, unset_hf_deepspeed_config
+ else:
+ from transformers.integrations import HfDeepSpeedConfig, unset_hf_deepspeed_config
+
+ unset_hf_deepspeed_config()
+ self.dschf = HfDeepSpeedConfig(ds_config) # keep this object alive # noqa
+
+ def is_zero3_init_enabled(self):
+ return self.zero3_init_flag
+
+ @contextmanager
+ def zero3_init_context_manager(self, enable=False):
+ old = self.zero3_init_flag
+ if old == enable:
+ yield
+ else:
+ self.zero3_init_flag = enable
+ self.dschf = None
+ self.set_deepspeed_weakref()
+ yield
+ self.zero3_init_flag = old
+ self.dschf = None
+ self.set_deepspeed_weakref()
+
+ def _deepspeed_config_checks(self):
+ env_variable_names_to_ignore = [
+ "ACCELERATE_GRADIENT_ACCUMULATION_STEPS",
+ "ACCELERATE_GRADIENT_CLIPPING",
+ "ACCELERATE_DEEPSPEED_ZERO_STAGE",
+ "ACCELERATE_DEEPSPEED_OFFLOAD_OPTIMIZER_DEVICE",
+ "ACCELERATE_DEEPSPEED_OFFLOAD_PARAM_DEVICE",
+ "ACCELERATE_DEEPSPEED_OFFLOAD_PARAM_NVME_PATH",
+ "ACCELERATE_DEEPSPEED_OFFLOAD_OPTIMIZER_NVME_PATH",
+ "ACCELERATE_DEEPSPEED_ZERO3_SAVE_16BIT_MODEL",
+ "ACCELERATE_MIXED_PRECISION",
+ ]
+ env_variable_names_to_ignore = [
+ name.replace("ACCELERATE_", "").replace("DEEPSPEED_", "").lower() for name in env_variable_names_to_ignore
+ ]
+
+ deepspeed_fields_from_accelerate_config = os.environ.get("ACCELERATE_CONFIG_DS_FIELDS", "").split(",")
+
+ if any(name in env_variable_names_to_ignore for name in deepspeed_fields_from_accelerate_config):
+ raise ValueError(
+ f"When using `deepspeed_config_file`, the following accelerate config variables will be ignored: {env_variable_names_to_ignore}.\n"
+ "Please specify them appropriately in the DeepSpeed config file.\n"
+ "If you are using an accelerate config file, remove others config variables mentioned in the above specified list.\n"
+ "The easiest method is to create a new config following the questionnaire via `accelerate config`.\n"
+ "It will only ask for the necessary config variables when using `deepspeed_config_file`."
+ )
+
+ def set_moe_leaf_modules(self, model):
+ if self.transformer_moe_cls_names is None:
+ self.transformer_moe_cls_names = os.environ.get("ACCELERATE_DEEPSPEED_MOE_LAYER_CLS_NAMES", None)
+ if self.transformer_moe_cls_names is not None:
+ if compare_versions("deepspeed", "<", "0.14.0"):
+ raise ImportError("DeepSpeed version must be >= 0.14.0 to use MOE support. Please update DeepSpeed.")
+ from deepspeed.utils import set_z3_leaf_modules
+
+ class_names = self.transformer_moe_cls_names.split(",")
+ transformer_moe_cls = []
+ for layer_class in class_names:
+ transformer_cls = get_module_class_from_name(model, layer_class)
+ if transformer_cls is None:
+ raise Exception(
+ f"Could not find a transformer layer class called '{layer_class}' to wrap in the model."
+ )
+ else:
+ transformer_moe_cls.append(transformer_cls)
+ set_z3_leaf_modules(model, transformer_moe_cls) # z3_leaf
+
+ def select(self, _from_accelerator_state: bool = False):
+ """
+ Sets the HfDeepSpeedWeakref to use the current deepspeed plugin configuration
+ """
+ if not _from_accelerator_state:
+ raise ValueError(
+ "A `DeepSpeedPlugin` object must be enabled manually by calling `AcceleratorState().enable_deepspeed_plugin(plugin_key)`."
+ )
+ self.set_deepspeed_weakref()
+ self._set_selected(True)
+
+ def _unselect(self):
+ self._set_selected(False)
+
+ def _set_selected(self, value: bool):
+ """
+ Private setter for the 'enabled' attribute.
+ """
+ self._selected = value
+
+ @property
+ def selected(self):
+ return self._selected
+
+ @selected.setter
+ def selected(self, value):
+ raise NotImplementedError(
+ "'enabled' can only be set through calling 'AcceleratorState().enable_deepspeed_plugin(key)'."
+ )
+
+
+@dataclass
+class FullyShardedDataParallelPlugin:
+ """
+ This plugin is used to enable fully sharded data parallelism.
+
+ Args:
+ fsdp_version (`int`, defaults to `1`):
+ The version of FSDP to use. Defaults to 1. If set to 2, launcher expects the config to be converted to
+ FSDP2 format.
+ sharding_strategy (`Union[str, torch.distributed.fsdp.ShardingStrategy]`, defaults to `'FULL_SHARD'`):
+ Sharding strategy to use. Should be either a `str` or an instance of
+ `torch.distributed.fsdp.fully_sharded_data_parallel.ShardingStrategy`. Is deprecated in favor of
+ `reshard_after_forward`.
+ reshard_after_forward (`Union[str, torch.distributed.fsdp.ShardingStrategy, bool]`, defaults to `'FULL_SHARD'` for `fsdp_version=1` and `True` for `fsdp_version=2`):
+ Sharding strategy to use. Should be a bool if `fsdp_version` is set to 2 else a `str` or an instance of
+ `torch.distributed.fsdp.fully_sharded_data_parallel.ShardingStrategy`.
+ backward_prefetch (`Union[str, torch.distributed.fsdp.BackwardPrefetch]`, defaults to `'NO_PREFETCH'`):
+ Backward prefetch strategy to use. Should be either a `str` or an instance of
+ `torch.distributed.fsdp.fully_sharded_data_parallel.BackwardPrefetch`.
+ mixed_precision_policy (`Optional[Union[dict, torch.distributed.fsdp.MixedPrecision, torch.distributed.fsdp.MixedPrecisionPolicy]]`, defaults to `None`):
+ A config to enable mixed precision training with FullyShardedDataParallel. If passing in a `dict`, it
+ should have the following keys: `param_dtype`, `reduce_dtype`, and `buffer_dtype`, can be an instance of
+ `torch.distributed.fsdp.MixedPrecisionPolicy` if `fsdp_version` is set to 2.
+ auto_wrap_policy (`Optional(Union[Callable, Literal["transformer_based_wrap", "size_based_wrap", "no_wrap"]]), defaults to `NO_WRAP`):
+ A callable or string specifying a policy to recursively wrap layers with FSDP. If a string, it must be one
+ of `transformer_based_wrap`, `size_based_wrap`, or `no_wrap`. See
+ `torch.distributed.fsdp.wrap.size_based_wrap_policy` for a direction on what it should look like.
+ cpu_offload (`Union[bool, torch.distributed.fsdp.CPUOffload, torch.distributed.fsdp.CPUOffloadPolicy]`, defaults to `False`):
+ Whether to offload parameters to CPU. Should be either a `bool` or an instance of
+ `torch.distributed.fsdp.fully_sharded_data_parallel.CPUOffload` or
+ `torch.distributed.fsdp.fully_sharded_data_parallel.CPUOffloadPolicy` if `fsdp_version` is set to 2.
+ ignored_modules (`Optional[Iterable[torch.nn.Module]]`, defaults to `None`):
+ A list of modules to ignore when wrapping with FSDP.
+ state_dict_type (`Union[str, torch.distributed.fsdp.StateDictType]`, defaults to `'FULL_STATE_DICT'`):
+ State dict type to use. If a string, it must be one of `full_state_dict`, `local_state_dict`, or
+ `sharded_state_dict`.
+ state_dict_config (`Optional[Union[torch.distributed.fsdp.FullStateDictConfig, torch.distributed.fsdp.ShardedStateDictConfig]`, defaults to `None`):
+ State dict config to use. Is determined based on the `state_dict_type` if not passed in.
+ optim_state_dict_config (`Optional[Union[torch.distributed.fsdp.FullOptimStateDictConfig, torch.distributed.fsdp.ShardedOptimStateDictConfig]`, defaults to `None`):
+ Optim state dict config to use. Is determined based on the `state_dict_type` if not passed in.
+ limit_all_gathers (`bool`, defaults to `True`):
+ Whether to have FSDP explicitly synchronizes the CPU thread to prevent too many in-flight all-gathers. This
+ bool only affects the sharded strategies that schedule all-gathers. Enabling this can help lower the number
+ of CUDA malloc retries.
+ use_orig_params (`bool`, defaults to `False`):
+ Whether to use the original parameters for the optimizer.
+ param_init_fn (`Optional[Callable[[torch.nn.Module], None]`, defaults to `None`):
+ A `Callable[torch.nn.Module] -> None` that specifies how modules that are currently on the meta device
+ should be initialized onto an actual device. Only applicable when `sync_module_states` is `True`. By
+ default is a `lambda` which calls `to_empty` on the module.
+ sync_module_states (`bool`, defaults to `False`):
+ Whether each individually wrapped FSDP unit should broadcast module parameters from rank 0 to ensure they
+ are the same across all ranks after initialization. Defaults to `False` unless `cpu_ram_efficient_loading`
+ is `True`, then will be forcibly enabled.
+ forward_prefetch (`bool`, defaults to `False`):
+ Whether to have FSDP explicitly prefetches the next upcoming all-gather while executing in the forward
+ pass. only use with Static graphs.
+ activation_checkpointing (`bool`, defaults to `False`):
+ A technique to reduce memory usage by clearing activations of certain layers and recomputing them during a
+ backward pass. Effectively, this trades extra computation time for reduced memory usage.
+ cpu_ram_efficient_loading (`bool`, defaults to `None`):
+ If True, only the first process loads the pretrained model checkoint while all other processes have empty
+ weights. Only applicable for Transformers. When using this, `sync_module_states` needs to be `True`.
+ transformer_cls_names_to_wrap (`Optional[List[str]]`, defaults to `None`):
+ A list of transformer layer class names to wrap. Only applicable when `auto_wrap_policy` is
+ `transformer_based_wrap`.
+ min_num_params (`Optional[int]`, defaults to `None`):
+ The minimum number of parameters a module must have to be wrapped. Only applicable when `auto_wrap_policy`
+ is `size_based_wrap`.
+ """
+
+ fsdp_version: int = field(
+ default=None,
+ metadata={
+ "help": "The version of FSDP to use. Defaults to 1. If set to 2, launcher expects the config to be converted to FSDP2 format."
+ },
+ )
+
+ sharding_strategy: Union[str, "torch.distributed.fsdp.ShardingStrategy"] = field(
+ default=None,
+ metadata={
+ "help": "Sharding strategy to use. Should be either a `str` or an instance of `torch.distributed.fsdp.fully_sharded_data_parallel.ShardingStrategy`. Defaults to 'FULL_SHARD'. Is deprecated in favor of `reshard_after_forward` "
+ },
+ )
+
+ reshard_after_forward: Union[str, "torch.distributed.fsdp.ShardingStrategy", bool] = field(
+ default=None,
+ metadata={
+ "help": "Sharding strategy to use. Should be a bool if `fsdp_version` is set to 2 else a `str` or an instance of `torch.distributed.fsdp.fully_sharded_data_parallel.ShardingStrategy`. Defaults to 'FULL_SHARD'"
+ },
+ )
+ backward_prefetch: Optional[Union[str, "torch.distributed.fsdp.BackwardPrefetch"]] = field(
+ default=None,
+ metadata={
+ "help": "Backward prefetch strategy to use. Should be either a `str` or an instance of `torch.distributed.fsdp.fully_sharded_data_parallel.BackwardPrefetch`. Defaults to 'NO_PREFETCH'. This becomes obsolete in FSDP2."
+ },
+ )
+ mixed_precision_policy: Optional[
+ Union[dict, "torch.distributed.fsdp.MixedPrecision", "torch.distributed.fsdp.MixedPrecisionPolicy"]
+ ] = field(
+ default=None,
+ metadata={
+ "help": "A config to enable mixed precision training with FullyShardedDataParallel. "
+ "If passing in a `dict`, it should have the following keys: `param_dtype`, `reduce_dtype`, and `buffer_dtype`."
+ "Can also be an instance of `torch.distributed.fsdp.MixedPrecisionPolicy` if `fsdp_version` is set to 2."
+ },
+ )
+ auto_wrap_policy: Optional[Union[Callable, Literal["transformer_based_wrap", "size_based_wrap", "no_wrap"]]] = (
+ field(
+ default=None,
+ metadata={
+ "help": "A callable or string specifying a policy to recursively wrap layers with FSDP. If a string, it must be one of `transformer_based_wrap`, `size_based_wrap`, or `no_wrap`. "
+ "Defaults to `NO_WRAP`. See `torch.distributed.fsdp.wrap.size_based_wrap_policy` for a direction on what it should look like"
+ },
+ )
+ )
+ cpu_offload: Union[bool, "torch.distributed.fsdp.CPUOffload", "torch.distributed.fsdp.CPUOffloadPolicy"] = field(
+ default=None,
+ metadata={
+ "help": "Whether to offload parameters to CPU. Should be either a `bool` or an instance of `torch.distributed.fsdp.fully_sharded_data_parallel.CPUOffload` or `torch.distributed.fsdp.fully_sharded_data_parallel.CPUOffloadPolicy` if `fsdp_version` is set to 2. Defaults to `False`"
+ },
+ )
+ ignored_modules: Optional[Iterable[torch.nn.Module]] = field(
+ default=None,
+ metadata={"help": "A list of modules to ignore when wrapping with FSDP."},
+ )
+
+ state_dict_type: Union[str, "torch.distributed.fsdp.StateDictType"] = field(
+ default=None,
+ metadata={
+ "help": "State dict type to use. If a string, it must be one of `full_state_dict`, `local_state_dict`, or `sharded_state_dict`. Defaults to `FULL_STATE_DICT`"
+ },
+ )
+ state_dict_config: Optional[
+ Union[
+ "torch.distributed.fsdp.FullStateDictConfig",
+ "torch.distributed.fsdp.ShardedStateDictConfig",
+ ]
+ ] = field(
+ default=None,
+ metadata={"help": "State dict config to use. Is determined based on the `state_dict_type` if not passed in."},
+ )
+ optim_state_dict_config: Optional[
+ Union["torch.distributed.fsdp.FullOptimStateDictConfig", "torch.distributed.fsdp.ShardedOptimStateDictConfig"]
+ ] = field(
+ default=None,
+ metadata={
+ "help": "Optim state dict config to use. Is determined based on the `state_dict_type` if not passed in."
+ },
+ )
+ limit_all_gathers: bool = field(
+ default=True,
+ metadata={
+ "help": "Whether to have FSDP explicitly synchronizes the CPU thread to prevent "
+ "too many in-flight all-gathers. This bool only affects the sharded strategies that schedule all-gathers. "
+ "Enabling this can help lower the number of CUDA malloc retries."
+ },
+ )
+ use_orig_params: Optional[bool] = field(
+ default=None,
+ metadata={
+ "help": "Whether to use the original parameters for the optimizer. Defaults to `False`. This becomes obsolete in FSDP2."
+ },
+ )
+ param_init_fn: Optional[Callable[[torch.nn.Module], None]] = field(
+ default=None,
+ metadata={
+ "help": "A Callable[torch.nn.Module] -> None that specifies how modules "
+ "that are currently on the meta device should be initialized onto an actual device. "
+ "Only applicable when `sync_module_states` is `True`. By default is a `lambda` which calls `to_empty` on the module."
+ },
+ )
+ sync_module_states: Optional[bool] = field(
+ default=None,
+ metadata={
+ "help": "Whether each individually wrapped FSDP unit should broadcast module parameters from rank 0 "
+ "to ensure they are the same across all ranks after initialization. Defaults to `False` unless "
+ "`cpu_ram_efficient_loading` is `True`, then will be forcibly enabled. This becomes obsolete in FSDP2."
+ },
+ )
+ forward_prefetch: bool = field(
+ default=None,
+ metadata={
+ "help": "Whether to have FSDP explicitly prefetches the next upcoming "
+ "all-gather while executing in the forward pass. only use with Static graphs. Defaults to `False`"
+ },
+ )
+ activation_checkpointing: bool = field(
+ default=None,
+ metadata={
+ "help": "A technique to reduce memory usage by clearing activations of "
+ "certain layers and recomputing them during a backward pass. Effectively, this trades extra computation time "
+ "for reduced memory usage. Defaults to `False`"
+ },
+ )
+ cpu_ram_efficient_loading: bool = field(
+ default=None,
+ metadata={
+ "help": "If True, only the first process loads the pretrained model checkoint while all other processes have empty weights. "
+ "Only applicable for 🤗 Transformers. When using this, `sync_module_states` needs to be `True`. Defaults to `False`."
+ },
+ )
+ transformer_cls_names_to_wrap: Optional[list[str]] = field(
+ default=None,
+ metadata={
+ "help": "A list of transformer layer class names to wrap. Only applicable when `auto_wrap_policy` is `transformer_based_wrap`."
+ },
+ )
+ min_num_params: Optional[int] = field(
+ default=None,
+ metadata={
+ "help": "The minimum number of parameters a module must have to be wrapped. Only applicable when `auto_wrap_policy` is `size_based_wrap`."
+ },
+ )
+
+ def __post_init__(self):
+ from torch.distributed.fsdp import (
+ BackwardPrefetch,
+ ShardingStrategy,
+ )
+
+ _fsdp2_warnings = set()
+
+ env_prefix = "FSDP_"
+ # Strategy: By default we should always assume that values are passed in, else we check the environment variables
+ if self.fsdp_version is None:
+ self.fsdp_version = int(os.environ.get(env_prefix + "VERSION", "1"))
+
+ if self.sharding_strategy is not None:
+ # We cannot properly detect all of the cases, as by default `args.fsdp_sharding_strategy` is set to `fully_shard`
+ # Therefore we issue a warning only if the user has explicitly set it inside their plugin
+ _fsdp2_warnings.add(
+ "sharding_strategy is deprecated in favor of reshard_after_forward. "
+ "This will be removed in a future version of Accelerate."
+ )
+ if self.fsdp_version == 1:
+ if self.sharding_strategy is None:
+ self.sharding_strategy = os.environ.get(env_prefix + "SHARDING_STRATEGY", "FULL_SHARD")
+ if isinstance(self.sharding_strategy, str):
+ if self.sharding_strategy.upper() in FSDP_SHARDING_STRATEGY:
+ self.sharding_strategy = FSDP_SHARDING_STRATEGY.index(self.sharding_strategy.upper()) + 1
+ if isinstance(self.sharding_strategy, int) or self.sharding_strategy.isdigit():
+ self.sharding_strategy = ShardingStrategy(int(self.sharding_strategy))
+ else:
+ self.sharding_strategy = ShardingStrategy[self.sharding_strategy.upper()]
+
+ # Fallback to `reshard_after_forward` in FSDP1 if `sharding_strategy` is not set
+ if self.reshard_after_forward is None and self.sharding_strategy is None:
+ reshard_after_forward = os.environ.get(
+ env_prefix + "RESHARD_AFTER_FORWARD", "true" if self.fsdp_version == 2 else "FULL_SHARD"
+ )
+ if self.fsdp_version == 2:
+ self.reshard_after_forward = str_to_bool(reshard_after_forward.lower(), to_bool=True)
+ else:
+ self.reshard_after_forward = reshard_after_forward
+ if isinstance(self.reshard_after_forward, str):
+ if self.fsdp_version == 2:
+ self.reshard_after_forward = str_to_bool(self.reshard_after_forward.lower(), to_bool=True)
+ else:
+ # We need to remap based on custom enum values for user readability
+ if self.reshard_after_forward.upper() in FSDP_SHARDING_STRATEGY:
+ self.reshard_after_forward = FSDP_SHARDING_STRATEGY.index(self.reshard_after_forward.upper()) + 1
+ if isinstance(self.reshard_after_forward, int) or self.reshard_after_forward.isdigit():
+ self.reshard_after_forward = ShardingStrategy(int(self.reshard_after_forward))
+ else:
+ self.reshard_after_forward = ShardingStrategy[self.reshard_after_forward.upper()]
+
+ if self.fsdp_version == 2 and not isinstance(self.reshard_after_forward, bool):
+ raise ValueError(
+ f"reshard_after_forward set to {self.reshard_after_forward}. This is not supported with FSDP2, please set to a `bool`"
+ )
+ if self.fsdp_version == 1 and isinstance(self.reshard_after_forward, bool):
+ raise ValueError(
+ f"reshard_after_forward set to {self.reshard_after_forward}. This is not supported with FSDP1, please set to a `str` or an instance of `torch.distributed.fsdp.fully_sharded_data_parallel.ShardingStrategy`"
+ )
+
+ if self.cpu_offload is None:
+ self.cpu_offload = str_to_bool(os.environ.get(env_prefix + "OFFLOAD_PARAMS", "False")) == 1
+
+ self.set_cpu_offload() # abstracted away to hide imports due to version checks
+ self.validate_cpu_offload()
+
+ if self.backward_prefetch is None:
+ self.backward_prefetch = os.environ.get(env_prefix + "BACKWARD_PREFETCH", None)
+ if isinstance(self.backward_prefetch, str) and self.backward_prefetch.upper() == "NO_PREFETCH":
+ self.backward_prefetch = None
+ if self.backward_prefetch is not None and not isinstance(self.backward_prefetch, BackwardPrefetch):
+ if isinstance(self.backward_prefetch, str) and self.backward_prefetch.upper() in FSDP_BACKWARD_PREFETCH:
+ self.backward_prefetch = FSDP_BACKWARD_PREFETCH.index(self.backward_prefetch.upper()) + 1
+ if isinstance(self.backward_prefetch, int) or self.backward_prefetch.isdigit():
+ self.backward_prefetch = BackwardPrefetch(int(self.backward_prefetch))
+ else:
+ self.backward_prefetch = BackwardPrefetch[self.backward_prefetch.upper()]
+ if self.fsdp_version == 2 and self.backward_prefetch is not None:
+ _fsdp2_warnings.add("backward_prefetch is not supported in FSDP2. Setting backward prefetch to None.")
+ self.backward_prefetch = None
+
+ self.set_state_dict_type()
+
+ if self.auto_wrap_policy is None:
+ self.auto_wrap_policy = os.environ.get(env_prefix + "AUTO_WRAP_POLICY", "NO_WRAP")
+ if isinstance(self.auto_wrap_policy, str):
+ if self.auto_wrap_policy.upper() not in FSDP_AUTO_WRAP_POLICY:
+ raise ValueError(
+ f"Invalid auto wrap policy: {self.auto_wrap_policy}. Must be one of {list(FSDP_AUTO_WRAP_POLICY.keys())}"
+ )
+ from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy, transformer_auto_wrap_policy
+
+ if self.auto_wrap_policy.upper() == "TRANSFORMER_BASED_WRAP":
+ self.auto_wrap_policy = transformer_auto_wrap_policy
+ if self.transformer_cls_names_to_wrap is None:
+ self.transformer_cls_names_to_wrap = os.environ.get(env_prefix + "TRANSFORMER_CLS_TO_WRAP", None)
+ if isinstance(self.transformer_cls_names_to_wrap, str):
+ self.transformer_cls_names_to_wrap = self.transformer_cls_names_to_wrap.split(",")
+ elif self.auto_wrap_policy.upper() == "SIZE_BASED_WRAP":
+ self.auto_wrap_policy = size_based_auto_wrap_policy
+ if self.min_num_params is None:
+ self.min_num_params = int(os.environ.get(env_prefix + "MIN_NUM_PARAMS", 0))
+ elif not isinstance(self.min_num_params, int):
+ raise ValueError(
+ f"`min_num_params` must be an integer. Got {self.min_num_params} of type {type(self.min_num_params)}"
+ )
+ elif self.auto_wrap_policy.upper() == "NO_WRAP":
+ self.auto_wrap_policy = None
+
+ if self.use_orig_params is None and self.fsdp_version == 1:
+ self.use_orig_params = str_to_bool(os.environ.get(env_prefix + "USE_ORIG_PARAMS", "False")) == 1
+ if self.fsdp_version == 2 and self.use_orig_params is not None:
+ _fsdp2_warnings.add("use_orig_params is obsolete in FSDP2, as FSDP2 always uses the original parameters.")
+ self.use_orig_params = None
+
+ if self.sync_module_states is None and self.fsdp_version == 1:
+ self.sync_module_states = str_to_bool(os.environ.get(env_prefix + "SYNC_MODULE_STATES", "False")) == 1
+ if self.fsdp_version == 2 and self.sync_module_states is not None:
+ _fsdp2_warnings.add(
+ "sync_module_states is obsolete in FSDP2, as it is not needed anymore."
+ "Setting sync_module_states to None."
+ )
+ self.sync_module_states = None
+
+ if self.forward_prefetch is None and self.fsdp_version == 1:
+ self.forward_prefetch = str_to_bool(os.environ.get(env_prefix + "FORWARD_PREFETCH", "False")) == 1
+ if self.fsdp_version == 2 and self.forward_prefetch is not None:
+ raise ValueError("forward_prefetch is not yet implemented in FSDP2, set to None or use `fsdp_version=1`")
+
+ if self.activation_checkpointing is None:
+ self.activation_checkpointing = (
+ str_to_bool(os.environ.get(env_prefix + "ACTIVATION_CHECKPOINTING", "False")) == 1
+ )
+
+ if self.cpu_ram_efficient_loading is None:
+ self.cpu_ram_efficient_loading = (
+ str_to_bool(os.environ.get(env_prefix + "CPU_RAM_EFFICIENT_LOADING", "False")) == 1
+ )
+ # There's no need to specify sync_module_states in FSDP2
+ if self.fsdp_version == 1 and self.cpu_ram_efficient_loading and not self.sync_module_states:
+ warnings.warn(
+ "sync_module_states cannot be False since efficient cpu ram loading enabled. "
+ "Setting sync_module_states to True."
+ )
+ self.sync_module_states = True
+
+ if isinstance(self.mixed_precision_policy, dict):
+ self.set_mixed_precision(self.mixed_precision_policy)
+ if self.mixed_precision_policy is not None:
+ self.validate_mixed_precision_policy()
+
+ if self.sync_module_states:
+ if is_npu_available():
+ device = torch.npu.current_device()
+ elif is_mlu_available():
+ device = torch.mlu.current_device()
+ elif is_musa_available():
+ device = torch.musa.current_device()
+ elif is_cuda_available():
+ device = torch.cuda.current_device()
+ elif is_xpu_available():
+ device = torch.xpu.current_device()
+ elif is_hpu_available():
+ device = torch.hpu.current_device()
+ else:
+ raise RuntimeError(
+ "There are currently no available devices found, must be one of 'XPU', 'CUDA', 'MLU', 'NPU', 'MUSA', or 'HPU'."
+ )
+ # Create a function that will be used to initialize the parameters of the model
+ # when using `sync_module_states`
+ self.param_init_fn = lambda x: x.to_empty(device=device, recurse=False)
+
+ # Single warning for all deprecation warnings due to FSDP2 conversion
+ if _fsdp2_warnings:
+ logger.warning("Multiple deprecation warnings due to FSDP2 conversion:\n".join(_fsdp2_warnings))
+
+ def set_state_dict_type(self, state_dict_type=None):
+ """
+ Set the state dict config based on the `StateDictType`.
+ """
+ from torch.distributed.fsdp.fully_sharded_data_parallel import (
+ FullOptimStateDictConfig,
+ FullStateDictConfig,
+ ShardedOptimStateDictConfig,
+ ShardedStateDictConfig,
+ StateDictType,
+ )
+
+ # Override the state_dict_type if provided, typical use case:
+ # user trains with sharded, but final save is with full
+ if state_dict_type is not None:
+ self.state_dict_type = state_dict_type
+
+ if self.state_dict_type is None:
+ self.state_dict_type = os.environ.get(
+ "FSDP_STATE_DICT_TYPE", "FULL_STATE_DICT" if self.fsdp_version == 1 else "SHARDED_STATE_DICT"
+ )
+ if isinstance(self.state_dict_type, str):
+ if self.state_dict_type.isdigit():
+ self.state_dict_type = StateDictType(int(self.state_dict_type))
+ else:
+ self.state_dict_type = StateDictType[self.state_dict_type.upper()]
+
+ if self.state_dict_type == StateDictType.FULL_STATE_DICT:
+ if self.state_dict_config is None:
+ self.state_dict_config = FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
+ if self.optim_state_dict_config is None:
+ self.optim_state_dict_config = FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
+ elif self.state_dict_type == StateDictType.SHARDED_STATE_DICT:
+ if self.state_dict_config is None:
+ self.state_dict_config = ShardedStateDictConfig(offload_to_cpu=True)
+ if self.optim_state_dict_config is None:
+ self.optim_state_dict_config = ShardedOptimStateDictConfig(offload_to_cpu=True)
+
+ # TODO(s1ro1): add support for FULL_STATE_DICT in FSDP2
+ if self.fsdp_version == 2 and self.state_dict_type != StateDictType.SHARDED_STATE_DICT:
+ raise ValueError(
+ "FSDP2 only supports SHARDED_STATE_DICT for now. "
+ "Please set `fsdp_state_dict_type` to `SHARDED_STATE_DICT`."
+ )
+
+ def set_auto_wrap_policy(self, model):
+ """
+ Given `model`, creates an `auto_wrap_policy` baesd on the passed in policy and if we can use the
+ `transformer_cls_to_wrap`
+ """
+ from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy, transformer_auto_wrap_policy
+
+ # First base off of `_no_split_modules`
+ no_split_modules = getattr(model, "_no_split_modules", None)
+ default_transformer_cls_names_to_wrap = list(no_split_modules) if no_split_modules is not None else []
+ if self.auto_wrap_policy == transformer_auto_wrap_policy:
+ if self.transformer_cls_names_to_wrap is None:
+ self.transformer_cls_names_to_wrap = default_transformer_cls_names_to_wrap
+ transformer_cls_to_wrap = set()
+ for layer_class in self.transformer_cls_names_to_wrap:
+ transformer_cls = get_module_class_from_name(model, layer_class)
+ if transformer_cls is None:
+ raise ValueError(f"Could not find the transformer layer class {layer_class} in the model.")
+ transformer_cls_to_wrap.add(transformer_cls)
+ # Finally we set the auto_wrap_policy to a callable
+ self.auto_wrap_policy = functools.partial(
+ self.auto_wrap_policy, transformer_layer_cls=transformer_cls_to_wrap
+ )
+
+ elif self.auto_wrap_policy == size_based_auto_wrap_policy:
+ # If zero, we silently ignore it.
+ if self.min_num_params > 0:
+ self.auto_wrap_policy = functools.partial(self.auto_wrap_policy, min_num_params=self.min_num_params)
+ else:
+ self.auto_wrap_policy = None
+
+ def set_mixed_precision(self, mixed_precision, buffer_autocast=False, override=False):
+ "Sets the mixed precision policy for FSDP"
+ mixed_precision_mapping = {
+ "fp8": torch.bfloat16,
+ "fp16": torch.float16,
+ "bf16": torch.bfloat16,
+ "fp32": torch.float32,
+ }
+ dtype = mixed_precision
+ if isinstance(mixed_precision, str):
+ dtype = mixed_precision_mapping.get(mixed_precision, None)
+ if dtype is None:
+ raise ValueError(
+ f"Invalid mixed precision: {mixed_precision}. Must be one of {list(mixed_precision_mapping.keys())}"
+ )
+ elif isinstance(mixed_precision, torch.dtype) and mixed_precision not in mixed_precision_mapping.values():
+ raise ValueError(
+ f"Invalid mixed precision: {mixed_precision}. Must be one of {list(mixed_precision_mapping.values())}"
+ )
+
+ buffer_type = torch.float32 if buffer_autocast else dtype
+
+ if self.fsdp_version == 1:
+ from torch.distributed.fsdp import MixedPrecision
+ elif self.fsdp_version == 2:
+ from torch.distributed.fsdp import MixedPrecisionPolicy as MixedPrecision
+
+ if override or self.mixed_precision_policy is None:
+ dtype_args = {"param_dtype": dtype, "reduce_dtype": dtype}
+ if self.fsdp_version == 1:
+ dtype_args["buffer_dtype"] = buffer_type
+ else:
+ dtype_args["output_dtype"] = dtype
+ # TODO(s1ro1): `cast_forward_inputs` for FSDP2?
+ self.mixed_precision_policy = MixedPrecision(**dtype_args)
+ elif isinstance(self.mixed_precision_policy, dict):
+ # Check for incompatible types
+ valid_keys = ["param_dtype", "reduce_dtype"] + (
+ ["buffer_dtype"] if self.fsdp_version == 1 else ["output_dtype"]
+ )
+ missing_keys = [k for k in valid_keys if k not in self.mixed_precision_policy]
+ invalid_values = [
+ k for k, v in self.mixed_precision_policy.items() if v not in mixed_precision_mapping.values()
+ ]
+ if missing_keys or invalid_values:
+ raise ValueError(
+ f"Invalid mixed precision policy: {self.mixed_precision_policy}. "
+ f"Must be a `dict` with keys {valid_keys}."
+ f"Values must be one of {list(mixed_precision_mapping.values())}"
+ )
+ self.mixed_precision_policy = MixedPrecision(**self.mixed_precision_policy)
+
+ def validate_mixed_precision_policy(self):
+ """
+ Validates the mixed precision policy, abstracted away to not bring in the imports if not needed.
+ """
+ if self.fsdp_version == 2:
+ from torch.distributed.fsdp import MixedPrecisionPolicy as MixedPrecision
+ else:
+ from torch.distributed.fsdp import MixedPrecision
+
+ if not isinstance(self.mixed_precision_policy, MixedPrecision):
+ required_type = (
+ "`torch.distributed.fsdp.MixedPrecisionPolicy`"
+ if self.fsdp_version == 2
+ else "`torch.distributed.fsdp.MixedPrecision`"
+ )
+ raise ValueError(f"mixed_precision_policy must be an instance of {required_type}.")
+
+ def set_cpu_offload(self):
+ if self.fsdp_version == 2:
+ from torch.distributed.fsdp import CPUOffloadPolicy, OffloadPolicy
+ else:
+ from torch.distributed.fsdp import CPUOffload
+
+ if isinstance(self.cpu_offload, bool):
+ if self.fsdp_version == 2:
+ if not self.cpu_offload:
+ self.cpu_offload = OffloadPolicy()
+ else:
+ self.cpu_offload = CPUOffloadPolicy()
+ else:
+ self.cpu_offload = CPUOffload(offload_params=self.cpu_offload)
+
+ def validate_cpu_offload(self):
+ if self.fsdp_version == 2:
+ from torch.distributed.fsdp import OffloadPolicy
+ else:
+ from torch.distributed.fsdp import CPUOffload
+
+ if self.fsdp_version == 2 and not isinstance(self.cpu_offload, OffloadPolicy):
+ raise ValueError(
+ f"`cpu_offload` must be an instance of `torch.distributed.fsdp.OffloadPolicy` in FSDP2, got {self.cpu_offload}"
+ )
+ if self.fsdp_version == 1 and not isinstance(self.cpu_offload, CPUOffload):
+ raise ValueError(
+ f"`cpu_offload` must be an instance of `torch.distributed.fsdp.CPUOffload` in FSDP1, got {self.cpu_offload}"
+ )
+
+
+@dataclass
+class TorchTensorParallelPlugin:
+ """
+ This plugin is used to enable tensor parallelism using PyTorch >= 2.0.
+ """
+
+ tp_size: int = field(
+ default=1,
+ metadata={"help": "tensor parallel size will be used in the device mesh preparation"},
+ )
+
+ # torch_device_mesh is fo type "torch.distributed.DeviceMesh"
+ torch_device_mesh: Optional["torch.distributed.DeviceMesh"] = field(default=None)
+
+ def __post_init__(self):
+ self.tp_size = self.tp_size if os.environ.get("TP_SIZE", "1") == "1" else int(os.environ.get("TP_SIZE", "1"))
+ if self.tp_size == 1:
+ raise ValueError("Provide TP degree > 1.")
+
+ if is_torch_version("<", BETA_TP_AVAILABLE_PYTORCH_VERSION):
+ raise ValueError(
+ f"Minimum PyTorch version {BETA_TP_AVAILABLE_PYTORCH_VERSION} needed to use tensor parallel."
+ )
+ from torch.distributed.device_mesh import init_device_mesh
+
+ # support for other devices has to be investigated
+ if is_hpu_available(init_hccl=True):
+ device = "hpu"
+ else:
+ device = "cuda"
+
+ mesh_dim_name = "tp"
+
+ self.torch_device_mesh = init_device_mesh(device, (self.tp_size,), mesh_dim_names=(mesh_dim_name,))
+
+
+@dataclass
+class MegatronLMPlugin:
+ """
+ Plugin for Megatron-LM to enable tensor, pipeline, sequence and data parallelism. Also to enable selective
+ activation recomputation and optimized fused kernels.
+
+ Args:
+ tp_degree (`int`, defaults to `None`):
+ Tensor parallelism degree.
+ pp_degree (`int`, defaults to `None`):
+ Pipeline parallelism degree.
+ num_micro_batches (`int`, defaults to `None`):
+ Number of micro-batches.
+ gradient_clipping (`float`, defaults to `None`):
+ Gradient clipping value based on global L2 Norm (0 to disable).
+ sequence_parallelism (`bool`, defaults to `None`):
+ Enable sequence parallelism.
+ recompute_activations (`bool`, defaults to `None`):
+ Enable selective activation recomputation.
+ use_distributed_optimizr (`bool`, defaults to `None`):
+ Enable distributed optimizer.
+ pipeline_model_parallel_split_rank (`int`, defaults to `None`):
+ Rank where encoder and decoder should be split.
+ num_layers_per_virtual_pipeline_stage (`int`, defaults to `None`):
+ Number of layers per virtual pipeline stage.
+ is_train_batch_min (`str`, defaults to `True`):
+ If both tran & eval dataloaders are specified, this will decide the `micro_batch_size`.
+ train_iters (`int`, defaults to `None`):
+ Total number of samples to train over all training runs. Note that either train-iters or train-samples
+ should be provided when using `MegatronLMDummyScheduler`.
+ train_samples (`int`, defaults to `None`):
+ Total number of samples to train over all training runs. Note that either train-iters or train-samples
+ should be provided when using `MegatronLMDummyScheduler`.
+ weight_decay_incr_style (`str`, defaults to `'constant'`):
+ Weight decay increment function. choices=["constant", "linear", "cosine"].
+ start_weight_decay (`float`, defaults to `None`):
+ Initial weight decay coefficient for L2 regularization.
+ end_weight_decay (`float`, defaults to `None`):
+ End of run weight decay coefficient for L2 regularization.
+ lr_decay_style (`str`, defaults to `'linear'`):
+ Learning rate decay function. choices=['constant', 'linear', 'cosine'].
+ lr_decay_iters (`int`, defaults to `None`):
+ Number of iterations for learning rate decay. If None defaults to `train_iters`.
+ lr_decay_samples (`int`, defaults to `None`):
+ Number of samples for learning rate decay. If None defaults to `train_samples`.
+ lr_warmup_iters (`int`, defaults to `None`):
+ Number of iterations to linearly warmup learning rate over.
+ lr_warmup_samples (`int`, defaults to `None`):
+ Number of samples to linearly warmup learning rate over.
+ lr_warmup_fraction (`float`, defaults to `None`):
+ Fraction of lr-warmup-(iters/samples) to linearly warmup learning rate over.
+ min_lr (`float`, defaults to `0`):
+ Minumum value for learning rate. The scheduler clip values below this threshold.
+ consumed_samples (`List`, defaults to `None`):
+ Number of samples consumed in the same order as the dataloaders to `accelerator.prepare` call.
+ no_wd_decay_cond (`Optional`, defaults to `None`):
+ Condition to disable weight decay.
+ scale_lr_cond (`Optional`, defaults to `None`):
+ Condition to scale learning rate.
+ lr_mult (`float`, defaults to `1.0`):
+ Learning rate multiplier.
+ megatron_dataset_flag (`bool`, defaults to `False`):
+ Whether the format of dataset follows Megatron-LM Indexed/Cached/MemoryMapped format.
+ seq_length (`int`, defaults to `None`):
+ Maximum sequence length to process.
+ encoder_seq_length (`int`, defaults to `None`):
+ Maximum sequence length to process for the encoder.
+ decoder_seq_length (`int`, defaults to `None`):
+ Maximum sequence length to process for the decoder.
+ tensorboard_dir (`str`, defaults to `None`):
+ Path to save tensorboard logs.
+ set_all_logging_options (`bool`, defaults to `False`):
+ Whether to set all logging options.
+ eval_iters (`int`, defaults to `100`):
+ Number of iterations to run for evaluation validation/test for.
+ eval_interval (`int`, defaults to `1000`):
+ Interval between running evaluation on validation set.
+ return_logits (`bool`, defaults to `False`):
+ Whether to return logits from the model.
+ custom_train_step_class (`Optional`, defaults to `None`):
+ Custom train step class.
+ custom_train_step_kwargs (`Optional`, defaults to `None`):
+ Custom train step kwargs.
+ custom_model_provider_function (`Optional`, defaults to `None`):
+ Custom model provider function.
+ custom_prepare_model_function (`Optional`, defaults to `None`):
+ Custom prepare model function.
+ custom_megatron_datasets_provider_function (`Optional`, defaults to `None`):
+ Custom megatron train_valid_test datasets provider function.
+ custom_get_batch_function (`Optional`, defaults to `None`):
+ Custom get batch function.
+ custom_loss_function (`Optional`, defaults to `None`):
+ Custom loss function.
+ other_megatron_args (`Optional`, defaults to `None`):
+ Other Megatron-LM arguments. Please refer Megatron-LM.
+ """
+
+ tp_degree: int = field(default=None, metadata={"help": "tensor parallelism degree."})
+ pp_degree: int = field(default=None, metadata={"help": "pipeline parallelism degree."})
+ num_micro_batches: int = field(default=None, metadata={"help": "number of micro-batches."})
+ gradient_clipping: float = field(
+ default=None, metadata={"help": "gradient clipping value based on global L2 Norm (0 to disable)"}
+ )
+ sequence_parallelism: bool = field(
+ default=None,
+ metadata={"help": "enable sequence parallelism"},
+ )
+ recompute_activations: bool = field(
+ default=None,
+ metadata={"help": "enable selective activation recomputation"},
+ )
+ use_distributed_optimizer: bool = field(
+ default=None,
+ metadata={"help": "enable distributed optimizer"},
+ )
+ pipeline_model_parallel_split_rank: int = field(
+ default=None, metadata={"help": "Rank where encoder and decoder should be split."}
+ )
+ num_layers_per_virtual_pipeline_stage: int = field(
+ default=None, metadata={"help": "Number of layers per virtual pipeline stage."}
+ )
+ is_train_batch_min: str = field(
+ default=True,
+ metadata={"help": "If both train & eval dataloaders are specified, this will decide the micro_batch_size"},
+ )
+ train_iters: int = field(
+ default=None,
+ metadata={
+ "help": "Total number of iterations to train over all training runs. "
+ "Note that either train-iters or train-samples should be provided when using `MegatronLMDummyScheduler`"
+ },
+ )
+ train_samples: int = field(
+ default=None,
+ metadata={
+ "help": "Total number of samples to train over all training runs. "
+ "Note that either train-iters or train-samples should be provided when using `MegatronLMDummyScheduler`"
+ },
+ )
+ weight_decay_incr_style: str = field(
+ default="constant",
+ metadata={"help": 'Weight decay increment function. choices=["constant", "linear", "cosine"]. '},
+ )
+ start_weight_decay: float = field(
+ default=None,
+ metadata={"help": "Initial weight decay coefficient for L2 regularization."},
+ )
+ end_weight_decay: float = field(
+ default=None,
+ metadata={"help": "End of run weight decay coefficient for L2 regularization."},
+ )
+ lr_decay_style: str = field(
+ default="linear",
+ metadata={"help": "Learning rate decay function. choices=['constant', 'linear', 'cosine']."},
+ )
+ lr_decay_iters: int = field(
+ default=None,
+ metadata={"help": "Number of iterations for learning rate decay. If None defaults to `train_iters`."},
+ )
+ lr_decay_samples: int = field(
+ default=None,
+ metadata={"help": "Number of samples for learning rate decay. If None defaults to `train_samples`."},
+ )
+ lr_warmup_iters: int = field(
+ default=None,
+ metadata={"help": "number of iterations to linearly warmup learning rate over."},
+ )
+ lr_warmup_samples: int = field(
+ default=None,
+ metadata={"help": "number of samples to linearly warmup learning rate over."},
+ )
+ lr_warmup_fraction: float = field(
+ default=None,
+ metadata={"help": "fraction of lr-warmup-(iters/samples) to linearly warmup learning rate over."},
+ )
+ min_lr: float = field(
+ default=0,
+ metadata={"help": "Minumum value for learning rate. The scheduler clip values below this threshold."},
+ )
+ consumed_samples: list[int] = field(
+ default=None,
+ metadata={
+ "help": "Number of samples consumed in the same order as the dataloaders to `accelerator.prepare` call."
+ },
+ )
+ no_wd_decay_cond: Optional[Callable] = field(default=None, metadata={"help": "Condition to disable weight decay."})
+ scale_lr_cond: Optional[Callable] = field(default=None, metadata={"help": "Condition to scale learning rate."})
+ lr_mult: float = field(default=1.0, metadata={"help": "Learning rate multiplier."})
+ megatron_dataset_flag: bool = field(
+ default=False,
+ metadata={"help": "Whether the format of dataset follows Megatron-LM Indexed/Cached/MemoryMapped format."},
+ )
+ seq_length: int = field(
+ default=None,
+ metadata={"help": "Maximum sequence length to process."},
+ )
+ encoder_seq_length: int = field(
+ default=None,
+ metadata={"help": "Maximum sequence length to process for the encoder."},
+ )
+ decoder_seq_length: int = field(
+ default=None,
+ metadata={"help": "Maximum sequence length to process for the decoder."},
+ )
+ tensorboard_dir: str = field(
+ default=None,
+ metadata={"help": "Path to save tensorboard logs."},
+ )
+ set_all_logging_options: bool = field(
+ default=False,
+ metadata={"help": "Whether to set all logging options."},
+ )
+ eval_iters: int = field(
+ default=100, metadata={"help": "Number of iterations to run for evaluation validation/test for."}
+ )
+ eval_interval: int = field(
+ default=1000, metadata={"help": "Interval between running evaluation on validation set."}
+ )
+ return_logits: bool = field(
+ default=False,
+ metadata={"help": "Whether to return logits from the model."},
+ )
+
+ # custom train step args
+ custom_train_step_class: Optional[Any] = field(
+ default=None,
+ metadata={"help": "Custom train step class."},
+ )
+ custom_train_step_kwargs: Optional[dict[str, Any]] = field(
+ default=None,
+ metadata={"help": "Custom train step kwargs."},
+ )
+
+ # custom model args
+ custom_model_provider_function: Optional[Callable] = field(
+ default=None,
+ metadata={"help": "Custom model provider function."},
+ )
+ custom_prepare_model_function: Optional[Callable] = field(
+ default=None,
+ metadata={"help": "Custom prepare model function."},
+ )
+ custom_megatron_datasets_provider_function: Optional[Callable] = field(
+ default=None,
+ metadata={"help": "Custom megatron train_valid_test datasets provider function."},
+ )
+ custom_get_batch_function: Optional[Callable] = field(
+ default=None,
+ metadata={"help": "Custom get batch function."},
+ )
+ custom_loss_function: Optional[Callable] = field(
+ default=None,
+ metadata={"help": "Custom loss function."},
+ )
+
+ # remaining args such as enabling Alibi/ROPE positional embeddings,
+ # wandb logging, Multi-Query Attention, etc.
+ other_megatron_args: Optional[dict[str, Any]] = field(
+ default=None,
+ metadata={"help": "Other Megatron-LM arguments. Please refer Megatron-LM"},
+ )
+
+ def __post_init__(self):
+ prefix = "MEGATRON_LM_"
+ if self.tp_degree is None:
+ self.tp_degree = int(os.environ.get(prefix + "TP_DEGREE", 1))
+ if self.pp_degree is None:
+ self.pp_degree = int(os.environ.get(prefix + "PP_DEGREE", 1))
+ if self.num_micro_batches is None:
+ self.num_micro_batches = int(os.environ.get(prefix + "NUM_MICRO_BATCHES", 1))
+ if self.gradient_clipping is None:
+ self.gradient_clipping = float(os.environ.get(prefix + "GRADIENT_CLIPPING", 1.0))
+ if self.recompute_activations is None:
+ self.recompute_activations = str_to_bool(os.environ.get(prefix + "RECOMPUTE_ACTIVATIONS", "False")) == 1
+ if self.use_distributed_optimizer is None:
+ self.use_distributed_optimizer = (
+ str_to_bool(os.environ.get(prefix + "USE_DISTRIBUTED_OPTIMIZER", "False")) == 1
+ )
+ if self.sequence_parallelism is None:
+ self.sequence_parallelism = str_to_bool(os.environ.get(prefix + "SEQUENCE_PARALLELISM", "False")) == 1
+
+ if self.pp_degree > 1 or self.use_distributed_optimizer:
+ self.DDP_impl = "local"
+ else:
+ self.DDP_impl = "torch"
+
+ if self.consumed_samples is not None:
+ if len(self.consumed_samples) == 1:
+ self.consumed_samples.extend([0, 0])
+ elif len(self.consumed_samples) == 2:
+ self.consumed_samples.append(0)
+
+ self.megatron_lm_default_args = {
+ "tensor_model_parallel_size": self.tp_degree,
+ "pipeline_model_parallel_size": self.pp_degree,
+ "pipeline_model_parallel_split_rank": self.pipeline_model_parallel_split_rank,
+ "num_layers_per_virtual_pipeline_stage": self.num_layers_per_virtual_pipeline_stage,
+ "DDP_impl": self.DDP_impl,
+ "use_distributed_optimizer": self.use_distributed_optimizer,
+ "sequence_parallel": self.sequence_parallelism,
+ "clip_grad": self.gradient_clipping,
+ "num_micro_batches": self.num_micro_batches,
+ "consumed_samples": self.consumed_samples,
+ "no_wd_decay_cond": self.no_wd_decay_cond,
+ "scale_lr_cond": self.scale_lr_cond,
+ "lr_mult": self.lr_mult,
+ "megatron_dataset_flag": self.megatron_dataset_flag,
+ "eval_iters": self.eval_iters,
+ "eval_interval": self.eval_interval,
+ }
+ if self.recompute_activations:
+ self.megatron_lm_default_args["recompute_granularity"] = "selective"
+ if self.tensorboard_dir is not None:
+ self.megatron_lm_default_args["tensorboard_dir"] = self.tensorboard_dir
+ if self.set_all_logging_options:
+ self.set_tensorboard_logging_options()
+ if self.other_megatron_args is not None:
+ self.megatron_lm_default_args.update(self.other_megatron_args)
+
+ def set_network_size_args(self, model, batch_data=None):
+ model_config_type = model.config.model_type.lower()
+ for model_type in MODEL_CONFIGS_TO_MEGATRON_PARSERS.keys():
+ if model_type in model_config_type:
+ MODEL_CONFIGS_TO_MEGATRON_PARSERS[model_type](self, model, batch_data)
+ return
+ raise ValueError(
+ f"Accelerate Megatron-LM integration not supports {model_config_type} model. "
+ "You can add your own model config parser."
+ )
+
+ def set_mixed_precision(self, mixed_precision):
+ if mixed_precision == "fp16":
+ self.megatron_lm_default_args["fp16"] = True
+ elif mixed_precision == "bf16":
+ self.megatron_lm_default_args["bf16"] = True
+ self.DDP_impl = "local"
+ self.megatron_lm_default_args["DDP_impl"] = self.DDP_impl
+
+ def set_training_args(self, micro_batch_size, dp_degree):
+ self.data_parallel_size = dp_degree
+ self.micro_batch_size = micro_batch_size
+ self.global_batch_size = dp_degree * micro_batch_size * self.num_micro_batches
+ self.megatron_lm_default_args["data_parallel_size"] = self.data_parallel_size
+ self.megatron_lm_default_args["micro_batch_size"] = self.micro_batch_size
+ self.megatron_lm_default_args["global_batch_size"] = self.global_batch_size
+
+ def set_optimizer_type(self, optimizer):
+ optimizer_name = optimizer.__class__.__name__.lower()
+ if "adam" in optimizer_name:
+ self.megatron_lm_default_args["optimizer"] = "adam"
+ self.megatron_lm_default_args["adam_beta1"] = optimizer.defaults["betas"][0]
+ self.megatron_lm_default_args["adam_beta2"] = optimizer.defaults["betas"][1]
+ self.megatron_lm_default_args["adam_eps"] = optimizer.defaults["eps"]
+ elif "sgd" in optimizer_name:
+ self.megatron_lm_default_args["optimizer"] = "sgd"
+ self.megatron_lm_default_args["sgd_momentum"] = optimizer.defaults["momentum"]
+ else:
+ raise ValueError(f"Optimizer {optimizer_name} is not supported by Megatron-LM")
+
+ self.megatron_lm_default_args["lr"] = optimizer.defaults["lr"]
+ self.megatron_lm_default_args["weight_decay"] = optimizer.defaults["weight_decay"]
+
+ def set_scheduler_args(self, scheduler):
+ if self.train_iters is None:
+ self.train_iters = scheduler.total_num_steps // self.megatron_lm_default_args["data_parallel_size"]
+ if self.train_samples is not None:
+ self.train_samples = None
+ warnings.warn(
+ "Ignoring `train_samples` as `train_iters` based on scheduler is being used for training."
+ )
+ if self.lr_warmup_iters is None:
+ self.lr_warmup_iters = scheduler.warmup_num_steps // self.megatron_lm_default_args["data_parallel_size"]
+ if self.lr_warmup_samples is not None:
+ warnings.warn(
+ "Ignoring `lr_warmup_samples` as `lr_warmup_iters` based on scheduler is being used for training."
+ )
+ self.lr_warmup_samples = 0
+
+ self.megatron_lm_default_args["train_iters"] = self.train_iters
+ self.megatron_lm_default_args["lr_warmup_iters"] = self.lr_warmup_iters
+ self.megatron_lm_default_args["train_samples"] = self.train_samples
+ self.megatron_lm_default_args["lr_warmup_samples"] = self.lr_warmup_samples
+ self.megatron_lm_default_args["lr_decay_iters"] = self.lr_decay_iters
+ self.megatron_lm_default_args["lr_decay_samples"] = self.lr_decay_samples
+ self.megatron_lm_default_args["lr_warmup_fraction"] = self.lr_warmup_fraction
+ self.megatron_lm_default_args["lr_decay_style"] = self.lr_decay_style
+ self.megatron_lm_default_args["weight_decay_incr_style"] = self.weight_decay_incr_style
+ self.megatron_lm_default_args["start_weight_decay"] = self.start_weight_decay
+ self.megatron_lm_default_args["end_weight_decay"] = self.end_weight_decay
+ self.megatron_lm_default_args["min_lr"] = self.min_lr
+
+ def set_tensorboard_logging_options(self):
+ from megatron.training.arguments import _add_logging_args
+
+ parser = argparse.ArgumentParser()
+ parser = _add_logging_args(parser)
+ logging_args = parser.parse_known_args()
+ self.dataset_args = vars(logging_args[0])
+ for key, value in self.dataset_args.items():
+ if key.startswith("log_"):
+ self.megatron_lm_default_args[key] = True
+ elif key.startswith("no_log_"):
+ self.megatron_lm_default_args[key.replace("no_", "")] = True
+
+
+MODEL_CONFIGS_TO_MEGATRON_PARSERS = {}
+
+
+def add_model_config_to_megatron_parser(model_type: str):
+ def add_model_config_parser_helper(func):
+ @functools.wraps(func)
+ def wrapper(*args, **kwargs):
+ return func(*args, **kwargs)
+
+ MODEL_CONFIGS_TO_MEGATRON_PARSERS[model_type] = func
+ return wrapper
+
+ return add_model_config_parser_helper
+
+
+@add_model_config_to_megatron_parser("megatron-bert")
+def parse_bert_config(megatron_lm_plugin, model, batch_data):
+ model_type_name = "bert"
+ num_layers = model.config.num_hidden_layers
+ hidden_size = model.config.hidden_size
+ num_attention_heads = model.config.num_attention_heads
+ max_position_embeddings = model.config.max_position_embeddings
+ num_labels = model.config.num_labels
+ orig_vocab_size = model.config.vocab_size
+ pretraining_flag = False
+ if "maskedlm" in model.__class__.__name__.lower():
+ pretraining_flag = True
+ if megatron_lm_plugin.seq_length is not None:
+ if megatron_lm_plugin.encoder_seq_length is not None:
+ warnings.warn("Both `seq_length` and `encoder_seq_length` are set. Using `encoder_seq_length`.")
+ megatron_lm_plugin.seq_length = megatron_lm_plugin.encoder_seq_length
+ elif megatron_lm_plugin.encoder_seq_length is not None:
+ megatron_lm_plugin.seq_length = megatron_lm_plugin.encoder_seq_length
+ elif batch_data is not None:
+ megatron_lm_plugin.seq_length = batch_data["input_ids"].shape[1]
+ else:
+ megatron_lm_plugin.seq_length = max_position_embeddings
+ megatron_lm_plugin.megatron_lm_default_args["seq_length"] = megatron_lm_plugin.seq_length
+ megatron_lm_plugin.megatron_lm_default_args["model_type_name"] = model_type_name
+ megatron_lm_plugin.megatron_lm_default_args["num_layers"] = num_layers
+ megatron_lm_plugin.megatron_lm_default_args["hidden_size"] = hidden_size
+ megatron_lm_plugin.megatron_lm_default_args["num_attention_heads"] = num_attention_heads
+ megatron_lm_plugin.megatron_lm_default_args["max_position_embeddings"] = max_position_embeddings
+ megatron_lm_plugin.megatron_lm_default_args["pretraining_flag"] = pretraining_flag
+ megatron_lm_plugin.megatron_lm_default_args["orig_vocab_size"] = orig_vocab_size
+ megatron_lm_plugin.megatron_lm_default_args["model_return_dict"] = model.config.return_dict
+ megatron_lm_plugin.megatron_lm_default_args["num_labels"] = num_labels
+
+
+@add_model_config_to_megatron_parser("gpt2")
+def parse_gpt2_config(megatron_lm_plugin, model, batch_data):
+ model_type_name = "gpt"
+ num_layers = model.config.n_layer
+ hidden_size = model.config.n_embd
+ num_attention_heads = model.config.n_head
+ max_position_embeddings = model.config.n_positions
+ orig_vocab_size = model.config.vocab_size
+ pretraining_flag = True
+ if megatron_lm_plugin.seq_length is not None:
+ if megatron_lm_plugin.decoder_seq_length is not None:
+ warnings.warn("Both `seq_length` and `decoder_seq_length` are set. Using `decoder_seq_length`.")
+ megatron_lm_plugin.seq_length = megatron_lm_plugin.decoder_seq_length
+ elif megatron_lm_plugin.decoder_seq_length is not None:
+ megatron_lm_plugin.seq_length = megatron_lm_plugin.decoder_seq_length
+ elif batch_data is not None:
+ megatron_lm_plugin.seq_length = batch_data["input_ids"].shape[1]
+ else:
+ megatron_lm_plugin.seq_length = max_position_embeddings
+ megatron_lm_plugin.megatron_lm_default_args["seq_length"] = megatron_lm_plugin.seq_length
+ megatron_lm_plugin.megatron_lm_default_args["return_logits"] = megatron_lm_plugin.return_logits
+ megatron_lm_plugin.megatron_lm_default_args["tokenizer_type"] = "GPT2BPETokenizer"
+ megatron_lm_plugin.megatron_lm_default_args["model_type_name"] = model_type_name
+ megatron_lm_plugin.megatron_lm_default_args["num_layers"] = num_layers
+ megatron_lm_plugin.megatron_lm_default_args["hidden_size"] = hidden_size
+ megatron_lm_plugin.megatron_lm_default_args["num_attention_heads"] = num_attention_heads
+ megatron_lm_plugin.megatron_lm_default_args["max_position_embeddings"] = max_position_embeddings
+ megatron_lm_plugin.megatron_lm_default_args["pretraining_flag"] = pretraining_flag
+ megatron_lm_plugin.megatron_lm_default_args["orig_vocab_size"] = orig_vocab_size
+ megatron_lm_plugin.megatron_lm_default_args["model_return_dict"] = model.config.return_dict
+
+
+@add_model_config_to_megatron_parser("t5")
+def parse_t5_config(megatron_lm_plugin, model, batch_data):
+ model_type_name = "t5"
+ num_layers = model.config.num_layers
+ hidden_size = model.config.d_model
+ num_attention_heads = model.config.num_heads
+ max_position_embeddings = model.config.n_positions if hasattr(model.config, "n_positions") else 1024
+ orig_vocab_size = model.config.vocab_size
+ pretraining_flag = True
+ if megatron_lm_plugin.encoder_seq_length is None:
+ if batch_data is not None:
+ megatron_lm_plugin.encoder_seq_length = batch_data["input_ids"].shape[1]
+ else:
+ megatron_lm_plugin.encoder_seq_length = max_position_embeddings
+ if megatron_lm_plugin.decoder_seq_length is None:
+ if batch_data is not None:
+ megatron_lm_plugin.decoder_seq_length = batch_data["labels"].shape[1]
+ else:
+ megatron_lm_plugin.decoder_seq_length = max_position_embeddings
+ megatron_lm_plugin.megatron_lm_default_args["encoder_seq_length"] = megatron_lm_plugin.encoder_seq_length
+ megatron_lm_plugin.megatron_lm_default_args["decoder_seq_length"] = megatron_lm_plugin.decoder_seq_length
+ megatron_lm_plugin.megatron_lm_default_args["model_type_name"] = model_type_name
+ megatron_lm_plugin.megatron_lm_default_args["num_layers"] = num_layers
+ megatron_lm_plugin.megatron_lm_default_args["hidden_size"] = hidden_size
+ megatron_lm_plugin.megatron_lm_default_args["num_attention_heads"] = num_attention_heads
+ megatron_lm_plugin.megatron_lm_default_args["max_position_embeddings"] = max_position_embeddings
+ megatron_lm_plugin.megatron_lm_default_args["pretraining_flag"] = pretraining_flag
+ megatron_lm_plugin.megatron_lm_default_args["orig_vocab_size"] = orig_vocab_size
+ megatron_lm_plugin.megatron_lm_default_args["model_return_dict"] = model.config.return_dict
+
+
+@add_model_config_to_megatron_parser("llama")
+def parse_llama_config(megatron_lm_plugin, model, batch_data):
+ model_type_name = "gpt"
+ num_layers = model.config.num_hidden_layers
+ pretraining_flag = True
+ hidden_size = model.config.hidden_size
+ num_attention_heads = model.config.num_attention_heads
+ orig_vocab_size = model.config.vocab_size
+
+ max_position_embeddings = model.config.max_position_embeddings
+ seq_length = getattr(model.config, "max_sequence_length", None)
+ if megatron_lm_plugin.seq_length is None:
+ if seq_length is not None:
+ megatron_lm_plugin.seq_length = seq_length
+ elif megatron_lm_plugin.decoder_seq_length is not None:
+ megatron_lm_plugin.seq_length = megatron_lm_plugin.decoder_seq_length
+ elif batch_data is not None:
+ megatron_lm_plugin.seq_length = batch_data["input_ids"].shape[1]
+ else:
+ megatron_lm_plugin.seq_length = max_position_embeddings
+
+ megatron_lm_plugin.megatron_lm_default_args["return_logits"] = megatron_lm_plugin.return_logits
+ megatron_lm_plugin.megatron_lm_default_args["tokenizer_type"] = "Llama2Tokenizer"
+ megatron_lm_plugin.megatron_lm_default_args["model_type_name"] = model_type_name
+ megatron_lm_plugin.megatron_lm_default_args["num_layers"] = num_layers
+ megatron_lm_plugin.megatron_lm_default_args["pretraining_flag"] = pretraining_flag
+ megatron_lm_plugin.megatron_lm_default_args["hidden_size"] = hidden_size
+ megatron_lm_plugin.megatron_lm_default_args["num_attention_heads"] = num_attention_heads
+ megatron_lm_plugin.megatron_lm_default_args["orig_vocab_size"] = orig_vocab_size
+ megatron_lm_plugin.megatron_lm_default_args["max_position_embeddings"] = max_position_embeddings
+ megatron_lm_plugin.megatron_lm_default_args["seq_length"] = megatron_lm_plugin.seq_length
+ megatron_lm_plugin.megatron_lm_default_args["model_return_dict"] = model.config.return_dict
+
+
+@dataclass
+class BnbQuantizationConfig:
+ """
+ A plugin to enable BitsAndBytes 4bit and 8bit quantization
+
+ Args:
+ load_in_8bit (`bool`, defaults to `False`):
+ Enable 8bit quantization.
+ llm_int8_threshold (`float`, defaults to `6.0`):
+ Value of the outliner threshold. Only relevant when `load_in_8bit=True`.
+ load_in_4_bit (`bool`, defaults to `False`):
+ Enable 4bit quantization.
+ bnb_4bit_quant_type (`str`, defaults to `fp4`):
+ Set the quantization data type in the `bnb.nn.Linear4Bit` layers. Options are {'fp4','np4'}.
+ bnb_4bit_use_double_quant (`bool`, defaults to `False`):
+ Enable nested quantization where the quantization constants from the first quantization are quantized
+ again.
+ bnb_4bit_compute_dtype (`bool`, defaults to `fp16`):
+ This sets the computational type which might be different than the input time. For example, inputs might be
+ fp32, but computation can be set to bf16 for speedups. Options are {'fp32','fp16','bf16'}.
+ torch_dtype (`torch.dtype`, defaults to `None`):
+ This sets the dtype of the remaining non quantized layers. `bitsandbytes` library suggests to set the value
+ to `torch.float16` for 8 bit model and use the same dtype as the compute dtype for 4 bit model.
+ skip_modules (`List[str]`, defaults to `None`):
+ An explicit list of the modules that we don't quantize. The dtype of these modules will be `torch_dtype`.
+ keep_in_fp32_modules (`List`, defaults to `None`):
+ An explicit list of the modules that we don't quantize. We keep them in `torch.float32`.
+ """
+
+ load_in_8bit: bool = field(default=False, metadata={"help": "enable 8bit quantization."})
+
+ llm_int8_threshold: float = field(
+ default=6.0, metadata={"help": "value of the outliner threshold. only relevant when load_in_8bit=True"}
+ )
+
+ load_in_4bit: bool = field(default=False, metadata={"help": "enable 4bit quantization."})
+
+ bnb_4bit_quant_type: str = field(
+ default="fp4",
+ metadata={
+ "help": "set the quantization data type in the `bnb.nn.Linear4Bit` layers. Options are {'fp4','nf4'}."
+ },
+ )
+
+ bnb_4bit_use_double_quant: bool = field(
+ default=False,
+ metadata={
+ "help": "enable nested quantization where the quantization constants from the first quantization are quantized again."
+ },
+ )
+
+ bnb_4bit_compute_dtype: str = field(
+ default="fp16",
+ metadata={
+ "help": "This sets the computational type which might be different than the input time. For example, inputs might be "
+ "fp32, but computation can be set to bf16 for speedups. Options are {'fp32','fp16','bf16'}."
+ },
+ )
+
+ torch_dtype: torch.dtype = field(
+ default=None,
+ metadata={
+ "help": "this sets the dtype of the remaining non quantized layers. `bitsandbytes` library suggests to set the value"
+ "to `torch.float16` for 8 bit model and use the same dtype as the compute dtype for 4 bit model "
+ },
+ )
+
+ skip_modules: list[str] = field(
+ default=None,
+ metadata={
+ "help": "an explicit list of the modules that we don't quantize. The dtype of these modules will be `torch_dtype`."
+ },
+ )
+
+ keep_in_fp32_modules: list[str] = field(
+ default=None,
+ metadata={"help": "an explicit list of the modules that we don't quantize. We keep them in `torch.float32`."},
+ )
+
+ def __post_init__(self):
+ """
+ Safety checker that arguments are correct - also replaces some NoneType arguments with their default values.
+ """
+ if not isinstance(self.load_in_8bit, bool):
+ raise ValueError("load_in_8bit must be a boolean")
+
+ if not isinstance(self.load_in_4bit, bool):
+ raise ValueError("load_in_4bit must be a boolean")
+
+ if self.load_in_4bit and self.load_in_8bit:
+ raise ValueError("load_in_4bit and load_in_8bit can't be both True")
+
+ if not self.load_in_4bit and not self.load_in_8bit:
+ raise ValueError("load_in_4bit and load_in_8bit can't be both False")
+
+ if not isinstance(self.llm_int8_threshold, (int, float)):
+ raise ValueError("llm_int8_threshold must be a float or an int")
+
+ if not isinstance(self.bnb_4bit_quant_type, str):
+ raise ValueError("bnb_4bit_quant_type must be a string")
+ elif self.bnb_4bit_quant_type not in ["fp4", "nf4"]:
+ raise ValueError(f"bnb_4bit_quant_type must be in ['fp4','nf4'] but found {self.bnb_4bit_quant_type}")
+
+ if not isinstance(self.bnb_4bit_use_double_quant, bool):
+ raise ValueError("bnb_4bit_use_double_quant must be a boolean")
+
+ if isinstance(self.bnb_4bit_compute_dtype, str):
+ if self.bnb_4bit_compute_dtype == "fp32":
+ self.bnb_4bit_compute_dtype = torch.float32
+ elif self.bnb_4bit_compute_dtype == "fp16":
+ self.bnb_4bit_compute_dtype = torch.float16
+ elif self.bnb_4bit_compute_dtype == "bf16":
+ self.bnb_4bit_compute_dtype = torch.bfloat16
+ else:
+ raise ValueError(
+ f"bnb_4bit_compute_dtype must be in ['fp32','fp16','bf16'] but found {self.bnb_4bit_compute_dtype}"
+ )
+ elif not isinstance(self.bnb_4bit_compute_dtype, torch.dtype):
+ raise ValueError("bnb_4bit_compute_dtype must be a string or a torch.dtype")
+
+ if self.skip_modules is not None and not isinstance(self.skip_modules, list):
+ raise ValueError("skip_modules must be a list of strings")
+
+ if self.keep_in_fp32_modules is not None and not isinstance(self.keep_in_fp32_modules, list):
+ raise ValueError("keep_in_fp_32_modules must be a list of strings")
+
+ if self.load_in_4bit:
+ self.target_dtype = CustomDtype.INT4
+
+ if self.load_in_8bit:
+ self.target_dtype = torch.int8
+
+ if self.load_in_4bit and self.llm_int8_threshold != 6.0:
+ warnings.warn("llm_int8_threshold can only be used for model loaded in 8bit")
+
+ if isinstance(self.torch_dtype, str):
+ if self.torch_dtype == "fp32":
+ self.torch_dtype = torch.float32
+ elif self.torch_dtype == "fp16":
+ self.torch_dtype = torch.float16
+ elif self.torch_dtype == "bf16":
+ self.torch_dtype = torch.bfloat16
+ else:
+ raise ValueError(f"torch_dtype must be in ['fp32','fp16','bf16'] but found {self.torch_dtype}")
+ if self.load_in_8bit and self.torch_dtype is None:
+ self.torch_dtype = torch.float16
+
+ if self.load_in_4bit and self.torch_dtype is None:
+ self.torch_dtype = self.bnb_4bit_compute_dtype
+
+ if not isinstance(self.torch_dtype, torch.dtype):
+ raise ValueError("torch_dtype must be a torch.dtype")
+
+
+def get_module_class_from_name(module, name):
+ """
+ Gets a class from a module by its name.
+
+ Args:
+ module (`torch.nn.Module`): The module to get the class from.
+ name (`str`): The name of the class.
+ """
+ modules_children = list(module.children())
+ if module.__class__.__name__ == name:
+ return module.__class__
+ elif len(modules_children) == 0:
+ return
+ else:
+ for child_module in modules_children:
+ module_class = get_module_class_from_name(child_module, name)
+ if module_class is not None:
+ return module_class
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/deepspeed.py b/venv/lib/python3.11/site-packages/accelerate/utils/deepspeed.py
new file mode 100644
index 0000000000000000000000000000000000000000..32e4d4842e91600f66c80d0f9e5cddfa0fa5dfcb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/deepspeed.py
@@ -0,0 +1,371 @@
+# Copyright 2021 The HuggingFace Team. All rights reserved.
+#
+# 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 base64
+import json
+import os
+from copy import deepcopy
+
+from torch import optim
+
+from ..optimizer import AcceleratedOptimizer
+from ..scheduler import AcceleratedScheduler
+from .dataclasses import DistributedType
+from .imports import is_bnb_available
+from .versions import compare_versions
+
+
+def map_pytorch_optim_to_deepspeed(optimizer):
+ """
+ Args:
+ optimizer: torch.optim.Optimizer
+
+ Returns the DeepSeedCPUOptimizer (deepspeed.ops) version of the optimizer.
+ """
+
+ defaults = {k: v for k, v in optimizer.defaults.items() if k in ["lr", "weight_decay"]}
+
+ # Select the DeepSpeedCPUOptimizer based on the original optimizer class.
+ # DeepSpeedCPUAdam is the default
+ from deepspeed.ops.adam import DeepSpeedCPUAdam
+
+ optimizer_class = DeepSpeedCPUAdam
+
+ # For DeepSpeedCPUAdam (adamw_mode)
+ if compare_versions("deepspeed", ">=", "0.3.1"):
+ defaults["adamw_mode"] = False
+ is_adaw = isinstance(optimizer, optim.AdamW)
+
+ if is_bnb_available() and not is_adaw:
+ import bitsandbytes.optim as bnb_opt
+
+ if isinstance(optimizer, (bnb_opt.AdamW, bnb_opt.AdamW32bit)):
+ try:
+ is_adaw = optimizer.optim_bits == 32
+ except AttributeError:
+ is_adaw = optimizer.args.optim_bits == 32
+ else:
+ is_adaw = False
+
+ if is_adaw:
+ defaults["adamw_mode"] = True
+
+ # For DeepSpeedCPUAdagrad
+ if compare_versions("deepspeed", ">=", "0.5.5"):
+ # Check if the optimizer is PyTorch's Adagrad.
+ is_ada = isinstance(optimizer, optim.Adagrad)
+ # If not, and bitsandbytes is available,
+ # # check if the optimizer is the 32-bit bitsandbytes Adagrad.
+ if is_bnb_available() and not is_ada:
+ import bitsandbytes.optim as bnb_opt
+
+ if isinstance(optimizer, (bnb_opt.Adagrad, bnb_opt.Adagrad32bit)):
+ try:
+ is_ada = optimizer.optim_bits == 32
+ except AttributeError:
+ is_ada = optimizer.args.optim_bits == 32
+ if is_ada:
+ from deepspeed.ops.adagrad import DeepSpeedCPUAdagrad
+
+ optimizer_class = DeepSpeedCPUAdagrad
+
+ # For DeepSpeedCPULion
+ if is_bnb_available(min_version="0.38.0") and compare_versions("deepspeed", ">=", "0.11.0"):
+ from bitsandbytes.optim import Lion, Lion32bit
+
+ if isinstance(optimizer, (Lion, Lion32bit)):
+ try:
+ is_bnb_32bits = optimizer.optim_bits == 32
+ except AttributeError:
+ is_bnb_32bits = optimizer.args.optim_bits == 32
+ if is_bnb_32bits:
+ from deepspeed.ops.lion import DeepSpeedCPULion
+
+ optimizer_class = DeepSpeedCPULion
+
+ return optimizer_class(optimizer.param_groups, **defaults)
+
+
+def get_active_deepspeed_plugin(state):
+ """
+ Returns the currently active DeepSpeedPlugin.
+
+ Raises:
+ ValueError: If DeepSpeed was not enabled and this function is called.
+ """
+ if state.distributed_type != DistributedType.DEEPSPEED:
+ raise ValueError(
+ "Couldn't retrieve the active `DeepSpeedPlugin` as none were enabled. "
+ "Please make sure that either `Accelerator` is configured for `deepspeed` "
+ "or make sure that the desired `DeepSpeedPlugin` has been enabled (`AcceleratorState().select_deepspeed_plugin(name)`) "
+ "before calling this function."
+ )
+ if not isinstance(state.deepspeed_plugins, dict):
+ return state.deepspeed_plugins
+ return next(plugin for plugin in state.deepspeed_plugins.values() if plugin.selected)
+
+
+class HfDeepSpeedConfig:
+ """
+ This object contains a DeepSpeed configuration dictionary and can be quickly queried for things like zero stage.
+
+ A `weakref` of this object is stored in the module's globals to be able to access the config from areas where
+ things like the Trainer object is not available (e.g. `from_pretrained` and `_get_resized_embeddings`). Therefore
+ it's important that this object remains alive while the program is still running.
+
+ [`Trainer`] uses the `HfTrainerDeepSpeedConfig` subclass instead. That subclass has logic to sync the configuration
+ with values of [`TrainingArguments`] by replacing special placeholder values: `"auto"`. Without this special logic
+ the DeepSpeed configuration is not modified in any way.
+
+ Args:
+ config_file_or_dict (`Union[str, Dict]`): path to DeepSpeed config file or dict.
+
+ """
+
+ def __init__(self, config_file_or_dict):
+ if isinstance(config_file_or_dict, dict):
+ # Don't modify user's data should they want to reuse it (e.g. in tests), because once we
+ # modified it, it will not be accepted here again, since `auto` values would have been overridden
+ config = deepcopy(config_file_or_dict)
+ elif os.path.exists(config_file_or_dict):
+ with open(config_file_or_dict, encoding="utf-8") as f:
+ config = json.load(f)
+ else:
+ try:
+ try:
+ # First try parsing as JSON directly
+ config = json.loads(config_file_or_dict)
+ except json.JSONDecodeError:
+ # If that fails, try base64 decoding
+ config_decoded = base64.urlsafe_b64decode(config_file_or_dict).decode("utf-8")
+ config = json.loads(config_decoded)
+ except (UnicodeDecodeError, AttributeError, ValueError):
+ raise ValueError(
+ f"Expected a string path to an existing deepspeed config, or a dictionary, or a base64 encoded string. Received: {config_file_or_dict}"
+ )
+
+ self.config = config
+
+ self.set_stage_and_offload()
+
+ def set_stage_and_offload(self):
+ # zero stage - this is done as early as possible, before model is created, to allow
+ # ``is_deepspeed_zero3_enabled`` query and getting to the early deepspeed config object
+ # during ``zero.Init()`` which needs to know the dtype, and some other hparams.
+ self._stage = self.get_value("zero_optimization.stage", -1)
+
+ # offload
+ self._offload = False
+ if self.is_zero2() or self.is_zero3():
+ offload_devices_valid = set(["cpu", "nvme"])
+ offload_devices = set(
+ [
+ self.get_value("zero_optimization.offload_optimizer.device"),
+ self.get_value("zero_optimization.offload_param.device"),
+ ]
+ )
+ if len(offload_devices & offload_devices_valid) > 0:
+ self._offload = True
+
+ def find_config_node(self, ds_key_long):
+ config = self.config
+
+ # find the config node of interest if it exists
+ nodes = ds_key_long.split(".")
+ ds_key = nodes.pop()
+ for node in nodes:
+ config = config.get(node)
+ if config is None:
+ return None, ds_key
+
+ return config, ds_key
+
+ def get_value(self, ds_key_long, default=None):
+ """
+ Returns the set value or `default` if no value is set
+ """
+ config, ds_key = self.find_config_node(ds_key_long)
+ if config is None:
+ return default
+ return config.get(ds_key, default)
+
+ def del_config_sub_tree(self, ds_key_long, must_exist=False):
+ """
+ Deletes a sub-section of the config file if it's found.
+
+ Unless `must_exist` is `True` the section doesn't have to exist.
+ """
+ config = self.config
+
+ # find the config node of interest if it exists
+ nodes = ds_key_long.split(".")
+ for node in nodes:
+ parent_config = config
+ config = config.get(node)
+ if config is None:
+ if must_exist:
+ raise ValueError(f"Can't find {ds_key_long} entry in the config: {self.config}")
+ else:
+ return
+
+ # if found remove it
+ if parent_config is not None:
+ parent_config.pop(node)
+
+ def is_true(self, ds_key_long):
+ """
+ Returns `True`/``False` only if the value is set, always `False` otherwise. So use this method to ask the very
+ specific question of whether the value is set to `True` (and it's not set to `False`` or isn't set).
+
+ """
+ value = self.get_value(ds_key_long)
+ return False if value is None else bool(value)
+
+ def is_false(self, ds_key_long):
+ """
+ Returns `True`/``False` only if the value is set, always `False` otherwise. So use this method to ask the very
+ specific question of whether the value is set to `False` (and it's not set to `True`` or isn't set).
+ """
+ value = self.get_value(ds_key_long)
+ return False if value is None else not bool(value)
+
+ def is_zero2(self):
+ return self._stage == 2
+
+ def is_zero3(self):
+ return self._stage == 3
+
+ def is_offload(self):
+ return self._offload
+
+
+class DeepSpeedEngineWrapper:
+ """
+ Internal wrapper for deepspeed.runtime.engine.DeepSpeedEngine. This is used to follow conventional training loop.
+
+ Args:
+ engine (deepspeed.runtime.engine.DeepSpeedEngine): deepspeed engine to wrap
+ """
+
+ def __init__(self, engine):
+ self.engine = engine
+
+ def backward(self, loss, **kwargs):
+ # runs backpropagation and handles mixed precision
+ self.engine.backward(loss, **kwargs)
+
+ # Deepspeed's `engine.step` performs the following operations:
+ # - gradient accumulation check
+ # - gradient clipping
+ # - optimizer step
+ # - zero grad
+ # - checking overflow
+ # - lr_scheduler step (only if engine.lr_scheduler is not None)
+ self.engine.step()
+ # and this plugin overrides the above calls with no-ops when Accelerate runs under
+ # Deepspeed, but allows normal functionality for non-Deepspeed cases thus enabling a simple
+ # training loop that works transparently under many training regimes.
+
+
+class DeepSpeedOptimizerWrapper(AcceleratedOptimizer):
+ """
+ Internal wrapper around a deepspeed optimizer.
+
+ Args:
+ optimizer (`torch.optim.optimizer.Optimizer`):
+ The optimizer to wrap.
+ """
+
+ def __init__(self, optimizer):
+ super().__init__(optimizer, device_placement=False, scaler=None)
+ self.__has_overflow__ = hasattr(self.optimizer, "overflow")
+
+ def zero_grad(self, set_to_none=None):
+ pass # `accelerator.backward(loss)` is doing that automatically. Therefore, its implementation is not needed
+
+ def step(self):
+ pass # `accelerator.backward(loss)` is doing that automatically. Therefore, its implementation is not needed
+
+ @property
+ def step_was_skipped(self):
+ """Whether or not the optimizer step was done, or skipped because of gradient overflow."""
+ if self.__has_overflow__:
+ return self.optimizer.overflow
+ return False
+
+
+class DeepSpeedSchedulerWrapper(AcceleratedScheduler):
+ """
+ Internal wrapper around a deepspeed scheduler.
+
+ Args:
+ scheduler (`torch.optim.lr_scheduler.LambdaLR`):
+ The scheduler to wrap.
+ optimizers (one or a list of `torch.optim.Optimizer`):
+ """
+
+ def __init__(self, scheduler, optimizers):
+ super().__init__(scheduler, optimizers)
+
+ def step(self):
+ pass # `accelerator.backward(loss)` is doing that automatically. Therefore, its implementation is not needed
+
+
+class DummyOptim:
+ """
+ Dummy optimizer presents model parameters or param groups, this is primarily used to follow conventional training
+ loop when optimizer config is specified in the deepspeed config file.
+
+ Args:
+ lr (float):
+ Learning rate.
+ params (iterable): iterable of parameters to optimize or dicts defining
+ parameter groups
+ weight_decay (float):
+ Weight decay.
+ **kwargs (additional keyword arguments, *optional*):
+ Other arguments.
+ """
+
+ def __init__(self, params, lr=0.001, weight_decay=0, **kwargs):
+ self.params = params
+ self.lr = lr
+ self.weight_decay = weight_decay
+ self.kwargs = kwargs
+
+
+class DummyScheduler:
+ """
+ Dummy scheduler presents model parameters or param groups, this is primarily used to follow conventional training
+ loop when scheduler config is specified in the deepspeed config file.
+
+ Args:
+ optimizer (`torch.optim.optimizer.Optimizer`):
+ The optimizer to wrap.
+ total_num_steps (int, *optional*):
+ Total number of steps.
+ warmup_num_steps (int, *optional*):
+ Number of steps for warmup.
+ lr_scheduler_callable (callable, *optional*):
+ A callable function that creates an LR Scheduler. It accepts only one argument `optimizer`.
+ **kwargs (additional keyword arguments, *optional*):
+ Other arguments.
+ """
+
+ def __init__(self, optimizer, total_num_steps=None, warmup_num_steps=0, lr_scheduler_callable=None, **kwargs):
+ self.optimizer = optimizer
+ self.total_num_steps = total_num_steps
+ self.warmup_num_steps = warmup_num_steps
+ self.lr_scheduler_callable = lr_scheduler_callable
+ self.kwargs = kwargs
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/environment.py b/venv/lib/python3.11/site-packages/accelerate/utils/environment.py
new file mode 100644
index 0000000000000000000000000000000000000000..c913702ef30290ed3d19c08a71213663297f32de
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/environment.py
@@ -0,0 +1,421 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 logging
+import math
+import os
+import platform
+import subprocess
+import sys
+from contextlib import contextmanager
+from dataclasses import dataclass, field
+from functools import lru_cache, wraps
+from shutil import which
+from typing import Optional, Union
+
+import torch
+from packaging.version import parse
+
+
+logger = logging.getLogger(__name__)
+
+
+def convert_dict_to_env_variables(current_env: dict):
+ """
+ Verifies that all keys and values in `current_env` do not contain illegal keys or values, and returns a list of
+ strings as the result.
+
+ Example:
+ ```python
+ >>> from accelerate.utils.environment import verify_env
+
+ >>> env = {"ACCELERATE_DEBUG_MODE": "1", "BAD_ENV_NAME": ">> valid_env_items = verify_env(env)
+ >>> print(valid_env_items)
+ ["ACCELERATE_DEBUG_MODE=1\n", "OTHER_ENV=2\n"]
+ ```
+ """
+ forbidden_chars = [";", "\n", "<", ">", " "]
+ valid_env_items = []
+ for key, value in current_env.items():
+ if all(char not in (key + value) for char in forbidden_chars) and len(key) >= 1 and len(value) >= 1:
+ valid_env_items.append(f"{key}={value}\n")
+ else:
+ logger.warning(f"WARNING: Skipping {key}={value} as it contains forbidden characters or missing values.")
+ return valid_env_items
+
+
+def str_to_bool(value, to_bool: bool = False) -> Union[int, bool]:
+ """
+ Converts a string representation of truth to `True` (1) or `False` (0).
+
+ True values are `y`, `yes`, `t`, `true`, `on`, and `1`; False value are `n`, `no`, `f`, `false`, `off`, and `0`;
+ """
+ value = value.lower()
+ if value in ("y", "yes", "t", "true", "on", "1"):
+ return 1 if not to_bool else True
+ elif value in ("n", "no", "f", "false", "off", "0"):
+ return 0 if not to_bool else False
+ else:
+ raise ValueError(f"invalid truth value {value}")
+
+
+def get_int_from_env(env_keys, default):
+ """Returns the first positive env value found in the `env_keys` list or the default."""
+ for e in env_keys:
+ val = int(os.environ.get(e, -1))
+ if val >= 0:
+ return val
+ return default
+
+
+def parse_flag_from_env(key, default=False):
+ """Returns truthy value for `key` from the env if available else the default."""
+ value = os.environ.get(key, str(default))
+ return str_to_bool(value) == 1 # As its name indicates `str_to_bool` actually returns an int...
+
+
+def parse_choice_from_env(key, default="no"):
+ value = os.environ.get(key, str(default))
+ return value
+
+
+def are_libraries_initialized(*library_names: str) -> list[str]:
+ """
+ Checks if any of `library_names` are imported in the environment. Will return any names that are.
+ """
+ return [lib_name for lib_name in library_names if lib_name in sys.modules.keys()]
+
+
+def _nvidia_smi():
+ """
+ Returns the right nvidia-smi command based on the system.
+ """
+ if platform.system() == "Windows":
+ # If platform is Windows and nvidia-smi can't be found in path
+ # try from systemd drive with default installation path
+ command = which("nvidia-smi")
+ if command is None:
+ command = f"{os.environ['systemdrive']}\\Program Files\\NVIDIA Corporation\\NVSMI\\nvidia-smi.exe"
+ else:
+ command = "nvidia-smi"
+ return command
+
+
+def get_gpu_info():
+ """
+ Gets GPU count and names using `nvidia-smi` instead of torch to not initialize CUDA.
+
+ Largely based on the `gputil` library.
+ """
+ # Returns as list of `n` GPUs and their names
+ output = subprocess.check_output(
+ [_nvidia_smi(), "--query-gpu=count,name", "--format=csv,noheader"], universal_newlines=True
+ )
+ output = output.strip()
+ gpus = output.split(os.linesep)
+ # Get names from output
+ gpu_count = len(gpus)
+ gpu_names = [gpu.split(",")[1].strip() for gpu in gpus]
+ return gpu_names, gpu_count
+
+
+def get_driver_version():
+ """
+ Returns the driver version
+
+ In the case of multiple GPUs, will return the first.
+ """
+ output = subprocess.check_output(
+ [_nvidia_smi(), "--query-gpu=driver_version", "--format=csv,noheader"], universal_newlines=True
+ )
+ output = output.strip()
+ return output.split(os.linesep)[0]
+
+
+def check_cuda_p2p_ib_support():
+ """
+ Checks if the devices being used have issues with P2P and IB communications, namely any consumer GPU hardware after
+ the 3090.
+
+ Noteably uses `nvidia-smi` instead of torch to not initialize CUDA.
+ """
+ try:
+ device_names, device_count = get_gpu_info()
+ # As new consumer GPUs get released, add them to `unsupported_devices``
+ unsupported_devices = {"RTX 40"}
+ if device_count > 1:
+ if any(
+ unsupported_device in device_name
+ for device_name in device_names
+ for unsupported_device in unsupported_devices
+ ):
+ # Check if they have the right driver version
+ acceptable_driver_version = "550.40.07"
+ current_driver_version = get_driver_version()
+ if parse(current_driver_version) < parse(acceptable_driver_version):
+ return False
+ return True
+ except Exception:
+ pass
+ return True
+
+
+@lru_cache
+def check_cuda_fp8_capability():
+ """
+ Checks if the current GPU available supports FP8.
+
+ Notably might initialize `torch.cuda` to check.
+ """
+
+ try:
+ # try to get the compute capability from nvidia-smi
+ output = subprocess.check_output(
+ [_nvidia_smi(), "--query-gpu=compute_capability", "--format=csv,noheader"], universal_newlines=True
+ )
+ output = output.strip()
+ # we take the first GPU's compute capability
+ compute_capability = tuple(map(int, output.split(os.linesep)[0].split(".")))
+ except Exception:
+ compute_capability = torch.cuda.get_device_capability()
+
+ return compute_capability >= (8, 9)
+
+
+@dataclass
+class CPUInformation:
+ """
+ Stores information about the CPU in a distributed environment. It contains the following attributes:
+ - rank: The rank of the current process.
+ - world_size: The total number of processes in the world.
+ - local_rank: The rank of the current process on the local node.
+ - local_world_size: The total number of processes on the local node.
+ """
+
+ rank: int = field(default=0, metadata={"help": "The rank of the current process."})
+ world_size: int = field(default=1, metadata={"help": "The total number of processes in the world."})
+ local_rank: int = field(default=0, metadata={"help": "The rank of the current process on the local node."})
+ local_world_size: int = field(default=1, metadata={"help": "The total number of processes on the local node."})
+
+
+def get_cpu_distributed_information() -> CPUInformation:
+ """
+ Returns various information about the environment in relation to CPU distributed training as a `CPUInformation`
+ dataclass.
+ """
+ information = {}
+ information["rank"] = get_int_from_env(["RANK", "PMI_RANK", "OMPI_COMM_WORLD_RANK", "MV2_COMM_WORLD_RANK"], 0)
+ information["world_size"] = get_int_from_env(
+ ["WORLD_SIZE", "PMI_SIZE", "OMPI_COMM_WORLD_SIZE", "MV2_COMM_WORLD_SIZE"], 1
+ )
+ information["local_rank"] = get_int_from_env(
+ ["LOCAL_RANK", "MPI_LOCALRANKID", "OMPI_COMM_WORLD_LOCAL_RANK", "MV2_COMM_WORLD_LOCAL_RANK"], 0
+ )
+ information["local_world_size"] = get_int_from_env(
+ ["LOCAL_WORLD_SIZE", "MPI_LOCALNRANKS", "OMPI_COMM_WORLD_LOCAL_SIZE", "MV2_COMM_WORLD_LOCAL_SIZE"],
+ 1,
+ )
+ return CPUInformation(**information)
+
+
+def override_numa_affinity(local_process_index: int, verbose: Optional[bool] = None) -> None:
+ """
+ Overrides whatever NUMA affinity is set for the current process. This is very taxing and requires recalculating the
+ affinity to set, ideally you should use `utils.environment.set_numa_affinity` instead.
+
+ Args:
+ local_process_index (int):
+ The index of the current process on the current server.
+ verbose (bool, *optional*):
+ Whether to log out the assignment of each CPU. If `ACCELERATE_DEBUG_MODE` is enabled, will default to True.
+ """
+ if verbose is None:
+ verbose = parse_flag_from_env("ACCELERATE_DEBUG_MODE", False)
+ if torch.cuda.is_available():
+ from accelerate.utils import is_pynvml_available
+
+ if not is_pynvml_available():
+ raise ImportError(
+ "To set CPU affinity on CUDA GPUs the `pynvml` package must be available. (`pip install pynvml`)"
+ )
+ import pynvml as nvml
+
+ # The below code is based on https://github.com/NVIDIA/DeepLearningExamples/blob/master/TensorFlow2/LanguageModeling/BERT/gpu_affinity.py
+ nvml.nvmlInit()
+ num_elements = math.ceil(os.cpu_count() / 64)
+ handle = nvml.nvmlDeviceGetHandleByIndex(local_process_index)
+ affinity_string = ""
+ for j in nvml.nvmlDeviceGetCpuAffinity(handle, num_elements):
+ # assume nvml returns list of 64 bit ints
+ affinity_string = f"{j:064b}{affinity_string}"
+ affinity_list = [int(x) for x in affinity_string]
+ affinity_list.reverse() # so core 0 is the 0th element
+ affinity_to_set = [i for i, e in enumerate(affinity_list) if e != 0]
+ os.sched_setaffinity(0, affinity_to_set)
+ if verbose:
+ cpu_cores = os.sched_getaffinity(0)
+ logger.info(f"Assigning {len(cpu_cores)} cpu cores to process {local_process_index}: {cpu_cores}")
+
+
+@lru_cache
+def set_numa_affinity(local_process_index: int, verbose: Optional[bool] = None) -> None:
+ """
+ Assigns the current process to a specific NUMA node. Ideally most efficient when having at least 2 cpus per node.
+
+ This result is cached between calls. If you want to override it, please use
+ `accelerate.utils.environment.override_numa_afifnity`.
+
+ Args:
+ local_process_index (int):
+ The index of the current process on the current server.
+ verbose (bool, *optional*):
+ Whether to print the new cpu cores assignment for each process. If `ACCELERATE_DEBUG_MODE` is enabled, will
+ default to True.
+ """
+ override_numa_affinity(local_process_index=local_process_index, verbose=verbose)
+
+
+@contextmanager
+def clear_environment():
+ """
+ A context manager that will temporarily clear environment variables.
+
+ When this context exits, the previous environment variables will be back.
+
+ Example:
+
+ ```python
+ >>> import os
+ >>> from accelerate.utils import clear_environment
+
+ >>> os.environ["FOO"] = "bar"
+ >>> with clear_environment():
+ ... print(os.environ)
+ ... os.environ["FOO"] = "new_bar"
+ ... print(os.environ["FOO"])
+ {}
+ new_bar
+
+ >>> print(os.environ["FOO"])
+ bar
+ ```
+ """
+ _old_os_environ = os.environ.copy()
+ os.environ.clear()
+
+ try:
+ yield
+ finally:
+ os.environ.clear() # clear any added keys,
+ os.environ.update(_old_os_environ) # then restore previous environment
+
+
+@contextmanager
+def patch_environment(**kwargs):
+ """
+ A context manager that will add each keyword argument passed to `os.environ` and remove them when exiting.
+
+ Will convert the values in `kwargs` to strings and upper-case all the keys.
+
+ Example:
+
+ ```python
+ >>> import os
+ >>> from accelerate.utils import patch_environment
+
+ >>> with patch_environment(FOO="bar"):
+ ... print(os.environ["FOO"]) # prints "bar"
+ >>> print(os.environ["FOO"]) # raises KeyError
+ ```
+ """
+ existing_vars = {}
+ for key, value in kwargs.items():
+ key = key.upper()
+ if key in os.environ:
+ existing_vars[key] = os.environ[key]
+ os.environ[key] = str(value)
+
+ try:
+ yield
+ finally:
+ for key in kwargs:
+ key = key.upper()
+ if key in existing_vars:
+ # restore previous value
+ os.environ[key] = existing_vars[key]
+ else:
+ os.environ.pop(key, None)
+
+
+def purge_accelerate_environment(func_or_cls):
+ """Decorator to clean up accelerate environment variables set by the decorated class or function.
+
+ In some circumstances, calling certain classes or functions can result in accelerate env vars being set and not
+ being cleaned up afterwards. As an example, when calling:
+
+ TrainingArguments(fp16=True, ...)
+
+ The following env var will be set:
+
+ ACCELERATE_MIXED_PRECISION=fp16
+
+ This can affect subsequent code, since the env var takes precedence over TrainingArguments(fp16=False). This is
+ especially relevant for unit testing, where we want to avoid the individual tests to have side effects on one
+ another. Decorate the unit test function or whole class with this decorator to ensure that after each test, the env
+ vars are cleaned up. This works for both unittest.TestCase and normal classes (pytest); it also works when
+ decorating the parent class.
+
+ """
+ prefix = "ACCELERATE_"
+
+ @contextmanager
+ def env_var_context():
+ # Store existing accelerate env vars
+ existing_vars = {k: v for k, v in os.environ.items() if k.startswith(prefix)}
+ try:
+ yield
+ finally:
+ # Restore original env vars or remove new ones
+ for key in [k for k in os.environ if k.startswith(prefix)]:
+ if key in existing_vars:
+ os.environ[key] = existing_vars[key]
+ else:
+ os.environ.pop(key, None)
+
+ def wrap_function(func):
+ @wraps(func)
+ def wrapper(*args, **kwargs):
+ with env_var_context():
+ return func(*args, **kwargs)
+
+ wrapper._accelerate_is_purged_environment_wrapped = True
+ return wrapper
+
+ if not isinstance(func_or_cls, type):
+ return wrap_function(func_or_cls)
+
+ # Handle classes by wrapping test methods
+ def wrap_test_methods(test_class_instance):
+ for name in dir(test_class_instance):
+ if name.startswith("test"):
+ method = getattr(test_class_instance, name)
+ if callable(method) and not hasattr(method, "_accelerate_is_purged_environment_wrapped"):
+ setattr(test_class_instance, name, wrap_function(method))
+ return test_class_instance
+
+ # Handle inheritance
+ wrap_test_methods(func_or_cls)
+ func_or_cls.__init_subclass__ = classmethod(lambda cls, **kw: wrap_test_methods(cls))
+ return func_or_cls
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/fsdp_utils.py b/venv/lib/python3.11/site-packages/accelerate/utils/fsdp_utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..c804fb427d7da582768eb894dc9eaf80a9799ada
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/fsdp_utils.py
@@ -0,0 +1,629 @@
+# Copyright 2023 The HuggingFace Team. All rights reserved.
+#
+# 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 functools
+import os
+import shutil
+import warnings
+from collections import defaultdict
+from contextlib import nullcontext
+from pathlib import Path
+from typing import Callable
+
+import torch
+
+from ..logging import get_logger
+from .constants import FSDP_MODEL_NAME, OPTIMIZER_NAME, SAFE_WEIGHTS_NAME, WEIGHTS_NAME
+from .dataclasses import get_module_class_from_name
+from .modeling import is_peft_model
+from .other import get_module_children_bottom_up, is_compiled_module, save
+from .versions import is_torch_version
+
+
+logger = get_logger(__name__)
+
+
+def enable_fsdp_ram_efficient_loading():
+ """
+ Enables RAM efficient loading of Hugging Face models for FSDP in the environment.
+ """
+ # Sets values for `transformers.modeling_utils.is_fsdp_enabled`
+ if "ACCELERATE_USE_FSDP" not in os.environ:
+ os.environ["ACCELERATE_USE_FSDP"] = "True"
+ os.environ["FSDP_CPU_RAM_EFFICIENT_LOADING"] = "True"
+
+
+def disable_fsdp_ram_efficient_loading():
+ """
+ Disables RAM efficient loading of Hugging Face models for FSDP in the environment.
+ """
+ os.environ["FSDP_CPU_RAM_EFFICIENT_LOADING"] = "False"
+
+
+def _get_model_state_dict(model, adapter_only=False):
+ if adapter_only and is_peft_model(model):
+ from peft import get_peft_model_state_dict
+
+ return get_peft_model_state_dict(model, adapter_name=model.active_adapter)
+ else:
+ return model.state_dict()
+
+
+def _set_model_state_dict(model, state_dict, adapter_only=False):
+ if adapter_only and is_peft_model(model):
+ from peft import set_peft_model_state_dict
+
+ return set_peft_model_state_dict(model, state_dict, adapter_name=model.active_adapter)
+ else:
+ return model.load_state_dict(state_dict)
+
+
+def save_fsdp_model(fsdp_plugin, accelerator, model, output_dir, model_index=0, adapter_only=False):
+ # Note: We import here to reduce import time from general modules, and isolate outside dependencies
+ import torch.distributed.checkpoint as dist_cp
+ from torch.distributed.checkpoint.default_planner import DefaultSavePlanner
+ from torch.distributed.fsdp.fully_sharded_data_parallel import FullyShardedDataParallel as FSDP
+ from torch.distributed.fsdp.fully_sharded_data_parallel import StateDictType
+
+ os.makedirs(output_dir, exist_ok=True)
+ if fsdp_plugin.state_dict_type == StateDictType.FULL_STATE_DICT and fsdp_plugin.fsdp_version == 1:
+ # FSDP raises error when single GPU is used with `offload_to_cpu=True` for FULL_STATE_DICT
+ # so, only enable it when num_processes>1
+ is_multi_process = accelerator.num_processes > 1
+ fsdp_plugin.state_dict_config.offload_to_cpu = is_multi_process
+ fsdp_plugin.state_dict_config.rank0_only = is_multi_process
+
+ ctx = (
+ FSDP.state_dict_type(
+ model, fsdp_plugin.state_dict_type, fsdp_plugin.state_dict_config, fsdp_plugin.optim_state_dict_config
+ )
+ if fsdp_plugin.fsdp_version == 1
+ else nullcontext()
+ )
+
+ with ctx:
+ state_dict = _get_model_state_dict(model, adapter_only=adapter_only)
+ if fsdp_plugin.state_dict_type == StateDictType.FULL_STATE_DICT:
+ weights_name = f"{FSDP_MODEL_NAME}.bin" if model_index == 0 else f"{FSDP_MODEL_NAME}_{model_index}.bin"
+ output_model_file = os.path.join(output_dir, weights_name)
+ if accelerator.process_index == 0:
+ logger.info(f"Saving model to {output_model_file}")
+ torch.save(state_dict, output_model_file)
+ logger.info(f"Model saved to {output_model_file}")
+ # Invariant: `LOCAL_STATE_DICT` is never possible with `FSDP2`
+ elif fsdp_plugin.state_dict_type == StateDictType.LOCAL_STATE_DICT:
+ weights_name = (
+ f"{FSDP_MODEL_NAME}_rank{accelerator.process_index}.bin"
+ if model_index == 0
+ else f"{FSDP_MODEL_NAME}_{model_index}_rank{accelerator.process_index}.bin"
+ )
+ output_model_file = os.path.join(output_dir, weights_name)
+ logger.info(f"Saving model to {output_model_file}")
+ torch.save(state_dict, output_model_file)
+ logger.info(f"Model saved to {output_model_file}")
+ elif fsdp_plugin.state_dict_type == StateDictType.SHARDED_STATE_DICT:
+ ckpt_dir = os.path.join(output_dir, f"{FSDP_MODEL_NAME}_{model_index}")
+ os.makedirs(ckpt_dir, exist_ok=True)
+ logger.info(f"Saving model to {ckpt_dir}")
+ state_dict = {"model": state_dict}
+
+ dist_cp.save_state_dict(
+ state_dict=state_dict,
+ storage_writer=dist_cp.FileSystemWriter(ckpt_dir),
+ planner=DefaultSavePlanner(),
+ )
+ logger.info(f"Model saved to {ckpt_dir}")
+
+
+def load_fsdp_model(fsdp_plugin, accelerator, model, input_dir, model_index=0, adapter_only=False):
+ # Note: We import here to reduce import time from general modules, and isolate outside dependencies
+ import torch.distributed.checkpoint as dist_cp
+ from torch.distributed.checkpoint.default_planner import DefaultLoadPlanner
+ from torch.distributed.fsdp.fully_sharded_data_parallel import FullyShardedDataParallel as FSDP
+ from torch.distributed.fsdp.fully_sharded_data_parallel import StateDictType
+
+ accelerator.wait_for_everyone()
+ if fsdp_plugin.state_dict_type == StateDictType.FULL_STATE_DICT and fsdp_plugin.fsdp_version == 1:
+ # FSDP raises error when single GPU is used with `offload_to_cpu=True` for FULL_STATE_DICT
+ # so, only enable it when num_processes>1
+ is_multi_process = accelerator.num_processes > 1
+ fsdp_plugin.state_dict_config.offload_to_cpu = is_multi_process
+ fsdp_plugin.state_dict_config.rank0_only = is_multi_process
+
+ ctx = (
+ FSDP.state_dict_type(
+ model, fsdp_plugin.state_dict_type, fsdp_plugin.state_dict_config, fsdp_plugin.optim_state_dict_config
+ )
+ if fsdp_plugin.fsdp_version == 1
+ else nullcontext()
+ )
+
+ with ctx:
+ if fsdp_plugin.state_dict_type == StateDictType.FULL_STATE_DICT:
+ if type(model) is not FSDP and accelerator.process_index != 0:
+ if not fsdp_plugin.sync_module_states and fsdp_plugin.fsdp_version == 1:
+ raise ValueError(
+ "Set the `sync_module_states` flag to `True` so that model states are synced across processes when "
+ "initializing FSDP object"
+ )
+ return
+ weights_name = f"{FSDP_MODEL_NAME}.bin" if model_index == 0 else f"{FSDP_MODEL_NAME}_{model_index}.bin"
+ input_model_file = os.path.join(input_dir, weights_name)
+ logger.info(f"Loading model from {input_model_file}")
+ state_dict = torch.load(input_model_file)
+ logger.info(f"Model loaded from {input_model_file}")
+ elif fsdp_plugin.state_dict_type == StateDictType.LOCAL_STATE_DICT:
+ weights_name = (
+ f"{FSDP_MODEL_NAME}_rank{accelerator.process_index}.bin"
+ if model_index == 0
+ else f"{FSDP_MODEL_NAME}_{model_index}_rank{accelerator.process_index}.bin"
+ )
+ input_model_file = os.path.join(input_dir, weights_name)
+ logger.info(f"Loading model from {input_model_file}")
+ state_dict = torch.load(input_model_file)
+ logger.info(f"Model loaded from {input_model_file}")
+ elif fsdp_plugin.state_dict_type == StateDictType.SHARDED_STATE_DICT:
+ ckpt_dir = (
+ os.path.join(input_dir, f"{FSDP_MODEL_NAME}_{model_index}")
+ if f"{FSDP_MODEL_NAME}" not in input_dir
+ else input_dir
+ )
+ logger.info(f"Loading model from {ckpt_dir}")
+ state_dict = {"model": _get_model_state_dict(model, adapter_only=adapter_only)}
+ dist_cp.load_state_dict(
+ state_dict=state_dict,
+ storage_reader=dist_cp.FileSystemReader(ckpt_dir),
+ planner=DefaultLoadPlanner(),
+ )
+ state_dict = state_dict["model"]
+ logger.info(f"Model loaded from {ckpt_dir}")
+
+ if fsdp_plugin.fsdp_version == 1:
+ load_result = _set_model_state_dict(model, state_dict, adapter_only=adapter_only)
+ else:
+ from torch.distributed.checkpoint.state_dict import set_model_state_dict
+
+ load_result = set_model_state_dict(model, state_dict)
+ return load_result
+
+
+def save_fsdp_optimizer(fsdp_plugin, accelerator, optimizer, model, output_dir, optimizer_index=0):
+ # Note: We import here to reduce import time from general modules, and isolate outside dependencies
+ import torch.distributed.checkpoint as dist_cp
+ from torch.distributed.checkpoint.default_planner import DefaultSavePlanner
+ from torch.distributed.fsdp.fully_sharded_data_parallel import FullyShardedDataParallel as FSDP
+ from torch.distributed.fsdp.fully_sharded_data_parallel import StateDictType
+
+ os.makedirs(output_dir, exist_ok=True)
+
+ ctx = (
+ FSDP.state_dict_type(
+ model, fsdp_plugin.state_dict_type, fsdp_plugin.state_dict_config, fsdp_plugin.optim_state_dict_config
+ )
+ if fsdp_plugin.fsdp_version == 1
+ else nullcontext()
+ )
+
+ with ctx:
+ if fsdp_plugin.fsdp_version == 1:
+ optim_state = FSDP.optim_state_dict(model, optimizer)
+ else:
+ optim_state = optimizer.state_dict()
+
+ if fsdp_plugin.state_dict_type == StateDictType.FULL_STATE_DICT:
+ if accelerator.process_index == 0:
+ optim_state_name = (
+ f"{OPTIMIZER_NAME}.bin" if optimizer_index == 0 else f"{OPTIMIZER_NAME}_{optimizer_index}.bin"
+ )
+ output_optimizer_file = os.path.join(output_dir, optim_state_name)
+ logger.info(f"Saving Optimizer state to {output_optimizer_file}")
+ torch.save(optim_state, output_optimizer_file)
+ logger.info(f"Optimizer state saved in {output_optimizer_file}")
+ else:
+ ckpt_dir = os.path.join(output_dir, f"{OPTIMIZER_NAME}_{optimizer_index}")
+ os.makedirs(ckpt_dir, exist_ok=True)
+ logger.info(f"Saving Optimizer state to {ckpt_dir}")
+ dist_cp.save_state_dict(
+ state_dict={"optimizer": optim_state},
+ storage_writer=dist_cp.FileSystemWriter(ckpt_dir),
+ planner=DefaultSavePlanner(),
+ )
+ logger.info(f"Optimizer state saved in {ckpt_dir}")
+
+
+def load_fsdp_optimizer(fsdp_plugin, accelerator, optimizer, model, input_dir, optimizer_index=0, adapter_only=False):
+ # Note: We import here to reduce import time from general modules, and isolate outside dependencies
+ import torch.distributed.checkpoint as dist_cp
+ from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_dict
+ from torch.distributed.fsdp.fully_sharded_data_parallel import FullyShardedDataParallel as FSDP
+ from torch.distributed.fsdp.fully_sharded_data_parallel import StateDictType
+
+ accelerator.wait_for_everyone()
+ ctx = (
+ FSDP.state_dict_type(
+ model, fsdp_plugin.state_dict_type, fsdp_plugin.state_dict_config, fsdp_plugin.optim_state_dict_config
+ )
+ if fsdp_plugin.fsdp_version == 1
+ else nullcontext()
+ )
+ with ctx:
+ if fsdp_plugin.state_dict_type == StateDictType.FULL_STATE_DICT:
+ optim_state = None
+ if accelerator.process_index == 0 or not fsdp_plugin.optim_state_dict_config.rank0_only:
+ optimizer_name = (
+ f"{OPTIMIZER_NAME}.bin" if optimizer_index == 0 else f"{OPTIMIZER_NAME}_{optimizer_index}.bin"
+ )
+ input_optimizer_file = os.path.join(input_dir, optimizer_name)
+ logger.info(f"Loading Optimizer state from {input_optimizer_file}")
+ optim_state = torch.load(input_optimizer_file)
+ logger.info(f"Optimizer state loaded from {input_optimizer_file}")
+ else:
+ ckpt_dir = (
+ os.path.join(input_dir, f"{OPTIMIZER_NAME}_{optimizer_index}")
+ if f"{OPTIMIZER_NAME}" not in input_dir
+ else input_dir
+ )
+ logger.info(f"Loading Optimizer from {ckpt_dir}")
+ if fsdp_plugin.fsdp_version == 1:
+ optim_state = load_sharded_optimizer_state_dict(
+ model_state_dict=_get_model_state_dict(model, adapter_only=adapter_only),
+ optimizer_key="optimizer",
+ storage_reader=dist_cp.FileSystemReader(ckpt_dir),
+ )
+ else:
+ optim_state = {"optimizer": optimizer.state_dict()}
+ dist_cp.load(
+ optim_state,
+ checkpoint_id=ckpt_dir,
+ storage_reader=dist_cp.FileSystemReader(ckpt_dir),
+ )
+ optim_state = optim_state["optimizer"]
+ logger.info(f"Optimizer loaded from {ckpt_dir}")
+ if fsdp_plugin.fsdp_version == 1:
+ flattened_osd = FSDP.optim_state_dict_to_load(model=model, optim=optimizer, optim_state_dict=optim_state)
+ optimizer.load_state_dict(flattened_osd)
+ else:
+ # we can't do `set_state_dict` here because it does a step of the optimizer which breaks with grad-scaler
+ # TODO(siro1): investigate
+ optimizer.load_state_dict(optim_state)
+
+
+def _distributed_checkpoint_to_merged_weights(checkpoint_dir: str, save_path: str, safe_serialization: bool = True):
+ """
+ Passthrough to `torch.distributed.checkpoint.format_utils.dcp_to_torch_save`
+
+ Will save under `save_path` as either `model.safetensors` or `pytorch_model.bin`.
+ """
+ # Note: We import here to reduce import time from general modules, and isolate outside dependencies
+ import torch.distributed.checkpoint as dist_cp
+ import torch.distributed.checkpoint.format_utils as dist_cp_format_utils
+
+ state_dict = {}
+ save_path = Path(save_path)
+ save_path.mkdir(exist_ok=True)
+ dist_cp_format_utils._load_state_dict(
+ state_dict,
+ storage_reader=dist_cp.FileSystemReader(checkpoint_dir),
+ planner=dist_cp_format_utils._EmptyStateDictLoadPlanner(),
+ no_dist=True,
+ )
+ save_path = save_path / SAFE_WEIGHTS_NAME if safe_serialization else save_path / WEIGHTS_NAME
+
+ # To handle if state is a dict like {model: {...}}
+ if len(state_dict.keys()) == 1:
+ state_dict = state_dict[list(state_dict)[0]]
+ save(state_dict, save_path, safe_serialization=safe_serialization)
+ return save_path
+
+
+def merge_fsdp_weights(
+ checkpoint_dir: str, output_path: str, safe_serialization: bool = True, remove_checkpoint_dir: bool = False
+):
+ """
+ Merge the weights from sharded FSDP model checkpoints into a single combined checkpoint. Should be used if
+ `SHARDED_STATE_DICT` was used for the model. Weights will be saved to `{output_path}/model.safetensors` if
+ `safe_serialization` else `pytorch_model.bin`.
+
+ Note: this is a CPU-bound process.
+
+ Args:
+ checkpoint_dir (`str`):
+ The directory containing the FSDP checkpoints (can be either the model or optimizer).
+ output_path (`str`):
+ The path to save the merged checkpoint.
+ safe_serialization (`bool`, *optional*, defaults to `True`):
+ Whether to save the merged weights with safetensors (recommended).
+ remove_checkpoint_dir (`bool`, *optional*, defaults to `False`):
+ Whether to remove the checkpoint directory after merging.
+ """
+ checkpoint_dir = Path(checkpoint_dir)
+ from accelerate.state import PartialState
+
+ if not is_torch_version(">=", "2.3.0"):
+ raise ValueError("`merge_fsdp_weights` requires PyTorch >= 2.3.0`")
+
+ # Verify that the checkpoint directory exists
+ if not checkpoint_dir.exists():
+ model_path_exists = (checkpoint_dir / "pytorch_model_fsdp_0").exists()
+ optimizer_path_exists = (checkpoint_dir / "optimizer_0").exists()
+ err = f"Tried to load from {checkpoint_dir} but couldn't find a valid metadata file."
+ if model_path_exists and optimizer_path_exists:
+ err += " However, potential model and optimizer checkpoint directories exist."
+ err += f"Please pass in either {checkpoint_dir}/pytorch_model_fsdp_0 or {checkpoint_dir}/optimizer_0"
+ err += "instead."
+ elif model_path_exists:
+ err += " However, a potential model checkpoint directory exists."
+ err += f"Please try passing in {checkpoint_dir}/pytorch_model_fsdp_0 instead."
+ elif optimizer_path_exists:
+ err += " However, a potential optimizer checkpoint directory exists."
+ err += f"Please try passing in {checkpoint_dir}/optimizer_0 instead."
+ raise ValueError(err)
+
+ # To setup `save` to work
+ state = PartialState()
+ if state.is_main_process:
+ logger.info(f"Merging FSDP weights from {checkpoint_dir}")
+ save_path = _distributed_checkpoint_to_merged_weights(checkpoint_dir, output_path, safe_serialization)
+ logger.info(f"Successfully merged FSDP weights and saved to {save_path}")
+ if remove_checkpoint_dir:
+ logger.info(f"Removing old checkpoint directory {checkpoint_dir}")
+ shutil.rmtree(checkpoint_dir)
+ state.wait_for_everyone()
+
+
+def ensure_weights_retied(param_init_fn, model: torch.nn.Module, device: torch.cuda.device):
+ _tied_names = getattr(model, "_tied_weights_keys", None)
+ if not _tied_names:
+ # if no tied names just passthrough
+ return param_init_fn
+
+ # get map of parameter instances to params.
+ # - needed for replacement later
+ _tied_params = {}
+ for name in _tied_names:
+ name = name.split(".")
+ name, param_name = ".".join(name[:-1]), name[-1]
+ mod = model.get_submodule(name)
+ param = getattr(mod, param_name)
+
+ _tied_params[id(param)] = None # placeholder for the param first
+
+ # build param_init_fn for the case with tied params
+ def param_init_fn_tied_param(module: torch.nn.Module):
+ # track which params to tie
+ # - usually only 1, but for completeness consider > 1
+ params_to_tie = defaultdict(list)
+ for n, param in module.named_parameters(recurse=False):
+ if id(param) in _tied_params:
+ params_to_tie[id(param)].append(n)
+
+ # call the param init fn, which potentially re-allocates the
+ # parameters
+ module = param_init_fn(module)
+
+ # search the parameters again and tie them up again
+ for id_key, _param_names in params_to_tie.items():
+ for param_name in _param_names:
+ param = _tied_params[id_key]
+ if param is None:
+ # everything will be tied to the first time the
+ # param is observed
+ _tied_params[id_key] = getattr(module, param_name)
+ else:
+ setattr(module, param_name, param) # tie
+
+ return module
+
+ return param_init_fn_tied_param
+
+
+def fsdp2_load_full_state_dict(accelerator, model: torch.nn.Module, full_sd: dict):
+ """
+ Loads the full state dict (could be only on rank 0) into the sharded model. This is done by broadcasting the
+ parameters from rank 0 to all other ranks. This function modifies the model in-place.
+
+ Args:
+ accelerator (`Accelerator`): The accelerator instance
+ model (`torch.nn.Module`): The model to load the state dict into
+ full_sd (`dict`): The full state dict to load, can only be on rank 0
+ """
+ import torch.distributed as dist
+ from torch.distributed.tensor import distribute_tensor
+
+ sharded_sd = model.state_dict()
+ if accelerator.is_main_process:
+ for (param_name, full_param), sharded_param in zip(full_sd.items(), sharded_sd.values()):
+ full_param = full_param.detach().cuda()
+ mesh = sharded_param.device_mesh
+ dist.broadcast(full_param, src=0, group=mesh.get_group())
+ sharded_tensor = distribute_tensor(full_param, mesh, sharded_param.placements)
+ sharded_sd[param_name] = sharded_tensor
+ else:
+ for param_name, sharded_param in sharded_sd.items():
+ full_tensor = torch.empty(sharded_param.size(), device="cuda", dtype=sharded_param.dtype)
+ mesh = sharded_param.device_mesh
+ dist.broadcast(full_tensor, src=0, group=mesh.get_group())
+ sharded_tensor = distribute_tensor(full_tensor, mesh, sharded_param.placements)
+ sharded_sd[param_name] = sharded_tensor
+
+ model.load_state_dict(sharded_sd)
+
+
+def fsdp2_switch_optimizer_parameters(optimizer: torch.optim.Optimizer, mapping: dict):
+ """
+ Switches the parameters of the optimizer to new ones (sharded parameters in usual case). This function modifies the
+ optimizer in-place.
+
+ Args:
+ optimizer (`torch.optim.Optimizer`): Optimizer instance which contains the original model parameters
+ mapping (`dict`): Mapping from the original parameter (specified by `data_ptr`) to the sharded parameter
+
+ Raises:
+ KeyError:
+ If a parameter in the optimizer couldn't be switched to its sharded version. This should never happen and
+ indicates a bug. If we kept the original params instead of raising, the training wouldn't be numerically
+ correct and weights wouldn't get updated.
+ """
+ try:
+ for param_group in optimizer.param_groups:
+ param_group["params"] = [mapping[p.data_ptr] for p in param_group["params"]]
+ except KeyError:
+ # This shouldn't ever happen, but we want to fail here else training wouldn't be numerically correct
+ # This basically means that we're missing a mapping from the original parameter to the sharded parameter
+ raise KeyError(
+ "A parameter in the optimizer couldn't be switched to its sharded version. This breaks the training. Please raise an issue on GitHub."
+ )
+
+
+def fsdp2_prepare_model(accelerator, model: torch.nn.Module) -> torch.nn.Module:
+ """Prepares the model for FSDP2 in-place. Also returns the model to avoid misuse of the original model.
+
+ Args:
+ accelerator (`Accelerator`): The accelerator instance
+ model (`torch.nn.Module`): The model to prepare
+
+ Returns:
+ `torch.nn.Module`: Prepared model
+ """
+ from torch.distributed.fsdp import FSDPModule, MixedPrecisionPolicy, fully_shard
+
+ is_type_fsdp = isinstance(model, FSDPModule) or (
+ is_compiled_module(model) and isinstance(model._orig_mod, FSDPModule)
+ )
+ if is_type_fsdp:
+ return model
+
+ fsdp2_plugin = accelerator.state.fsdp_plugin
+
+ original_sd = model.state_dict()
+
+ from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy, transformer_auto_wrap_policy
+
+ # We need the `auto_wrap_policy` original type to create a custom poilicy function for sharding
+ # This is because `fully_shard` doesn't support old auto wrap policies, rather we have to imitate the behaviour
+ auto_wrap_policy_type = None
+ if fsdp2_plugin.auto_wrap_policy is transformer_auto_wrap_policy:
+ auto_wrap_policy_type = "transformer"
+ elif fsdp2_plugin.auto_wrap_policy is size_based_auto_wrap_policy:
+ auto_wrap_policy_type = "size"
+
+ # We set `auto_wrap_policy` to `functools.partial` to avoid creating it again
+ # This is because of `apply_activation_checkpointing` which will can reuse this function
+ fsdp2_plugin.set_auto_wrap_policy(model)
+
+ if fsdp2_plugin.activation_checkpointing:
+ from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
+ CheckpointImpl,
+ apply_activation_checkpointing,
+ checkpoint_wrapper,
+ )
+
+ # Apply activation checkpointing before applying `fully_shard`
+ apply_activation_checkpointing(
+ model,
+ checkpoint_wrapper_fn=functools.partial(
+ checkpoint_wrapper,
+ checkpoint_impl=CheckpointImpl.NO_REENTRANT,
+ ),
+ auto_wrap_policy=fsdp2_plugin.auto_wrap_policy,
+ )
+
+ fsdp2_kwargs = {
+ "reshard_after_forward": fsdp2_plugin.reshard_after_forward,
+ "offload_policy": fsdp2_plugin.cpu_offload,
+ # `fully_shard` doesn't accept `None` in case of `MixedPrecisionPolicy`
+ "mp_policy": fsdp2_plugin.mixed_precision_policy or MixedPrecisionPolicy(),
+ }
+
+ auto_wrap_policy = fsdp2_prepare_auto_wrap_policy(fsdp2_plugin, auto_wrap_policy_type, model)
+ if auto_wrap_policy is not None:
+ # We skip the model itself, as that one is always wrapped
+ for module in get_module_children_bottom_up(model)[:-1]:
+ if auto_wrap_policy(module):
+ fully_shard(module, **fsdp2_kwargs)
+
+ fully_shard(model, **fsdp2_kwargs)
+
+ if fsdp2_plugin.cpu_ram_efficient_loading:
+ # If `cpu_ram_efficient_loading` is enabled, only rank 0 loads the weights
+ # Other ranks have an empty model on `meta` device, so we need to distribute the weights properly
+ fsdp2_load_full_state_dict(accelerator, model, original_sd)
+
+ if accelerator.mixed_precision != "no" and model.dtype != torch.float32:
+ # We upcast the model according to `deepspeed`'s implementation
+ # More info about this can be found in `accelerator.py:prepare_model`s FSDP1 section
+ model = model.to(torch.float32)
+ if accelerator.is_main_process:
+ # TODO(siro1): Add a warning for each parameter that was upcasted
+ warnings.warn(
+ "FSDP upcast of low precision parameters to fp32 (since mixed_precision != 'no') may affect the precision of model checkpoints."
+ )
+ return model
+
+
+def fsdp2_prepare_auto_wrap_policy(
+ fsdp2_plugin, auto_wrap_policy_type: str, model: torch.nn.Module
+) -> Callable[[torch.nn.Module], bool]:
+ """Prepares the auto wrap policy based on its type, done to mimic the behaviour of FSDP1 auto wrap policy.
+
+ Args:
+ fsdp2_plugin (`FullyShardedDataParallelPlugin`):
+ Instance of `FullyShardedDataParallelPlugin` containing the configuration options
+ auto_wrap_policy_type (`str`):
+ Either `transformer` or `size`
+ model (`torch.nn.Module`):
+ The model to wrap
+
+ Returns:
+ `Callable[[torch.nn.Module], bool]`:
+ The auto wrap policy function to be applied to the model
+ """
+ if auto_wrap_policy_type == "transformer":
+ no_split_modules = model._no_split_modules
+ if no_split_modules is None:
+ no_split_modules = []
+ transformer_cls_names_to_wrap = list(no_split_modules)
+ if fsdp2_plugin.transformer_cls_names_to_wrap is not None:
+ transformer_cls_names_to_wrap = fsdp2_plugin.transformer_cls_names_to_wrap
+ transformer_cls_to_wrap = set()
+
+ for layer_class in transformer_cls_names_to_wrap:
+ transformer_cls = get_module_class_from_name(model, layer_class)
+ if transformer_cls is None:
+ raise ValueError(f"Could not find the transformer layer class {layer_class} in the model.")
+ transformer_cls_to_wrap.add(transformer_cls)
+
+ def policy(module: torch.nn.Module) -> bool:
+ if fsdp2_plugin.transformer_cls_names_to_wrap is None:
+ return False
+ return isinstance(module, tuple(transformer_cls_to_wrap))
+
+ elif auto_wrap_policy_type == "size":
+
+ def policy(module: torch.nn.Module) -> bool:
+ return module.numel() > fsdp2_plugin.min_num_params
+ else:
+ return None
+
+ return policy
+
+
+def get_fsdp2_grad_scaler(**kwargs):
+ """
+ Returns a `GradScaler` for FSDP2, as the current implementation of `get_grad_scaler` doesn't accept other args. We
+ need this as current `get_grad_scaler` accepts only `distributed_type` as arg, which doesn't differentiate between
+ FSDP1 and FSDP2
+ """
+ from torch.amp.grad_scaler import GradScaler
+
+ return GradScaler(**kwargs)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/imports.py b/venv/lib/python3.11/site-packages/accelerate/utils/imports.py
new file mode 100644
index 0000000000000000000000000000000000000000..26fda1bc732cd74b06f02a3de48b9020de192c88
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/imports.py
@@ -0,0 +1,546 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 importlib
+import importlib.metadata
+import os
+import warnings
+from functools import lru_cache, wraps
+
+import torch
+from packaging import version
+from packaging.version import parse
+
+from .environment import parse_flag_from_env, patch_environment, str_to_bool
+from .versions import compare_versions, is_torch_version
+
+
+# Try to run Torch native job in an environment with TorchXLA installed by setting this value to 0.
+USE_TORCH_XLA = parse_flag_from_env("USE_TORCH_XLA", default=True)
+
+_torch_xla_available = False
+if USE_TORCH_XLA:
+ try:
+ import torch_xla.core.xla_model as xm # noqa: F401
+ import torch_xla.runtime
+
+ _torch_xla_available = True
+ except ImportError:
+ pass
+
+# Keep it for is_tpu_available. It will be removed along with is_tpu_available.
+_tpu_available = _torch_xla_available
+
+# Cache this result has it's a C FFI call which can be pretty time-consuming
+_torch_distributed_available = torch.distributed.is_available()
+
+
+def _is_package_available(pkg_name, metadata_name=None):
+ # Check we're not importing a "pkg_name" directory somewhere but the actual library by trying to grab the version
+ package_exists = importlib.util.find_spec(pkg_name) is not None
+ if package_exists:
+ try:
+ # Some libraries have different names in the metadata
+ _ = importlib.metadata.metadata(pkg_name if metadata_name is None else metadata_name)
+ return True
+ except importlib.metadata.PackageNotFoundError:
+ return False
+
+
+def is_torch_distributed_available() -> bool:
+ return _torch_distributed_available
+
+
+def is_xccl_available():
+ # Currently IPEX uses custom "ccl" distributed backend. Return False for "xccl"
+ # here to avoid collisions.
+ if is_ipex_available():
+ return False
+ # TODO: switch to is_torch_version() once torch 2.7 will be released
+ if version.parse(torch.__version__).release >= version.parse("2.7").release:
+ return torch.distributed.distributed_c10d.is_xccl_available()
+ return False
+
+
+def is_ccl_available():
+ try:
+ pass
+ except ImportError:
+ print(
+ "Intel(R) oneCCL Bindings for PyTorch* is required to run DDP on Intel(R) GPUs, but it is not"
+ " detected. If you see \"ValueError: Invalid backend: 'ccl'\" error, please install Intel(R) oneCCL"
+ " Bindings for PyTorch*."
+ )
+ return (
+ importlib.util.find_spec("torch_ccl") is not None
+ or importlib.util.find_spec("oneccl_bindings_for_pytorch") is not None
+ )
+
+
+def get_ccl_version():
+ return importlib.metadata.version("oneccl_bind_pt")
+
+
+def is_import_timer_available():
+ return _is_package_available("import_timer")
+
+
+def is_pynvml_available():
+ return _is_package_available("pynvml") or _is_package_available("pynvml", "nvidia-ml-py")
+
+
+def is_pytest_available():
+ return _is_package_available("pytest")
+
+
+def is_msamp_available():
+ return _is_package_available("msamp", "ms-amp")
+
+
+def is_schedulefree_available():
+ return _is_package_available("schedulefree")
+
+
+def is_transformer_engine_available():
+ if is_hpu_available():
+ return _is_package_available("intel_transformer_engine", "intel-transformer-engine")
+ else:
+ return _is_package_available("transformer_engine", "transformer-engine")
+
+
+def is_lomo_available():
+ return _is_package_available("lomo_optim")
+
+
+def is_cuda_available():
+ """
+ Checks if `cuda` is available via an `nvml-based` check which won't trigger the drivers and leave cuda
+ uninitialized.
+ """
+ with patch_environment(PYTORCH_NVML_BASED_CUDA_CHECK="1"):
+ available = torch.cuda.is_available()
+
+ return available
+
+
+@lru_cache
+def is_torch_xla_available(check_is_tpu=False, check_is_gpu=False):
+ """
+ Check if `torch_xla` is available. To train a native pytorch job in an environment with torch xla installed, set
+ the USE_TORCH_XLA to false.
+ """
+ assert not (check_is_tpu and check_is_gpu), "The check_is_tpu and check_is_gpu cannot both be true."
+
+ if not _torch_xla_available:
+ return False
+ elif check_is_gpu:
+ return torch_xla.runtime.device_type() in ["GPU", "CUDA"]
+ elif check_is_tpu:
+ return torch_xla.runtime.device_type() == "TPU"
+
+ return True
+
+
+def is_torchao_available():
+ package_exists = _is_package_available("torchao")
+ if package_exists:
+ torchao_version = version.parse(importlib.metadata.version("torchao"))
+ return compare_versions(torchao_version, ">=", "0.6.1")
+ return False
+
+
+def is_deepspeed_available():
+ return _is_package_available("deepspeed")
+
+
+def is_pippy_available():
+ return is_torch_version(">=", "2.4.0")
+
+
+def is_bf16_available(ignore_tpu=False):
+ "Checks if bf16 is supported, optionally ignoring the TPU"
+ if is_torch_xla_available(check_is_tpu=True):
+ return not ignore_tpu
+ if is_cuda_available():
+ return torch.cuda.is_bf16_supported()
+ if is_mlu_available():
+ return torch.mlu.is_bf16_supported()
+ if is_mps_available():
+ return False
+ return True
+
+
+def is_fp16_available():
+ "Checks if fp16 is supported"
+ if is_habana_gaudi1():
+ return False
+
+ return True
+
+
+def is_fp8_available():
+ "Checks if fp8 is supported"
+ return is_msamp_available() or is_transformer_engine_available() or is_torchao_available()
+
+
+def is_4bit_bnb_available():
+ package_exists = _is_package_available("bitsandbytes")
+ if package_exists:
+ bnb_version = version.parse(importlib.metadata.version("bitsandbytes"))
+ return compare_versions(bnb_version, ">=", "0.39.0")
+ return False
+
+
+def is_8bit_bnb_available():
+ package_exists = _is_package_available("bitsandbytes")
+ if package_exists:
+ bnb_version = version.parse(importlib.metadata.version("bitsandbytes"))
+ return compare_versions(bnb_version, ">=", "0.37.2")
+ return False
+
+
+def is_bnb_available(min_version=None):
+ package_exists = _is_package_available("bitsandbytes")
+ if package_exists and min_version is not None:
+ bnb_version = version.parse(importlib.metadata.version("bitsandbytes"))
+ return compare_versions(bnb_version, ">=", min_version)
+ else:
+ return package_exists
+
+
+def is_bitsandbytes_multi_backend_available():
+ if not is_bnb_available():
+ return False
+ import bitsandbytes as bnb
+
+ return "multi_backend" in getattr(bnb, "features", set())
+
+
+def is_torchvision_available():
+ return _is_package_available("torchvision")
+
+
+def is_megatron_lm_available():
+ if str_to_bool(os.environ.get("ACCELERATE_USE_MEGATRON_LM", "False")) == 1:
+ if importlib.util.find_spec("megatron") is not None:
+ try:
+ megatron_version = parse(importlib.metadata.version("megatron-core"))
+ if compare_versions(megatron_version, ">=", "0.8.0"):
+ return importlib.util.find_spec(".training", "megatron")
+ except Exception as e:
+ warnings.warn(f"Parse Megatron version failed. Exception:{e}")
+ return False
+
+
+def is_transformers_available():
+ return _is_package_available("transformers")
+
+
+def is_datasets_available():
+ return _is_package_available("datasets")
+
+
+def is_peft_available():
+ return _is_package_available("peft")
+
+
+def is_timm_available():
+ return _is_package_available("timm")
+
+
+def is_triton_available():
+ if is_xpu_available():
+ return _is_package_available("triton", "pytorch-triton-xpu")
+ return _is_package_available("triton")
+
+
+def is_aim_available():
+ package_exists = _is_package_available("aim")
+ if package_exists:
+ aim_version = version.parse(importlib.metadata.version("aim"))
+ return compare_versions(aim_version, "<", "4.0.0")
+ return False
+
+
+def is_tensorboard_available():
+ return _is_package_available("tensorboard") or _is_package_available("tensorboardX")
+
+
+def is_wandb_available():
+ return _is_package_available("wandb")
+
+
+def is_comet_ml_available():
+ return _is_package_available("comet_ml")
+
+
+def is_boto3_available():
+ return _is_package_available("boto3")
+
+
+def is_rich_available():
+ if _is_package_available("rich"):
+ return parse_flag_from_env("ACCELERATE_ENABLE_RICH", False)
+ return False
+
+
+def is_sagemaker_available():
+ return _is_package_available("sagemaker")
+
+
+def is_tqdm_available():
+ return _is_package_available("tqdm")
+
+
+def is_clearml_available():
+ return _is_package_available("clearml")
+
+
+def is_pandas_available():
+ return _is_package_available("pandas")
+
+
+def is_matplotlib_available():
+ return _is_package_available("matplotlib")
+
+
+def is_mlflow_available():
+ if _is_package_available("mlflow"):
+ return True
+
+ if importlib.util.find_spec("mlflow") is not None:
+ try:
+ _ = importlib.metadata.metadata("mlflow-skinny")
+ return True
+ except importlib.metadata.PackageNotFoundError:
+ return False
+ return False
+
+
+def is_mps_available(min_version="1.12"):
+ "Checks if MPS device is available. The minimum version required is 1.12."
+ # With torch 1.12, you can use torch.backends.mps
+ # With torch 2.0.0, you can use torch.mps
+ return is_torch_version(">=", min_version) and torch.backends.mps.is_available() and torch.backends.mps.is_built()
+
+
+def is_ipex_available():
+ "Checks if ipex is installed."
+
+ def get_major_and_minor_from_version(full_version):
+ return str(version.parse(full_version).major) + "." + str(version.parse(full_version).minor)
+
+ _torch_version = importlib.metadata.version("torch")
+ if importlib.util.find_spec("intel_extension_for_pytorch") is None:
+ return False
+ _ipex_version = "N/A"
+ try:
+ _ipex_version = importlib.metadata.version("intel_extension_for_pytorch")
+ except importlib.metadata.PackageNotFoundError:
+ return False
+ torch_major_and_minor = get_major_and_minor_from_version(_torch_version)
+ ipex_major_and_minor = get_major_and_minor_from_version(_ipex_version)
+ if torch_major_and_minor != ipex_major_and_minor:
+ warnings.warn(
+ f"Intel Extension for PyTorch {ipex_major_and_minor} needs to work with PyTorch {ipex_major_and_minor}.*,"
+ f" but PyTorch {_torch_version} is found. Please switch to the matching version and run again."
+ )
+ return False
+ return True
+
+
+@lru_cache
+def is_mlu_available(check_device=False):
+ """
+ Checks if `mlu` is available via an `cndev-based` check which won't trigger the drivers and leave mlu
+ uninitialized.
+ """
+ if importlib.util.find_spec("torch_mlu") is None:
+ return False
+
+ import torch_mlu # noqa: F401
+
+ with patch_environment(PYTORCH_CNDEV_BASED_MLU_CHECK="1"):
+ available = torch.mlu.is_available()
+
+ return available
+
+
+@lru_cache
+def is_musa_available(check_device=False):
+ "Checks if `torch_musa` is installed and potentially if a MUSA is in the environment"
+ if importlib.util.find_spec("torch_musa") is None:
+ return False
+
+ import torch_musa # noqa: F401
+
+ if check_device:
+ try:
+ # Will raise a RuntimeError if no MUSA is found
+ _ = torch.musa.device_count()
+ return torch.musa.is_available()
+ except RuntimeError:
+ return False
+ return hasattr(torch, "musa") and torch.musa.is_available()
+
+
+@lru_cache
+def is_npu_available(check_device=False):
+ "Checks if `torch_npu` is installed and potentially if a NPU is in the environment"
+ if importlib.util.find_spec("torch_npu") is None:
+ return False
+
+ import torch_npu # noqa: F401
+
+ if check_device:
+ try:
+ # Will raise a RuntimeError if no NPU is found
+ _ = torch.npu.device_count()
+ return torch.npu.is_available()
+ except RuntimeError:
+ return False
+ return hasattr(torch, "npu") and torch.npu.is_available()
+
+
+@lru_cache
+def is_sdaa_available(check_device=False):
+ "Checks if `torch_sdaa` is installed and potentially if a SDAA is in the environment"
+ if importlib.util.find_spec("torch_sdaa") is None:
+ return False
+
+ import torch_sdaa # noqa: F401
+
+ if check_device:
+ try:
+ # Will raise a RuntimeError if no NPU is found
+ _ = torch.sdaa.device_count()
+ return torch.sdaa.is_available()
+ except RuntimeError:
+ return False
+ return hasattr(torch, "sdaa") and torch.sdaa.is_available()
+
+
+@lru_cache
+def is_hpu_available(init_hccl=False):
+ "Checks if `torch.hpu` is installed and potentially if a HPU is in the environment"
+ if (
+ importlib.util.find_spec("habana_frameworks") is None
+ or importlib.util.find_spec("habana_frameworks.torch") is None
+ ):
+ return False
+
+ import habana_frameworks.torch # noqa: F401
+
+ if init_hccl:
+ import habana_frameworks.torch.distributed.hccl as hccl # noqa: F401
+
+ return hasattr(torch, "hpu") and torch.hpu.is_available()
+
+
+def is_habana_gaudi1():
+ if is_hpu_available():
+ import habana_frameworks.torch.utils.experimental as htexp # noqa: F401
+
+ if htexp._get_device_type() == htexp.synDeviceType.synDeviceGaudi:
+ return True
+
+ return False
+
+
+@lru_cache
+def is_xpu_available(check_device=False):
+ """
+ Checks if XPU acceleration is available either via `intel_extension_for_pytorch` or via stock PyTorch (>=2.4) and
+ potentially if a XPU is in the environment
+ """
+
+ if is_ipex_available():
+ import intel_extension_for_pytorch # noqa: F401
+ else:
+ if is_torch_version("<=", "2.3"):
+ return False
+
+ if check_device:
+ try:
+ # Will raise a RuntimeError if no XPU is found
+ _ = torch.xpu.device_count()
+ return torch.xpu.is_available()
+ except RuntimeError:
+ return False
+ return hasattr(torch, "xpu") and torch.xpu.is_available()
+
+
+def is_dvclive_available():
+ return _is_package_available("dvclive")
+
+
+def is_torchdata_available():
+ return _is_package_available("torchdata")
+
+
+# TODO: Remove this function once stateful_dataloader is a stable feature in torchdata.
+def is_torchdata_stateful_dataloader_available():
+ package_exists = _is_package_available("torchdata")
+ if package_exists:
+ torchdata_version = version.parse(importlib.metadata.version("torchdata"))
+ return compare_versions(torchdata_version, ">=", "0.8.0")
+ return False
+
+
+def torchao_required(func):
+ """
+ A decorator that ensures the decorated function is only called when torchao is available.
+ """
+
+ @wraps(func)
+ def wrapper(*args, **kwargs):
+ if not is_torchao_available():
+ raise ImportError(
+ "`torchao` is not available, please install it before calling this function via `pip install torchao`."
+ )
+ return func(*args, **kwargs)
+
+ return wrapper
+
+
+# TODO: Rework this into `utils.deepspeed` and migrate the "core" chunks into `accelerate.deepspeed`
+def deepspeed_required(func):
+ """
+ A decorator that ensures the decorated function is only called when deepspeed is enabled.
+ """
+
+ @wraps(func)
+ def wrapper(*args, **kwargs):
+ from accelerate.state import AcceleratorState
+ from accelerate.utils.dataclasses import DistributedType
+
+ if AcceleratorState._shared_state != {} and AcceleratorState().distributed_type != DistributedType.DEEPSPEED:
+ raise ValueError(
+ "DeepSpeed is not enabled, please make sure that an `Accelerator` is configured for `deepspeed` "
+ "before calling this function."
+ )
+ return func(*args, **kwargs)
+
+ return wrapper
+
+
+def is_weights_only_available():
+ # Weights only with allowlist was added in 2.4.0
+ # ref: https://github.com/pytorch/pytorch/pull/124331
+ return is_torch_version(">=", "2.4.0")
+
+
+def is_numpy_available(min_version="1.25.0"):
+ numpy_version = parse(importlib.metadata.version("numpy"))
+ return compare_versions(numpy_version, ">=", min_version)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/launch.py b/venv/lib/python3.11/site-packages/accelerate/utils/launch.py
new file mode 100644
index 0000000000000000000000000000000000000000..ba5bf182ffd462094380abe3a3d1a268631eca6c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/launch.py
@@ -0,0 +1,709 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 argparse
+import os
+import subprocess
+import sys
+from ast import literal_eval
+from shutil import which
+from typing import Any
+
+import torch
+
+from ..commands.config.config_args import SageMakerConfig
+from ..utils import (
+ DynamoBackend,
+ PrecisionType,
+ is_fp8_available,
+ is_hpu_available,
+ is_ipex_available,
+ is_mlu_available,
+ is_musa_available,
+ is_npu_available,
+ is_sdaa_available,
+ is_torch_xla_available,
+ is_xpu_available,
+)
+from ..utils.constants import DEEPSPEED_MULTINODE_LAUNCHERS
+from ..utils.other import is_port_in_use, merge_dicts
+from ..utils.versions import compare_versions
+from .dataclasses import DistributedType, SageMakerDistributedType
+
+
+def _filter_args(args, parser, default_args=[]):
+ """
+ Filters out all `accelerate` specific args
+ """
+ new_args, _ = parser.parse_known_args(default_args)
+ for key, value in vars(args).items():
+ if key in vars(new_args).keys():
+ setattr(new_args, key, value)
+ return new_args
+
+
+def _get_mpirun_args():
+ """
+ Determines the executable and argument names for mpirun, based on the type of install. The supported MPI programs
+ are: OpenMPI, Intel MPI, or MVAPICH.
+
+ Returns: Program name and arg names for hostfile, num processes, and processes per node
+ """
+ # Find the MPI program name
+ mpi_apps = [x for x in ["mpirun", "mpiexec"] if which(x)]
+
+ if len(mpi_apps) == 0:
+ raise OSError("mpirun or mpiexec were not found. Ensure that Intel MPI, Open MPI, or MVAPICH are installed.")
+
+ # Call the app with the --version flag to determine which MPI app is installed
+ mpi_app = mpi_apps[0]
+ mpirun_version = subprocess.check_output([mpi_app, "--version"])
+
+ if b"Open MPI" in mpirun_version:
+ return mpi_app, "--hostfile", "-n", "--npernode", "--bind-to"
+ else:
+ # Intel MPI and MVAPICH both use the same arg names
+ return mpi_app, "-f", "-n", "-ppn", ""
+
+
+def setup_fp8_env(args: argparse.Namespace, current_env: dict[str, str]):
+ """
+ Setup the FP8 environment variables.
+ """
+ prefix = "ACCELERATE_"
+ for arg in vars(args):
+ if arg.startswith("fp8_"):
+ value = getattr(args, arg)
+ if value is not None:
+ if arg == "fp8_override_linear_precision":
+ current_env[prefix + "FP8_OVERRIDE_FPROP"] = value[0]
+ current_env[prefix + "FP8_OVERRIDE_DGRAD"] = value[1]
+ current_env[prefix + "FP8_OVERRIDE_WGRAD"] = value[2]
+ else:
+ current_env[f"{prefix}{arg.upper()}"] = str(getattr(args, arg))
+ return current_env
+
+
+def prepare_simple_launcher_cmd_env(args: argparse.Namespace) -> tuple[list[str], dict[str, str]]:
+ """
+ Prepares and returns the command list and an environment with the correct simple launcher environment variables.
+ """
+ cmd = []
+ if args.no_python and args.module:
+ raise ValueError("--module and --no_python cannot be used together")
+
+ if args.mpirun_hostfile is not None:
+ mpi_app_name, hostfile_arg, num_proc_arg, proc_per_node_arg, bind_to_arg = _get_mpirun_args()
+ mpirun_ccl = getattr(args, "mpirun_ccl", None)
+ bind_to = getattr(args, "bind-to", "socket")
+ num_machines = args.num_machines
+ num_processes = getattr(args, "num_processes", None)
+ nproc_per_node = str(num_processes // num_machines) if num_processes and num_machines else "1"
+ cmd += [
+ mpi_app_name,
+ hostfile_arg,
+ args.mpirun_hostfile,
+ proc_per_node_arg,
+ nproc_per_node,
+ ]
+ if num_processes:
+ cmd += [num_proc_arg, str(num_processes)]
+ if bind_to_arg:
+ cmd += [bind_to_arg, bind_to]
+ if not args.no_python:
+ cmd.append(sys.executable)
+ if args.module:
+ cmd.append("-m")
+ cmd.append(args.training_script)
+ cmd.extend(args.training_script_args)
+
+ current_env = os.environ.copy()
+ current_env["ACCELERATE_USE_CPU"] = str(args.cpu or args.use_cpu)
+ if args.debug:
+ current_env["ACCELERATE_DEBUG_MODE"] = "true"
+ if args.gpu_ids != "all" and args.gpu_ids is not None:
+ if is_xpu_available():
+ current_env["ZE_AFFINITY_MASK"] = args.gpu_ids
+ elif is_mlu_available():
+ current_env["MLU_VISIBLE_DEVICES"] = args.gpu_ids
+ elif is_sdaa_available():
+ current_env["SDAA_VISIBLE_DEVICES"] = args.gpu_ids
+ elif is_musa_available():
+ current_env["MUSA_VISIBLE_DEVICES"] = args.gpu_ids
+ elif is_npu_available():
+ current_env["ASCEND_RT_VISIBLE_DEVICES"] = args.gpu_ids
+ elif is_hpu_available():
+ current_env["HABANA_VISIBLE_MODULES"] = args.gpu_ids
+ else:
+ current_env["CUDA_VISIBLE_DEVICES"] = args.gpu_ids
+ if args.num_machines > 1:
+ current_env["MASTER_ADDR"] = args.main_process_ip
+ current_env["MASTER_PORT"] = str(args.main_process_port)
+
+ if args.mpirun_hostfile is not None:
+ current_env["CCL_WORKER_COUNT"] = str(mpirun_ccl)
+ elif args.num_processes > 1:
+ current_env["MASTER_ADDR"] = args.main_process_ip if args.main_process_ip is not None else "127.0.0.1"
+ current_env["MASTER_PORT"] = str(args.main_process_port) if args.main_process_port is not None else "29500"
+
+ try:
+ mixed_precision = PrecisionType(args.mixed_precision.lower())
+ except ValueError:
+ raise ValueError(
+ f"Unknown mixed_precision mode: {args.mixed_precision.lower()}. Choose between {PrecisionType.list()}."
+ )
+
+ current_env["ACCELERATE_MIXED_PRECISION"] = str(mixed_precision)
+ if args.mixed_precision.lower() == "fp8":
+ if not is_fp8_available():
+ raise RuntimeError(
+ "FP8 is not available on this machine. Please ensure that either Transformer Engine or MSAMP is installed."
+ )
+ current_env = setup_fp8_env(args, current_env)
+
+ try:
+ dynamo_backend = DynamoBackend(args.dynamo_backend.upper())
+ except ValueError:
+ raise ValueError(
+ f"Unknown dynamo backend: {args.dynamo_backend.upper()}. Choose between {DynamoBackend.list()}."
+ )
+ current_env["ACCELERATE_DYNAMO_BACKEND"] = dynamo_backend.value
+ current_env["ACCELERATE_DYNAMO_MODE"] = args.dynamo_mode
+ current_env["ACCELERATE_DYNAMO_USE_FULLGRAPH"] = str(args.dynamo_use_fullgraph)
+ current_env["ACCELERATE_DYNAMO_USE_DYNAMIC"] = str(args.dynamo_use_dynamic)
+
+ current_env["OMP_NUM_THREADS"] = str(args.num_cpu_threads_per_process)
+ if is_ipex_available():
+ current_env["ACCELERATE_USE_IPEX"] = str(args.ipex).lower()
+ if args.enable_cpu_affinity:
+ current_env["ACCELERATE_CPU_AFFINITY"] = "1"
+ return cmd, current_env
+
+
+def prepare_multi_gpu_env(args: argparse.Namespace) -> dict[str, str]:
+ """
+ Prepares and returns an environment with the correct multi-GPU environment variables.
+ """
+ num_processes = args.num_processes
+ num_machines = args.num_machines
+ main_process_ip = args.main_process_ip
+ main_process_port = args.main_process_port
+ if num_machines > 1:
+ args.nproc_per_node = str(num_processes // num_machines)
+ args.nnodes = str(num_machines)
+ args.node_rank = int(args.machine_rank)
+ if getattr(args, "same_network", False):
+ args.master_addr = str(main_process_ip)
+ args.master_port = str(main_process_port)
+ else:
+ args.rdzv_endpoint = f"{main_process_ip}:{main_process_port}"
+ else:
+ args.nproc_per_node = str(num_processes)
+ if main_process_port is not None:
+ args.master_port = str(main_process_port)
+
+ if main_process_port is None:
+ main_process_port = 29500
+
+ # only need to check port availability in main process, in case we have to start multiple launchers on the same machine
+ # for some reasons like splitting log files.
+ need_port_check = num_machines <= 1 or int(args.machine_rank) == 0
+ if need_port_check and is_port_in_use(main_process_port):
+ raise ConnectionError(
+ f"Tried to launch distributed communication on port `{main_process_port}`, but another process is utilizing it. "
+ "Please specify a different port (such as using the `--main_process_port` flag or specifying a different `main_process_port` in your config file)"
+ " and rerun your script. To automatically use the next open port (on a single node), you can set this to `0`."
+ )
+
+ if args.module and args.no_python:
+ raise ValueError("--module and --no_python cannot be used together")
+ elif args.module:
+ args.module = True
+ elif args.no_python:
+ args.no_python = True
+
+ current_env = os.environ.copy()
+ if args.debug:
+ current_env["ACCELERATE_DEBUG_MODE"] = "true"
+ gpu_ids = getattr(args, "gpu_ids", "all")
+ if gpu_ids != "all" and args.gpu_ids is not None:
+ if is_xpu_available():
+ current_env["ZE_AFFINITY_MASK"] = gpu_ids
+ elif is_mlu_available():
+ current_env["MLU_VISIBLE_DEVICES"] = gpu_ids
+ elif is_sdaa_available():
+ current_env["SDAA_VISIBLE_DEVICES"] = gpu_ids
+ elif is_musa_available():
+ current_env["MUSA_VISIBLE_DEVICES"] = gpu_ids
+ elif is_npu_available():
+ current_env["ASCEND_RT_VISIBLE_DEVICES"] = gpu_ids
+ elif is_hpu_available():
+ current_env["HABANA_VISIBLE_MODULES"] = gpu_ids
+ else:
+ current_env["CUDA_VISIBLE_DEVICES"] = gpu_ids
+ mixed_precision = args.mixed_precision.lower()
+ try:
+ mixed_precision = PrecisionType(mixed_precision)
+ except ValueError:
+ raise ValueError(f"Unknown mixed_precision mode: {mixed_precision}. Choose between {PrecisionType.list()}.")
+
+ current_env["ACCELERATE_MIXED_PRECISION"] = str(mixed_precision)
+ if args.mixed_precision.lower() == "fp8":
+ if not is_fp8_available():
+ raise RuntimeError(
+ "FP8 is not available on this machine. Please ensure that either Transformer Engine or MSAMP is installed."
+ )
+ current_env = setup_fp8_env(args, current_env)
+
+ try:
+ dynamo_backend = DynamoBackend(args.dynamo_backend.upper())
+ except ValueError:
+ raise ValueError(
+ f"Unknown dynamo backend: {args.dynamo_backend.upper()}. Choose between {DynamoBackend.list()}."
+ )
+ current_env["ACCELERATE_DYNAMO_BACKEND"] = dynamo_backend.value
+ current_env["ACCELERATE_DYNAMO_MODE"] = args.dynamo_mode
+ current_env["ACCELERATE_DYNAMO_USE_FULLGRAPH"] = str(args.dynamo_use_fullgraph)
+ current_env["ACCELERATE_DYNAMO_USE_DYNAMIC"] = str(args.dynamo_use_dynamic)
+
+ if args.use_fsdp:
+ current_env["ACCELERATE_USE_FSDP"] = "true"
+ if args.fsdp_cpu_ram_efficient_loading and not args.fsdp_sync_module_states:
+ raise ValueError("When using `--fsdp_cpu_ram_efficient_loading` set `--fsdp_sync_module_states` to `True`")
+
+ current_env["FSDP_VERSION"] = str(args.fsdp_version) if hasattr(args, "fsdp_version") else "1"
+
+ # For backwards compatibility, we support this in launched scripts,
+ # however, we do not ask users for this in `accelerate config` CLI
+ current_env["FSDP_SHARDING_STRATEGY"] = str(args.fsdp_sharding_strategy)
+
+ current_env["FSDP_RESHARD_AFTER_FORWARD"] = str(args.fsdp_reshard_after_forward).lower()
+ current_env["FSDP_OFFLOAD_PARAMS"] = str(args.fsdp_offload_params).lower()
+ current_env["FSDP_MIN_NUM_PARAMS"] = str(args.fsdp_min_num_params)
+ if args.fsdp_auto_wrap_policy is not None:
+ current_env["FSDP_AUTO_WRAP_POLICY"] = str(args.fsdp_auto_wrap_policy)
+ if args.fsdp_transformer_layer_cls_to_wrap is not None:
+ current_env["FSDP_TRANSFORMER_CLS_TO_WRAP"] = str(args.fsdp_transformer_layer_cls_to_wrap)
+ if args.fsdp_backward_prefetch is not None:
+ current_env["FSDP_BACKWARD_PREFETCH"] = str(args.fsdp_backward_prefetch)
+ if args.fsdp_state_dict_type is not None:
+ current_env["FSDP_STATE_DICT_TYPE"] = str(args.fsdp_state_dict_type)
+ current_env["FSDP_FORWARD_PREFETCH"] = str(args.fsdp_forward_prefetch).lower()
+ current_env["FSDP_USE_ORIG_PARAMS"] = str(args.fsdp_use_orig_params).lower()
+ current_env["FSDP_CPU_RAM_EFFICIENT_LOADING"] = str(args.fsdp_cpu_ram_efficient_loading).lower()
+ current_env["FSDP_SYNC_MODULE_STATES"] = str(args.fsdp_sync_module_states).lower()
+ current_env["FSDP_ACTIVATION_CHECKPOINTING"] = str(args.fsdp_activation_checkpointing).lower()
+
+ if args.use_tp:
+ current_env["ACCELERATE_USE_TP"] = "true"
+ current_env["TP_SIZE"] = str(args.tp_size)
+
+ if args.use_megatron_lm:
+ prefix = "MEGATRON_LM_"
+ current_env["ACCELERATE_USE_MEGATRON_LM"] = "true"
+ current_env[prefix + "TP_DEGREE"] = str(args.megatron_lm_tp_degree)
+ current_env[prefix + "PP_DEGREE"] = str(args.megatron_lm_pp_degree)
+ current_env[prefix + "GRADIENT_CLIPPING"] = str(args.megatron_lm_gradient_clipping)
+ if args.megatron_lm_num_micro_batches is not None:
+ current_env[prefix + "NUM_MICRO_BATCHES"] = str(args.megatron_lm_num_micro_batches)
+ if args.megatron_lm_sequence_parallelism is not None:
+ current_env[prefix + "SEQUENCE_PARALLELISM"] = str(args.megatron_lm_sequence_parallelism)
+ if args.megatron_lm_recompute_activations is not None:
+ current_env[prefix + "RECOMPUTE_ACTIVATIONS"] = str(args.megatron_lm_recompute_activations)
+ if args.megatron_lm_use_distributed_optimizer is not None:
+ current_env[prefix + "USE_DISTRIBUTED_OPTIMIZER"] = str(args.megatron_lm_use_distributed_optimizer)
+
+ current_env["OMP_NUM_THREADS"] = str(args.num_cpu_threads_per_process)
+ if args.enable_cpu_affinity:
+ current_env["ACCELERATE_CPU_AFFINITY"] = "1"
+ return current_env
+
+
+def prepare_deepspeed_cmd_env(args: argparse.Namespace) -> tuple[list[str], dict[str, str]]:
+ """
+ Prepares and returns the command list and an environment with the correct DeepSpeed environment variables.
+ """
+ num_processes = args.num_processes
+ num_machines = args.num_machines
+ main_process_ip = args.main_process_ip
+ main_process_port = args.main_process_port
+ cmd = None
+
+ # make sure launcher is not None
+ if args.deepspeed_multinode_launcher is None:
+ # set to default pdsh
+ args.deepspeed_multinode_launcher = DEEPSPEED_MULTINODE_LAUNCHERS[0]
+
+ if num_machines > 1 and args.deepspeed_multinode_launcher != DEEPSPEED_MULTINODE_LAUNCHERS[1]:
+ cmd = ["deepspeed"]
+ cmd.extend(["--hostfile", str(args.deepspeed_hostfile)])
+ if args.deepspeed_multinode_launcher == "nossh":
+ if compare_versions("deepspeed", "<", "0.14.5"):
+ raise ValueError("nossh launcher requires DeepSpeed >= 0.14.5")
+ cmd.extend(["--node_rank", str(args.machine_rank), "--no_ssh"])
+ else:
+ cmd.extend(["--no_local_rank", "--launcher", str(args.deepspeed_multinode_launcher)])
+ if args.deepspeed_exclusion_filter is not None:
+ cmd.extend(
+ [
+ "--exclude",
+ str(args.deepspeed_exclusion_filter),
+ ]
+ )
+ elif args.deepspeed_inclusion_filter is not None:
+ cmd.extend(
+ [
+ "--include",
+ str(args.deepspeed_inclusion_filter),
+ ]
+ )
+ else:
+ cmd.extend(["--num_gpus", str(args.num_processes // args.num_machines)])
+ if main_process_ip:
+ cmd.extend(["--master_addr", str(main_process_ip)])
+ cmd.extend(["--master_port", str(main_process_port)])
+ if args.module and args.no_python:
+ raise ValueError("--module and --no_python cannot be used together")
+ elif args.module:
+ cmd.append("--module")
+ elif args.no_python:
+ cmd.append("--no_python")
+ cmd.append(args.training_script)
+ cmd.extend(args.training_script_args)
+ elif num_machines > 1 and args.deepspeed_multinode_launcher == DEEPSPEED_MULTINODE_LAUNCHERS[1]:
+ args.nproc_per_node = str(num_processes // num_machines)
+ args.nnodes = str(num_machines)
+ args.node_rank = int(args.machine_rank)
+ if getattr(args, "same_network", False):
+ args.master_addr = str(main_process_ip)
+ args.master_port = str(main_process_port)
+ else:
+ args.rdzv_endpoint = f"{main_process_ip}:{main_process_port}"
+ else:
+ args.nproc_per_node = str(num_processes)
+ if main_process_port is not None:
+ args.master_port = str(main_process_port)
+
+ if main_process_port is None:
+ main_process_port = 29500
+
+ # only need to check port availability in main process, in case we have to start multiple launchers on the same machine
+ # for some reasons like splitting log files.
+ need_port_check = num_machines <= 1 or int(args.machine_rank) == 0
+ if need_port_check and is_port_in_use(main_process_port):
+ raise ConnectionError(
+ f"Tried to launch distributed communication on port `{main_process_port}`, but another process is utilizing it. "
+ "Please specify a different port (such as using the `--main_process_port` flag or specifying a different `main_process_port` in your config file)"
+ " and rerun your script. To automatically use the next open port (on a single node), you can set this to `0`."
+ )
+
+ if args.module and args.no_python:
+ raise ValueError("--module and --no_python cannot be used together")
+ elif args.module:
+ args.module = True
+ elif args.no_python:
+ args.no_python = True
+
+ current_env = os.environ.copy()
+ if args.debug:
+ current_env["ACCELERATE_DEBUG_MODE"] = "true"
+ gpu_ids = getattr(args, "gpu_ids", "all")
+ if gpu_ids != "all" and args.gpu_ids is not None:
+ if is_xpu_available():
+ current_env["ZE_AFFINITY_MASK"] = gpu_ids
+ elif is_mlu_available():
+ current_env["MLU_VISIBLE_DEVICES"] = gpu_ids
+ elif is_sdaa_available():
+ current_env["SDAA_VISIBLE_DEVICES"] = gpu_ids
+ elif is_musa_available():
+ current_env["MUSA_VISIBLE_DEVICES"] = gpu_ids
+ elif is_npu_available():
+ current_env["ASCEND_RT_VISIBLE_DEVICES"] = gpu_ids
+ elif is_hpu_available():
+ current_env["HABANA_VISIBLE_MODULES"] = gpu_ids
+ else:
+ current_env["CUDA_VISIBLE_DEVICES"] = gpu_ids
+ try:
+ mixed_precision = PrecisionType(args.mixed_precision.lower())
+ except ValueError:
+ raise ValueError(
+ f"Unknown mixed_precision mode: {args.mixed_precision.lower()}. Choose between {PrecisionType.list()}."
+ )
+
+ current_env["PYTHONPATH"] = env_var_path_add("PYTHONPATH", os.path.abspath("."))
+ current_env["ACCELERATE_MIXED_PRECISION"] = str(mixed_precision)
+ if args.mixed_precision.lower() == "fp8":
+ if not is_fp8_available():
+ raise RuntimeError(
+ "FP8 is not available on this machine. Please ensure that either Transformer Engine or MSAMP is installed."
+ )
+ current_env = setup_fp8_env(args, current_env)
+ current_env["ACCELERATE_CONFIG_DS_FIELDS"] = str(args.deepspeed_fields_from_accelerate_config).lower()
+ current_env["ACCELERATE_USE_DEEPSPEED"] = "true"
+ if args.zero_stage is not None:
+ current_env["ACCELERATE_DEEPSPEED_ZERO_STAGE"] = str(args.zero_stage)
+ if args.gradient_accumulation_steps is not None:
+ current_env["ACCELERATE_GRADIENT_ACCUMULATION_STEPS"] = str(args.gradient_accumulation_steps)
+ if args.gradient_clipping is not None:
+ current_env["ACCELERATE_GRADIENT_CLIPPING"] = str(args.gradient_clipping).lower()
+ if args.offload_optimizer_device is not None:
+ current_env["ACCELERATE_DEEPSPEED_OFFLOAD_OPTIMIZER_DEVICE"] = str(args.offload_optimizer_device).lower()
+ if args.offload_param_device is not None:
+ current_env["ACCELERATE_DEEPSPEED_OFFLOAD_PARAM_DEVICE"] = str(args.offload_param_device).lower()
+ if args.zero3_init_flag is not None:
+ current_env["ACCELERATE_DEEPSPEED_ZERO3_INIT"] = str(args.zero3_init_flag).lower()
+ if args.zero3_save_16bit_model is not None:
+ current_env["ACCELERATE_DEEPSPEED_ZERO3_SAVE_16BIT_MODEL"] = str(args.zero3_save_16bit_model).lower()
+ if args.deepspeed_config_file is not None:
+ current_env["ACCELERATE_DEEPSPEED_CONFIG_FILE"] = str(args.deepspeed_config_file)
+ if args.enable_cpu_affinity:
+ current_env["ACCELERATE_CPU_AFFINITY"] = "1"
+ if args.deepspeed_moe_layer_cls_names is not None:
+ current_env["ACCELERATE_DEEPSPEED_MOE_LAYER_CLS_NAMES"] = str(args.deepspeed_moe_layer_cls_names)
+ return cmd, current_env
+
+
+def prepare_tpu(
+ args: argparse.Namespace, current_env: dict[str, str], pod: bool = False
+) -> tuple[argparse.Namespace, dict[str, str]]:
+ """
+ Prepares and returns an environment with the correct TPU environment variables.
+ """
+ if args.mixed_precision == "bf16" and is_torch_xla_available(check_is_tpu=True):
+ if args.downcast_bf16:
+ current_env["XLA_DOWNCAST_BF16"] = "1"
+ else:
+ current_env["XLA_USE_BF16"] = "1"
+ if args.debug:
+ current_env["ACCELERATE_DEBUG_MODE"] = "true"
+ if pod:
+ # Take explicit args and set them up for XLA
+ args.vm = args.tpu_vm
+ args.tpu = args.tpu_name
+ return args, current_env
+
+
+def _convert_nargs_to_dict(nargs: list[str]) -> dict[str, str]:
+ if len(nargs) < 0:
+ return {}
+ # helper function to infer type for argsparser
+
+ def _infer_type(s):
+ try:
+ s = float(s)
+
+ if s // 1 == s:
+ return int(s)
+ return s
+ except ValueError:
+ return s
+
+ parser = argparse.ArgumentParser()
+ _, unknown = parser.parse_known_args(nargs)
+ for index, argument in enumerate(unknown):
+ if argument.startswith(("-", "--")):
+ action = None
+ if index + 1 < len(unknown): # checks if next index would be in list
+ if unknown[index + 1].startswith(("-", "--")): # checks if next element is an key
+ # raise an error if element is store_true or store_false
+ raise ValueError(
+ "SageMaker doesn’t support argparse actions for `store_true` or `store_false`. Please define explicit types"
+ )
+ else: # raise an error if last element is store_true or store_false
+ raise ValueError(
+ "SageMaker doesn’t support argparse actions for `store_true` or `store_false`. Please define explicit types"
+ )
+ # adds argument to parser based on action_store true
+ if action is None:
+ parser.add_argument(argument, type=_infer_type)
+ else:
+ parser.add_argument(argument, action=action)
+
+ return {
+ key: (literal_eval(value) if value in ("True", "False") else value)
+ for key, value in parser.parse_args(nargs).__dict__.items()
+ }
+
+
+def prepare_sagemager_args_inputs(
+ sagemaker_config: SageMakerConfig, args: argparse.Namespace
+) -> tuple[argparse.Namespace, dict[str, Any]]:
+ # configure environment
+ print("Configuring Amazon SageMaker environment")
+ os.environ["AWS_DEFAULT_REGION"] = sagemaker_config.region
+
+ # configure credentials
+ if sagemaker_config.profile is not None:
+ os.environ["AWS_PROFILE"] = sagemaker_config.profile
+ elif args.aws_access_key_id is not None and args.aws_secret_access_key is not None:
+ os.environ["AWS_ACCESS_KEY_ID"] = args.aws_access_key_id
+ os.environ["AWS_SECRET_ACCESS_KEY"] = args.aws_secret_access_key
+ else:
+ raise OSError("You need to provide an aws_access_key_id and aws_secret_access_key when not using aws_profile")
+
+ # extract needed arguments
+ source_dir = os.path.dirname(args.training_script)
+ if not source_dir: # checks if string is empty
+ source_dir = "."
+ entry_point = os.path.basename(args.training_script)
+ if not entry_point.endswith(".py"):
+ raise ValueError(f'Your training script should be a python script and not "{entry_point}"')
+
+ print("Converting Arguments to Hyperparameters")
+ hyperparameters = _convert_nargs_to_dict(args.training_script_args)
+
+ try:
+ mixed_precision = PrecisionType(args.mixed_precision.lower())
+ except ValueError:
+ raise ValueError(
+ f"Unknown mixed_precision mode: {args.mixed_precision.lower()}. Choose between {PrecisionType.list()}."
+ )
+
+ try:
+ dynamo_backend = DynamoBackend(args.dynamo_backend.upper())
+ except ValueError:
+ raise ValueError(
+ f"Unknown dynamo backend: {args.dynamo_backend.upper()}. Choose between {DynamoBackend.list()}."
+ )
+
+ # Environment variables to be set for use during training job
+ environment = {
+ "ACCELERATE_USE_SAGEMAKER": "true",
+ "ACCELERATE_MIXED_PRECISION": str(mixed_precision),
+ "ACCELERATE_DYNAMO_BACKEND": dynamo_backend.value,
+ "ACCELERATE_DYNAMO_MODE": args.dynamo_mode,
+ "ACCELERATE_DYNAMO_USE_FULLGRAPH": str(args.dynamo_use_fullgraph),
+ "ACCELERATE_DYNAMO_USE_DYNAMIC": str(args.dynamo_use_dynamic),
+ "ACCELERATE_SAGEMAKER_DISTRIBUTED_TYPE": sagemaker_config.distributed_type.value,
+ }
+ if args.mixed_precision.lower() == "fp8":
+ if not is_fp8_available():
+ raise RuntimeError(
+ "FP8 is not available on this machine. Please ensure that either Transformer Engine or MSAMP is installed."
+ )
+ environment = setup_fp8_env(args, environment)
+ # configure distribution set up
+ distribution = None
+ if sagemaker_config.distributed_type == SageMakerDistributedType.DATA_PARALLEL:
+ distribution = {"smdistributed": {"dataparallel": {"enabled": True}}}
+
+ # configure sagemaker inputs
+ sagemaker_inputs = None
+ if sagemaker_config.sagemaker_inputs_file is not None:
+ print(f"Loading SageMaker Inputs from {sagemaker_config.sagemaker_inputs_file} file")
+ sagemaker_inputs = {}
+ with open(sagemaker_config.sagemaker_inputs_file) as file:
+ for i, line in enumerate(file):
+ if i == 0:
+ continue
+ l = line.split("\t")
+ sagemaker_inputs[l[0]] = l[1].strip()
+ print(f"Loaded SageMaker Inputs: {sagemaker_inputs}")
+
+ # configure sagemaker metrics
+ sagemaker_metrics = None
+ if sagemaker_config.sagemaker_metrics_file is not None:
+ print(f"Loading SageMaker Metrics from {sagemaker_config.sagemaker_metrics_file} file")
+ sagemaker_metrics = []
+ with open(sagemaker_config.sagemaker_metrics_file) as file:
+ for i, line in enumerate(file):
+ if i == 0:
+ continue
+ l = line.split("\t")
+ metric_dict = {
+ "Name": l[0],
+ "Regex": l[1].strip(),
+ }
+ sagemaker_metrics.append(metric_dict)
+ print(f"Loaded SageMaker Metrics: {sagemaker_metrics}")
+
+ # configure session
+ print("Creating Estimator")
+ args = {
+ "image_uri": sagemaker_config.image_uri,
+ "entry_point": entry_point,
+ "source_dir": source_dir,
+ "role": sagemaker_config.iam_role_name,
+ "transformers_version": sagemaker_config.transformers_version,
+ "pytorch_version": sagemaker_config.pytorch_version,
+ "py_version": sagemaker_config.py_version,
+ "base_job_name": sagemaker_config.base_job_name,
+ "instance_count": sagemaker_config.num_machines,
+ "instance_type": sagemaker_config.ec2_instance_type,
+ "debugger_hook_config": False,
+ "distribution": distribution,
+ "hyperparameters": hyperparameters,
+ "environment": environment,
+ "metric_definitions": sagemaker_metrics,
+ }
+
+ if sagemaker_config.additional_args is not None:
+ args = merge_dicts(sagemaker_config.additional_args, args)
+ return args, sagemaker_inputs
+
+
+def env_var_path_add(env_var_name, path_to_add):
+ """
+ Extends a path-based environment variable's value with a new path and returns the updated value. It's up to the
+ caller to set it in os.environ.
+ """
+ paths = [p for p in os.environ.get(env_var_name, "").split(":") if len(p) > 0]
+ paths.append(str(path_to_add))
+ return ":".join(paths)
+
+
+class PrepareForLaunch:
+ """
+ Prepare a function that will launched in a distributed setup.
+
+ Args:
+ launcher (`Callable`):
+ The function to launch.
+ distributed_type ([`~state.DistributedType`]):
+ The distributed type to prepare for.
+ debug (`bool`, *optional*, defaults to `False`):
+ Whether or not this is a debug launch.
+ """
+
+ def __init__(self, launcher, distributed_type="NO", debug=False):
+ self.launcher = launcher
+ self.distributed_type = DistributedType(distributed_type)
+ self.debug = debug
+
+ def __call__(self, index, *args):
+ if self.debug:
+ world_size = int(os.environ.get("WORLD_SIZE"))
+ rdv_file = os.environ.get("ACCELERATE_DEBUG_RDV_FILE")
+ torch.distributed.init_process_group(
+ "gloo",
+ rank=index,
+ store=torch.distributed.FileStore(rdv_file, world_size),
+ world_size=world_size,
+ )
+ elif self.distributed_type in (
+ DistributedType.MULTI_GPU,
+ DistributedType.MULTI_MLU,
+ DistributedType.MULTI_MUSA,
+ DistributedType.MULTI_NPU,
+ DistributedType.MULTI_XPU,
+ DistributedType.MULTI_CPU,
+ ):
+ # Prepare the environment for torch.distributed
+ os.environ["LOCAL_RANK"] = str(index)
+ nproc = int(os.environ.get("NPROC", 1))
+ node_rank = int(os.environ.get("NODE_RANK", 0))
+ os.environ["RANK"] = str(nproc * node_rank + index)
+
+ os.environ["FORK_LAUNCHED"] = str(1)
+ self.launcher(*args)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/megatron_lm.py b/venv/lib/python3.11/site-packages/accelerate/utils/megatron_lm.py
new file mode 100644
index 0000000000000000000000000000000000000000..c867d6007bcdd6b252c9852cc87c2919ca225891
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/megatron_lm.py
@@ -0,0 +1,1424 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 argparse
+import math
+import os
+from abc import ABC
+from functools import partial
+
+import torch
+import torch.nn.functional as F
+from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss
+from torch.nn.parallel.distributed import DistributedDataParallel as torchDDP
+
+from ..optimizer import AcceleratedOptimizer
+from ..scheduler import AcceleratedScheduler
+from .imports import is_megatron_lm_available
+from .operations import recursively_apply, send_to_device
+
+
+if is_megatron_lm_available():
+ from megatron.core import mpu, tensor_parallel
+ from megatron.core.distributed import DistributedDataParallel as LocalDDP
+ from megatron.core.distributed import finalize_model_grads
+ from megatron.core.enums import ModelType
+ from megatron.core.num_microbatches_calculator import get_num_microbatches
+ from megatron.core.optimizer import get_megatron_optimizer
+ from megatron.core.parallel_state import get_tensor_model_parallel_group, get_tensor_model_parallel_src_rank
+ from megatron.core.pipeline_parallel import get_forward_backward_func
+ from megatron.core.utils import get_model_config
+ from megatron.inference.text_generation.communication import broadcast_int_list, broadcast_tensor
+ from megatron.inference.text_generation.generation import (
+ beam_search_and_return_on_first_stage,
+ generate_tokens_probs_and_return_on_first_stage,
+ )
+ from megatron.legacy.data.dataset_utils import build_train_valid_test_datasets
+ from megatron.legacy.model import BertModel, Float16Module, GPTModel, T5Model
+ from megatron.legacy.model.classification import Classification
+ from megatron.training import (
+ get_args,
+ get_tensorboard_writer,
+ get_tokenizer,
+ print_rank_last,
+ )
+ from megatron.training.arguments import (
+ _add_data_args,
+ _add_validation_args,
+ core_transformer_config_from_args,
+ parse_args,
+ validate_args,
+ )
+ from megatron.training.checkpointing import load_args_from_checkpoint, load_checkpoint, save_checkpoint
+ from megatron.training.global_vars import set_global_variables
+ from megatron.training.initialize import (
+ _compile_dependencies,
+ _init_autoresume,
+ _initialize_distributed,
+ _set_random_seed,
+ set_jit_fusion_options,
+ write_args_to_tensorboard,
+ )
+ from megatron.training.tokenizer.tokenizer import _vocab_size_with_padding
+ from megatron.training.training import (
+ build_train_valid_test_data_iterators,
+ get_optimizer_param_scheduler,
+ num_floating_point_operations,
+ setup_model_and_optimizer,
+ train_step,
+ training_log,
+ )
+ from megatron.training.utils import (
+ average_losses_across_data_parallel_group,
+ calc_params_l2_norm,
+ get_ltor_masks_and_position_ids,
+ unwrap_model,
+ )
+
+
+# model utilities
+def model_provider_func(pre_process=True, post_process=True, add_encoder=True, add_decoder=True):
+ """Build the model."""
+ args = get_args()
+ mode = "pre-training" if args.pretraining_flag else "fine-tuning"
+ if args.rank == 0:
+ print(f"Building {args.model_type_name} model in the {mode} mode.")
+ print(
+ "The Megatron LM model weights are initialized at random in `accelerator.prepare`. "
+ "Please use `accelerator.load_checkpoint` to load a pre-trained checkpoint matching the distributed setup."
+ )
+ config = core_transformer_config_from_args(args)
+ if args.model_type_name == "bert":
+ if args.pretraining_flag:
+ num_tokentypes = 2 if args.bert_binary_head else 0
+ model = BertModel(
+ config=config,
+ num_tokentypes=num_tokentypes,
+ add_binary_head=args.bert_binary_head,
+ parallel_output=True,
+ pre_process=pre_process,
+ post_process=post_process,
+ )
+ else:
+ model = Classification(
+ config=config,
+ num_classes=args.num_labels,
+ num_tokentypes=2,
+ pre_process=pre_process,
+ post_process=post_process,
+ )
+ elif args.model_type_name == "gpt":
+ model = GPTModel(
+ config=config,
+ num_tokentypes=0,
+ parallel_output=True,
+ pre_process=pre_process,
+ post_process=post_process,
+ )
+ elif args.model_type_name == "t5":
+ model = T5Model(
+ config=config,
+ num_tokentypes=0,
+ parallel_output=True,
+ pre_process=pre_process,
+ post_process=post_process,
+ add_encoder=add_encoder,
+ add_decoder=add_decoder,
+ )
+ else:
+ raise ValueError(f"Unsupported model type: {args.model_type_name}")
+ return model
+
+
+def prepare_model_optimizer_scheduler(accelerator):
+ accelerator.print("Preparing model optimizer scheduler")
+ args = get_args()
+ if accelerator.state.megatron_lm_plugin.custom_prepare_model_function is not None:
+ if accelerator.state.megatron_lm_plugin.custom_model_provider_function is None:
+ raise ValueError(
+ "You must provide a `custom_model_provider_function` when using a `custom_prepare_model_function`."
+ )
+ custom_model_provider_func = accelerator.state.megatron_lm_plugin.custom_model_provider_function
+ model = accelerator.state.megatron_lm_plugin.custom_prepare_model_function(custom_model_provider_func)
+ optimizer = prepare_optimizer(accelerator, model)
+ scheduler = prepare_scheduler(accelerator, optimizer, scheduler=None)
+ else:
+ model_type = ModelType.encoder_or_decoder
+ if args.model_type_name == "t5":
+ model_type = ModelType.encoder_and_decoder
+ model_provider_func_ = model_provider_func
+ if accelerator.state.megatron_lm_plugin.custom_model_provider_function is not None:
+ model_provider_func_ = accelerator.state.megatron_lm_plugin.custom_model_provider_function
+ (model, optimizer, scheduler) = setup_model_and_optimizer(
+ model_provider_func_,
+ model_type,
+ no_wd_decay_cond=args.no_wd_decay_cond,
+ scale_lr_cond=args.scale_lr_cond,
+ lr_mult=args.lr_mult,
+ )
+ args.model_len = len(model)
+ return model, optimizer, scheduler
+
+
+# dataloader utilities
+class MegatronLMDummyDataLoader:
+ """
+ Dummy dataloader presents model parameters or param groups, this is primarily used to follow conventional training
+
+ Args:
+ **dataset_kwargs: Megatron data arguments.
+ """
+
+ def __init__(self, **dataset_kwargs):
+ parser = argparse.ArgumentParser()
+ parser = _add_data_args(parser)
+ parser = _add_validation_args(parser)
+ data_args = parser.parse_known_args()
+ self.dataset_args = vars(data_args[0])
+ self.dataset_args.update(dataset_kwargs)
+ self.dataset_args["megatron_dataset_flag"] = True
+
+ def set_megatron_data_args(self):
+ args = get_args()
+ for key, value in self.dataset_args.items():
+ old_value = getattr(args, key, "")
+ if old_value != value:
+ print(
+ f"WARNING: MegatronLMDummyDataLoader overriding arguments for {key}:{old_value} with {key}:{value}"
+ )
+ setattr(args, key, value)
+
+ def get_train_valid_test_datasets_provider(self, accelerator):
+ def train_valid_test_datasets_provider(train_val_test_num_samples):
+ """Build train, valid, and test datasets."""
+ args = get_args()
+ dataset_args = {
+ "data_prefix": args.data_path if isinstance(args.data_path, (list, tuple)) else [args.data_path],
+ "splits_string": args.split,
+ "train_valid_test_num_samples": train_val_test_num_samples,
+ "seed": args.seed,
+ }
+ if args.model_type_name == "bert":
+ dataset_args.update(
+ {
+ "max_seq_length": args.seq_length,
+ "binary_head": args.bert_binary_head,
+ }
+ )
+ elif args.model_type_name == "gpt":
+ dataset_args.update(
+ {
+ "max_seq_length": args.seq_length,
+ }
+ )
+ elif args.model_type_name == "t5":
+ dataset_args.update(
+ {
+ "max_seq_length": args.encoder_seq_length,
+ "max_seq_length_dec": args.decoder_seq_length,
+ "dataset_type": "t5",
+ }
+ )
+ else:
+ raise ValueError(f"Unsupported model type: {args.model_type_name}")
+ train_ds, valid_ds, test_ds = build_train_valid_test_datasets(**dataset_args)
+ return train_ds, valid_ds, test_ds
+
+ if accelerator.state.megatron_lm_plugin.custom_megatron_datasets_provider_function is not None:
+ return accelerator.state.megatron_lm_plugin.custom_megatron_datasets_provider_function
+ try:
+ args = get_args()
+ # Use '--no-use-pep517 -e' to pip install nvidia's megatron from source
+ if args.model_type_name == "bert":
+ from pretrain_bert import train_valid_test_datasets_provider
+
+ train_valid_test_datasets_provider.is_distributed = True
+ return train_valid_test_datasets_provider
+ elif args.model_type_name == "gpt":
+ from pretrain_gpt import train_valid_test_datasets_provider
+
+ train_valid_test_datasets_provider.is_distributed = True
+ return train_valid_test_datasets_provider
+ elif args.model_type_name == "t5":
+ from pretrain_t5 import train_valid_test_datasets_provider
+
+ train_valid_test_datasets_provider.is_distributed = True
+ return train_valid_test_datasets_provider
+ except ImportError:
+ pass
+ return train_valid_test_datasets_provider
+
+ def build_train_valid_test_data_iterators(self, accelerator):
+ args = get_args()
+
+ train_valid_test_dataset_provider = self.get_train_valid_test_datasets_provider(accelerator)
+ if args.virtual_pipeline_model_parallel_size is not None:
+ train_data_iterator = []
+ valid_data_iterator = []
+ test_data_iterator = []
+ for i in range(getattr(args, "model_len", 0)):
+ mpu.set_virtual_pipeline_model_parallel_rank(i)
+ iterators = build_train_valid_test_data_iterators(train_valid_test_dataset_provider)
+ train_data_iterator.append(iterators[0])
+ valid_data_iterator.append(iterators[1])
+ test_data_iterator.append(iterators[2])
+ else:
+ train_data_iterator, valid_data_iterator, test_data_iterator = build_train_valid_test_data_iterators(
+ train_valid_test_dataset_provider
+ )
+
+ return train_data_iterator, valid_data_iterator, test_data_iterator
+
+
+def _handle_megatron_data_iterator(accelerator, data_iterator):
+ class DummyMegatronDataloader:
+ def __iter__(self):
+ return self
+
+ def __next__(self):
+ return {}
+
+ is_data_iterator_empty = data_iterator is None
+ is_src_data_iterator_empty = torch.tensor(is_data_iterator_empty, dtype=torch.bool, device=accelerator.device)
+ torch.distributed.broadcast(
+ is_src_data_iterator_empty, get_tensor_model_parallel_src_rank(), group=get_tensor_model_parallel_group()
+ )
+ if not is_src_data_iterator_empty and is_data_iterator_empty:
+ return DummyMegatronDataloader()
+ return data_iterator
+
+
+def prepare_data_loader(accelerator, dataloader):
+ accelerator.print("Preparing dataloader")
+ args = get_args()
+ if not args.megatron_dataset_flag:
+ from ..data_loader import _PYTORCH_DATALOADER_KWARGS, prepare_data_loader
+
+ micro_batch_size = args.micro_batch_size * args.num_micro_batches
+ kwargs = {k: getattr(dataloader, k, _PYTORCH_DATALOADER_KWARGS[k]) for k in _PYTORCH_DATALOADER_KWARGS}
+ if kwargs["batch_size"] is None:
+ if isinstance(kwargs["sampler"], torch.utils.data.BatchSampler):
+ kwargs["sampler"].batch_size = micro_batch_size
+ else:
+ del kwargs["sampler"]
+ del kwargs["shuffle"]
+ del kwargs["batch_size"]
+ kwargs["batch_sampler"].batch_size = micro_batch_size
+ else:
+ del kwargs["batch_sampler"]
+ kwargs["batch_size"] = micro_batch_size
+
+ dataloader = torch.utils.data.DataLoader(dataloader.dataset, **kwargs)
+ # split_batches:
+ # Megatron only needs to fetch different data between different dp groups,
+ # and does not need to split the data within the dp group.
+ return prepare_data_loader(
+ dataloader,
+ accelerator.device,
+ num_processes=mpu.get_data_parallel_world_size(),
+ process_index=mpu.get_data_parallel_rank(),
+ split_batches=False,
+ put_on_device=True,
+ rng_types=accelerator.rng_types.copy(),
+ dispatch_batches=accelerator.dispatch_batches,
+ )
+ else:
+ if args.consumed_samples is not None:
+ (
+ args.consumed_train_samples,
+ args.consumed_valid_samples,
+ args.consumed_test_samples,
+ ) = args.consumed_samples
+ else:
+ args.consumed_train_samples, args.consumed_valid_samples, args.consumed_test_samples = 0, 0, 0
+ args.micro_batch_size = args.micro_batch_size * args.num_micro_batches
+ # In order to be compatible with data in transform format,
+ # it needs to increase the size of mbs first,
+ # and then split the large batch data into some mbs.
+ (
+ train_data_iterator,
+ valid_data_iterator,
+ test_data_iterator,
+ ) = dataloader.build_train_valid_test_data_iterators(accelerator)
+ args.micro_batch_size = args.micro_batch_size // args.num_micro_batches
+
+ train_data_iterator = _handle_megatron_data_iterator(
+ accelerator=accelerator, data_iterator=train_data_iterator
+ )
+ valid_data_iterator = _handle_megatron_data_iterator(
+ accelerator=accelerator, data_iterator=valid_data_iterator
+ )
+ test_data_iterator = _handle_megatron_data_iterator(accelerator=accelerator, data_iterator=test_data_iterator)
+
+ return train_data_iterator, valid_data_iterator, test_data_iterator
+
+
+# optimizer utilities
+class MegatronLMOptimizerWrapper(AcceleratedOptimizer):
+ def __init__(self, optimizer):
+ super().__init__(optimizer, device_placement=False, scaler=None)
+
+ def zero_grad(self, set_to_none=None):
+ pass # `model(**batch)` is doing that automatically. Therefore, it's implementation is not needed
+
+ def step(self):
+ pass # `model(**batch)` is doing that automatically. Therefore, it's implementation is not needed
+
+ @property
+ def step_was_skipped(self):
+ """Whether or not the optimizer step was done, or skipped because of gradient overflow."""
+ return self.optimizer.skipped_iter
+
+
+def prepare_optimizer(accelerator, model):
+ accelerator.print("Preparing optimizer")
+ args = get_args()
+ return get_megatron_optimizer(model, args.no_wd_decay_cond, args.scale_lr_cond, args.lr_mult)
+
+
+# scheduler utilities
+class MegatronLMDummyScheduler:
+ """
+ Dummy scheduler presents model parameters or param groups, this is primarily used to follow conventional training
+ loop when scheduler config is specified in the deepspeed config file.
+
+ Args:
+ optimizer (`torch.optim.optimizer.Optimizer`):
+ The optimizer to wrap.
+ total_num_steps (int):
+ Total number of steps.
+ warmup_num_steps (int):
+ Number of steps for warmup.
+ **kwargs (additional keyword arguments, *optional*):
+ Other arguments.
+ """
+
+ def __init__(self, optimizer, total_num_steps=None, warmup_num_steps=0, **kwargs):
+ self.optimizer = optimizer
+ self.total_num_steps = total_num_steps
+ self.warmup_num_steps = warmup_num_steps
+ self.kwargs = kwargs
+
+
+class MegatronLMSchedulerWrapper(AcceleratedScheduler):
+ def __init__(self, scheduler, optimizers):
+ super().__init__(scheduler, optimizers)
+
+ def step(self, *args, **kwargs):
+ return # `model(**batch)` is doing that automatically. Therefore, it's implementation is not needed
+
+
+def prepare_scheduler(accelerator, optimizer, scheduler):
+ accelerator.print("Preparing scheduler")
+ scheduler = get_optimizer_param_scheduler(optimizer)
+ return scheduler
+
+
+class AbstractTrainStep(ABC):
+ """Abstract class for batching, forward pass and loss handler."""
+
+ def __init__(self, name):
+ super().__init__()
+ self.name = name
+
+ def get_batch_func(self, accelerator, megatron_dataset_flag):
+ pass
+
+ def get_forward_step_func(self):
+ pass
+
+ def get_loss_func(self, accelerator):
+ pass
+
+
+class BertTrainStep(AbstractTrainStep):
+ """
+ Bert train step class.
+
+ Args:
+ args (`argparse.Namespace`): Megatron-LM arguments.
+ """
+
+ def __init__(self, accelerator, args):
+ super().__init__("BertTrainStep")
+ self.get_batch = self.get_batch_func(accelerator, args.megatron_dataset_flag)
+ self.loss_func = self.get_loss_func(accelerator, args.pretraining_flag, args.num_labels)
+ self.forward_step = self.get_forward_step_func(args.pretraining_flag, args.bert_binary_head)
+ if not args.model_return_dict:
+ self.model_output_class = None
+ else:
+ from transformers.modeling_outputs import SequenceClassifierOutput
+
+ self.model_output_class = SequenceClassifierOutput
+
+ def get_batch_func(self, accelerator, megatron_dataset_flag):
+ def get_batch_megatron(data_iterator):
+ """Build the batch."""
+
+ # Items and their type.
+ keys = ["text", "types", "labels", "is_random", "loss_mask", "padding_mask"]
+ datatype = torch.int64
+
+ # Broadcast data.
+ if data_iterator is not None:
+ data = next(data_iterator)
+ else:
+ data = None
+ data_b = tensor_parallel.broadcast_data(keys, data, datatype)
+
+ # Unpack.
+ tokens = data_b["text"].long()
+ types = data_b["types"].long()
+ sentence_order = data_b["is_random"].long()
+ loss_mask = data_b["loss_mask"].float()
+ lm_labels = data_b["labels"].long()
+ padding_mask = data_b["padding_mask"].long()
+
+ return tokens, types, sentence_order, loss_mask, lm_labels, padding_mask
+
+ def get_batch_transformer(data_iterator):
+ """Build the batch."""
+ data = next(data_iterator)
+ data = send_to_device(data, torch.cuda.current_device())
+
+ # Unpack.
+ tokens = data["input_ids"].long()
+ padding_mask = data["attention_mask"].long()
+ if "token_type_ids" in data:
+ types = data["token_type_ids"].long()
+ else:
+ types = None
+ if "labels" in data:
+ lm_labels = data["labels"].long()
+ loss_mask = (data["labels"] != -100).to(torch.float)
+ else:
+ lm_labels = None
+ loss_mask = None
+ if "next_sentence_label" in data:
+ sentence_order = data["next_sentence_label"].long()
+ else:
+ sentence_order = None
+
+ return tokens, types, sentence_order, loss_mask, lm_labels, padding_mask
+
+ if accelerator.state.megatron_lm_plugin.custom_get_batch_function is not None:
+ return accelerator.state.megatron_lm_plugin.custom_get_batch_function
+ if megatron_dataset_flag:
+ try:
+ # Use '--no-use-pep517 -e' to pip install nvidia's megatron from source
+ from pretrain_bert import get_batch
+
+ return get_batch
+ except ImportError:
+ pass
+ return get_batch_megatron
+ else:
+ return get_batch_transformer
+
+ def get_loss_func(self, accelerator, pretraining_flag, num_labels):
+ def loss_func_pretrain(loss_mask, sentence_order, output_tensor):
+ lm_loss_, sop_logits = output_tensor
+
+ lm_loss_ = lm_loss_.float()
+ loss_mask = loss_mask.float()
+ lm_loss = torch.sum(lm_loss_.view(-1) * loss_mask.reshape(-1)) / loss_mask.sum()
+
+ if sop_logits is not None:
+ sop_loss = F.cross_entropy(sop_logits.view(-1, 2).float(), sentence_order.view(-1), ignore_index=-1)
+ sop_loss = sop_loss.float()
+ loss = lm_loss + sop_loss
+ averaged_losses = average_losses_across_data_parallel_group([lm_loss, sop_loss])
+ return loss, {"lm loss": averaged_losses[0], "sop loss": averaged_losses[1]}
+
+ else:
+ loss = lm_loss
+ averaged_losses = average_losses_across_data_parallel_group([lm_loss])
+ return loss, {"lm loss": averaged_losses[0]}
+
+ def loss_func_finetune(labels, logits):
+ if num_labels == 1:
+ # We are doing regression
+ loss_fct = MSELoss()
+ loss = loss_fct(logits.view(-1), labels.view(-1))
+ elif self.num_labels > 1 and (labels.dtype in (torch.long, torch.int)):
+ loss_fct = CrossEntropyLoss()
+ loss = loss_fct(logits.view(-1, num_labels), labels.view(-1))
+ else:
+ loss_fct = BCEWithLogitsLoss()
+ loss = loss_fct(logits, labels)
+ averaged_losses = average_losses_across_data_parallel_group([loss])
+ return loss, {"loss": averaged_losses[0]}
+
+ if accelerator.state.megatron_lm_plugin.custom_loss_function is not None:
+ return accelerator.state.megatron_lm_plugin.custom_loss_function
+ if pretraining_flag:
+ return loss_func_pretrain
+ else:
+ return loss_func_finetune
+
+ def get_forward_step_func(self, pretraining_flag, bert_binary_head):
+ def forward_step(data_iterator, model):
+ """Forward step."""
+ tokens, types, sentence_order, loss_mask, labels, padding_mask = self.get_batch(data_iterator)
+ if not bert_binary_head:
+ types = None
+ # Forward pass through the model.
+ if pretraining_flag:
+ output_tensor = model(tokens, padding_mask, tokentype_ids=types, lm_labels=labels)
+ return output_tensor, partial(self.loss_func, loss_mask, sentence_order)
+ else:
+ logits = model(tokens, padding_mask, tokentype_ids=types)
+ return logits, partial(self.loss_func, labels)
+
+ return forward_step
+
+
+class GPTTrainStep(AbstractTrainStep):
+ """
+ GPT train step class.
+
+ Args:
+ args (`argparse.Namespace`): Megatron-LM arguments.
+ """
+
+ def __init__(self, accelerator, args):
+ super().__init__("GPTTrainStep")
+ self.get_batch = self.get_batch_func(accelerator, args.megatron_dataset_flag)
+ self.loss_func = self.get_loss_func(accelerator)
+ self.forward_step = self.get_forward_step_func()
+ self.eod_token = args.padded_vocab_size - 1
+ if args.vocab_file is not None:
+ tokenizer = get_tokenizer()
+ self.eod_token = tokenizer.eod
+ self.reset_position_ids = args.reset_position_ids
+ self.reset_attention_mask = args.reset_attention_mask
+ self.eod_mask_loss = args.eod_mask_loss
+ if not args.model_return_dict:
+ self.model_output_class = None
+ else:
+ from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions
+
+ self.model_output_class = CausalLMOutputWithCrossAttentions
+
+ def get_batch_func(self, accelerator, megatron_dataset_flag):
+ def get_batch_megatron(data_iterator):
+ """Generate a batch"""
+ # Items and their type.
+ keys = ["text"]
+ datatype = torch.int64
+
+ # Broadcast data.
+ if data_iterator is not None:
+ data = next(data_iterator)
+ else:
+ data = None
+ data_b = tensor_parallel.broadcast_data(keys, data, datatype)
+
+ # Unpack.
+ tokens_ = data_b["text"].long()
+ labels = tokens_[:, 1:].contiguous()
+ tokens = tokens_[:, :-1].contiguous()
+
+ # Get the masks and postition ids.
+ attention_mask, loss_mask, position_ids = get_ltor_masks_and_position_ids(
+ tokens, self.eod_token, self.reset_position_ids, self.reset_attention_mask, self.eod_mask_loss
+ )
+
+ return tokens, labels, loss_mask, attention_mask, position_ids
+
+ def get_batch_transformer(data_iterator):
+ data = next(data_iterator)
+ data = {"input_ids": data["input_ids"]}
+ data = send_to_device(data, torch.cuda.current_device())
+
+ tokens_ = data["input_ids"].long()
+ padding = torch.zeros((tokens_.shape[0], 1), dtype=tokens_.dtype, device=tokens_.device) + self.eod_token
+ tokens_ = torch.concat([tokens_, padding], dim=1)
+ labels = tokens_[:, 1:].contiguous()
+ tokens = tokens_[:, :-1].contiguous()
+ # Get the masks and postition ids.
+ attention_mask, loss_mask, position_ids = get_ltor_masks_and_position_ids(
+ tokens, self.eod_token, self.reset_position_ids, self.reset_attention_mask, True
+ )
+ return tokens, labels, loss_mask, attention_mask, position_ids
+
+ if accelerator.state.megatron_lm_plugin.custom_get_batch_function is not None:
+ return accelerator.state.megatron_lm_plugin.custom_get_batch_function
+ if megatron_dataset_flag:
+ try:
+ # Use '--no-use-pep517 -e' to pip install nvidia's megatron from source
+ from pretrain_gpt import get_batch
+
+ return get_batch
+ except ImportError:
+ pass
+ return get_batch_megatron
+ else:
+ return get_batch_transformer
+
+ def get_loss_func(self, accelerator):
+ args = get_args()
+
+ def loss_func(loss_mask, output_tensor):
+ if args.return_logits:
+ losses, logits = output_tensor
+ else:
+ losses = output_tensor
+ losses = losses.float()
+ loss_mask = loss_mask.view(-1).float()
+ if args.context_parallel_size > 1:
+ loss = torch.cat([torch.sum(losses.view(-1) * loss_mask).view(1), loss_mask.sum().view(1)])
+ torch.distributed.all_reduce(loss, group=mpu.get_context_parallel_group())
+ loss = loss[0] / loss[1]
+ else:
+ loss = torch.sum(losses.view(-1) * loss_mask) / loss_mask.sum()
+
+ # Check individual rank losses are not NaN prior to DP all-reduce.
+ if args.check_for_nan_in_loss_and_grad:
+ global_rank = torch.distributed.get_rank()
+ assert not loss.isnan(), (
+ f"Rank {global_rank}: found NaN in local forward loss calculation. "
+ f"Device: {torch.cuda.current_device()}, node: {os.uname()[1]}"
+ )
+
+ # Reduce loss for logging.
+ averaged_loss = average_losses_across_data_parallel_group([loss])
+
+ output_dict = {"lm loss": averaged_loss[0]}
+ if args.return_logits:
+ output_dict.update({"logits": logits})
+ return loss, output_dict
+
+ if accelerator.state.megatron_lm_plugin.custom_loss_function is not None:
+ return accelerator.state.megatron_lm_plugin.custom_loss_function
+ return loss_func
+
+ def get_forward_step_func(self):
+ def forward_step(data_iterator, model):
+ """Forward step."""
+ # Get the batch.
+ tokens, labels, loss_mask, attention_mask, position_ids = self.get_batch(data_iterator)
+ output_tensor = model(tokens, position_ids, attention_mask, labels=labels)
+
+ return output_tensor, partial(self.loss_func, loss_mask)
+
+ return forward_step
+
+
+class T5TrainStep(AbstractTrainStep):
+ """
+ T5 train step class.
+
+ Args:
+ args (`argparse.Namespace`): Megatron-LM arguments.
+ """
+
+ def __init__(self, accelerator, args):
+ super().__init__("T5TrainStep")
+ self.get_batch = self.get_batch_func(accelerator, args.megatron_dataset_flag)
+ self.loss_func = self.get_loss_func(accelerator)
+ self.forward_step = self.get_forward_step_func()
+ if not args.model_return_dict:
+ self.model_output_class = None
+ else:
+ from transformers.modeling_outputs import Seq2SeqLMOutput
+
+ self.model_output_class = Seq2SeqLMOutput
+
+ @staticmethod
+ def attn_mask_postprocess(attention_mask):
+ # We create a 3D attention mask from a 2D tensor mask.
+ # [b, 1, s]
+ attention_mask_b1s = attention_mask.unsqueeze(1)
+ # [b, s, 1]
+ attention_mask_bs1 = attention_mask.unsqueeze(2)
+ # [b, s, s]
+ attention_mask_bss = attention_mask_b1s * attention_mask_bs1
+ # Convert attention mask to binary:
+ extended_attention_mask = attention_mask_bss < 0.5
+ return extended_attention_mask
+
+ @staticmethod
+ def get_decoder_mask(seq_length, device):
+ attention_mask = torch.tril(torch.ones((1, seq_length, seq_length), device=device))
+ attention_mask = attention_mask < 0.5
+ return attention_mask
+
+ @staticmethod
+ def get_enc_dec_mask(attention_mask, dec_seq_length, device):
+ batch_size, _ = attention_mask.shape
+ # We create a 3D attention mask from a 2D tensor mask.
+ # [b, 1, s]
+ attention_mask_b1s = attention_mask.unsqueeze(1)
+ # [b, s, 1]
+ attention_mask_bs1 = torch.ones((batch_size, dec_seq_length, 1), device=device)
+ attention_mask_bss = attention_mask_bs1 * attention_mask_b1s
+ extended_attention_mask = attention_mask_bss < 0.5
+ return extended_attention_mask
+
+ def get_batch_func(self, accelerator, megatron_dataset_flag):
+ def get_batch_megatron(data_iterator):
+ """Build the batch."""
+
+ keys = ["text_enc", "text_dec", "labels", "loss_mask", "enc_mask", "dec_mask", "enc_dec_mask"]
+ datatype = torch.int64
+
+ # Broadcast data.
+ if data_iterator is not None:
+ data = next(data_iterator)
+ else:
+ data = None
+ data_b = tensor_parallel.broadcast_data(keys, data, datatype)
+
+ # Unpack.
+ tokens_enc = data_b["text_enc"].long()
+ tokens_dec = data_b["text_dec"].long()
+ labels = data_b["labels"].long()
+ loss_mask = data_b["loss_mask"].float()
+
+ enc_mask = data_b["enc_mask"] < 0.5
+ dec_mask = data_b["dec_mask"] < 0.5
+ enc_dec_mask = data_b["enc_dec_mask"] < 0.5
+
+ return tokens_enc, tokens_dec, loss_mask, labels, enc_mask, dec_mask, enc_dec_mask
+
+ def get_batch_transformer(data_iterator):
+ """Build the batch."""
+ data = next(data_iterator)
+ data = send_to_device(data, torch.cuda.current_device())
+
+ tokens_enc = data["input_ids"].long()
+ labels = data["labels"].long()
+ loss_mask = (labels != -100).to(torch.float)
+ if "decoder_input_ids" in data:
+ tokens_dec = data["decoder_input_ids"].long()
+ else:
+ tokens_dec = labels.new_zeros(labels.shape, device=labels.device, dtype=torch.long)
+ tokens_dec[..., 1:] = labels[..., :-1].clone()
+ tokens_dec[..., 0] = 0
+ tokens_dec.masked_fill_(tokens_dec == -100, 0)
+ enc_mask = T5TrainStep.attn_mask_postprocess(data["attention_mask"].long())
+ dec_mask = T5TrainStep.get_decoder_mask(tokens_dec.shape[1], tokens_dec.device)
+ enc_dec_mask = T5TrainStep.get_enc_dec_mask(
+ data["attention_mask"].long(), tokens_dec.shape[1], tokens_dec.device
+ )
+
+ return tokens_enc, tokens_dec, loss_mask, labels, enc_mask, dec_mask, enc_dec_mask
+
+ if accelerator.state.megatron_lm_plugin.custom_get_batch_function is not None:
+ return accelerator.state.megatron_lm_plugin.custom_get_batch_function
+ if megatron_dataset_flag:
+ try:
+ # Use '--no-use-pep517 -e' to pip install nvidia's megatron from source
+ from pretrain_t5 import get_batch
+
+ return get_batch
+ except ImportError:
+ pass
+ return get_batch_megatron
+ else:
+ return get_batch_transformer
+
+ def get_loss_func(self, accelerator):
+ def loss_func(loss_mask, output_tensor):
+ lm_loss_ = output_tensor.float()
+ lm_loss = torch.sum(lm_loss_.view(-1) * loss_mask.reshape(-1)) / loss_mask.sum()
+
+ loss = lm_loss
+ averaged_losses = average_losses_across_data_parallel_group([lm_loss])
+
+ return loss, {"lm loss": averaged_losses[0]}
+
+ if accelerator.state.megatron_lm_plugin.custom_loss_function is not None:
+ return accelerator.state.megatron_lm_plugin.custom_loss_function
+ return loss_func
+
+ def get_forward_step_func(self):
+ def forward_step(data_iterator, model):
+ """Forward step."""
+ # Get the batch.
+ tokens_enc, tokens_dec, loss_mask, lm_labels, enc_mask, dec_mask, enc_dec_mask = self.get_batch(
+ data_iterator
+ )
+ # Forward model lm_labels
+ output_tensor = model(
+ tokens_enc, tokens_dec, enc_mask, dec_mask, enc_dec_mask, tokentype_ids=None, lm_labels=lm_labels
+ )
+
+ return output_tensor, partial(self.loss_func, loss_mask)
+
+ return forward_step
+
+
+def finish_mpu_init():
+ # torch.distributed initialization
+ args = get_args()
+ # Pytorch distributed.
+ _initialize_distributed()
+
+ # Random seeds for reproducibility.
+ if args.rank == 0:
+ print(f"> setting random seeds to {args.seed} ...")
+ _set_random_seed(args.seed, args.data_parallel_random_init)
+
+
+# intialize megatron setup
+def initialize(accelerator, extra_args_provider=None, args_defaults={}):
+ accelerator.print("Initializing Megatron-LM")
+ assert torch.cuda.is_available(), "Megatron requires CUDA."
+
+ # Parse arguments
+ args = parse_args(extra_args_provider, ignore_unknown_args=True)
+
+ # Set defaults
+ for key, value in args_defaults.items():
+ if getattr(args, key, None) is not None:
+ if args.rank == 0:
+ print(
+ f"WARNING: overriding default arguments for {key}:{getattr(args, key)} with {key}:{value}",
+ flush=True,
+ )
+ setattr(args, key, value)
+
+ if args.use_checkpoint_args or args_defaults.get("use_checkpoint_args", False):
+ assert args.load is not None, "--use-checkpoints-args requires --load argument"
+ load_args_from_checkpoint(args)
+
+ validate_args(args)
+
+ # set global args, build tokenizer, and set adlr-autoresume,
+ # tensorboard-writer, and timers.
+ set_global_variables(args)
+
+ # Megatron's MPU is the master. Complete initialization right away.
+ finish_mpu_init()
+
+ # Autoresume.
+ _init_autoresume()
+
+ # Compile dependencies.
+ _compile_dependencies()
+
+ # Set pytorch JIT layer fusion options and warmup JIT functions.
+ set_jit_fusion_options()
+ args = get_args()
+ if getattr(args, "padded_vocab_size", None) is None:
+ args.padded_vocab_size = _vocab_size_with_padding(args.orig_vocab_size, args)
+ if args.model_type_name == "bert" and args.pretraining_flag and args.num_labels == 2:
+ args.bert_binary_head = True
+ else:
+ args.bert_binary_head = False
+ args.iteration = 0
+
+
+class MegatronEngine(torch.nn.Module):
+ """
+ Megatron-LM model wrapper
+
+ Args:
+ accelerator (:class:`~accelerate.Accelerator`): The accelerator object to use.
+ model: Megatron-LM model
+ optimizer: Megatron-LM optimizer
+ lr_scheduler: Megatron-LM lr scheduler
+ """
+
+ def __init__(self, accelerator, model, optimizer, scheduler):
+ super().__init__()
+ self.module = model
+ self.base_model = model[0]
+ self.optimizer = optimizer
+ self.scheduler = scheduler
+ args = get_args()
+ if accelerator.state.megatron_lm_plugin.custom_train_step_class is not None:
+ self.train_step_handler = accelerator.state.megatron_lm_plugin.custom_train_step_class(
+ args, **accelerator.state.megatron_lm_plugin.custom_train_step_kwargs
+ )
+ elif args.model_type_name == "bert":
+ self.train_step_handler = BertTrainStep(accelerator, args)
+ elif args.model_type_name == "gpt":
+ self.train_step_handler = GPTTrainStep(accelerator, args)
+ elif args.model_type_name == "t5":
+ self.train_step_handler = T5TrainStep(accelerator, args)
+ else:
+ raise ValueError(f"Unsupported model type: {args.model_type_name}")
+ self.optimizer.skipped_iter = False
+
+ # Tracking loss.
+ self.total_loss_dict = {}
+ self.eval_total_loss_dict = {}
+ self.iteration = 0
+ self.report_memory_flag = True
+ self.num_floating_point_operations_so_far = 0
+ self.module_config = None
+ if args.tensorboard_dir is not None:
+ write_args_to_tensorboard()
+
+ def get_module_config(self):
+ args = get_args()
+ config = get_model_config(self.module[0])
+ # Setup some training config params
+ config.grad_scale_func = self.optimizer.scale_loss
+ if isinstance(self.module[0], LocalDDP) and args.overlap_grad_reduce:
+ assert config.no_sync_func is None, (
+ "When overlap_grad_reduce is True, config.no_sync_func must be None; "
+ "a custom no_sync_func is not supported when overlapping grad-reduce"
+ )
+ config.no_sync_func = [model_chunk.no_sync for model_chunk in self.module]
+ if len(self.module) == 1:
+ config.no_sync_func = config.no_sync_func[0]
+ if args.delay_grad_reduce:
+ config.grad_sync_func = [model_chunk.start_grad_sync for model_chunk in self.module]
+ if len(self.module) == 1:
+ config.grad_sync_func = config.grad_sync_func[0]
+ if args.overlap_param_gather and args.delay_param_gather:
+ config.param_sync_func = [
+ lambda x: self.optimizer.finish_param_sync(model_index, x) for model_index in range(len(self.module))
+ ]
+ if len(self.module) == 1:
+ config.param_sync_func = config.param_sync_func[0]
+ config.finalize_model_grads_func = finalize_model_grads
+ return config
+
+ def train(self):
+ for model_module in self.module:
+ model_module.train()
+
+ if self.module_config is None:
+ self.module_config = self.get_module_config()
+
+ self.log_eval_results()
+
+ def eval(self):
+ for model_module in self.module:
+ model_module.eval()
+
+ if self.module_config is None:
+ self.module_config = self.get_module_config()
+
+ def get_batch_data_iterator(self, batch_data):
+ args = get_args()
+ data_chunks = []
+ if len(batch_data) > 0:
+ if args.num_micro_batches > 1:
+ for i in range(0, args.num_micro_batches):
+ data_chunks.append(
+ {
+ k: v[i * args.micro_batch_size : (i + 1) * args.micro_batch_size]
+ for k, v in batch_data.items()
+ }
+ )
+ else:
+ data_chunks = [batch_data]
+
+ if len(self.module) > 1:
+ batch_data_iterator = (
+ [iter(data_chunks) for _ in range(len(self.module))]
+ if len(batch_data) > 0
+ else [None] * len(self.module)
+ )
+ else:
+ batch_data_iterator = iter(data_chunks) if len(batch_data) > 0 else None
+ return batch_data_iterator
+
+ def train_step(self, **batch_data):
+ """
+ Training step for Megatron-LM
+
+ Args:
+ batch_data (:obj:`dict`): The batch data to train on.
+ """
+
+ batch_data_iterator = self.get_batch_data_iterator(batch_data)
+
+ loss_reduced, skipped_iter, grad_norm, num_zeros_in_grad = train_step(
+ forward_step_func=self.train_step_handler.forward_step,
+ data_iterator=batch_data_iterator,
+ model=self.module,
+ optimizer=self.optimizer,
+ opt_param_scheduler=self.scheduler,
+ config=self.module_config,
+ )
+
+ self.optimizer.skipped_iter = skipped_iter == 1
+
+ return loss_reduced, skipped_iter, grad_norm, num_zeros_in_grad
+
+ def eval_step(self, **batch_data):
+ """
+ Evaluation step for Megatron-LM
+
+ Args:
+ batch_data (:obj:`dict`): The batch data to evaluate on.
+ """
+
+ args = get_args()
+ batch_data_iterator = self.get_batch_data_iterator(batch_data)
+ forward_backward_func = get_forward_backward_func()
+ loss_dicts = forward_backward_func(
+ forward_step_func=self.train_step_handler.forward_step,
+ data_iterator=batch_data_iterator,
+ model=self.module,
+ num_microbatches=get_num_microbatches(),
+ seq_length=args.seq_length,
+ micro_batch_size=args.micro_batch_size,
+ forward_only=True,
+ )
+ # Empty unused memory
+ if args.empty_unused_memory_level >= 1:
+ torch.cuda.empty_cache()
+
+ args.consumed_valid_samples += (
+ mpu.get_data_parallel_world_size() * args.micro_batch_size * get_num_microbatches()
+ )
+
+ if mpu.is_pipeline_last_stage(ignore_virtual=True):
+ # Average loss across microbatches.
+ loss_reduced = {}
+ for key in loss_dicts[0]:
+ losses_reduced_for_key = [x[key] for x in loss_dicts]
+ if len(losses_reduced_for_key[0].shape) == 0:
+ loss_reduced[key] = sum(losses_reduced_for_key) / len(losses_reduced_for_key)
+ else:
+ loss_reduced[key] = torch.concat(losses_reduced_for_key)
+ return loss_reduced
+ return {}
+
+ def forward(self, **batch_data):
+ # During training, we use train_step()
+ # model(**batch_data) performs following operations by delegating it to `self.train_step`:
+ # 1. Prepare **batch_data for Tendor, Pipeline and Model Parallelism
+ # 2. Set grad to zero.
+ # 3. forward pass and backward pass using Pipeline Parallelism
+ # 4. Empty unused memory.
+ # 5. Reduce gradients.
+ # 6. Update parameters.
+ # 7. Gather params when using Distributed Optimizer (Data Parallelism).
+ # 8. Update learning rate if scheduler is specified.
+ # 9. Empty unused memory.
+ # 10. Average loss across microbatches and across DP ranks.
+ #
+ # During evaluation, we use eval_step()
+ args = get_args()
+ if self.module[0].training:
+ loss_dict, skipped_iter, grad_norm, num_zeros_in_grad = self.train_step(**batch_data)
+ self.iteration += 1
+ batch_size = mpu.get_data_parallel_world_size() * args.micro_batch_size * get_num_microbatches()
+ args.consumed_train_samples += batch_size
+ self.num_floating_point_operations_so_far += num_floating_point_operations(args, batch_size)
+ if args.tensorboard_dir is not None:
+ # Logging.
+ loss_scale = self.optimizer.get_loss_scale().item()
+ params_norm = None
+ if args.log_params_norm:
+ params_norm = calc_params_l2_norm(self.model)
+ self.report_memory_flag = training_log(
+ loss_dict,
+ self.total_loss_dict,
+ self.optimizer.param_groups[0]["lr"],
+ self.iteration,
+ loss_scale,
+ self.report_memory_flag,
+ skipped_iter,
+ grad_norm,
+ params_norm,
+ num_zeros_in_grad,
+ )
+ else:
+ loss_dict = self.eval_step(**batch_data)
+ if args.tensorboard_dir is not None:
+ for key in loss_dict:
+ self.eval_total_loss_dict[key] = (
+ self.eval_total_loss_dict.get(key, torch.cuda.FloatTensor([0.0])) + loss_dict[key]
+ )
+ self.eval_total_loss_dict[key + "_num_iters"] = self.eval_total_loss_dict.get(
+ key + "_num_iters", torch.cuda.FloatTensor([0.0])
+ ) + torch.cuda.FloatTensor([1.0])
+
+ loss = torch.tensor(0.0, device=torch.cuda.current_device())
+ for key in loss_dict:
+ if len(loss_dict[key].shape) == 0:
+ loss += loss_dict[key]
+
+ logits = None
+ if "logits" in loss_dict:
+ logits = loss_dict["logits"]
+ if self.train_step_handler.model_output_class is not None:
+ return self.train_step_handler.model_output_class(loss=loss, logits=logits)
+ return loss
+
+ def log_eval_results(self):
+ args = get_args()
+ if args.tensorboard_dir is None or self.iteration == 0:
+ return
+ args = get_args()
+ writer = get_tensorboard_writer()
+ string = f"validation loss at iteration {self.iteration} | "
+ for key in self.eval_total_loss_dict:
+ if key.endswith("_num_iters"):
+ continue
+ value = self.eval_total_loss_dict[key] / self.eval_total_loss_dict[key + "_num_iters"]
+ string += f"{key} value: {value} | "
+ ppl = math.exp(min(20, value.item()))
+ if args.pretraining_flag:
+ string += f"{key} PPL: {ppl} | "
+ if writer:
+ writer.add_scalar(f"{key} validation", value.item(), self.iteration)
+ if args.pretraining_flag:
+ writer.add_scalar(f"{key} validation ppl", ppl, self.iteration)
+
+ length = len(string) + 1
+ print_rank_last("-" * length)
+ print_rank_last(string)
+ print_rank_last("-" * length)
+ self.eval_total_loss_dict = {}
+
+ def save_checkpoint(self, output_dir):
+ self.log_eval_results()
+ args = get_args()
+ args.save = output_dir
+ torch.distributed.barrier()
+ save_checkpoint(
+ self.iteration,
+ self.module,
+ self.optimizer,
+ self.scheduler,
+ num_floating_point_operations_so_far=self.num_floating_point_operations_so_far,
+ )
+ torch.distributed.barrier()
+
+ def load_checkpoint(self, input_dir):
+ args = get_args()
+ args.load = input_dir
+ args.consumed_train_samples = 0
+ args.consumed_valid_samples = 0
+ torch.distributed.barrier()
+ iteration, num_floating_point_operations_so_far = load_checkpoint(self.module, self.optimizer, self.scheduler)
+ torch.distributed.barrier()
+ self.iteration = iteration
+ self.num_floating_point_operations_so_far = num_floating_point_operations_so_far
+ if args.fp16 and self.iteration == 0:
+ self.optimizer.reload_model_params()
+
+ def megatron_generate(
+ self,
+ inputs,
+ attention_mask=None,
+ max_length=None,
+ max_new_tokens=None,
+ num_beams=None,
+ temperature=None,
+ top_k=None,
+ top_p=None,
+ length_penalty=None,
+ **kwargs,
+ ):
+ """
+ Generate method for GPT2 model. This method is used for inference. Supports both greedy and beam search along
+ with sampling. Refer the Megatron-LM repo for more details
+
+ Args:
+ inputs (torch.Tensor): input ids
+ attention_mask (torch.Tensor, optional): attention mask. Defaults to None.
+ max_length (int, optional): max length of the generated sequence. Defaults to None.
+ Either this or max_new_tokens should be provided.
+ max_new_tokens (int, optional): max number of tokens to be generated. Defaults to None.
+ Either this or max_length should be provided.
+ num_beams (int, optional): number of beams to use for beam search. Defaults to None.
+ temperature (float, optional): temperature for sampling. Defaults to 1.0.
+ top_k (int, optional): top k tokens to consider for sampling. Defaults to 0.0.
+ top_p (float, optional): tokens in top p probability are considered for sampling. Defaults to 0.0.
+ length_penalty (float, optional): length penalty for beam search. Defaults to None.
+ kwargs: additional key-value arguments
+ """
+
+ # checking if required arguments are passed
+ args = get_args()
+ if args.model_type_name != "gpt":
+ raise NotImplementedError("Generate method is not implemented for this model")
+
+ if args.data_parallel_size > 1:
+ raise ValueError("Generate method requires data parallelism to be 1")
+
+ if args.sequence_parallel:
+ raise ValueError("Generate method requires sequence parallelism to be False")
+
+ if args.recompute_granularity is not None:
+ raise ValueError("Checkpoint activations cannot be set for inference")
+
+ if args.vocab_file is None:
+ raise ValueError("Vocab file is required for inference")
+
+ # Prepare inputs
+ if max_length is None and max_new_tokens is None:
+ raise ValueError("`max_length` or `max_new_tokens` are required for inference")
+
+ if temperature is None:
+ temperature = 1.0
+ elif not (0.0 < temperature <= 100.0):
+ raise ValueError("temperature must be a positive number less than or equal to 100.0")
+
+ if top_k is None:
+ top_k = 0
+ elif not (0 <= top_k <= 1000):
+ raise ValueError("top_k must be a positive number less than or equal to 1000")
+
+ if top_p is None:
+ top_p = 0.0
+ elif top_p > 0.0 and top_k > 0.0:
+ raise ValueError("top_p and top_k sampling cannot be set together")
+ else:
+ if not (0.0 <= top_p <= 1.0):
+ raise ValueError("top_p must be less than or equal to 1.0")
+
+ top_p_decay = kwargs.get("top_p_decay", 0.0)
+ if not (0.0 <= top_p_decay <= 1.0):
+ raise ValueError("top_p_decay must be less than or equal to 1.0")
+
+ top_p_bound = kwargs.get("top_p_bound", 0.0)
+ if not (0.0 <= top_p_bound <= 1.0):
+ raise ValueError("top_p_bound must be less than or equal to 1.0")
+
+ add_BOS = kwargs.get("add_BOS", False)
+ if not (isinstance(add_BOS, bool)):
+ raise ValueError("add_BOS must be a boolean")
+
+ beam_width = num_beams
+ if beam_width is not None:
+ if not isinstance(beam_width, int):
+ raise ValueError("beam_width must be an integer")
+ if beam_width < 1:
+ raise ValueError("beam_width must be greater than 0")
+ if inputs.shape[0] > 1:
+ return "When doing beam_search, batch size must be 1"
+
+ tokenizer = get_tokenizer()
+
+ stop_token = kwargs.get("stop_token", tokenizer.eod)
+ if stop_token is not None:
+ if not isinstance(stop_token, int):
+ raise ValueError("stop_token must be an integer")
+
+ if length_penalty is None:
+ length_penalty = 1.0
+
+ sizes_list = None
+ prompts_tokens_tensor = None
+ prompts_length_tensor = None
+ if torch.distributed.get_rank() == 0:
+ # Get the prompts length.
+ if attention_mask is None:
+ prompts_length_tensor = torch.cuda.LongTensor([inputs.shape[1]] * inputs.shape[0])
+ else:
+ prompts_length_tensor = attention_mask.sum(axis=-1).cuda()
+
+ if max_new_tokens is None:
+ max_new_tokens = max_length - inputs.shape[1]
+ if max_new_tokens <= 0:
+ raise ValueError("max_new_tokens must be greater than 0")
+
+ if add_BOS:
+ max_length = max_new_tokens + inputs.shape[1] + 1
+ # making sure that `max_length` is a multiple of 4 to leverage fused kernels
+ max_length = 4 * math.ceil(max_length / 4)
+ max_new_tokens = max_length - (inputs.shape[1] + 1)
+ padding = torch.cuda.LongTensor([[tokenizer.eod] * max_new_tokens] * inputs.shape[0])
+ prompts_tokens_tensor = torch.concat(
+ [torch.unsqueeze(padding[:, 0], axis=-1), inputs.cuda(), padding], axis=-1
+ )
+ else:
+ # making sure that `max_length` is a multiple of 4 to leverage fused kernels
+ max_length = max_new_tokens + inputs.shape[1]
+ max_length = 4 * math.ceil(max_length / 4)
+ max_new_tokens = max_length - inputs.shape[1]
+ padding = torch.cuda.LongTensor([[tokenizer.eod] * max_new_tokens] * inputs.shape[0])
+ prompts_tokens_tensor = torch.concat([inputs.cuda(), padding], axis=-1)
+
+ # We need the sizes of these tensors for the boradcast
+ sizes_list = [
+ prompts_tokens_tensor.size(0), # Batch size
+ prompts_tokens_tensor.size(1),
+ ] # Sequence lenght
+
+ # First, broadcast the sizes.
+ sizes_tensor = broadcast_int_list(2, int_list=sizes_list, rank=0)
+
+ # Now that we have the sizes, we can boradcast the tokens
+ # and length tensors.
+ sizes = sizes_tensor.tolist()
+ context_tokens_tensor = broadcast_tensor(sizes, torch.int64, tensor=prompts_tokens_tensor, rank=0)
+ context_length_tensor = broadcast_tensor(sizes[0], torch.int64, tensor=prompts_length_tensor, rank=0)
+
+ # Run the inference
+ random_seed = kwargs.get("random_seed", 0)
+ torch.random.manual_seed(random_seed)
+ unwrapped_model = unwrap_model(self.base_model, (torchDDP, LocalDDP, Float16Module))
+ if beam_width is not None:
+ tokens, _ = beam_search_and_return_on_first_stage(
+ unwrapped_model,
+ context_tokens_tensor,
+ context_length_tensor,
+ beam_width,
+ stop_token=stop_token,
+ num_return_gen=1,
+ length_penalty=length_penalty,
+ )
+ else:
+ tokens, _, _ = generate_tokens_probs_and_return_on_first_stage(
+ unwrapped_model,
+ context_tokens_tensor,
+ context_length_tensor,
+ return_output_log_probs=False,
+ top_k=top_k,
+ top_p=top_p,
+ top_p_decay=top_p_decay,
+ top_p_bound=top_p_bound,
+ temperature=temperature,
+ use_eod_token_for_early_termination=True,
+ )
+ return tokens
+
+
+# other utilities
+def avg_losses_across_data_parallel_group(losses):
+ """
+ Average losses across data parallel group.
+
+ Args:
+ losses (List[Tensor]): List of losses to average across data parallel group.
+ """
+
+ return average_losses_across_data_parallel_group(losses)
+
+
+def gather_across_data_parallel_groups(tensor):
+ """
+ Recursively gather tensor in a nested list/tuple/dictionary of tensors from data parallel ranks.
+
+ Args:
+ tensor (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to gather across data parallel ranks.
+
+ """
+
+ def _gpu_gather_one(tensor):
+ if tensor.ndim == 0:
+ tensor = tensor.clone()[None]
+ output_tensors = [
+ torch.empty_like(tensor)
+ for _ in range(torch.distributed.get_world_size(group=mpu.get_data_parallel_group()))
+ ]
+ torch.distributed.all_gather(output_tensors, tensor, group=mpu.get_data_parallel_group())
+ return torch.cat(output_tensors, dim=0)
+
+ return recursively_apply(_gpu_gather_one, tensor, error_on_other_type=True)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/memory.py b/venv/lib/python3.11/site-packages/accelerate/utils/memory.py
new file mode 100644
index 0000000000000000000000000000000000000000..b06b027dff57d47a41c8708fd15ec07d572c4228
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/memory.py
@@ -0,0 +1,199 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+"""
+A collection of utilities for ensuring that training can always occur. Heavily influenced by the
+[toma](https://github.com/BlackHC/toma) library.
+"""
+
+import functools
+import gc
+import importlib
+import inspect
+import warnings
+
+import torch
+from packaging import version
+
+from .imports import (
+ is_cuda_available,
+ is_hpu_available,
+ is_ipex_available,
+ is_mlu_available,
+ is_mps_available,
+ is_musa_available,
+ is_npu_available,
+ is_sdaa_available,
+ is_xpu_available,
+)
+from .versions import compare_versions
+
+
+def clear_device_cache(garbage_collection=False):
+ """
+ Clears the device cache by calling `torch.{backend}.empty_cache`. Can also run `gc.collect()`, but do note that
+ this is a *considerable* slowdown and should be used sparingly.
+ """
+ if garbage_collection:
+ gc.collect()
+
+ if is_xpu_available():
+ torch.xpu.empty_cache()
+ elif is_mlu_available():
+ torch.mlu.empty_cache()
+ elif is_sdaa_available():
+ torch.sdaa.empty_cache()
+ elif is_musa_available():
+ torch.musa.empty_cache()
+ elif is_npu_available():
+ torch.npu.empty_cache()
+ elif is_mps_available(min_version="2.0"):
+ torch.mps.empty_cache()
+ elif is_cuda_available():
+ torch.cuda.empty_cache()
+ elif is_hpu_available():
+ # torch.hpu.empty_cache() # not available on hpu as it reserves all device memory for the current process
+ pass
+
+
+def release_memory(*objects):
+ """
+ Releases memory from `objects` by setting them to `None` and calls `gc.collect()` and `torch.cuda.empty_cache()`.
+ Returned objects should be reassigned to the same variables.
+
+ Args:
+ objects (`Iterable`):
+ An iterable of objects
+ Returns:
+ A list of `None` objects to replace `objects`
+
+ Example:
+
+ ```python
+ >>> import torch
+ >>> from accelerate.utils import release_memory
+
+ >>> a = torch.ones(1000, 1000).cuda()
+ >>> b = torch.ones(1000, 1000).cuda()
+ >>> a, b = release_memory(a, b)
+ ```
+ """
+ if not isinstance(objects, list):
+ objects = list(objects)
+ for i in range(len(objects)):
+ objects[i] = None
+ clear_device_cache(garbage_collection=True)
+ return objects
+
+
+def should_reduce_batch_size(exception: Exception) -> bool:
+ """
+ Checks if `exception` relates to CUDA out-of-memory, XPU out-of-memory, CUDNN not supported, or CPU out-of-memory
+
+ Args:
+ exception (`Exception`):
+ An exception
+ """
+ _statements = [
+ " out of memory.", # OOM for CUDA, HIP, XPU
+ "cuDNN error: CUDNN_STATUS_NOT_SUPPORTED.", # CUDNN SNAFU
+ "DefaultCPUAllocator: can't allocate memory", # CPU OOM
+ "FATAL ERROR :: MODULE:PT_DEVMEM Allocation failed", # HPU OOM
+ ]
+ if isinstance(exception, RuntimeError) and len(exception.args) == 1:
+ return any(err in exception.args[0] for err in _statements)
+ return False
+
+
+def find_executable_batch_size(function: callable = None, starting_batch_size: int = 128):
+ """
+ A basic decorator that will try to execute `function`. If it fails from exceptions related to out-of-memory or
+ CUDNN, the batch size is cut in half and passed to `function`
+
+ `function` must take in a `batch_size` parameter as its first argument.
+
+ Args:
+ function (`callable`, *optional*):
+ A function to wrap
+ starting_batch_size (`int`, *optional*):
+ The batch size to try and fit into memory
+
+ Example:
+
+ ```python
+ >>> from accelerate.utils import find_executable_batch_size
+
+
+ >>> @find_executable_batch_size(starting_batch_size=128)
+ ... def train(batch_size, model, optimizer):
+ ... ...
+
+
+ >>> train(model, optimizer)
+ ```
+ """
+ if function is None:
+ return functools.partial(find_executable_batch_size, starting_batch_size=starting_batch_size)
+
+ batch_size = starting_batch_size
+
+ def decorator(*args, **kwargs):
+ nonlocal batch_size
+ clear_device_cache(garbage_collection=True)
+ params = list(inspect.signature(function).parameters.keys())
+ # Guard against user error
+ if len(params) < (len(args) + 1):
+ arg_str = ", ".join([f"{arg}={value}" for arg, value in zip(params[1:], args[1:])])
+ raise TypeError(
+ f"Batch size was passed into `{function.__name__}` as the first argument when called."
+ f"Remove this as the decorator already does so: `{function.__name__}({arg_str})`"
+ )
+ while True:
+ if batch_size == 0:
+ raise RuntimeError("No executable batch size found, reached zero.")
+ try:
+ return function(batch_size, *args, **kwargs)
+ except Exception as e:
+ if should_reduce_batch_size(e):
+ clear_device_cache(garbage_collection=True)
+ batch_size //= 2
+ else:
+ raise
+
+ return decorator
+
+
+def get_xpu_available_memory(device_index: int):
+ if is_ipex_available():
+ ipex_version = version.parse(importlib.metadata.version("intel_extension_for_pytorch"))
+ if compare_versions(ipex_version, ">=", "2.5"):
+ from intel_extension_for_pytorch.xpu import mem_get_info
+
+ return mem_get_info(device_index)[0]
+ elif version.parse(torch.__version__).release >= version.parse("2.6").release:
+ # torch.xpu.mem_get_info API is available starting from PyTorch 2.6
+ # It further requires PyTorch built with the SYCL runtime which supports API
+ # to query available device memory. If not available, exception will be
+ # raised. Version of SYCL runtime used to build PyTorch is being reported
+ # with print(torch.version.xpu) and corresponds to the version of Intel DPC++
+ # SYCL compiler. First version to support required feature is 20250001.
+ try:
+ return torch.xpu.mem_get_info(device_index)[0]
+ except Exception:
+ pass
+
+ warnings.warn(
+ "The XPU `mem_get_info` API is available in IPEX version >=2.5 or PyTorch >=2.6. The current returned available memory is incorrect. Please consider upgrading your IPEX or PyTorch version."
+ )
+ return torch.xpu.max_memory_allocated(device_index)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/modeling.py b/venv/lib/python3.11/site-packages/accelerate/utils/modeling.py
new file mode 100644
index 0000000000000000000000000000000000000000..651c427c1669a5ba0a1c36d88cdeca5be0a4d03f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/modeling.py
@@ -0,0 +1,2140 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 contextlib
+import gc
+import importlib
+import inspect
+import json
+import logging
+import os
+import re
+import shutil
+import tempfile
+import warnings
+from collections import OrderedDict, defaultdict
+from typing import Optional, Union
+
+import torch
+import torch.nn as nn
+
+from ..state import AcceleratorState
+from .constants import SAFE_WEIGHTS_NAME, WEIGHTS_NAME
+from .dataclasses import AutocastKwargs, CustomDtype, DistributedType
+from .imports import (
+ is_hpu_available,
+ is_mlu_available,
+ is_mps_available,
+ is_musa_available,
+ is_npu_available,
+ is_peft_available,
+ is_sdaa_available,
+ is_torch_xla_available,
+ is_xpu_available,
+)
+from .memory import clear_device_cache, get_xpu_available_memory
+from .offload import load_offloaded_weight, offload_weight, save_offload_index
+from .tqdm import is_tqdm_available, tqdm
+from .versions import compare_versions, is_torch_version
+
+
+if is_npu_available(check_device=False):
+ import torch_npu # noqa: F401
+
+if is_mlu_available(check_device=False):
+ import torch_mlu # noqa: F401
+
+if is_sdaa_available(check_device=False):
+ import torch_sdaa # noqa: F401
+
+if is_musa_available(check_device=False):
+ import torch_musa # noqa: F401
+
+from safetensors import safe_open
+from safetensors.torch import load_file as safe_load_file
+
+
+WEIGHTS_INDEX_NAME = "pytorch_model.bin.index.json"
+
+logger = logging.getLogger(__name__)
+
+
+def is_peft_model(model):
+ from .other import extract_model_from_parallel
+
+ if is_peft_available():
+ from peft import PeftModel
+
+ return is_peft_available() and isinstance(extract_model_from_parallel(model), PeftModel)
+
+
+def check_device_same(first_device, second_device):
+ """
+ Utility method to check if two `torch` devices are similar. When dealing with CUDA devices, torch throws `False`
+ for `torch.device("cuda") == torch.device("cuda:0")` whereas they should be the same
+
+ Args:
+ first_device (`torch.device`):
+ First device to check
+ second_device (`torch.device`):
+ Second device to check
+ """
+ if first_device.type != second_device.type:
+ return False
+
+ if first_device.type != "cpu" and first_device.index is None:
+ # In case the first_device is a cuda device and have
+ # the index attribute set to `None`, default it to `0`
+ first_device = torch.device(first_device.type, index=0)
+
+ if second_device.type != "cpu" and second_device.index is None:
+ # In case the second_device is a cuda device and have
+ # the index attribute set to `None`, default it to `0`
+ second_device = torch.device(second_device.type, index=0)
+
+ return first_device == second_device
+
+
+def convert_file_size_to_int(size: Union[int, str]):
+ """
+ Converts a size expressed as a string with digits an unit (like `"5MB"`) to an integer (in bytes).
+
+ Args:
+ size (`int` or `str`): The size to convert. Will be directly returned if an `int`.
+
+ Example:
+
+ ```py
+ >>> convert_file_size_to_int("1MiB")
+ 1048576
+ ```
+ """
+ mem_size = -1
+ err_msg = (
+ f"`size` {size} is not in a valid format. Use an integer for bytes, or a string with an unit (like '5.0GB')."
+ )
+ try:
+ if isinstance(size, int):
+ mem_size = size
+ elif size.upper().endswith("GIB"):
+ mem_size = int(float(size[:-3]) * (2**30))
+ elif size.upper().endswith("MIB"):
+ mem_size = int(float(size[:-3]) * (2**20))
+ elif size.upper().endswith("KIB"):
+ mem_size = int(float(size[:-3]) * (2**10))
+ elif size.upper().endswith("GB"):
+ int_size = int(float(size[:-2]) * (10**9))
+ mem_size = int_size // 8 if size.endswith("b") else int_size
+ elif size.upper().endswith("MB"):
+ int_size = int(float(size[:-2]) * (10**6))
+ mem_size = int_size // 8 if size.endswith("b") else int_size
+ elif size.upper().endswith("KB"):
+ int_size = int(float(size[:-2]) * (10**3))
+ mem_size = int_size // 8 if size.endswith("b") else int_size
+ except ValueError:
+ raise ValueError(err_msg)
+
+ if mem_size < 0:
+ raise ValueError(err_msg)
+ return mem_size
+
+
+def dtype_byte_size(dtype: torch.dtype):
+ """
+ Returns the size (in bytes) occupied by one parameter of type `dtype`.
+
+ Example:
+
+ ```py
+ >>> dtype_byte_size(torch.float32)
+ 4
+ ```
+ """
+ if dtype == torch.bool:
+ return 1 / 8
+ elif dtype == CustomDtype.INT2:
+ return 1 / 4
+ elif dtype == CustomDtype.INT4:
+ return 1 / 2
+ elif dtype == CustomDtype.FP8:
+ return 1
+ elif is_torch_version(">=", "2.1.0") and dtype == torch.float8_e4m3fn:
+ return 1
+ bit_search = re.search(r"[^\d](\d+)$", str(dtype))
+ if bit_search is None:
+ raise ValueError(f"`dtype` is not a valid dtype: {dtype}.")
+ bit_size = int(bit_search.groups()[0])
+ return bit_size // 8
+
+
+def id_tensor_storage(tensor: torch.Tensor) -> tuple[torch.device, int, int]:
+ """
+ Unique identifier to a tensor storage. Multiple different tensors can share the same underlying storage. For
+ example, "meta" tensors all share the same storage, and thus their identifier will all be equal. This identifier is
+ guaranteed to be unique and constant for this tensor's storage during its lifetime. Two tensor storages with
+ non-overlapping lifetimes may have the same id.
+ """
+ _SIZE = {
+ torch.int64: 8,
+ torch.float32: 4,
+ torch.int32: 4,
+ torch.bfloat16: 2,
+ torch.float16: 2,
+ torch.int16: 2,
+ torch.uint8: 1,
+ torch.int8: 1,
+ torch.bool: 1,
+ torch.float64: 8,
+ }
+ try:
+ storage_ptr = tensor.untyped_storage().data_ptr()
+ storage_size = tensor.untyped_storage().nbytes()
+ except Exception:
+ try:
+ # Fallback for torch==1.10
+ storage_ptr = tensor.storage().data_ptr()
+ storage_size = tensor.storage().size() * _SIZE[tensor.dtype]
+ except NotImplementedError:
+ # Fallback for meta storage
+ storage_ptr = 0
+ # On torch >=2.0 this is the tensor size
+ storage_size = tensor.nelement() * _SIZE[tensor.dtype]
+
+ return tensor.device, storage_ptr, storage_size
+
+
+def set_module_tensor_to_device(
+ module: nn.Module,
+ tensor_name: str,
+ device: Union[int, str, torch.device],
+ value: Optional[torch.Tensor] = None,
+ dtype: Optional[Union[str, torch.dtype]] = None,
+ fp16_statistics: Optional[torch.HalfTensor] = None,
+ tied_params_map: Optional[dict[int, dict[torch.device, torch.Tensor]]] = None,
+):
+ """
+ A helper function to set a given tensor (parameter of buffer) of a module on a specific device (note that doing
+ `param.to(device)` creates a new tensor not linked to the parameter, which is why we need this function).
+
+ Args:
+ module (`torch.nn.Module`):
+ The module in which the tensor we want to move lives.
+ tensor_name (`str`):
+ The full name of the parameter/buffer.
+ device (`int`, `str` or `torch.device`):
+ The device on which to set the tensor.
+ value (`torch.Tensor`, *optional*):
+ The value of the tensor (useful when going from the meta device to any other device).
+ dtype (`torch.dtype`, *optional*):
+ If passed along the value of the parameter will be cast to this `dtype`. Otherwise, `value` will be cast to
+ the dtype of the existing parameter in the model.
+ fp16_statistics (`torch.HalfTensor`, *optional*):
+ The list of fp16 statistics to set on the module, used for 8 bit model serialization.
+ tied_params_map (Dict[int, Dict[torch.device, torch.Tensor]], *optional*, defaults to `None`):
+ A map of current data pointers to dictionaries of devices to already dispatched tied weights. For a given
+ execution device, this parameter is useful to reuse the first available pointer of a shared weight on the
+ device for all others, instead of duplicating memory.
+ """
+ # Recurse if needed
+ if "." in tensor_name:
+ splits = tensor_name.split(".")
+ for split in splits[:-1]:
+ new_module = getattr(module, split)
+ if new_module is None:
+ raise ValueError(f"{module} has no attribute {split}.")
+ module = new_module
+ tensor_name = splits[-1]
+
+ if tensor_name not in module._parameters and tensor_name not in module._buffers:
+ raise ValueError(f"{module} does not have a parameter or a buffer named {tensor_name}.")
+ is_buffer = tensor_name in module._buffers
+ old_value = getattr(module, tensor_name)
+
+ # Treat the case where old_value (or a custom `value`, typically offloaded to RAM/disk) belongs to a tied group, and one of the weight
+ # in the tied group has already been dispatched to the device, by avoiding reallocating memory on the device and just copying the pointer.
+ if (
+ value is not None
+ and tied_params_map is not None
+ and value.data_ptr() in tied_params_map
+ and device in tied_params_map[value.data_ptr()]
+ ):
+ module._parameters[tensor_name] = tied_params_map[value.data_ptr()][device]
+ return
+ elif (
+ tied_params_map is not None
+ and old_value.data_ptr() in tied_params_map
+ and device in tied_params_map[old_value.data_ptr()]
+ ):
+ module._parameters[tensor_name] = tied_params_map[old_value.data_ptr()][device]
+ return
+
+ if old_value.device == torch.device("meta") and device not in ["meta", torch.device("meta")] and value is None:
+ raise ValueError(f"{tensor_name} is on the meta device, we need a `value` to put in on {device}.")
+
+ param = module._parameters[tensor_name] if tensor_name in module._parameters else None
+ param_cls = type(param)
+
+ if value is not None:
+ # We can expect mismatches when using bnb 4bit since Params4bit will reshape and pack the weights.
+ # In other cases, we want to make sure we're not loading checkpoints that do not match the config.
+ if old_value.shape != value.shape and param_cls.__name__ != "Params4bit":
+ raise ValueError(
+ f'Trying to set a tensor of shape {value.shape} in "{tensor_name}" (which has shape {old_value.shape}), this looks incorrect.'
+ )
+
+ if dtype is None:
+ # For compatibility with PyTorch load_state_dict which converts state dict dtype to existing dtype in model
+ value = value.to(old_value.dtype)
+ elif not str(value.dtype).startswith(("torch.uint", "torch.int", "torch.bool")):
+ value = value.to(dtype)
+
+ device_quantization = None
+ with torch.no_grad():
+ # leave it on cpu first before moving them to cuda
+ # # fix the case where the device is meta, we don't want to put it on cpu because there is no data =0
+ if (
+ param is not None
+ and param.device.type != "cuda"
+ and torch.device(device).type == "cuda"
+ and param_cls.__name__ in ["Int8Params", "FP4Params", "Params4bit"]
+ ):
+ device_quantization = device
+ device = "cpu"
+ # `torch.Tensor.to()` is not supported by `torch_npu` (see this [issue](https://github.com/Ascend/pytorch/issues/16)).
+ if isinstance(device, int):
+ if is_npu_available():
+ device = f"npu:{device}"
+ elif is_mlu_available():
+ device = f"mlu:{device}"
+ elif is_sdaa_available():
+ device = f"sdaa:{device}"
+ elif is_musa_available():
+ device = f"musa:{device}"
+ elif is_hpu_available():
+ device = "hpu"
+ if "xpu" in str(device) and not is_xpu_available():
+ raise ValueError(f'{device} is not available, you should use device="cpu" instead')
+ if value is None:
+ new_value = old_value.to(device)
+ if dtype is not None and device in ["meta", torch.device("meta")]:
+ if not str(old_value.dtype).startswith(("torch.uint", "torch.int", "torch.bool")):
+ new_value = new_value.to(dtype)
+
+ if not is_buffer:
+ module._parameters[tensor_name] = param_cls(new_value, requires_grad=old_value.requires_grad)
+ elif isinstance(value, torch.Tensor):
+ new_value = value.to(device)
+ else:
+ new_value = torch.tensor(value, device=device)
+ if device_quantization is not None:
+ device = device_quantization
+ if is_buffer:
+ module._buffers[tensor_name] = new_value
+ elif value is not None or not check_device_same(torch.device(device), module._parameters[tensor_name].device):
+ param_cls = type(module._parameters[tensor_name])
+ kwargs = module._parameters[tensor_name].__dict__
+ if param_cls.__name__ in ["Int8Params", "FP4Params", "Params4bit"]:
+ if param_cls.__name__ == "Int8Params" and new_value.dtype == torch.float32:
+ # downcast to fp16 if any - needed for 8bit serialization
+ new_value = new_value.to(torch.float16)
+ # quantize module that are going to stay on the cpu so that we offload quantized weights
+ if device == "cpu" and param_cls.__name__ == "Int8Params":
+ new_value = param_cls(new_value, requires_grad=old_value.requires_grad, **kwargs).to(0).to("cpu")
+ new_value.CB = new_value.CB.to("cpu")
+ new_value.SCB = new_value.SCB.to("cpu")
+ else:
+ new_value = param_cls(new_value, requires_grad=old_value.requires_grad, **kwargs).to(device)
+ elif param_cls.__name__ in ["QTensor", "QBitsTensor"]:
+ new_value = torch.nn.Parameter(new_value, requires_grad=old_value.requires_grad).to(device)
+ elif param_cls.__name__ in ["AffineQuantizedTensor"]:
+ if importlib.util.find_spec("torchao") is not None and compare_versions("torchao", ">=", "0.7.0"):
+ # TorchAO v0.7.0 made layout_tensor an internal private variable and exposed tensor_impl
+ args = (new_value.tensor_impl,)
+ else:
+ args = (new_value.layout_tensor,)
+ args += (
+ new_value.block_size,
+ new_value.shape,
+ new_value.quant_min,
+ new_value.quant_max,
+ new_value.zero_point_domain,
+ )
+ new_value = torch.nn.Parameter(param_cls(*args), requires_grad=old_value.requires_grad).to(device)
+ else:
+ new_value = param_cls(new_value, requires_grad=old_value.requires_grad).to(device)
+
+ module._parameters[tensor_name] = new_value
+ if fp16_statistics is not None:
+ module._parameters[tensor_name].SCB = fp16_statistics.to(device)
+ del fp16_statistics
+ # as we put the weight to meta, it doesn't have SCB attr anymore. make sure that it is not a meta weight
+ if (
+ module.__class__.__name__ == "Linear8bitLt"
+ and getattr(module.weight, "SCB", None) is None
+ and str(module.weight.device) != "meta"
+ ):
+ # quantize only if necessary
+ device_index = torch.device(device).index if torch.device(device).type == "cuda" else None
+ if not getattr(module.weight, "SCB", None) and device_index is not None:
+ if module.bias is not None and module.bias.device.type != "meta":
+ # if a bias exists, we need to wait until the bias is set on the correct device
+ module = module.cuda(device_index)
+ elif module.bias is None:
+ # if no bias exists, we can quantize right away
+ module = module.cuda(device_index)
+ elif (
+ module.__class__.__name__ == "Linear4bit"
+ and getattr(module.weight, "quant_state", None) is None
+ and str(module.weight.device) != "meta"
+ ):
+ # quantize only if necessary
+ device_index = torch.device(device).index if torch.device(device).type == "cuda" else None
+ if not getattr(module.weight, "quant_state", None) and device_index is not None:
+ module.weight = module.weight.cuda(device_index)
+ # clean pre and post foward hook
+ if device != "cpu":
+ clear_device_cache()
+
+ # When handling tied weights, we update tied_params_map to keep track of the tied weights that have already been allocated on the device in
+ # order to avoid duplicating memory, see above.
+ if (
+ tied_params_map is not None
+ and old_value.data_ptr() in tied_params_map
+ and device not in tied_params_map[old_value.data_ptr()]
+ ):
+ tied_params_map[old_value.data_ptr()][device] = new_value
+ elif (
+ value is not None
+ and tied_params_map is not None
+ and value.data_ptr() in tied_params_map
+ and device not in tied_params_map[value.data_ptr()]
+ ):
+ tied_params_map[value.data_ptr()][device] = new_value
+
+
+def named_module_tensors(
+ module: nn.Module, include_buffers: bool = True, recurse: bool = False, remove_non_persistent: bool = False
+):
+ """
+ A helper function that gathers all the tensors (parameters + buffers) of a given module. If `include_buffers=True`
+ it's the same as doing `module.named_parameters(recurse=recurse) + module.named_buffers(recurse=recurse)`.
+
+ Args:
+ module (`torch.nn.Module`):
+ The module we want the tensors on.
+ include_buffer (`bool`, *optional*, defaults to `True`):
+ Whether or not to include the buffers in the result.
+ recurse (`bool`, *optional`, defaults to `False`):
+ Whether or not to go look in every submodule or just return the direct parameters and buffers.
+ remove_non_persistent (`bool`, *optional*, defaults to `False`):
+ Whether or not to remove the non persistent buffer from the buffers. Useful only when include_buffers =
+ True
+ """
+ yield from module.named_parameters(recurse=recurse)
+
+ if include_buffers:
+ non_persistent_buffers = set()
+ if remove_non_persistent:
+ non_persistent_buffers = get_non_persistent_buffers(module, recurse=recurse)
+ for named_buffer in module.named_buffers(recurse=recurse):
+ name, _ = named_buffer
+ if name not in non_persistent_buffers:
+ yield named_buffer
+
+
+def get_non_persistent_buffers(module: nn.Module, recurse: bool = False):
+ """
+ Gather all non persistent buffers of a given modules into a set
+
+ Args:
+ module (`nn.Module`):
+ The module we want the non persistent buffers on.
+ recurse (`bool`, *optional*, defaults to `False`):
+ Whether or not to go look in every submodule or just return the direct non persistent buffers.
+ """
+
+ non_persistent_buffers_set = module._non_persistent_buffers_set
+ if recurse:
+ for _, m in module.named_modules():
+ non_persistent_buffers_set |= m._non_persistent_buffers_set
+
+ return non_persistent_buffers_set
+
+
+class FindTiedParametersResult(list):
+ """
+ This is a subclass of a list to handle backward compatibility for Transformers. Do not rely on the fact this is not
+ a list or on the `values` method as in the future this will be removed.
+ """
+
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+
+ def values(self):
+ warnings.warn(
+ "The 'values' method of FindTiedParametersResult is deprecated and will be removed in Accelerate v1.3.0. ",
+ FutureWarning,
+ )
+ return sum([x[1:] for x in self], [])
+
+
+def check_tied_parameters_in_config(model: nn.Module):
+ """
+ Check if there is any indication in the given model that some weights should be tied.
+
+ Args:
+ model (`torch.nn.Module`): The model to inspect
+
+ Returns:
+ bool: True if the model needs to have tied weights
+ """
+
+ # based on model.tie_weights() method
+ has_tied_word_embedding = False
+ has_tied_encoder_decoder = False
+ has_tied_module = False
+
+ if "PreTrainedModel" in [c.__name__ for c in inspect.getmro(model.__class__)]:
+ has_tied_word_embedding = (
+ hasattr(model, "config")
+ and getattr(model.config, "tie_word_embeddings", False)
+ and model.get_output_embeddings()
+ )
+ has_tied_encoder_decoder = (
+ hasattr(model, "config")
+ and getattr(model.config, "is_encoder_decoder", False)
+ and getattr(model.config, "tie_encoder_decoder", False)
+ )
+ has_tied_module = any(hasattr(module, "_tie_weights") for module in model.modules())
+
+ return any([has_tied_word_embedding, has_tied_encoder_decoder, has_tied_module])
+
+
+def _get_param_device(param, device_map):
+ if param in device_map:
+ return device_map[param]
+ parent_param = ".".join(param.split(".")[:-1])
+ if parent_param == param:
+ raise ValueError(f"The `device_map` does not contain the module {param}.")
+ else:
+ return _get_param_device(parent_param, device_map)
+
+
+def check_tied_parameters_on_same_device(tied_params, device_map):
+ """
+ Check if tied parameters are on the same device
+
+ Args:
+ tied_params (`List[List[str]]`):
+ A list of lists of parameter names being all tied together.
+
+ device_map (`Dict[str, Union[int, str, torch.device]]`):
+ A map that specifies where each submodule should go.
+
+ """
+ for tie_param in tied_params:
+ tie_param_devices = {}
+ for param in tie_param:
+ tie_param_devices[param] = _get_param_device(param, device_map)
+ if len(set(tie_param_devices.values())) > 1:
+ logger.warn(
+ f"Tied parameters are on different devices: {tie_param_devices}. "
+ "Please modify your custom device map or set `device_map='auto'`. "
+ )
+
+
+def find_tied_parameters(model: torch.nn.Module, **kwargs):
+ """
+ Find the tied parameters in a given model.
+
+
+
+ The signature accepts keyword arguments, but they are for the recursive part of this function and you should ignore
+ them.
+
+
+
+ Args:
+ model (`torch.nn.Module`): The model to inspect.
+
+ Returns:
+ List[List[str]]: A list of lists of parameter names being all tied together.
+
+ Example:
+
+ ```py
+ >>> from collections import OrderedDict
+ >>> import torch.nn as nn
+
+ >>> model = nn.Sequential(OrderedDict([("linear1", nn.Linear(4, 4)), ("linear2", nn.Linear(4, 4))]))
+ >>> model.linear2.weight = model.linear1.weight
+ >>> find_tied_parameters(model)
+ [['linear1.weight', 'linear2.weight']]
+ ```
+ """
+
+ # get ALL model parameters and their names
+ all_named_parameters = {name: param for name, param in model.named_parameters(remove_duplicate=False)}
+
+ # get ONLY unique named parameters,
+ # if parameter is tied and have multiple names, it will be included only once
+ no_duplicate_named_parameters = {name: param for name, param in model.named_parameters(remove_duplicate=True)}
+
+ # the difference of the two sets will give us the tied parameters
+ tied_param_names = set(all_named_parameters.keys()) - set(no_duplicate_named_parameters.keys())
+
+ # 'tied_param_names' contains the names of parameters that are tied in the model, but we do not know
+ # which names refer to the same parameter. To identify this, we need to group them together.
+ tied_param_groups = {}
+ for tied_param_name in tied_param_names:
+ tied_param = all_named_parameters[tied_param_name]
+ for param_name, param in no_duplicate_named_parameters.items():
+ # compare if parameters are the same, if so, group their names together
+ if param is tied_param:
+ if param_name not in tied_param_groups:
+ tied_param_groups[param_name] = []
+ tied_param_groups[param_name].append(tied_param_name)
+
+ return FindTiedParametersResult([sorted([weight] + list(set(tied))) for weight, tied in tied_param_groups.items()])
+
+
+def retie_parameters(model, tied_params):
+ """
+ Reties tied parameters in a given model if the link was broken (for instance when adding hooks).
+
+ Args:
+ model (`torch.nn.Module`):
+ The model in which to retie parameters.
+ tied_params (`List[List[str]]`):
+ A mapping parameter name to tied parameter name as obtained by `find_tied_parameters`.
+ """
+ for tied_group in tied_params:
+ param_to_tie = None
+ # two loops : the first one to set param_to_tie , the second one to change the values of tied_group
+ for param_name in tied_group:
+ module = model
+ splits = param_name.split(".")
+ for split in splits[:-1]:
+ module = getattr(module, split)
+ param = getattr(module, splits[-1])
+ if param_to_tie is None and param.device != torch.device("meta"):
+ param_to_tie = param
+ break
+ if param_to_tie is not None:
+ for param_name in tied_group:
+ module = model
+ splits = param_name.split(".")
+ for split in splits[:-1]:
+ module = getattr(module, split)
+ setattr(module, splits[-1], param_to_tie)
+
+
+def _get_proper_dtype(dtype: Union[str, torch.device]) -> torch.dtype:
+ """
+ Just does torch.dtype(dtype) if necessary.
+ """
+ if isinstance(dtype, str):
+ # We accept "torch.float16" or just "float16"
+ dtype = dtype.replace("torch.", "")
+ dtype = getattr(torch, dtype)
+ return dtype
+
+
+def compute_module_sizes(
+ model: nn.Module,
+ dtype: Optional[Union[str, torch.device]] = None,
+ special_dtypes: Optional[dict[str, Union[str, torch.device]]] = None,
+ buffers_only: bool = False,
+):
+ """
+ Compute the size of each submodule of a given model.
+ """
+ if dtype is not None:
+ dtype = _get_proper_dtype(dtype)
+ dtype_size = dtype_byte_size(dtype)
+ if special_dtypes is not None:
+ special_dtypes = {key: _get_proper_dtype(dtyp) for key, dtyp in special_dtypes.items()}
+ special_dtypes_size = {key: dtype_byte_size(dtyp) for key, dtyp in special_dtypes.items()}
+ module_sizes = defaultdict(int)
+
+ module_list = []
+
+ if not buffers_only:
+ module_list = named_module_tensors(model, recurse=True)
+ else:
+ module_list = model.named_buffers(recurse=True)
+
+ for name, tensor in module_list:
+ if special_dtypes is not None and name in special_dtypes:
+ size = tensor.numel() * special_dtypes_size[name]
+ elif dtype is None:
+ size = tensor.numel() * dtype_byte_size(tensor.dtype)
+ elif str(tensor.dtype).startswith(("torch.uint", "torch.int", "torch.bool")):
+ # According to the code in set_module_tensor_to_device, these types won't be converted
+ # so use their original size here
+ size = tensor.numel() * dtype_byte_size(tensor.dtype)
+ else:
+ size = tensor.numel() * min(dtype_size, dtype_byte_size(tensor.dtype))
+ name_parts = name.split(".")
+ for idx in range(len(name_parts) + 1):
+ module_sizes[".".join(name_parts[:idx])] += size
+
+ return module_sizes
+
+
+def compute_module_total_buffer_size(
+ model: nn.Module,
+ dtype: Optional[Union[str, torch.device]] = None,
+ special_dtypes: Optional[dict[str, Union[str, torch.device]]] = None,
+):
+ """
+ Compute the total size of buffers in each submodule of a given model.
+ """
+ module_sizes = compute_module_sizes(model, dtype=dtype, special_dtypes=special_dtypes, buffers_only=True)
+ return module_sizes.get("", 0)
+
+
+def get_max_layer_size(
+ modules: list[tuple[str, torch.nn.Module]], module_sizes: dict[str, int], no_split_module_classes: list[str]
+):
+ """
+ Utility function that will scan a list of named modules and return the maximum size used by one full layer. The
+ definition of a layer being:
+ - a module with no direct children (just parameters and buffers)
+ - a module whose class name is in the list `no_split_module_classes`
+
+ Args:
+ modules (`List[Tuple[str, torch.nn.Module]]`):
+ The list of named modules where we want to determine the maximum layer size.
+ module_sizes (`Dict[str, int]`):
+ A dictionary mapping each layer name to its size (as generated by `compute_module_sizes`).
+ no_split_module_classes (`List[str]`):
+ A list of class names for layers we don't want to be split.
+
+ Returns:
+ `Tuple[int, List[str]]`: The maximum size of a layer with the list of layer names realizing that maximum size.
+ """
+ max_size = 0
+ layer_names = []
+ modules_to_treat = modules.copy()
+ while len(modules_to_treat) > 0:
+ module_name, module = modules_to_treat.pop(0)
+ modules_children = list(module.named_children()) if isinstance(module, torch.nn.Module) else []
+ if len(modules_children) == 0 or module.__class__.__name__ in no_split_module_classes:
+ # No splitting this one so we compare to the max_size
+ size = module_sizes[module_name]
+ if size > max_size:
+ max_size = size
+ layer_names = [module_name]
+ elif size == max_size:
+ layer_names.append(module_name)
+ else:
+ modules_to_treat = [(f"{module_name}.{n}", v) for n, v in modules_children] + modules_to_treat
+ return max_size, layer_names
+
+
+def get_max_memory(max_memory: Optional[dict[Union[int, str], Union[int, str]]] = None):
+ """
+ Get the maximum memory available if nothing is passed, converts string to int otherwise.
+ """
+ import psutil
+
+ if max_memory is None:
+ max_memory = {}
+ # Make sure CUDA is initialized on each GPU to have the right memory info.
+ if is_npu_available():
+ for i in range(torch.npu.device_count()):
+ try:
+ _ = torch.tensor(0, device=torch.device("npu", i))
+ max_memory[i] = torch.npu.mem_get_info(i)[0]
+ except Exception:
+ logger.info(f"Device {i} seems unavailable, Proceeding to check subsequent devices.")
+ continue
+ elif is_mlu_available():
+ for i in range(torch.mlu.device_count()):
+ try:
+ _ = torch.tensor(0, device=torch.device("mlu", i))
+ max_memory[i] = torch.mlu.mem_get_info(i)[0]
+ except Exception:
+ logger.info(f"Device {i} seems unavailable, Proceeding to check subsequent devices.")
+ continue
+ elif is_sdaa_available():
+ for i in range(torch.sdaa.device_count()):
+ try:
+ _ = torch.tensor(0, device=torch.device("sdaa", i))
+ max_memory[i] = torch.sdaa.mem_get_info(i)[0]
+ except Exception:
+ logger.info(f"Device {i} seems unavailable, Proceeding to check subsequent devices.")
+ continue
+ elif is_musa_available():
+ for i in range(torch.musa.device_count()):
+ try:
+ _ = torch.tensor(0, device=torch.device("musa", i))
+ max_memory[i] = torch.musa.mem_get_info(i)[0]
+ except Exception:
+ logger.info(f"Device {i} seems unavailable, Proceeding to check subsequent devices.")
+ continue
+ elif is_xpu_available():
+ for i in range(torch.xpu.device_count()):
+ try:
+ _ = torch.tensor(0, device=torch.device("xpu", i))
+ max_memory[i] = get_xpu_available_memory(i)
+ except Exception:
+ logger.info(f"Device {i} seems unavailable, Proceeding to check subsequent devices.")
+ continue
+ elif is_hpu_available():
+ for i in range(torch.hpu.device_count()):
+ try:
+ _ = torch.tensor(0, device=torch.device("hpu", i))
+ max_memory[i] = torch.hpu.mem_get_info(i)[0]
+ except Exception:
+ logger.info(f"Device {i} seems unavailable, Proceeding to check subsequent devices.")
+ continue
+ else:
+ for i in range(torch.cuda.device_count()):
+ try:
+ _ = torch.tensor([0], device=i)
+ max_memory[i] = torch.cuda.mem_get_info(i)[0]
+ except Exception:
+ logger.info(f"Device {i} seems unavailable, Proceeding to check subsequent devices.")
+ continue
+ # allocate everything in the mps device as the RAM is shared
+ if is_mps_available():
+ max_memory["mps"] = psutil.virtual_memory().available
+ else:
+ max_memory["cpu"] = psutil.virtual_memory().available
+ return max_memory
+
+ for key in max_memory:
+ if isinstance(max_memory[key], str):
+ max_memory[key] = convert_file_size_to_int(max_memory[key])
+
+ # Need to sort the device by type to make sure that we allocate the gpu first.
+ # As gpu/npu/xpu are represented by int, we need to sort them first.
+ gpu_devices = [k for k in max_memory.keys() if isinstance(k, int)]
+ gpu_devices.sort()
+ # check if gpu/npu/xpu devices are available and if not, throw a warning
+ if is_npu_available():
+ num_devices = torch.npu.device_count()
+ elif is_mlu_available():
+ num_devices = torch.mlu.device_count()
+ elif is_sdaa_available():
+ num_devices = torch.sdaa.device_count()
+ elif is_musa_available():
+ num_devices = torch.musa.device_count()
+ elif is_xpu_available():
+ num_devices = torch.xpu.device_count()
+ elif is_hpu_available():
+ num_devices = torch.hpu.device_count()
+ else:
+ num_devices = torch.cuda.device_count()
+ for device in gpu_devices:
+ if device >= num_devices or device < 0:
+ logger.warning(f"Device {device} is not available, available devices are {list(range(num_devices))}")
+ # Add the other devices in the preset order if they are available
+ all_devices = gpu_devices + [k for k in ["mps", "cpu", "disk"] if k in max_memory.keys()]
+ # Raise an error if a device is not recognized
+ for k in max_memory.keys():
+ if k not in all_devices:
+ raise ValueError(
+ f"Device {k} is not recognized, available devices are integers(for GPU/XPU), 'mps', 'cpu' and 'disk'"
+ )
+ max_memory = {k: max_memory[k] for k in all_devices}
+
+ return max_memory
+
+
+def clean_device_map(device_map: dict[str, Union[int, str, torch.device]], module_name: str = ""):
+ """
+ Cleans a device_map by grouping all submodules that go on the same device together.
+ """
+ # Get the value of the current module and if there is only one split across several keys, regroup it.
+ prefix = "" if module_name == "" else f"{module_name}."
+ values = [v for k, v in device_map.items() if k.startswith(prefix)]
+ if len(set(values)) == 1 and len(values) > 1:
+ for k in [k for k in device_map if k.startswith(prefix)]:
+ del device_map[k]
+ device_map[module_name] = values[0]
+
+ # Recurse over the children
+ children_modules = [k for k in device_map.keys() if k.startswith(prefix) and len(k) > len(module_name)]
+ idx = len(module_name.split(".")) + 1 if len(module_name) > 0 else 1
+ children_modules = set(".".join(k.split(".")[:idx]) for k in children_modules)
+ for child in children_modules:
+ clean_device_map(device_map, module_name=child)
+
+ return device_map
+
+
+def load_offloaded_weights(model, index, offload_folder):
+ """
+ Loads the weights from the offload folder into the model.
+
+ Args:
+ model (`torch.nn.Module`):
+ The model to load the weights into.
+ index (`dict`):
+ A dictionary containing the parameter name and its metadata for each parameter that was offloaded from the
+ model.
+ offload_folder (`str`):
+ The folder where the offloaded weights are stored.
+ """
+ if index is None or len(index) == 0:
+ # Nothing to do
+ return
+ for param_name, metadata in index.items():
+ if "SCB" in param_name:
+ continue
+ fp16_statistics = None
+ if "weight" in param_name and param_name.replace("weight", "SCB") in index.keys():
+ weight_name = param_name.replace("weight", "SCB")
+ fp16_statistics = load_offloaded_weight(
+ os.path.join(offload_folder, f"{weight_name}.dat"), index[weight_name]
+ )
+ tensor_file = os.path.join(offload_folder, f"{param_name}.dat")
+ weight = load_offloaded_weight(tensor_file, metadata)
+ set_module_tensor_to_device(model, param_name, "cpu", value=weight, fp16_statistics=fp16_statistics)
+
+
+def get_module_leaves(module_sizes):
+ module_children = {}
+ for module in module_sizes:
+ if module == "" or "." not in module:
+ continue
+ parent = module.rsplit(".", 1)[0]
+ module_children[parent] = module_children.get(parent, 0) + 1
+ leaves = [module for module in module_sizes if module_children.get(module, 0) == 0 and module != ""]
+ return leaves
+
+
+def get_balanced_memory(
+ model: nn.Module,
+ max_memory: Optional[dict[Union[int, str], Union[int, str]]] = None,
+ no_split_module_classes: Optional[list[str]] = None,
+ dtype: Optional[Union[str, torch.dtype]] = None,
+ special_dtypes: Optional[dict[str, Union[str, torch.device]]] = None,
+ low_zero: bool = False,
+):
+ """
+ Compute a `max_memory` dictionary for [`infer_auto_device_map`] that will balance the use of each available GPU.
+
+
+
+ All computation is done analyzing sizes and dtypes of the model parameters. As a result, the model can be on the
+ meta device (as it would if initialized within the `init_empty_weights` context manager).
+
+
+
+ Args:
+ model (`torch.nn.Module`):
+ The model to analyze.
+ max_memory (`Dict`, *optional*):
+ A dictionary device identifier to maximum memory. Will default to the maximum memory available if unset.
+ Example: `max_memory={0: "1GB"}`.
+ no_split_module_classes (`List[str]`, *optional*):
+ A list of layer class names that should never be split across device (for instance any layer that has a
+ residual connection).
+ dtype (`str` or `torch.dtype`, *optional*):
+ If provided, the weights will be converted to that type when loaded.
+ special_dtypes (`Dict[str, Union[str, torch.device]]`, *optional*):
+ If provided, special dtypes to consider for some specific weights (will override dtype used as default for
+ all weights).
+ low_zero (`bool`, *optional*):
+ Minimizes the number of weights on GPU 0, which is convenient when it's used for other operations (like the
+ Transformers generate function).
+ """
+ # Get default / clean up max_memory
+ user_not_set_max_memory = max_memory is None
+ max_memory = get_max_memory(max_memory)
+
+ if is_npu_available():
+ expected_device_type = "npu"
+ elif is_mlu_available():
+ expected_device_type = "mlu"
+ elif is_sdaa_available():
+ expected_device_type = "sdaa"
+ elif is_musa_available():
+ expected_device_type = "musa"
+ elif is_xpu_available():
+ expected_device_type = "xpu"
+ elif is_hpu_available():
+ expected_device_type = "hpu"
+ elif is_mps_available():
+ expected_device_type = "mps"
+ else:
+ expected_device_type = "cuda"
+ num_devices = len([d for d in max_memory if torch.device(d).type == expected_device_type and max_memory[d] > 0])
+
+ if num_devices == 0:
+ return max_memory
+
+ if num_devices == 1:
+ # We cannot do low_zero on just one GPU, but we will still reserve some memory for the buffer
+ low_zero = False
+ # If user just asked us to handle memory usage, we should avoid OOM
+ if user_not_set_max_memory:
+ for key in max_memory.keys():
+ if isinstance(key, int):
+ max_memory[key] *= 0.9 # 90% is a good compromise
+ logger.info(
+ f"We will use 90% of the memory on device {key} for storing the model, and 10% for the buffer to avoid OOM. "
+ "You can set `max_memory` in to a higher value to use more memory (at your own risk)."
+ )
+ break # only one device
+
+ module_sizes = compute_module_sizes(model, dtype=dtype, special_dtypes=special_dtypes)
+ per_gpu = module_sizes[""] // (num_devices - 1 if low_zero else num_devices)
+
+ # We can't just set the memory to model_size // num_devices as it will end being too small: each GPU will get
+ # slightly less layers and some layers will end up offload at the end. So this function computes a buffer size to
+ # add which is the biggest of:
+ # - the size of no split block (if applicable)
+ # - the mean of the layer sizes
+ if no_split_module_classes is None:
+ no_split_module_classes = []
+ elif not isinstance(no_split_module_classes, (list, tuple)):
+ no_split_module_classes = [no_split_module_classes]
+
+ # Identify the size of the no_split_block modules
+ if len(no_split_module_classes) > 0:
+ no_split_children = {}
+ for name, size in module_sizes.items():
+ if name == "":
+ continue
+ submodule = model
+ for submodule_name in name.split("."):
+ submodule = getattr(submodule, submodule_name)
+ class_name = submodule.__class__.__name__
+ if class_name in no_split_module_classes and class_name not in no_split_children:
+ no_split_children[class_name] = size
+
+ if set(no_split_children.keys()) == set(no_split_module_classes):
+ break
+ buffer = max(no_split_children.values()) if len(no_split_children) > 0 else 0
+ else:
+ buffer = 0
+
+ # Compute mean of final modules. In the first dict of module sizes, leaves are the parameters
+ leaves = get_module_leaves(module_sizes)
+ module_sizes = {n: v for n, v in module_sizes.items() if n not in leaves}
+ # Once removed, leaves are the final modules.
+ leaves = get_module_leaves(module_sizes)
+ mean_leaves = int(sum([module_sizes[n] for n in leaves]) / max(len(leaves), 1))
+ buffer = int(1.25 * max(buffer, mean_leaves))
+ per_gpu += buffer
+
+ # Sorted list of GPUs id (we may have some gpu ids not included in the our max_memory list - let's ignore them)
+ gpus_idx_list = list(
+ sorted(
+ device_id for device_id, device_mem in max_memory.items() if isinstance(device_id, int) and device_mem > 0
+ )
+ )
+ # The last device is left with max_memory just in case the buffer is not enough.
+ for idx in gpus_idx_list[:-1]:
+ max_memory[idx] = min(max_memory[0] if low_zero and idx == 0 else per_gpu, max_memory[idx])
+
+ if low_zero:
+ min_zero = max(0, module_sizes[""] - sum([max_memory[i] for i in range(1, num_devices)]))
+ max_memory[0] = min(min_zero, max_memory[0])
+
+ return max_memory
+
+
+def calculate_maximum_sizes(model: torch.nn.Module):
+ "Computes the total size of the model and its largest layer"
+ sizes = compute_module_sizes(model)
+ # `transformers` models store this information for us
+ no_split_modules = getattr(model, "_no_split_modules", None)
+ if no_split_modules is None:
+ no_split_modules = []
+
+ modules_to_treat = (
+ list(model.named_parameters(recurse=False))
+ + list(model.named_children())
+ + list(model.named_buffers(recurse=False))
+ )
+ largest_layer = get_max_layer_size(modules_to_treat, sizes, no_split_modules)
+ total_size = sizes[""]
+ return total_size, largest_layer
+
+
+def _init_infer_auto_device_map(
+ model: nn.Module,
+ max_memory: Optional[dict[Union[int, str], Union[int, str]]] = None,
+ no_split_module_classes: Optional[list[str]] = None,
+ dtype: Optional[Union[str, torch.dtype]] = None,
+ special_dtypes: Optional[dict[str, Union[str, torch.device]]] = None,
+) -> tuple[
+ list[Union[int, str]],
+ dict[Union[int, str], Union[int, str]],
+ list[Union[int, str]],
+ list[int],
+ dict[str, int],
+ list[list[str]],
+ list[str],
+ list[tuple[str, nn.Module]],
+]:
+ """
+ Initialize variables required for computing the device map for model allocation.
+ """
+ max_memory = get_max_memory(max_memory)
+ if no_split_module_classes is None:
+ no_split_module_classes = []
+ elif not isinstance(no_split_module_classes, (list, tuple)):
+ no_split_module_classes = [no_split_module_classes]
+
+ devices = list(max_memory.keys())
+ if "disk" not in devices:
+ devices.append("disk")
+ gpus = [device for device in devices if device not in ["cpu", "disk"]]
+
+ # Devices that need to keep space for a potential offloaded layer.
+ if "mps" in gpus:
+ main_devices = ["mps"]
+ elif len(gpus) > 0:
+ main_devices = [gpus[0], "cpu"]
+ else:
+ main_devices = ["cpu"]
+
+ module_sizes = compute_module_sizes(model, dtype=dtype, special_dtypes=special_dtypes)
+ tied_parameters = find_tied_parameters(model)
+
+ if check_tied_parameters_in_config(model) and len(tied_parameters) == 0:
+ logger.warn(
+ "The model weights are not tied. Please use the `tie_weights` method before using the `infer_auto_device` function."
+ )
+
+ # Direct submodules and parameters
+ modules_to_treat = (
+ list(model.named_parameters(recurse=False))
+ + list(model.named_children())
+ + list(model.named_buffers(recurse=False))
+ )
+
+ return (
+ devices,
+ max_memory,
+ main_devices,
+ gpus,
+ module_sizes,
+ tied_parameters,
+ no_split_module_classes,
+ modules_to_treat,
+ )
+
+
+def get_module_size_with_ties(
+ tied_params,
+ module_size,
+ module_sizes,
+ modules_to_treat,
+) -> tuple[int, list[str], list[nn.Module]]:
+ """
+ Calculate the total size of a module, including its tied parameters.
+
+ Args:
+ tied_params (`List[str]`): The list of tied parameters.
+ module_size (`int`): The size of the module without tied parameters.
+ module_sizes (`Dict[str, int]`): A dictionary mapping each layer name to its size.
+ modules_to_treat (`List[Tuple[str, nn.Module]]`): The list of named modules to treat.
+
+ Returns:
+ `Tuple[int, List[str], List[nn.Module]]`: The total size of the module, the names of the tied modules, and the
+ tied modules.
+ """
+ if len(tied_params) < 1:
+ return module_size, [], []
+ tied_module_names = []
+ tied_modules = []
+
+ for tied_param in tied_params:
+ tied_module_index = [i for i, (n, _) in enumerate(modules_to_treat) if tied_param.startswith(n + ".")][0]
+ tied_module_names.append(modules_to_treat[tied_module_index][0])
+ tied_modules.append(modules_to_treat[tied_module_index][1])
+
+ module_size_with_ties = module_size
+ for tied_param, tied_module_name in zip(tied_params, tied_module_names):
+ module_size_with_ties += module_sizes[tied_module_name] - module_sizes[tied_param]
+
+ return module_size_with_ties, tied_module_names, tied_modules
+
+
+def fallback_allocate(
+ modules: list[tuple[str, nn.Module]],
+ module_sizes: dict[str, int],
+ size_limit: Union[int, str],
+ no_split_module_classes: Optional[list[str]] = None,
+ tied_parameters: Optional[list[list[str]]] = None,
+) -> tuple[Optional[str], Optional[nn.Module], list[tuple[str, nn.Module]]]:
+ """
+ Find a module that fits in the size limit using BFS and return it with its name and the remaining modules.
+
+ Args:
+ modules (`List[Tuple[str, nn.Module]]`):
+ The list of named modules to search in.
+ module_sizes (`Dict[str, int]`):
+ A dictionary mapping each layer name to its size (as generated by `compute_module_sizes`).
+ size_limit (`Union[int, str]`):
+ The maximum size a module can have.
+ no_split_module_classes (`Optional[List[str]]`, *optional*):
+ A list of class names for layers we don't want to be split.
+ tied_parameters (`Optional[List[List[str]]`, *optional*):
+ A list of lists of parameter names being all tied together.
+
+ Returns:
+ `Tuple[Optional[str], Optional[nn.Module], List[Tuple[str, nn.Module]]]`: A tuple containing:
+ - The name of the module that fits within the size limit.
+ - The module itself.
+ - The list of remaining modules after the found module is removed.
+ """
+ try:
+ size_limit = convert_file_size_to_int(size_limit)
+ except ValueError:
+ return None, None, modules
+
+ if no_split_module_classes is None:
+ no_split_module_classes = []
+
+ if tied_parameters is None:
+ tied_parameters = []
+
+ modules_to_search = modules.copy()
+ module_found = False
+
+ while modules_to_search:
+ name, module = modules_to_search.pop(0)
+
+ tied_param_groups = [
+ tied_group
+ for tied_group in tied_parameters
+ if any(name + "." in k + "." for k in tied_group) and not all(name + "." in k + "." for k in tied_group)
+ ]
+
+ tied_params = sum(
+ [[p for p in tied_group if name + "." not in p + "."] for tied_group in tied_param_groups], []
+ )
+
+ module_size_with_ties, _, _ = get_module_size_with_ties(
+ tied_params, module_sizes[name], module_sizes, modules_to_search
+ )
+
+ # If the module fits in the size limit, we found it.
+ if module_size_with_ties <= size_limit:
+ module_found = True
+ break
+
+ # The module is too big, we need to split it if possible.
+ modules_children = (
+ []
+ if isinstance(module, nn.Parameter) or isinstance(module, torch.Tensor)
+ else list(module.named_children())
+ )
+
+ # Split fails, move to the next module
+ if len(modules_children) == 0 or module.__class__.__name__ in no_split_module_classes:
+ continue
+
+ # split is possible, add the children to the list of modules to search
+ modules_children = list(module.named_parameters(recurse=False)) + modules_children
+ modules_to_search = [(f"{name}.{n}", v) for n, v in modules_children] + modules_to_search
+
+ if not module_found:
+ return None, None, modules
+
+ # Prepare the module list for removal of the found module
+ current_names = [n for n, _ in modules]
+ dot_idx = [i for i, c in enumerate(name) if c == "."]
+
+ for dot_index in dot_idx:
+ parent_name = name[:dot_index]
+ if parent_name in current_names:
+ parent_module_idx = current_names.index(parent_name)
+ _, parent_module = modules[parent_module_idx]
+ module_children = list(parent_module.named_parameters(recurse=False)) + list(
+ parent_module.named_children()
+ )
+ modules = (
+ modules[:parent_module_idx]
+ + [(f"{parent_name}.{n}", v) for n, v in module_children]
+ + modules[parent_module_idx + 1 :]
+ )
+ current_names = [n for n, _ in modules]
+
+ # Now the target module should be directly in the list
+ target_idx = current_names.index(name)
+ name, module = modules.pop(target_idx)
+
+ return name, module, modules
+
+
+def infer_auto_device_map(
+ model: nn.Module,
+ max_memory: Optional[dict[Union[int, str], Union[int, str]]] = None,
+ no_split_module_classes: Optional[list[str]] = None,
+ dtype: Optional[Union[str, torch.dtype]] = None,
+ special_dtypes: Optional[dict[str, Union[str, torch.dtype]]] = None,
+ verbose: bool = False,
+ clean_result: bool = True,
+ offload_buffers: bool = False,
+ fallback_allocation: bool = False,
+):
+ """
+ Compute a device map for a given model giving priority to GPUs, then offload on CPU and finally offload to disk,
+ such that:
+ - we don't exceed the memory available of any of the GPU.
+ - if offload to the CPU is needed, there is always room left on GPU 0 to put back the layer offloaded on CPU that
+ has the largest size.
+ - if offload to the CPU is needed,we don't exceed the RAM available on the CPU.
+ - if offload to the disk is needed, there is always room left on the CPU to put back the layer offloaded on disk
+ that has the largest size.
+
+
+
+ All computation is done analyzing sizes and dtypes of the model parameters. As a result, the model can be on the
+ meta device (as it would if initialized within the `init_empty_weights` context manager).
+
+
+
+ Args:
+ model (`torch.nn.Module`):
+ The model to analyze.
+ max_memory (`Dict`, *optional*):
+ A dictionary device identifier to maximum memory. Will default to the maximum memory available if unset.
+ Example: `max_memory={0: "1GB"}`.
+ no_split_module_classes (`List[str]`, *optional*):
+ A list of layer class names that should never be split across device (for instance any layer that has a
+ residual connection).
+ dtype (`str` or `torch.dtype`, *optional*):
+ If provided, the weights will be converted to that type when loaded.
+ special_dtypes (`Dict[str, Union[str, torch.device]]`, *optional*):
+ If provided, special dtypes to consider for some specific weights (will override dtype used as default for
+ all weights).
+ verbose (`bool`, *optional*, defaults to `False`):
+ Whether or not to provide debugging statements as the function builds the device_map.
+ clean_result (`bool`, *optional*, defaults to `True`):
+ Clean the resulting device_map by grouping all submodules that go on the same device together.
+ offload_buffers (`bool`, *optional*, defaults to `False`):
+ In the layers that are offloaded on the CPU or the hard drive, whether or not to offload the buffers as
+ well as the parameters.
+ fallback_allocation (`bool`, *optional*, defaults to `False`):
+ When regular allocation fails, try to allocate a module that fits in the size limit using BFS.
+ """
+
+ # Initialize the variables
+ (
+ devices,
+ max_memory,
+ main_devices,
+ gpus,
+ module_sizes,
+ tied_parameters,
+ no_split_module_classes,
+ modules_to_treat,
+ ) = _init_infer_auto_device_map(model, max_memory, no_split_module_classes, dtype, special_dtypes)
+
+ device_map = OrderedDict()
+ current_device = 0
+ device_memory_used = {device: 0 for device in devices}
+ device_buffer_sizes = {}
+ device_minimum_assignment_memory = {}
+
+ # Initialize maximum largest layer, to know which space to keep in memory
+ max_layer_size, max_layer_names = get_max_layer_size(modules_to_treat, module_sizes, no_split_module_classes)
+
+ # Ready ? This is going to be a bit messy.
+ while len(modules_to_treat) > 0:
+ name, module = modules_to_treat.pop(0)
+ if verbose:
+ print(f"\nTreating module {name}.")
+ # Max size in the remaining layers may have changed since we took one, so we maybe update it.
+ max_layer_names = [n for n in max_layer_names if n != name and not n.startswith(name + ".")]
+ if len(max_layer_names) == 0:
+ max_layer_size, max_layer_names = get_max_layer_size(
+ [(n, m) for n, m in modules_to_treat if isinstance(m, torch.nn.Module)],
+ module_sizes,
+ no_split_module_classes,
+ )
+ # Assess size needed
+ module_size = module_sizes[name]
+
+ # We keep relevant tied parameters only: one of the tied parameters in the group is inside the current module
+ # and the other is not.
+ # Note: If we are currently processing the name `compute.weight`, an other parameter named
+ # e.g. `compute.weight_submodule.parameter`
+ # needs to be considered outside the current module, hence the check with additional dots.
+ tied_param_groups = [
+ tied_group
+ for tied_group in tied_parameters
+ if any(name + "." in k + "." for k in tied_group) and not all(name + "." in k + "." for k in tied_group)
+ ]
+
+ if verbose and len(tied_param_groups) > 0:
+ print(f" Found the relevant tied param groups {tied_param_groups}")
+
+ # Then we keep track of all the parameters that are tied to the current module, but not in the current module
+ tied_params = sum(
+ [[p for p in tied_group if name + "." not in p + "."] for tied_group in tied_param_groups], []
+ )
+
+ if verbose and len(tied_params) > 0:
+ print(f" So those parameters need to be taken into account {tied_params}")
+
+ device = devices[current_device]
+ current_max_size = max_memory[device] if device != "disk" else None
+ current_memory_reserved = 0
+ # Reduce max size available by the largest layer.
+ if devices[current_device] in main_devices:
+ current_max_size = current_max_size - max_layer_size
+ current_memory_reserved = max_layer_size
+
+ module_size_with_ties, tied_module_names, tied_modules = get_module_size_with_ties(
+ tied_params, module_size, module_sizes, modules_to_treat
+ )
+
+ # The module and its tied modules fit on the current device.
+ if current_max_size is None or device_memory_used[device] + module_size_with_ties <= current_max_size:
+ if verbose:
+ output = f"Putting {name}"
+
+ if tied_module_names:
+ output += f" and {tied_module_names}"
+ else:
+ output += f" (size={module_size})"
+
+ if current_max_size is not None:
+ output += f" (available={current_max_size - device_memory_used[device]})"
+
+ output += f" on {device}."
+ print(output)
+
+ device_memory_used[device] += module_size_with_ties
+
+ # Assign the primary module to the device.
+ device_map[name] = device
+
+ # Assign tied modules if any.
+ for tied_module_name in tied_module_names:
+ if tied_module_name in [m[0] for m in modules_to_treat]:
+ # Find the index of the tied module in the list
+ tied_module_index = next(i for i, (n, _) in enumerate(modules_to_treat) if n == tied_module_name)
+ # Remove the tied module from the list to prevent reprocessing
+ modules_to_treat.pop(tied_module_index)
+
+ # Assign the tied module to the device
+ device_map[tied_module_name] = device
+
+ # Buffer Handling
+ if not offload_buffers and isinstance(module, nn.Module):
+ # Compute the total buffer size for the module
+ current_buffer_size = compute_module_total_buffer_size(
+ module, dtype=dtype, special_dtypes=special_dtypes
+ )
+ # Update the buffer size on the device
+ device_buffer_sizes[device] = device_buffer_sizes.get(device, 0) + current_buffer_size
+
+ continue
+
+ # The current module itself fits, so we try to split the tied modules.
+ if len(tied_params) > 0 and device_memory_used[device] + module_size <= current_max_size:
+ # can we split one of the tied modules to make it smaller or do we need to go on the next device?
+ if verbose:
+ print(
+ f"Not enough space on {devices[current_device]} to put {name} and {tied_module_names} (space "
+ f"available {current_max_size - device_memory_used[device]}, needed size {module_size_with_ties})."
+ )
+ split_happened = False
+ for tied_module_name, tied_module in zip(tied_module_names, tied_modules):
+ tied_module_children = list(tied_module.named_children())
+ if len(tied_module_children) == 0 or tied_module.__class__.__name__ in no_split_module_classes:
+ # can't break this one.
+ continue
+
+ if verbose:
+ print(f"Splitting {tied_module_name}.")
+ tied_module_children = list(tied_module.named_parameters(recurse=False)) + tied_module_children
+ tied_module_children = [(f"{tied_module_name}.{n}", v) for n, v in tied_module_children]
+ tied_module_index = [i for i, (n, _) in enumerate(modules_to_treat) if n == tied_module_name][0]
+
+ modules_to_treat = (
+ [(name, module)]
+ + modules_to_treat[:tied_module_index]
+ + tied_module_children
+ + modules_to_treat[tied_module_index + 1 :]
+ )
+ # Update the max layer size.
+ max_layer_size, max_layer_names = get_max_layer_size(
+ [(n, m) for n, m in modules_to_treat if isinstance(m, torch.nn.Module)],
+ module_sizes,
+ no_split_module_classes,
+ )
+ split_happened = True
+ break
+
+ if split_happened:
+ continue
+
+ # If the tied module is not split, we go to the next device
+ if verbose:
+ print("None of the tied module can be split, going to the next device.")
+
+ # The current module itself doesn't fit, so we have to split it or go to the next device.
+ if device_memory_used[device] + module_size >= current_max_size:
+ # Split or not split?
+ modules_children = (
+ []
+ if isinstance(module, nn.Parameter) or isinstance(module, torch.Tensor)
+ else list(module.named_children())
+ )
+ if verbose:
+ print(
+ f"Not enough space on {devices[current_device]} to put {name} (space available "
+ f"{current_max_size - device_memory_used[device]}, module size {module_size})."
+ )
+ if len(modules_children) == 0 or module.__class__.__name__ in no_split_module_classes:
+ # -> no split, we go to the next device
+ if verbose:
+ print("This module cannot be split, going to the next device.")
+
+ else:
+ # -> split, we replace the module studied by its children + parameters
+ if verbose:
+ print(f"Splitting {name}.")
+ modules_children = list(module.named_parameters(recurse=False)) + modules_children
+ modules_to_treat = [(f"{name}.{n}", v) for n, v in modules_children] + modules_to_treat
+ # Update the max layer size.
+ max_layer_size, max_layer_names = get_max_layer_size(
+ [(n, m) for n, m in modules_to_treat if isinstance(m, torch.nn.Module)],
+ module_sizes,
+ no_split_module_classes,
+ )
+ continue
+
+ # If no module is assigned to the current device, we attempt to allocate a fallback module
+ # if fallback_allocation is enabled.
+ if device_memory_used[device] == 0 and fallback_allocation and device != "disk":
+ # We try to allocate a module that fits in the size limit using BFS.
+ # Recompute the current max size as we need to consider the current module as well.
+ current_max_size = max_memory[device] - max(max_layer_size, module_size_with_ties)
+
+ fallback_module_name, fallback_module, remaining_modules = fallback_allocate(
+ modules_to_treat,
+ module_sizes,
+ current_max_size - device_memory_used[device],
+ no_split_module_classes,
+ tied_parameters,
+ )
+ # use the next iteration to put the fallback module on the next device to avoid code duplication
+ if fallback_module is not None:
+ modules_to_treat = [(fallback_module_name, fallback_module)] + [(name, module)] + remaining_modules
+ continue
+
+ if device_memory_used[device] == 0:
+ device_minimum_assignment_memory[device] = module_size_with_ties + current_memory_reserved
+
+ # Neither the current module nor any tied modules can be split, so we move to the next device.
+ device_memory_used[device] = device_memory_used[device] + current_memory_reserved
+ current_device += 1
+ modules_to_treat = [(name, module)] + modules_to_treat
+
+ device_memory_used = {device: mem for device, mem in device_memory_used.items() if mem > 0}
+
+ if clean_result:
+ device_map = clean_device_map(device_map)
+
+ non_gpu_buffer_size = device_buffer_sizes.get("cpu", 0) + device_buffer_sizes.get("disk", 0)
+ if non_gpu_buffer_size > 0 and not offload_buffers:
+ is_buffer_fit_any_gpu = False
+ for gpu_device, gpu_max_memory in max_memory.items():
+ if gpu_device == "cpu" or gpu_device == "disk":
+ continue
+
+ if not is_buffer_fit_any_gpu:
+ gpu_memory_used = device_memory_used.get(gpu_device, 0)
+
+ if gpu_max_memory >= non_gpu_buffer_size + gpu_memory_used:
+ is_buffer_fit_any_gpu = True
+
+ if len(gpus) > 0 and not is_buffer_fit_any_gpu:
+ warnings.warn(
+ f"Current model requires {non_gpu_buffer_size} bytes of buffer for offloaded layers, which seems does "
+ f"not fit any GPU's remaining memory. If you are experiencing a OOM later, please consider using "
+ f"offload_buffers=True."
+ )
+
+ if device_minimum_assignment_memory:
+ devices_info = "\n".join(
+ f" - {device}: {mem} bytes required" for device, mem in device_minimum_assignment_memory.items()
+ )
+ logger.info(
+ f"Based on the current allocation process, no modules could be assigned to the following devices due to "
+ f"insufficient memory:\n"
+ f"{devices_info}\n"
+ f"These minimum requirements are specific to this allocation attempt and may vary. Consider increasing "
+ f"the available memory for these devices to at least the specified minimum, or adjusting the model config."
+ )
+ return device_map
+
+
+def check_device_map(model: nn.Module, device_map: dict[str, Union[int, str, torch.device]]):
+ """
+ Checks a device map covers everything in a given model.
+
+ Args:
+ model (`torch.nn.Module`): The model to check the device map against.
+ device_map (`Dict[str, Union[int, str, torch.device]]`): The device map to check.
+ """
+ all_model_tensors = [name for name, _ in model.state_dict().items()]
+ for module_name in device_map.keys():
+ if module_name == "":
+ all_model_tensors.clear()
+ break
+ else:
+ all_model_tensors = [
+ name
+ for name in all_model_tensors
+ if not name == module_name and not name.startswith(module_name + ".")
+ ]
+ if len(all_model_tensors) > 0:
+ non_covered_params = ", ".join(all_model_tensors)
+ raise ValueError(
+ f"The device_map provided does not give any device for the following parameters: {non_covered_params}"
+ )
+
+
+def load_state_dict(checkpoint_file, device_map=None):
+ """
+ Load a checkpoint from a given file. If the checkpoint is in the safetensors format and a device map is passed, the
+ weights can be fast-loaded directly on the GPU.
+
+ Args:
+ checkpoint_file (`str`): The path to the checkpoint to load.
+ device_map (`Dict[str, Union[int, str, torch.device]]`, *optional*):
+ A map that specifies where each submodule should go. It doesn't need to be refined to each parameter/buffer
+ name, once a given module name is inside, every submodule of it will be sent to the same device.
+ """
+ if checkpoint_file.endswith(".safetensors"):
+ with safe_open(checkpoint_file, framework="pt") as f:
+ metadata = f.metadata()
+ weight_names = f.keys()
+
+ if metadata is None:
+ logger.warn(
+ f"The safetensors archive passed at {checkpoint_file} does not contain metadata. "
+ "Make sure to save your model with the `save_pretrained` method. Defaulting to 'pt' metadata."
+ )
+ metadata = {"format": "pt"}
+
+ if metadata.get("format") not in ["pt", "tf", "flax"]:
+ raise OSError(
+ f"The safetensors archive passed at {checkpoint_file} does not contain the valid metadata. Make sure "
+ "you save your model with the `save_pretrained` method."
+ )
+ elif metadata["format"] != "pt":
+ raise ValueError(f"The checkpoint passed was saved with {metadata['format']}, we need a the pt format.")
+ if device_map is None:
+ return safe_load_file(checkpoint_file)
+ else:
+ # if we only have one device we can load everything directly
+ if len(set(device_map.values())) == 1:
+ device = list(device_map.values())[0]
+ target_device = device
+ if isinstance(device, int):
+ if is_npu_available():
+ target_device = f"npu:{device}"
+ elif is_hpu_available():
+ target_device = "hpu"
+
+ return safe_load_file(checkpoint_file, device=target_device)
+
+ devices = list(set(device_map.values()) - {"disk"})
+ # cpu device should always exist as fallback option
+ if "cpu" not in devices:
+ devices.append("cpu")
+
+ # For each device, get the weights that go there
+ device_weights = {device: [] for device in devices}
+ for module_name, device in device_map.items():
+ if device in devices:
+ device_weights[device].extend(
+ [k for k in weight_names if k == module_name or k.startswith(module_name + ".")]
+ )
+
+ # all weights that haven't defined a device should be loaded on CPU
+ device_weights["cpu"].extend([k for k in weight_names if k not in sum(device_weights.values(), [])])
+ tensors = {}
+ if is_tqdm_available():
+ progress_bar = tqdm(
+ main_process_only=False,
+ total=sum([len(device_weights[device]) for device in devices]),
+ unit="w",
+ smoothing=0,
+ leave=False,
+ )
+ else:
+ progress_bar = None
+ for device in devices:
+ target_device = device
+ if isinstance(device, int):
+ if is_npu_available():
+ target_device = f"npu:{device}"
+ elif is_hpu_available():
+ target_device = "hpu"
+
+ with safe_open(checkpoint_file, framework="pt", device=target_device) as f:
+ for key in device_weights[device]:
+ if progress_bar is not None:
+ progress_bar.set_postfix(dev=device, refresh=False)
+ progress_bar.set_description(key)
+ tensors[key] = f.get_tensor(key)
+ if progress_bar is not None:
+ progress_bar.update()
+ if progress_bar is not None:
+ progress_bar.close()
+
+ return tensors
+ else:
+ return torch.load(checkpoint_file, map_location=torch.device("cpu"))
+
+
+def get_state_dict_offloaded_model(model: nn.Module):
+ """
+ Returns the state dictionary for an offloaded model via iterative onloading
+
+ Args:
+ model (`torch.nn.Module`):
+ The offloaded model we want to save
+ """
+
+ state_dict = {}
+ placeholders = set()
+ for name, module in model.named_modules():
+ if name == "":
+ continue
+
+ try:
+ with align_module_device(module, "cpu"):
+ module_state_dict = module.state_dict()
+ except MemoryError:
+ raise MemoryError("Offloaded module must fit in CPU memory to call save_model!") from None
+
+ for key in module_state_dict:
+ # ignore placeholder parameters that are still on the meta device
+ if module_state_dict[key].device == torch.device("meta"):
+ placeholders.add(name + f".{key}")
+ continue
+ params = module_state_dict[key]
+ state_dict[name + f".{key}"] = params.to("cpu") # move buffers to cpu
+ for key in placeholders.copy():
+ if key in state_dict:
+ placeholders.remove(key)
+ if placeholders:
+ logger.warning(f"The following tensors were not saved because they were still on meta device: {placeholders}")
+
+ return state_dict
+
+
+def get_state_dict_from_offload(
+ module: nn.Module,
+ module_name: str,
+ state_dict: dict[str, Union[str, torch.tensor]],
+ device_to_put_offload: Union[int, str, torch.device] = "cpu",
+):
+ """
+ Retrieve the state dictionary (with parameters) from an offloaded module and load into a specified device (defaults
+ to cpu).
+
+ Args:
+ module: (`torch.nn.Module`):
+ The module we want to retrieve a state dictionary from
+ module_name: (`str`):
+ The name of the module of interest
+ state_dict (`Dict[str, Union[int, str, torch.device]]`):
+ Dictionary of {module names: parameters}
+ device_to_put_offload (`Union[int, str, torch.device]`):
+ Device to load offloaded parameters into, defaults to the cpu.
+ """
+
+ root = module_name[: module_name.rfind(".")] # module name without .weight or .bias
+
+ # do not move parameters if the module is not offloaded
+ if not has_offloaded_params(module):
+ device_to_put_offload = None
+
+ # assign the device to which the offloaded parameters will be sent
+ with align_module_device(module, device_to_put_offload):
+ for m_key, params in module.state_dict().items():
+ if (root + f".{m_key}") in state_dict:
+ state_dict[root + f".{m_key}"] = params
+
+ return state_dict
+
+
+def load_checkpoint_in_model(
+ model: nn.Module,
+ checkpoint: Union[str, os.PathLike],
+ device_map: Optional[dict[str, Union[int, str, torch.device]]] = None,
+ offload_folder: Optional[Union[str, os.PathLike]] = None,
+ dtype: Optional[Union[str, torch.dtype]] = None,
+ offload_state_dict: bool = False,
+ offload_buffers: bool = False,
+ keep_in_fp32_modules: list[str] = None,
+ offload_8bit_bnb: bool = False,
+ strict: bool = False,
+):
+ """
+ Loads a (potentially sharded) checkpoint inside a model, potentially sending weights to a given device as they are
+ loaded.
+
+
+
+ Once loaded across devices, you still need to call [`dispatch_model`] on your model to make it able to run. To
+ group the checkpoint loading and dispatch in one single call, use [`load_checkpoint_and_dispatch`].
+
+
+
+ Args:
+ model (`torch.nn.Module`):
+ The model in which we want to load a checkpoint.
+ checkpoint (`str` or `os.PathLike`):
+ The folder checkpoint to load. It can be:
+ - a path to a file containing a whole model state dict
+ - a path to a `.json` file containing the index to a sharded checkpoint
+ - a path to a folder containing a unique `.index.json` file and the shards of a checkpoint.
+ - a path to a folder containing a unique pytorch_model.bin or a model.safetensors file.
+ device_map (`Dict[str, Union[int, str, torch.device]]`, *optional*):
+ A map that specifies where each submodule should go. It doesn't need to be refined to each parameter/buffer
+ name, once a given module name is inside, every submodule of it will be sent to the same device.
+ offload_folder (`str` or `os.PathLike`, *optional*):
+ If the `device_map` contains any value `"disk"`, the folder where we will offload weights.
+ dtype (`str` or `torch.dtype`, *optional*):
+ If provided, the weights will be converted to that type when loaded.
+ offload_state_dict (`bool`, *optional*, defaults to `False`):
+ If `True`, will temporarily offload the CPU state dict on the hard drive to avoid getting out of CPU RAM if
+ the weight of the CPU state dict + the biggest shard does not fit.
+ offload_buffers (`bool`, *optional*, defaults to `False`):
+ Whether or not to include the buffers in the weights offloaded to disk.
+ keep_in_fp32_modules(`List[str]`, *optional*):
+ A list of the modules that we keep in `torch.float32` dtype.
+ offload_8bit_bnb (`bool`, *optional*):
+ Whether or not to enable offload of 8-bit modules on cpu/disk.
+ strict (`bool`, *optional*, defaults to `False`):
+ Whether to strictly enforce that the keys in the checkpoint state_dict match the keys of the model's
+ state_dict.
+
+ """
+ if offload_8bit_bnb:
+ from .bnb import quantize_and_offload_8bit
+
+ tied_params = find_tied_parameters(model)
+
+ if check_tied_parameters_in_config(model) and len(tied_params) == 0:
+ logger.warn(
+ "The model weights are not tied. Please use the `tie_weights` method before using the `infer_auto_device` function."
+ )
+ if device_map is not None:
+ check_tied_parameters_on_same_device(tied_params, device_map)
+
+ if offload_folder is None and device_map is not None and "disk" in device_map.values():
+ raise ValueError(
+ "At least one of the model submodule will be offloaded to disk, please pass along an `offload_folder`."
+ )
+ elif offload_folder is not None and device_map is not None and "disk" in device_map.values():
+ os.makedirs(offload_folder, exist_ok=True)
+
+ if isinstance(dtype, str):
+ # We accept "torch.float16" or just "float16"
+ dtype = dtype.replace("torch.", "")
+ dtype = getattr(torch, dtype)
+
+ checkpoint_files = None
+ index_filename = None
+ if os.path.isfile(checkpoint):
+ if str(checkpoint).endswith(".json"):
+ index_filename = checkpoint
+ else:
+ checkpoint_files = [checkpoint]
+ elif os.path.isdir(checkpoint):
+ # check if the whole state dict is present
+ potential_state_bin = [f for f in os.listdir(checkpoint) if f == WEIGHTS_NAME]
+ potential_state_safetensor = [f for f in os.listdir(checkpoint) if f == SAFE_WEIGHTS_NAME]
+ if len(potential_state_bin) == 1:
+ checkpoint_files = [os.path.join(checkpoint, potential_state_bin[0])]
+ elif len(potential_state_safetensor) == 1:
+ checkpoint_files = [os.path.join(checkpoint, potential_state_safetensor[0])]
+ else:
+ # otherwise check for sharded checkpoints
+ potential_index = [f for f in os.listdir(checkpoint) if f.endswith(".index.json")]
+ if len(potential_index) == 0:
+ raise ValueError(
+ f"{checkpoint} is not a folder containing a `.index.json` file or a {WEIGHTS_NAME} or a {SAFE_WEIGHTS_NAME} file"
+ )
+ elif len(potential_index) == 1:
+ index_filename = os.path.join(checkpoint, potential_index[0])
+ else:
+ raise ValueError(
+ f"{checkpoint} containing more than one `.index.json` file, delete the irrelevant ones."
+ )
+ else:
+ raise ValueError(
+ "`checkpoint` should be the path to a file containing a whole state dict, or the index of a sharded "
+ f"checkpoint, or a folder containing a sharded checkpoint or the whole state dict, but got {checkpoint}."
+ )
+
+ if index_filename is not None:
+ checkpoint_folder = os.path.split(index_filename)[0]
+ with open(index_filename) as f:
+ index = json.loads(f.read())
+
+ if "weight_map" in index:
+ index = index["weight_map"]
+ checkpoint_files = sorted(list(set(index.values())))
+ checkpoint_files = [os.path.join(checkpoint_folder, f) for f in checkpoint_files]
+
+ # Logic for missing/unexepected keys goes here.
+
+ offload_index = {}
+ if offload_state_dict:
+ state_dict_folder = tempfile.mkdtemp()
+ state_dict_index = {}
+
+ unexpected_keys = set()
+ model_keys = set(model.state_dict().keys())
+ buffer_names = [name for name, _ in model.named_buffers()]
+ for checkpoint_file in checkpoint_files:
+ loaded_checkpoint = load_state_dict(checkpoint_file, device_map=device_map)
+ if device_map is None:
+ model.load_state_dict(loaded_checkpoint, strict=strict)
+ unexpected_keys.update(set(loaded_checkpoint.keys()) - model_keys)
+ else:
+ for param_name, param in loaded_checkpoint.items():
+ # skip SCB parameter (for 8-bit serialization)
+ if "SCB" in param_name:
+ continue
+
+ if param_name not in model_keys:
+ unexpected_keys.add(param_name)
+ if not strict:
+ continue # Skip loading this parameter.
+
+ module_name = param_name
+
+ while len(module_name) > 0 and module_name not in device_map:
+ module_name = ".".join(module_name.split(".")[:-1])
+ if module_name == "" and "" not in device_map:
+ # TODO: group all errors and raise at the end.
+ raise ValueError(f"{param_name} doesn't have any device set.")
+ param_device = device_map[module_name]
+ new_dtype = dtype
+ if dtype is not None and torch.is_floating_point(param):
+ if keep_in_fp32_modules is not None and dtype == torch.float16:
+ proceed = False
+ for key in keep_in_fp32_modules:
+ if ((key in param_name) and (key + "." in param_name)) or key == param_name:
+ proceed = True
+ break
+ if proceed:
+ new_dtype = torch.float32
+
+ if "weight" in param_name and param_name.replace("weight", "SCB") in loaded_checkpoint.keys():
+ if param.dtype == torch.int8:
+ fp16_statistics = loaded_checkpoint[param_name.replace("weight", "SCB")]
+ else:
+ fp16_statistics = None
+
+ if param_device == "disk":
+ if offload_buffers or param_name not in buffer_names:
+ if new_dtype is None:
+ new_dtype = param.dtype
+ if offload_8bit_bnb:
+ quantize_and_offload_8bit(
+ model, param, param_name, new_dtype, offload_folder, offload_index, fp16_statistics
+ )
+ continue
+ else:
+ set_module_tensor_to_device(model, param_name, "meta", dtype=new_dtype)
+ offload_weight(param, param_name, offload_folder, index=offload_index)
+ elif param_device == "cpu" and offload_state_dict:
+ if new_dtype is None:
+ new_dtype = param.dtype
+ if offload_8bit_bnb:
+ quantize_and_offload_8bit(
+ model, param, param_name, new_dtype, state_dict_folder, state_dict_index, fp16_statistics
+ )
+ else:
+ set_module_tensor_to_device(model, param_name, "meta", dtype=new_dtype)
+ offload_weight(param, param_name, state_dict_folder, index=state_dict_index)
+ else:
+ set_module_tensor_to_device(
+ model,
+ param_name,
+ param_device,
+ value=param,
+ dtype=new_dtype,
+ fp16_statistics=fp16_statistics,
+ )
+
+ # Force Python to clean up.
+ del loaded_checkpoint
+ gc.collect()
+
+ if not strict and len(unexpected_keys) > 0:
+ logger.warning(
+ f"Some weights of the model checkpoint at {checkpoint} were not used when"
+ f" initializing {model.__class__.__name__}: {unexpected_keys}. This may or may not be an issue - make sure that the checkpoint does not have unnecessary parameters, or that the model definition correctly corresponds to the checkpoint."
+ )
+
+ save_offload_index(offload_index, offload_folder)
+
+ # Load back offloaded state dict on CPU
+ if offload_state_dict:
+ load_offloaded_weights(model, state_dict_index, state_dict_folder)
+ shutil.rmtree(state_dict_folder)
+
+ retie_parameters(model, tied_params)
+
+
+def get_mixed_precision_context_manager(native_amp: bool = False, autocast_kwargs: AutocastKwargs = None):
+ """
+ Return a context manager for autocasting mixed precision
+
+ Args:
+ native_amp (`bool`, *optional*, defaults to False):
+ Whether mixed precision is actually enabled.
+ cache_enabled (`bool`, *optional*, defaults to True):
+ Whether the weight cache inside autocast should be enabled.
+ """
+ state = AcceleratorState()
+ if autocast_kwargs is None:
+ autocast_kwargs = {}
+ else:
+ autocast_kwargs = autocast_kwargs.to_kwargs()
+ if native_amp:
+ device_type = (
+ "cuda"
+ if (state.distributed_type == DistributedType.XLA and is_torch_xla_available(check_is_gpu=True))
+ else state.device.type
+ )
+ if state.mixed_precision == "fp16":
+ return torch.autocast(device_type=device_type, dtype=torch.float16, **autocast_kwargs)
+ elif state.mixed_precision in ["bf16", "fp8"] and state.distributed_type in [
+ DistributedType.NO,
+ DistributedType.MULTI_CPU,
+ DistributedType.MULTI_GPU,
+ DistributedType.MULTI_MLU,
+ DistributedType.MULTI_SDAA,
+ DistributedType.MULTI_MUSA,
+ DistributedType.MULTI_NPU,
+ DistributedType.MULTI_XPU,
+ DistributedType.MULTI_HPU,
+ DistributedType.FSDP,
+ DistributedType.XLA,
+ ]:
+ return torch.autocast(device_type=device_type, dtype=torch.bfloat16, **autocast_kwargs)
+ else:
+ return torch.autocast(device_type=device_type, **autocast_kwargs)
+ else:
+ return contextlib.nullcontext()
+
+
+def get_grad_scaler(distributed_type: DistributedType = None, **kwargs):
+ """
+ A generic helper which will initialize the correct `GradScaler` implementation based on the environment and return
+ it.
+
+ Args:
+ distributed_type (`DistributedType`, *optional*, defaults to None):
+ The type of distributed environment.
+ kwargs:
+ Additional arguments for the utilized `GradScaler` constructor.
+ """
+ if distributed_type == DistributedType.FSDP:
+ from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
+
+ return ShardedGradScaler(**kwargs)
+ if is_torch_xla_available(check_is_gpu=True):
+ import torch_xla.amp as xamp
+
+ return xamp.GradScaler(**kwargs)
+ elif is_mlu_available():
+ return torch.mlu.amp.GradScaler(**kwargs)
+ elif is_sdaa_available():
+ return torch.sdaa.amp.GradScaler(**kwargs)
+ elif is_musa_available():
+ return torch.musa.amp.GradScaler(**kwargs)
+ elif is_npu_available():
+ return torch.npu.amp.GradScaler(**kwargs)
+ elif is_hpu_available():
+ return torch.amp.GradScaler("hpu", **kwargs)
+ elif is_xpu_available():
+ return torch.amp.GradScaler("xpu", **kwargs)
+ else:
+ if is_torch_version(">=", "2.3"):
+ return torch.amp.GradScaler("cuda", **kwargs)
+ else:
+ return torch.cuda.amp.GradScaler(**kwargs)
+
+
+def has_offloaded_params(module: torch.nn.Module) -> bool:
+ """
+ Checks if a module has offloaded parameters by checking if the given module has a AlignDevicesHook attached with
+ offloading enabled
+
+ Args:
+ module (`torch.nn.Module`): The module to check for an offload hook.
+
+ Returns:
+ bool: `True` if the module has an offload hook and offloading is enabled, `False` otherwise.
+ """
+ from ..hooks import AlignDevicesHook # avoid circular import
+
+ return hasattr(module, "_hf_hook") and isinstance(module._hf_hook, AlignDevicesHook) and module._hf_hook.offload
+
+
+@contextlib.contextmanager
+def align_module_device(module: torch.nn.Module, execution_device: Optional[torch.device] = None):
+ """
+ Context manager that moves a module's parameters to the specified execution device.
+
+ Args:
+ module (`torch.nn.Module`):
+ Module with parameters to align.
+ execution_device (`torch.device`, *optional*):
+ If provided, overrides the module's execution device within the context. Otherwise, use hook execution
+ device or pass
+ """
+ if has_offloaded_params(module):
+ if execution_device is not None:
+ original_device = module._hf_hook.execution_device
+ module._hf_hook.execution_device = execution_device
+
+ try:
+ module._hf_hook.pre_forward(module)
+ yield
+ finally:
+ module._hf_hook.post_forward(module, None)
+ if execution_device is not None:
+ module._hf_hook.execution_device = original_device
+
+ elif execution_device is not None:
+ devices = {name: param.device for name, param in module.named_parameters(recurse=False)}
+ try:
+ for name in devices:
+ set_module_tensor_to_device(module, name, execution_device)
+ yield
+ finally:
+ for name, device in devices.items():
+ set_module_tensor_to_device(module, name, device)
+
+ else:
+ yield
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/offload.py b/venv/lib/python3.11/site-packages/accelerate/utils/offload.py
new file mode 100644
index 0000000000000000000000000000000000000000..d8bff7dc6ad41ffc7f14555261d115d51b76ccec
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/offload.py
@@ -0,0 +1,213 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 json
+import os
+from collections.abc import Mapping
+from typing import Optional, Union
+
+import numpy as np
+import torch
+from safetensors import safe_open
+
+
+def offload_weight(weight, weight_name, offload_folder, index=None):
+ dtype = None
+ # Check the string instead of the dtype to be compatible with versions of PyTorch that don't have bfloat16.
+ if str(weight.dtype) == "torch.bfloat16":
+ # Need to reinterpret the underlined data as int16 since NumPy does not handle bfloat16s.
+ weight = weight.view(torch.int16)
+ dtype = "bfloat16"
+ array = weight.cpu().numpy()
+ tensor_file = os.path.join(offload_folder, f"{weight_name}.dat")
+ if index is not None:
+ if dtype is None:
+ dtype = str(array.dtype)
+ index[weight_name] = {"dtype": dtype, "shape": list(array.shape)}
+ if array.ndim == 0:
+ array = array[None]
+ file_array = np.memmap(tensor_file, dtype=array.dtype, mode="w+", shape=array.shape)
+ file_array[:] = array[:]
+ file_array.flush()
+ return index
+
+
+def load_offloaded_weight(weight_file, weight_info):
+ shape = tuple(weight_info["shape"])
+ if shape == ():
+ # NumPy memory-mapped arrays can't have 0 dims so it was saved as 1d tensor
+ shape = (1,)
+
+ dtype = weight_info["dtype"]
+ if dtype == "bfloat16":
+ # NumPy does not support bfloat16 so this was saved as a int16
+ dtype = "int16"
+
+ weight = np.memmap(weight_file, dtype=dtype, shape=shape, mode="r")
+
+ if len(weight_info["shape"]) == 0:
+ weight = weight[0]
+ weight = torch.tensor(weight)
+ if weight_info["dtype"] == "bfloat16":
+ weight = weight.view(torch.bfloat16)
+
+ return weight
+
+
+def save_offload_index(index, offload_folder):
+ if index is None or len(index) == 0:
+ # Nothing to save
+ return
+
+ offload_index_file = os.path.join(offload_folder, "index.json")
+ if os.path.isfile(offload_index_file):
+ with open(offload_index_file, encoding="utf-8") as f:
+ current_index = json.load(f)
+ else:
+ current_index = {}
+ current_index.update(index)
+
+ with open(offload_index_file, "w", encoding="utf-8") as f:
+ json.dump(current_index, f, indent=2)
+
+
+def offload_state_dict(save_dir: Union[str, os.PathLike], state_dict: dict[str, torch.Tensor]):
+ """
+ Offload a state dict in a given folder.
+
+ Args:
+ save_dir (`str` or `os.PathLike`):
+ The directory in which to offload the state dict.
+ state_dict (`Dict[str, torch.Tensor]`):
+ The dictionary of tensors to offload.
+ """
+ os.makedirs(save_dir, exist_ok=True)
+ index = {}
+ for name, parameter in state_dict.items():
+ index = offload_weight(parameter, name, save_dir, index=index)
+
+ # Update index
+ save_offload_index(index, save_dir)
+
+
+class PrefixedDataset(Mapping):
+ """
+ Will access keys in a given dataset by adding a prefix.
+
+ Args:
+ dataset (`Mapping`): Any map with string keys.
+ prefix (`str`): A prefix to add when trying to access any element in the underlying dataset.
+ """
+
+ def __init__(self, dataset: Mapping, prefix: str):
+ self.dataset = dataset
+ self.prefix = prefix
+
+ def __getitem__(self, key):
+ return self.dataset[f"{self.prefix}{key}"]
+
+ def __iter__(self):
+ return iter([key for key in self.dataset if key.startswith(self.prefix)])
+
+ def __len__(self):
+ return len(self.dataset)
+
+
+class OffloadedWeightsLoader(Mapping):
+ """
+ A collection that loads weights stored in a given state dict or memory-mapped on disk.
+
+ Args:
+ state_dict (`Dict[str, torch.Tensor]`, *optional*):
+ A dictionary parameter name to tensor.
+ save_folder (`str` or `os.PathLike`, *optional*):
+ The directory in which the weights are stored (by `offload_state_dict` for instance).
+ index (`Dict`, *optional*):
+ A dictionary from weight name to their information (`dtype`/ `shape` or safetensors filename). Will default
+ to the index saved in `save_folder`.
+ """
+
+ def __init__(
+ self,
+ state_dict: dict[str, torch.Tensor] = None,
+ save_folder: Optional[Union[str, os.PathLike]] = None,
+ index: Mapping = None,
+ device=None,
+ ):
+ if state_dict is None and save_folder is None and index is None:
+ raise ValueError("Need either a `state_dict`, a `save_folder` or an `index` containing offloaded weights.")
+
+ self.state_dict = {} if state_dict is None else state_dict
+ self.save_folder = save_folder
+ if index is None and save_folder is not None:
+ with open(os.path.join(save_folder, "index.json")) as f:
+ index = json.load(f)
+ self.index = {} if index is None else index
+ self.all_keys = list(self.state_dict.keys())
+ self.all_keys.extend([key for key in self.index if key not in self.all_keys])
+ self.device = device
+
+ def __getitem__(self, key: str):
+ # State dict gets priority
+ if key in self.state_dict:
+ return self.state_dict[key]
+ weight_info = self.index[key]
+ if weight_info.get("safetensors_file") is not None:
+ device = "cpu" if self.device is None else self.device
+ tensor = None
+ try:
+ with safe_open(weight_info["safetensors_file"], framework="pt", device=device) as f:
+ tensor = f.get_tensor(weight_info.get("weight_name", key))
+ except TypeError:
+ # if failed to get_tensor on the device, such as bf16 on mps, try to load it on CPU first
+ with safe_open(weight_info["safetensors_file"], framework="pt", device="cpu") as f:
+ tensor = f.get_tensor(weight_info.get("weight_name", key))
+
+ if "dtype" in weight_info:
+ tensor = tensor.to(getattr(torch, weight_info["dtype"]))
+
+ if tensor.device != torch.device(device):
+ tensor = tensor.to(device)
+ return tensor
+
+ weight_file = os.path.join(self.save_folder, f"{key}.dat")
+ return load_offloaded_weight(weight_file, weight_info)
+
+ def __iter__(self):
+ return iter(self.all_keys)
+
+ def __len__(self):
+ return len(self.all_keys)
+
+
+def extract_submodules_state_dict(state_dict: dict[str, torch.Tensor], submodule_names: list[str]):
+ """
+ Extract the sub state-dict corresponding to a list of given submodules.
+
+ Args:
+ state_dict (`Dict[str, torch.Tensor]`): The state dict to extract from.
+ submodule_names (`List[str]`): The list of submodule names we want to extract.
+ """
+ result = {}
+ for module_name in submodule_names:
+ # We want to catch module_name parameter (module_name.xxx) or potentially module_name, but not any of the
+ # submodules that could being like module_name (transformers.h.1 and transformers.h.10 for instance)
+ result.update(
+ {
+ key: param
+ for key, param in state_dict.items()
+ if key == module_name or key.startswith(module_name + ".")
+ }
+ )
+ return result
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/operations.py b/venv/lib/python3.11/site-packages/accelerate/utils/operations.py
new file mode 100644
index 0000000000000000000000000000000000000000..4b402f447c6e7b71fed21b872088a90a6d1ce9b8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/operations.py
@@ -0,0 +1,862 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+"""
+A set of basic tensor ops compatible with tpu, gpu, and multigpu
+"""
+
+import pickle
+import warnings
+from collections.abc import Mapping
+from contextlib import contextmanager, nullcontext
+from functools import update_wrapper, wraps
+from typing import Any
+
+import torch
+
+from ..state import AcceleratorState, PartialState
+from .constants import TORCH_DISTRIBUTED_OPERATION_TYPES
+from .dataclasses import DistributedType, TensorInformation
+from .imports import (
+ is_npu_available,
+ is_torch_distributed_available,
+ is_torch_xla_available,
+)
+
+
+if is_torch_xla_available():
+ import torch_xla.core.xla_model as xm
+
+if is_torch_distributed_available():
+ from torch.distributed import ReduceOp
+
+
+def is_torch_tensor(tensor):
+ return isinstance(tensor, torch.Tensor)
+
+
+def is_torch_xpu_tensor(tensor):
+ return isinstance(
+ tensor,
+ torch.xpu.FloatTensor,
+ torch.xpu.ByteTensor,
+ torch.xpu.IntTensor,
+ torch.xpu.LongTensor,
+ torch.xpu.HalfTensor,
+ torch.xpu.DoubleTensor,
+ torch.xpu.BFloat16Tensor,
+ )
+
+
+def is_tensor_information(tensor_info):
+ return isinstance(tensor_info, TensorInformation)
+
+
+def is_namedtuple(data):
+ """
+ Checks if `data` is a `namedtuple` or not. Can have false positives, but only if a user is trying to mimic a
+ `namedtuple` perfectly.
+ """
+ return isinstance(data, tuple) and hasattr(data, "_asdict") and hasattr(data, "_fields")
+
+
+def honor_type(obj, generator):
+ """
+ Cast a generator to the same type as obj (list, tuple, or namedtuple)
+ """
+ # Some objects may not be able to instantiate from a generator directly
+ if is_namedtuple(obj):
+ return type(obj)(*list(generator))
+ else:
+ return type(obj)(generator)
+
+
+def recursively_apply(func, data, *args, test_type=is_torch_tensor, error_on_other_type=False, **kwargs):
+ """
+ Recursively apply a function on a data structure that is a nested list/tuple/dictionary of a given base type.
+
+ Args:
+ func (`callable`):
+ The function to recursively apply.
+ data (nested list/tuple/dictionary of `main_type`):
+ The data on which to apply `func`
+ *args:
+ Positional arguments that will be passed to `func` when applied on the unpacked data.
+ main_type (`type`, *optional*, defaults to `torch.Tensor`):
+ The base type of the objects to which apply `func`.
+ error_on_other_type (`bool`, *optional*, defaults to `False`):
+ Whether to return an error or not if after unpacking `data`, we get on an object that is not of type
+ `main_type`. If `False`, the function will leave objects of types different than `main_type` unchanged.
+ **kwargs (additional keyword arguments, *optional*):
+ Keyword arguments that will be passed to `func` when applied on the unpacked data.
+
+ Returns:
+ The same data structure as `data` with `func` applied to every object of type `main_type`.
+ """
+ if isinstance(data, (tuple, list)):
+ return honor_type(
+ data,
+ (
+ recursively_apply(
+ func, o, *args, test_type=test_type, error_on_other_type=error_on_other_type, **kwargs
+ )
+ for o in data
+ ),
+ )
+ elif isinstance(data, Mapping):
+ return type(data)(
+ {
+ k: recursively_apply(
+ func, v, *args, test_type=test_type, error_on_other_type=error_on_other_type, **kwargs
+ )
+ for k, v in data.items()
+ }
+ )
+ elif test_type(data):
+ return func(data, *args, **kwargs)
+ elif error_on_other_type:
+ raise TypeError(
+ f"Unsupported types ({type(data)}) passed to `{func.__name__}`. Only nested list/tuple/dicts of "
+ f"objects that are valid for `{test_type.__name__}` should be passed."
+ )
+ return data
+
+
+def send_to_device(tensor, device, non_blocking=False, skip_keys=None):
+ """
+ Recursively sends the elements in a nested list/tuple/dictionary of tensors to a given device.
+
+ Args:
+ tensor (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to send to a given device.
+ device (`torch.device`):
+ The device to send the data to.
+
+ Returns:
+ The same data structure as `tensor` with all tensors sent to the proper device.
+ """
+ if is_torch_tensor(tensor) or hasattr(tensor, "to"):
+ # `torch.Tensor.to("npu")` could not find context when called for the first time (see this [issue](https://gitee.com/ascend/pytorch/issues/I8KECW?from=project-issue)).
+ if device == "npu":
+ device = "npu:0"
+ try:
+ return tensor.to(device, non_blocking=non_blocking)
+ except TypeError: # .to() doesn't accept non_blocking as kwarg
+ return tensor.to(device)
+ except AssertionError as error:
+ # `torch.Tensor.to()` is not supported by `torch_npu` (see this [issue](https://github.com/Ascend/pytorch/issues/16)).
+ # This call is inside the try-block since is_npu_available is not supported by torch.compile.
+ if is_npu_available():
+ if isinstance(device, int):
+ device = f"npu:{device}"
+ else:
+ raise error
+ try:
+ return tensor.to(device, non_blocking=non_blocking)
+ except TypeError: # .to() doesn't accept non_blocking as kwarg
+ return tensor.to(device)
+ elif isinstance(tensor, (tuple, list)):
+ return honor_type(
+ tensor, (send_to_device(t, device, non_blocking=non_blocking, skip_keys=skip_keys) for t in tensor)
+ )
+ elif isinstance(tensor, Mapping):
+ if isinstance(skip_keys, str):
+ skip_keys = [skip_keys]
+ elif skip_keys is None:
+ skip_keys = []
+ return type(tensor)(
+ {
+ k: t if k in skip_keys else send_to_device(t, device, non_blocking=non_blocking, skip_keys=skip_keys)
+ for k, t in tensor.items()
+ }
+ )
+ else:
+ return tensor
+
+
+def get_data_structure(data):
+ """
+ Recursively gathers the information needed to rebuild a nested list/tuple/dictionary of tensors.
+
+ Args:
+ data (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to send to analyze.
+
+ Returns:
+ The same data structure as `data` with [`~utils.TensorInformation`] instead of tensors.
+ """
+
+ def _get_data_structure(tensor):
+ return TensorInformation(shape=tensor.shape, dtype=tensor.dtype)
+
+ return recursively_apply(_get_data_structure, data)
+
+
+def get_shape(data):
+ """
+ Recursively gathers the shape of a nested list/tuple/dictionary of tensors as a list.
+
+ Args:
+ data (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to send to analyze.
+
+ Returns:
+ The same data structure as `data` with lists of tensor shapes instead of tensors.
+ """
+
+ def _get_shape(tensor):
+ return list(tensor.shape)
+
+ return recursively_apply(_get_shape, data)
+
+
+def initialize_tensors(data_structure):
+ """
+ Recursively initializes tensors from a nested list/tuple/dictionary of [`~utils.TensorInformation`].
+
+ Returns:
+ The same data structure as `data` with tensors instead of [`~utils.TensorInformation`].
+ """
+
+ def _initialize_tensor(tensor_info):
+ return torch.empty(*tensor_info.shape, dtype=tensor_info.dtype)
+
+ return recursively_apply(_initialize_tensor, data_structure, test_type=is_tensor_information)
+
+
+def find_batch_size(data):
+ """
+ Recursively finds the batch size in a nested list/tuple/dictionary of lists of tensors.
+
+ Args:
+ data (nested list/tuple/dictionary of `torch.Tensor`): The data from which to find the batch size.
+
+ Returns:
+ `int`: The batch size.
+ """
+ if isinstance(data, (tuple, list, Mapping)) and (len(data) == 0):
+ raise ValueError(f"Cannot find the batch size from empty {type(data)}.")
+
+ if isinstance(data, (tuple, list)):
+ return find_batch_size(data[0])
+ elif isinstance(data, Mapping):
+ for k in data.keys():
+ return find_batch_size(data[k])
+ elif not isinstance(data, torch.Tensor):
+ raise TypeError(f"Can only find the batch size of tensors but got {type(data)}.")
+ return data.shape[0]
+
+
+def ignorant_find_batch_size(data):
+ """
+ Same as [`utils.operations.find_batch_size`] except will ignore if `ValueError` and `TypeErrors` are raised
+
+ Args:
+ data (nested list/tuple/dictionary of `torch.Tensor`): The data from which to find the batch size.
+
+ Returns:
+ `int`: The batch size.
+ """
+ try:
+ return find_batch_size(data)
+ except (ValueError, TypeError):
+ pass
+ return None
+
+
+def listify(data):
+ """
+ Recursively finds tensors in a nested list/tuple/dictionary and converts them to a list of numbers.
+
+ Args:
+ data (nested list/tuple/dictionary of `torch.Tensor`): The data from which to convert to regular numbers.
+
+ Returns:
+ The same data structure as `data` with lists of numbers instead of `torch.Tensor`.
+ """
+
+ def _convert_to_list(tensor):
+ tensor = tensor.detach().cpu()
+ if tensor.dtype == torch.bfloat16:
+ # As of Numpy 1.21.4, NumPy does not support bfloat16 (see
+ # https://github.com/numpy/numpy/blob/a47ecdea856986cd60eabbd53265c2ca5916ad5d/doc/source/user/basics.types.rst ).
+ # Until Numpy adds bfloat16, we must convert float32.
+ tensor = tensor.to(torch.float32)
+ return tensor.tolist()
+
+ return recursively_apply(_convert_to_list, data)
+
+
+def _tpu_gather(tensor):
+ def _tpu_gather_one(tensor):
+ if tensor.ndim == 0:
+ tensor = tensor.clone()[None]
+
+ # Can only gather contiguous tensors
+ if not tensor.is_contiguous():
+ tensor = tensor.contiguous()
+ return xm.all_gather(tensor)
+
+ res = recursively_apply(_tpu_gather_one, tensor, error_on_other_type=True)
+ xm.mark_step()
+ return res
+
+
+def _gpu_gather(tensor):
+ state = PartialState()
+ gather_op = torch.distributed.all_gather_into_tensor
+
+ def _gpu_gather_one(tensor):
+ if tensor.ndim == 0:
+ tensor = tensor.clone()[None]
+
+ # Can only gather contiguous tensors
+ if not tensor.is_contiguous():
+ tensor = tensor.contiguous()
+
+ if state.backend is not None and state.backend != "gloo":
+ # We use `empty` as `all_gather_into_tensor` slightly
+ # differs from `all_gather` for better efficiency,
+ # and we rely on the number of items in the tensor
+ # rather than its direct shape
+ output_tensors = torch.empty(
+ state.num_processes * tensor.numel(),
+ dtype=tensor.dtype,
+ device=state.device,
+ )
+ gather_op(output_tensors, tensor)
+ return output_tensors.view(-1, *tensor.size()[1:])
+ else:
+ # a backend of `None` is always CPU
+ # also gloo does not support `all_gather_into_tensor`,
+ # which will result in a larger memory overhead for the op
+ output_tensors = [torch.empty_like(tensor) for _ in range(state.num_processes)]
+ torch.distributed.all_gather(output_tensors, tensor)
+ return torch.cat(output_tensors, dim=0)
+
+ return recursively_apply(_gpu_gather_one, tensor, error_on_other_type=True)
+
+
+class DistributedOperationException(Exception):
+ """
+ An exception class for distributed operations. Raised if the operation cannot be performed due to the shape of the
+ tensors.
+ """
+
+ pass
+
+
+def verify_operation(function):
+ """
+ Verifies that `tensor` is the same shape across all processes. Only ran if `PartialState().debug` is `True`.
+ """
+
+ @wraps(function)
+ def wrapper(*args, **kwargs):
+ if PartialState().distributed_type == DistributedType.NO or not PartialState().debug:
+ return function(*args, **kwargs)
+ operation = f"{function.__module__}.{function.__name__}"
+ if "tensor" in kwargs:
+ tensor = kwargs["tensor"]
+ else:
+ tensor = args[0]
+ if PartialState().device.type != find_device(tensor).type:
+ raise DistributedOperationException(
+ f"One or more of the tensors passed to {operation} were not on the {tensor.device.type} while the `Accelerator` is configured for {PartialState().device.type}. "
+ f"Please move it to the {PartialState().device.type} before calling {operation}."
+ )
+ shapes = get_shape(tensor)
+ output = gather_object([shapes])
+ if output[0] is not None:
+ are_same = output.count(output[0]) == len(output)
+ if not are_same:
+ process_shape_str = "\n - ".join([f"Process {i}: {shape}" for i, shape in enumerate(output)])
+ raise DistributedOperationException(
+ f"Cannot apply desired operation due to shape mismatches. "
+ "All shapes across devices must be valid."
+ f"\n\nOperation: `{operation}`\nInput shapes:\n - {process_shape_str}"
+ )
+ return function(*args, **kwargs)
+
+ return wrapper
+
+
+def chained_operation(function):
+ """
+ Checks that `verify_operation` failed and if so reports a more helpful error chaining the existing
+ `DistributedOperationException`.
+ """
+
+ @wraps(function)
+ def wrapper(*args, **kwargs):
+ try:
+ return function(*args, **kwargs)
+ except DistributedOperationException as e:
+ operation = f"{function.__module__}.{function.__name__}"
+ raise DistributedOperationException(
+ f"Error found while calling `{operation}`. Please see the earlier error for more details."
+ ) from e
+
+ return wrapper
+
+
+@verify_operation
+def gather(tensor):
+ """
+ Recursively gather tensor in a nested list/tuple/dictionary of tensors from all devices.
+
+ Args:
+ tensor (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to gather.
+
+ Returns:
+ The same data structure as `tensor` with all tensors sent to the proper device.
+ """
+ if PartialState().distributed_type == DistributedType.XLA:
+ return _tpu_gather(tensor)
+ elif PartialState().distributed_type in TORCH_DISTRIBUTED_OPERATION_TYPES:
+ return _gpu_gather(tensor)
+ else:
+ return tensor
+
+
+def _gpu_gather_object(object: Any):
+ output_objects = [None for _ in range(PartialState().num_processes)]
+ torch.distributed.all_gather_object(output_objects, object)
+ # all_gather_object returns a list of lists, so we need to flatten it
+ return [x for y in output_objects for x in y]
+
+
+def gather_object(object: Any):
+ """
+ Recursively gather object in a nested list/tuple/dictionary of objects from all devices.
+
+ Args:
+ object (nested list/tuple/dictionary of picklable object):
+ The data to gather.
+
+ Returns:
+ The same data structure as `object` with all the objects sent to every device.
+ """
+ if PartialState().distributed_type == DistributedType.XLA:
+ raise NotImplementedError("gather objects in TPU is not supported")
+ elif PartialState().distributed_type in TORCH_DISTRIBUTED_OPERATION_TYPES:
+ return _gpu_gather_object(object)
+ else:
+ return object
+
+
+def _gpu_broadcast(data, src=0):
+ def _gpu_broadcast_one(tensor, src=0):
+ torch.distributed.broadcast(tensor, src=src)
+ return tensor
+
+ return recursively_apply(_gpu_broadcast_one, data, error_on_other_type=True, src=src)
+
+
+def _tpu_broadcast(tensor, src=0, name="broadcast tensor"):
+ if isinstance(tensor, (list, tuple)):
+ return honor_type(tensor, (_tpu_broadcast(t, name=f"{name}_{i}") for i, t in enumerate(tensor)))
+ elif isinstance(tensor, Mapping):
+ return type(tensor)({k: _tpu_broadcast(v, name=f"{name}_{k}") for k, v in tensor.items()})
+ return xm.mesh_reduce(name, tensor, lambda x: x[src])
+
+
+TENSOR_TYPE_TO_INT = {
+ torch.float: 1,
+ torch.double: 2,
+ torch.half: 3,
+ torch.bfloat16: 4,
+ torch.uint8: 5,
+ torch.int8: 6,
+ torch.int16: 7,
+ torch.int32: 8,
+ torch.int64: 9,
+ torch.bool: 10,
+}
+
+TENSOR_INT_TO_DTYPE = {v: k for k, v in TENSOR_TYPE_TO_INT.items()}
+
+
+def gather_tensor_shape(tensor):
+ """
+ Grabs the shape of `tensor` only available on one process and returns a tensor of its shape
+ """
+ # Allocate 80 bytes to store the shape
+ max_tensor_dimension = 2**20
+ state = PartialState()
+ base_tensor = torch.empty(max_tensor_dimension, dtype=torch.int, device=state.device)
+
+ # Since PyTorch can't just send a tensor to another GPU without
+ # knowing its size, we store the size of the tensor with data
+ # in an allocation
+ if tensor is not None:
+ shape = tensor.shape
+ tensor_dtype = TENSOR_TYPE_TO_INT[tensor.dtype]
+ base_tensor[: len(shape) + 1] = torch.tensor(list(shape) + [tensor_dtype], dtype=int)
+ # Perform a reduction to copy the size data onto all GPUs
+ base_tensor = reduce(base_tensor, reduction="sum")
+ base_tensor = base_tensor[base_tensor.nonzero()]
+ # The last non-zero data contains the coded dtype the source tensor is
+ dtype = int(base_tensor[-1:][0])
+ base_tensor = base_tensor[:-1]
+ return base_tensor, dtype
+
+
+def copy_tensor_to_devices(tensor=None) -> torch.Tensor:
+ """
+ Copys a tensor that only exists on a single device and broadcasts it to other devices. Differs from `broadcast` as
+ each worker doesn't need to know its shape when used (and tensor can be `None`)
+
+ Args:
+ tensor (`torch.tensor`):
+ The tensor that should be sent to all devices. Must only have it be defined on a single device, the rest
+ should be `None`.
+ """
+ state = PartialState()
+ shape, dtype = gather_tensor_shape(tensor)
+ if tensor is None:
+ tensor = torch.zeros(shape, dtype=TENSOR_INT_TO_DTYPE[dtype]).to(state.device)
+ return reduce(tensor, reduction="sum")
+
+
+@verify_operation
+def broadcast(tensor, from_process: int = 0):
+ """
+ Recursively broadcast tensor in a nested list/tuple/dictionary of tensors to all devices.
+
+ Args:
+ tensor (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to gather.
+ from_process (`int`, *optional*, defaults to 0):
+ The process from which to send the data
+
+ Returns:
+ The same data structure as `tensor` with all tensors broadcasted to the proper device.
+ """
+ if PartialState().distributed_type == DistributedType.XLA:
+ return _tpu_broadcast(tensor, src=from_process, name="accelerate.utils.broadcast")
+ elif PartialState().distributed_type in TORCH_DISTRIBUTED_OPERATION_TYPES:
+ return _gpu_broadcast(tensor, src=from_process)
+ else:
+ return tensor
+
+
+def broadcast_object_list(object_list, from_process: int = 0):
+ """
+ Broadcast a list of picklable objects form one process to the others.
+
+ Args:
+ object_list (list of picklable objects):
+ The list of objects to broadcast. This list will be modified inplace.
+ from_process (`int`, *optional*, defaults to 0):
+ The process from which to send the data.
+
+ Returns:
+ The same list containing the objects from process 0.
+ """
+ if PartialState().distributed_type == DistributedType.XLA:
+ for i, obj in enumerate(object_list):
+ object_list[i] = xm.mesh_reduce("accelerate.utils.broadcast_object_list", obj, lambda x: x[from_process])
+ elif PartialState().distributed_type in TORCH_DISTRIBUTED_OPERATION_TYPES:
+ torch.distributed.broadcast_object_list(object_list, src=from_process)
+ return object_list
+
+
+def slice_tensors(data, tensor_slice, process_index=None, num_processes=None):
+ """
+ Recursively takes a slice in a nested list/tuple/dictionary of tensors.
+
+ Args:
+ data (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to slice.
+ tensor_slice (`slice`):
+ The slice to take.
+
+ Returns:
+ The same data structure as `data` with all the tensors slices.
+ """
+
+ def _slice_tensor(tensor, tensor_slice):
+ return tensor[tensor_slice]
+
+ return recursively_apply(_slice_tensor, data, tensor_slice)
+
+
+def concatenate(data, dim=0):
+ """
+ Recursively concatenate the tensors in a nested list/tuple/dictionary of lists of tensors with the same shape.
+
+ Args:
+ data (nested list/tuple/dictionary of lists of tensors `torch.Tensor`):
+ The data to concatenate.
+ dim (`int`, *optional*, defaults to 0):
+ The dimension on which to concatenate.
+
+ Returns:
+ The same data structure as `data` with all the tensors concatenated.
+ """
+ if isinstance(data[0], (tuple, list)):
+ return honor_type(data[0], (concatenate([d[i] for d in data], dim=dim) for i in range(len(data[0]))))
+ elif isinstance(data[0], Mapping):
+ return type(data[0])({k: concatenate([d[k] for d in data], dim=dim) for k in data[0].keys()})
+ elif not isinstance(data[0], torch.Tensor):
+ raise TypeError(f"Can only concatenate tensors but got {type(data[0])}")
+ return torch.cat(data, dim=dim)
+
+
+class CannotPadNestedTensorWarning(UserWarning):
+ pass
+
+
+@chained_operation
+def pad_across_processes(tensor, dim=0, pad_index=0, pad_first=False):
+ """
+ Recursively pad the tensors in a nested list/tuple/dictionary of tensors from all devices to the same size so they
+ can safely be gathered.
+
+ Args:
+ tensor (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to gather.
+ dim (`int`, *optional*, defaults to 0):
+ The dimension on which to pad.
+ pad_index (`int`, *optional*, defaults to 0):
+ The value with which to pad.
+ pad_first (`bool`, *optional*, defaults to `False`):
+ Whether to pad at the beginning or the end.
+ """
+
+ def _pad_across_processes(tensor, dim=0, pad_index=0, pad_first=False):
+ if getattr(tensor, "is_nested", False):
+ warnings.warn(
+ "Cannot pad nested tensors without more information. Leaving unprocessed.",
+ CannotPadNestedTensorWarning,
+ )
+ return tensor
+ if dim >= len(tensor.shape) or dim < -len(tensor.shape):
+ return tensor
+ # Convert negative dimensions to non-negative
+ if dim < 0:
+ dim += len(tensor.shape)
+
+ # Gather all sizes
+ size = torch.tensor(tensor.shape, device=tensor.device)[None]
+ sizes = gather(size).cpu()
+ # Then pad to the maximum size
+ max_size = max(s[dim] for s in sizes)
+ if max_size == tensor.shape[dim]:
+ return tensor
+
+ old_size = tensor.shape
+ new_size = list(old_size)
+ new_size[dim] = max_size
+ new_tensor = tensor.new_zeros(tuple(new_size)) + pad_index
+ if pad_first:
+ indices = tuple(
+ slice(max_size - old_size[dim], max_size) if i == dim else slice(None) for i in range(len(new_size))
+ )
+ else:
+ indices = tuple(slice(0, old_size[dim]) if i == dim else slice(None) for i in range(len(new_size)))
+ new_tensor[indices] = tensor
+ return new_tensor
+
+ return recursively_apply(
+ _pad_across_processes, tensor, error_on_other_type=True, dim=dim, pad_index=pad_index, pad_first=pad_first
+ )
+
+
+def pad_input_tensors(tensor, batch_size, num_processes, dim=0):
+ """
+ Takes a `tensor` of arbitrary size and pads it so that it can work given `num_processes` needed dimensions.
+
+ New tensors are just the last input repeated.
+
+ E.g.:
+ Tensor: ([3,4,4]) Num processes: 4 Expected result shape: ([4,4,4])
+
+ """
+
+ def _pad_input_tensors(tensor, batch_size, num_processes, dim=0):
+ remainder = batch_size // num_processes
+ last_inputs = batch_size - (remainder * num_processes)
+ if batch_size // num_processes == 0:
+ to_pad = num_processes - batch_size
+ else:
+ to_pad = num_processes - (batch_size // num_processes)
+ # In the rare case that `to_pad` is negative,
+ # we need to pad the last inputs - the found `to_pad`
+ if last_inputs > to_pad & to_pad < 1:
+ to_pad = last_inputs - to_pad
+ old_size = tensor.shape
+ new_size = list(old_size)
+ new_size[0] = batch_size + to_pad
+ new_tensor = tensor.new_zeros(tuple(new_size))
+ indices = tuple(slice(0, old_size[dim]) if i == dim else slice(None) for i in range(len(new_size)))
+ new_tensor[indices] = tensor
+ return new_tensor
+
+ return recursively_apply(
+ _pad_input_tensors,
+ tensor,
+ error_on_other_type=True,
+ batch_size=batch_size,
+ num_processes=num_processes,
+ dim=dim,
+ )
+
+
+@verify_operation
+def reduce(tensor, reduction="mean", scale=1.0):
+ """
+ Recursively reduce the tensors in a nested list/tuple/dictionary of lists of tensors across all processes by the
+ mean of a given operation.
+
+ Args:
+ tensor (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to reduce.
+ reduction (`str`, *optional*, defaults to `"mean"`):
+ A reduction method. Can be of "mean", "sum", or "none"
+ scale (`float`, *optional*):
+ A default scaling value to be applied after the reduce, only valied on XLA.
+
+ Returns:
+ The same data structure as `data` with all the tensors reduced.
+ """
+
+ def _reduce_across_processes(tensor, reduction="mean", scale=1.0):
+ state = PartialState()
+ cloned_tensor = tensor.clone()
+ if state.distributed_type == DistributedType.NO:
+ return cloned_tensor
+ if state.distributed_type == DistributedType.XLA:
+ # Some processes may have different HLO graphs than other
+ # processes, for example in the breakpoint API
+ # accelerator.set_trigger(). Use mark_step to make HLOs
+ # the same on all processes.
+ xm.mark_step()
+ xm.all_reduce(xm.REDUCE_SUM, [cloned_tensor], scale)
+ xm.mark_step()
+ elif state.distributed_type.value in TORCH_DISTRIBUTED_OPERATION_TYPES:
+ torch.distributed.all_reduce(cloned_tensor, ReduceOp.SUM)
+ if reduction == "mean":
+ cloned_tensor /= state.num_processes
+ return cloned_tensor
+
+ return recursively_apply(
+ _reduce_across_processes, tensor, error_on_other_type=True, reduction=reduction, scale=scale
+ )
+
+
+def convert_to_fp32(tensor):
+ """
+ Recursively converts the elements nested list/tuple/dictionary of tensors in FP16/BF16 precision to FP32.
+
+ Args:
+ tensor (nested list/tuple/dictionary of `torch.Tensor`):
+ The data to convert from FP16/BF16 to FP32.
+
+ Returns:
+ The same data structure as `tensor` with all tensors that were in FP16/BF16 precision converted to FP32.
+ """
+
+ def _convert_to_fp32(tensor):
+ return tensor.float()
+
+ def _is_fp16_bf16_tensor(tensor):
+ return (is_torch_tensor(tensor) or hasattr(tensor, "dtype")) and tensor.dtype in (
+ torch.float16,
+ torch.bfloat16,
+ )
+
+ return recursively_apply(_convert_to_fp32, tensor, test_type=_is_fp16_bf16_tensor)
+
+
+class ConvertOutputsToFp32:
+ """
+ Decorator to apply to a function outputing tensors (like a model forward pass) that ensures the outputs in FP16
+ precision will be convert back to FP32.
+
+ Args:
+ model_forward (`Callable`):
+ The function which outputs we want to treat.
+
+ Returns:
+ The same function as `model_forward` but with converted outputs.
+ """
+
+ def __init__(self, model_forward):
+ self.model_forward = model_forward
+ update_wrapper(self, model_forward)
+
+ def __call__(self, *args, **kwargs):
+ return convert_to_fp32(self.model_forward(*args, **kwargs))
+
+ def __getstate__(self):
+ raise pickle.PicklingError(
+ "Cannot pickle a prepared model with automatic mixed precision, please unwrap the model with `Accelerator.unwrap_model(model)` before pickling it."
+ )
+
+
+def convert_outputs_to_fp32(model_forward):
+ model_forward = ConvertOutputsToFp32(model_forward)
+
+ def forward(*args, **kwargs):
+ return model_forward(*args, **kwargs)
+
+ # To act like a decorator so that it can be popped when doing `extract_model_from_parallel`
+ forward.__wrapped__ = model_forward
+
+ return forward
+
+
+def find_device(data):
+ """
+ Finds the device on which a nested dict/list/tuple of tensors lies (assuming they are all on the same device).
+
+ Args:
+ (nested list/tuple/dictionary of `torch.Tensor`): The data we want to know the device of.
+ """
+ if isinstance(data, Mapping):
+ for obj in data.values():
+ device = find_device(obj)
+ if device is not None:
+ return device
+ elif isinstance(data, (tuple, list)):
+ for obj in data:
+ device = find_device(obj)
+ if device is not None:
+ return device
+ elif isinstance(data, torch.Tensor):
+ return data.device
+
+
+@contextmanager
+def GatheredParameters(params, modifier_rank=None, fwd_module=None, enabled=True):
+ """
+ Wrapper around `deepspeed.runtime.zero.GatheredParameters`, but if Zero-3 is not enabled, will be a no-op context
+ manager.
+ """
+ # We need to use the `AcceleratorState` here since it has access to the deepspeed plugin
+ if AcceleratorState().distributed_type != DistributedType.DEEPSPEED or (
+ AcceleratorState().deepspeed_plugin is not None
+ and not AcceleratorState().deepspeed_plugin.is_zero3_init_enabled()
+ ):
+ gather_param_context = nullcontext()
+ else:
+ import deepspeed
+
+ gather_param_context = deepspeed.zero.GatheredParameters(
+ params, modifier_rank=modifier_rank, fwd_module=fwd_module, enabled=enabled
+ )
+ with gather_param_context:
+ yield
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/other.py b/venv/lib/python3.11/site-packages/accelerate/utils/other.py
new file mode 100644
index 0000000000000000000000000000000000000000..301a4fe6ab72a579d7efbfd9ff5d7bc4c526f882
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/other.py
@@ -0,0 +1,373 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 collections
+import platform
+import re
+import socket
+from codecs import encode
+from collections import OrderedDict
+from functools import partial, reduce
+from types import MethodType
+
+import numpy as np
+import torch
+from packaging.version import Version
+from safetensors.torch import save_file as safe_save_file
+
+from ..commands.config.default import write_basic_config # noqa: F401
+from ..logging import get_logger
+from ..state import PartialState
+from .constants import FSDP_PYTORCH_VERSION
+from .dataclasses import DistributedType
+from .imports import (
+ is_deepspeed_available,
+ is_numpy_available,
+ is_torch_distributed_available,
+ is_torch_xla_available,
+ is_weights_only_available,
+)
+from .modeling import id_tensor_storage
+from .transformer_engine import convert_model
+from .versions import is_torch_version
+
+
+logger = get_logger(__name__)
+
+
+if is_torch_xla_available():
+ import torch_xla.core.xla_model as xm
+
+
+def is_compiled_module(module):
+ """
+ Check whether the module was compiled with torch.compile()
+ """
+ if not hasattr(torch, "_dynamo"):
+ return False
+ return isinstance(module, torch._dynamo.eval_frame.OptimizedModule)
+
+
+def extract_model_from_parallel(
+ model, keep_fp32_wrapper: bool = True, keep_torch_compile: bool = True, recursive: bool = False
+):
+ """
+ Extract a model from its distributed containers.
+
+ Args:
+ model (`torch.nn.Module`):
+ The model to extract.
+ keep_fp32_wrapper (`bool`, *optional*):
+ Whether to remove mixed precision hooks from the model.
+ keep_torch_compile (`bool`, *optional*):
+ Whether to unwrap compiled model.
+ recursive (`bool`, *optional*, defaults to `False`):
+ Whether to recursively extract all cases of `module.module` from `model` as well as unwrap child sublayers
+ recursively, not just the top-level distributed containers.
+
+ Returns:
+ `torch.nn.Module`: The extracted model.
+ """
+ options = (torch.nn.parallel.DistributedDataParallel, torch.nn.DataParallel)
+
+ is_compiled = is_compiled_module(model)
+ if is_compiled:
+ compiled_model = model
+ model = model._orig_mod
+
+ if is_deepspeed_available():
+ from deepspeed import DeepSpeedEngine
+
+ options += (DeepSpeedEngine,)
+
+ if is_torch_version(">=", FSDP_PYTORCH_VERSION) and is_torch_distributed_available():
+ from torch.distributed.fsdp.fully_sharded_data_parallel import FullyShardedDataParallel as FSDP
+
+ options += (FSDP,)
+
+ while isinstance(model, options):
+ model = model.module
+
+ if recursive:
+ # This is needed in cases such as using FSDPv2 on XLA
+ def _recursive_unwrap(module):
+ # Wrapped modules are standardly wrapped as `module`, similar to the cases earlier
+ # with DDP, DataParallel, DeepSpeed, and FSDP
+ if hasattr(module, "module"):
+ unwrapped_module = _recursive_unwrap(module.module)
+ else:
+ unwrapped_module = module
+ # Next unwrap child sublayers recursively
+ for name, child in unwrapped_module.named_children():
+ setattr(unwrapped_module, name, _recursive_unwrap(child))
+ return unwrapped_module
+
+ # Start with top-level
+ model = _recursive_unwrap(model)
+
+ if not keep_fp32_wrapper:
+ forward = model.forward
+ original_forward = model.__dict__.pop("_original_forward", None)
+ if original_forward is not None:
+ while hasattr(forward, "__wrapped__"):
+ forward = forward.__wrapped__
+ if forward == original_forward:
+ break
+ model.forward = MethodType(forward, model)
+ if getattr(model, "_converted_to_transformer_engine", False):
+ convert_model(model, to_transformer_engine=False)
+
+ if keep_torch_compile and is_compiled:
+ compiled_model._orig_mod = model
+ model = compiled_model
+
+ return model
+
+
+def wait_for_everyone():
+ """
+ Introduces a blocking point in the script, making sure all processes have reached this point before continuing.
+
+
+
+ Make sure all processes will reach this instruction otherwise one of your processes will hang forever.
+
+
+ """
+ PartialState().wait_for_everyone()
+
+
+def clean_state_dict_for_safetensors(state_dict: dict):
+ """
+ Cleans the state dictionary from a model and removes tensor aliasing if present.
+
+ Args:
+ state_dict (`dict`):
+ The state dictionary from a model
+ """
+ ptrs = collections.defaultdict(list)
+ # When bnb serialization is used, weights in state dict can be strings
+ for name, tensor in state_dict.items():
+ if not isinstance(tensor, str):
+ ptrs[id_tensor_storage(tensor)].append(name)
+
+ # These are all pointers of tensors with shared memory
+ shared_ptrs = {ptr: names for ptr, names in ptrs.items() if len(names) > 1}
+ warn_names = set()
+ for names in shared_ptrs.values():
+ # When not all duplicates have been cleaned, we still remove those keys but put a clear warning.
+ # If the link between tensors was done at runtime then `from_pretrained` will not get
+ # the key back leading to random tensor. A proper warning will be shown
+ # during reload (if applicable), but since the file is not necessarily compatible with
+ # the config, better show a proper warning.
+ found_names = [name for name in names if name in state_dict]
+ warn_names.update(found_names[1:])
+ for name in found_names[1:]:
+ del state_dict[name]
+ if len(warn_names) > 0:
+ logger.warning(
+ f"Removed shared tensor {warn_names} while saving. This should be OK, but check by verifying that you don't receive any warning while reloading",
+ )
+ state_dict = {k: v.contiguous() if isinstance(v, torch.Tensor) else v for k, v in state_dict.items()}
+ return state_dict
+
+
+def save(obj, f, save_on_each_node: bool = False, safe_serialization: bool = False):
+ """
+ Save the data to disk. Use in place of `torch.save()`.
+
+ Args:
+ obj:
+ The data to save
+ f:
+ The file (or file-like object) to use to save the data
+ save_on_each_node (`bool`, *optional*, defaults to `False`):
+ Whether to only save on the global main process
+ safe_serialization (`bool`, *optional*, defaults to `False`):
+ Whether to save `obj` using `safetensors` or the traditional PyTorch way (that uses `pickle`).
+ """
+ # When TorchXLA is enabled, it's necessary to transfer all data to the CPU before saving.
+ # Another issue arises with `id_tensor_storage`, which treats all XLA tensors as identical.
+ # If tensors remain on XLA, calling `clean_state_dict_for_safetensors` will result in only
+ # one XLA tensor remaining.
+ if PartialState().distributed_type == DistributedType.XLA:
+ obj = xm._maybe_convert_to_cpu(obj)
+ # Check if it's a model and remove duplicates
+ if safe_serialization:
+ save_func = partial(safe_save_file, metadata={"format": "pt"})
+ if isinstance(obj, OrderedDict):
+ obj = clean_state_dict_for_safetensors(obj)
+ else:
+ save_func = torch.save
+
+ if PartialState().is_main_process and not save_on_each_node:
+ save_func(obj, f)
+ elif PartialState().is_local_main_process and save_on_each_node:
+ save_func(obj, f)
+
+
+# The following are considered "safe" globals to reconstruct various types of objects when using `weights_only=True`
+# These should be added and then removed after loading in the file
+np_core = np._core if is_numpy_available("2.0.0") else np.core
+TORCH_SAFE_GLOBALS = [
+ # numpy arrays are just numbers, not objects, so we can reconstruct them safely
+ np_core.multiarray._reconstruct,
+ np.ndarray,
+ # The following are needed for the RNG states
+ encode,
+ np.dtype,
+]
+
+if is_numpy_available("1.25.0"):
+ TORCH_SAFE_GLOBALS.append(np.dtypes.UInt32DType)
+
+
+def load(f, map_location=None, **kwargs):
+ """
+ Compatible drop-in replacement of `torch.load()` which allows for `weights_only` to be used if `torch` version is
+ 2.4.0 or higher. Otherwise will ignore the kwarg.
+
+ Will also add (and then remove) an exception for numpy arrays
+
+ Args:
+ f:
+ The file (or file-like object) to use to load the data
+ map_location:
+ a function, `torch.device`, string or a dict specifying how to remap storage locations
+ **kwargs:
+ Additional keyword arguments to pass to `torch.load()`.
+ """
+ try:
+ if is_weights_only_available():
+ old_safe_globals = torch.serialization.get_safe_globals()
+ if "weights_only" not in kwargs:
+ kwargs["weights_only"] = True
+ torch.serialization.add_safe_globals(TORCH_SAFE_GLOBALS)
+ else:
+ kwargs.pop("weights_only", None)
+ loaded_obj = torch.load(f, map_location=map_location, **kwargs)
+ finally:
+ if is_weights_only_available():
+ torch.serialization.clear_safe_globals()
+ if old_safe_globals:
+ torch.serialization.add_safe_globals(old_safe_globals)
+ return loaded_obj
+
+
+def get_pretty_name(obj):
+ """
+ Gets a pretty name from `obj`.
+ """
+ if not hasattr(obj, "__qualname__") and not hasattr(obj, "__name__"):
+ obj = getattr(obj, "__class__", obj)
+ if hasattr(obj, "__qualname__"):
+ return obj.__qualname__
+ if hasattr(obj, "__name__"):
+ return obj.__name__
+ return str(obj)
+
+
+def merge_dicts(source, destination):
+ """
+ Recursively merges two dictionaries.
+
+ Args:
+ source (`dict`): The dictionary to merge into `destination`.
+ destination (`dict`): The dictionary to merge `source` into.
+ """
+ for key, value in source.items():
+ if isinstance(value, dict):
+ node = destination.setdefault(key, {})
+ merge_dicts(value, node)
+ else:
+ destination[key] = value
+
+ return destination
+
+
+def is_port_in_use(port: int = None) -> bool:
+ """
+ Checks if a port is in use on `localhost`. Useful for checking if multiple `accelerate launch` commands have been
+ run and need to see if the port is already in use.
+ """
+ if port is None:
+ port = 29500
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
+ return s.connect_ex(("localhost", port)) == 0
+
+
+def convert_bytes(size):
+ "Converts `size` from bytes to the largest possible unit"
+ for x in ["bytes", "KB", "MB", "GB", "TB"]:
+ if size < 1024.0:
+ return f"{round(size, 2)} {x}"
+ size /= 1024.0
+
+ return f"{round(size, 2)} PB"
+
+
+def check_os_kernel():
+ """Warns if the kernel version is below the recommended minimum on Linux."""
+ # see issue #1929
+ info = platform.uname()
+ system = info.system
+ if system != "Linux":
+ return
+
+ _, version, *_ = re.split(r"(\d+\.\d+\.\d+)", info.release)
+ min_version = "5.5.0"
+ if Version(version) < Version(min_version):
+ msg = (
+ f"Detected kernel version {version}, which is below the recommended minimum of {min_version}; this can "
+ "cause the process to hang. It is recommended to upgrade the kernel to the minimum version or higher."
+ )
+ logger.warning(msg, main_process_only=True)
+
+
+def recursive_getattr(obj, attr: str):
+ """
+ Recursive `getattr`.
+
+ Args:
+ obj:
+ A class instance holding the attribute.
+ attr (`str`):
+ The attribute that is to be retrieved, e.g. 'attribute1.attribute2'.
+ """
+
+ def _getattr(obj, attr):
+ return getattr(obj, attr)
+
+ return reduce(_getattr, [obj] + attr.split("."))
+
+
+def get_module_children_bottom_up(model: torch.nn.Module) -> list[torch.nn.Module]:
+ """Traverse the model in bottom-up order and return the children modules in that order.
+
+ Args:
+ model (`torch.nn.Module`): the model to get the children of
+
+ Returns:
+ `list[torch.nn.Module]`: a list of children modules of `model` in bottom-up order. The last element is the
+ `model` itself.
+ """
+ stack = [model]
+ ordered_modules = []
+ while stack:
+ current_module = stack.pop()
+ for _, attr in current_module.named_children():
+ if isinstance(attr, torch.nn.Module):
+ stack.append(attr)
+ ordered_modules.append(current_module)
+ return ordered_modules[::-1]
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/random.py b/venv/lib/python3.11/site-packages/accelerate/utils/random.py
new file mode 100644
index 0000000000000000000000000000000000000000..9dceb598cacc1c1d17b198e0a7c19789ab3b9f39
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/random.py
@@ -0,0 +1,156 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 random
+from typing import Optional, Union
+
+import numpy as np
+import torch
+
+from ..state import AcceleratorState
+from .constants import CUDA_DISTRIBUTED_TYPES
+from .dataclasses import DistributedType, RNGType
+from .imports import (
+ is_hpu_available,
+ is_mlu_available,
+ is_musa_available,
+ is_npu_available,
+ is_sdaa_available,
+ is_torch_xla_available,
+ is_xpu_available,
+)
+
+
+if is_torch_xla_available():
+ import torch_xla.core.xla_model as xm
+
+
+def set_seed(seed: int, device_specific: bool = False, deterministic: bool = False):
+ """
+ Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.
+
+ Args:
+ seed (`int`):
+ The seed to set.
+ device_specific (`bool`, *optional*, defaults to `False`):
+ Whether to differ the seed on each device slightly with `self.process_index`.
+ deterministic (`bool`, *optional*, defaults to `False`):
+ Whether to use deterministic algorithms where available. Can slow down training.
+ """
+ if device_specific:
+ seed += AcceleratorState().process_index
+ random.seed(seed)
+ np.random.seed(seed)
+ torch.manual_seed(seed)
+ if is_xpu_available():
+ torch.xpu.manual_seed_all(seed)
+ elif is_npu_available():
+ torch.npu.manual_seed_all(seed)
+ elif is_mlu_available():
+ torch.mlu.manual_seed_all(seed)
+ elif is_sdaa_available():
+ torch.sdaa.manual_seed_all(seed)
+ elif is_musa_available():
+ torch.musa.manual_seed_all(seed)
+ elif is_hpu_available():
+ torch.hpu.manual_seed_all(seed)
+ else:
+ torch.cuda.manual_seed_all(seed)
+ # ^^ safe to call this function even if cuda is not available
+ if is_torch_xla_available():
+ xm.set_rng_state(seed)
+
+ if deterministic:
+ torch.use_deterministic_algorithms(True)
+
+
+def synchronize_rng_state(rng_type: Optional[RNGType] = None, generator: Optional[torch.Generator] = None):
+ # Get the proper rng state
+ if rng_type == RNGType.TORCH:
+ rng_state = torch.get_rng_state()
+ elif rng_type == RNGType.CUDA:
+ rng_state = torch.cuda.get_rng_state()
+ elif rng_type == RNGType.XLA:
+ assert is_torch_xla_available(), "Can't synchronize XLA seeds as torch_xla is unavailable."
+ rng_state = torch.tensor(xm.get_rng_state())
+ elif rng_type == RNGType.NPU:
+ assert is_npu_available(), "Can't synchronize NPU seeds on an environment without NPUs."
+ rng_state = torch.npu.get_rng_state()
+ elif rng_type == RNGType.MLU:
+ assert is_mlu_available(), "Can't synchronize MLU seeds on an environment without MLUs."
+ rng_state = torch.mlu.get_rng_state()
+ elif rng_type == RNGType.SDAA:
+ assert is_sdaa_available(), "Can't synchronize SDAA seeds on an environment without SDAAs."
+ rng_state = torch.sdaa.get_rng_state()
+ elif rng_type == RNGType.MUSA:
+ assert is_musa_available(), "Can't synchronize MUSA seeds on an environment without MUSAs."
+ rng_state = torch.musa.get_rng_state()
+ elif rng_type == RNGType.XPU:
+ assert is_xpu_available(), "Can't synchronize XPU seeds on an environment without XPUs."
+ rng_state = torch.xpu.get_rng_state()
+ elif rng_type == RNGType.HPU:
+ assert is_hpu_available(), "Can't synchronize HPU seeds on an environment without HPUs."
+ rng_state = torch.hpu.get_rng_state()
+ elif rng_type == RNGType.GENERATOR:
+ assert generator is not None, "Need a generator to synchronize its seed."
+ rng_state = generator.get_state()
+
+ # Broadcast the rng state from device 0 to other devices
+ state = AcceleratorState()
+ if state.distributed_type == DistributedType.XLA:
+ rng_state = rng_state.to(xm.xla_device())
+ xm.collective_broadcast([rng_state])
+ xm.mark_step()
+ rng_state = rng_state.cpu()
+ elif (
+ state.distributed_type in CUDA_DISTRIBUTED_TYPES
+ or state.distributed_type == DistributedType.MULTI_MLU
+ or state.distributed_type == DistributedType.MULTI_SDAA
+ or state.distributed_type == DistributedType.MULTI_MUSA
+ or state.distributed_type == DistributedType.MULTI_NPU
+ or state.distributed_type == DistributedType.MULTI_XPU
+ or state.distributed_type == DistributedType.MULTI_HPU
+ ):
+ rng_state = rng_state.to(state.device)
+ torch.distributed.broadcast(rng_state, 0)
+ rng_state = rng_state.cpu()
+ elif state.distributed_type == DistributedType.MULTI_CPU:
+ torch.distributed.broadcast(rng_state, 0)
+
+ # Set the broadcast rng state
+ if rng_type == RNGType.TORCH:
+ torch.set_rng_state(rng_state)
+ elif rng_type == RNGType.CUDA:
+ torch.cuda.set_rng_state(rng_state)
+ elif rng_type == RNGType.NPU:
+ torch.npu.set_rng_state(rng_state)
+ elif rng_type == RNGType.MLU:
+ torch.mlu.set_rng_state(rng_state)
+ elif rng_type == RNGType.SDAA:
+ torch.sdaa.set_rng_state(rng_state)
+ elif rng_type == RNGType.MUSA:
+ torch.musa.set_rng_state(rng_state)
+ elif rng_type == RNGType.XPU:
+ torch.xpu.set_rng_state(rng_state)
+ elif rng_state == RNGType.HPU:
+ torch.hpu.set_rng_state(rng_state)
+ elif rng_type == RNGType.XLA:
+ xm.set_rng_state(rng_state.item())
+ elif rng_type == RNGType.GENERATOR:
+ generator.set_state(rng_state)
+
+
+def synchronize_rng_states(rng_types: list[Union[str, RNGType]], generator: Optional[torch.Generator] = None):
+ for rng_type in rng_types:
+ synchronize_rng_state(RNGType(rng_type), generator=generator)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/rich.py b/venv/lib/python3.11/site-packages/accelerate/utils/rich.py
new file mode 100644
index 0000000000000000000000000000000000000000..2d48661b7fcef92ef1168b74cc275c6d3ccc67a1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/rich.py
@@ -0,0 +1,24 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+from .imports import is_rich_available
+
+
+if is_rich_available():
+ from rich.traceback import install
+
+ install(show_locals=False)
+
+else:
+ raise ModuleNotFoundError("To use the rich extension, install rich with `pip install rich`")
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/torch_xla.py b/venv/lib/python3.11/site-packages/accelerate/utils/torch_xla.py
new file mode 100644
index 0000000000000000000000000000000000000000..140133926c2f88d39c70f5a9f46a08f88bed36da
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/torch_xla.py
@@ -0,0 +1,51 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 importlib.metadata
+import subprocess
+import sys
+
+
+def install_xla(upgrade: bool = False):
+ """
+ Helper function to install appropriate xla wheels based on the `torch` version in Google Colaboratory.
+
+ Args:
+ upgrade (`bool`, *optional*, defaults to `False`):
+ Whether to upgrade `torch` and install the latest `torch_xla` wheels.
+
+ Example:
+
+ ```python
+ >>> from accelerate.utils import install_xla
+
+ >>> install_xla(upgrade=True)
+ ```
+ """
+ in_colab = False
+ if "IPython" in sys.modules:
+ in_colab = "google.colab" in str(sys.modules["IPython"].get_ipython())
+
+ if in_colab:
+ if upgrade:
+ torch_install_cmd = ["pip", "install", "-U", "torch"]
+ subprocess.run(torch_install_cmd, check=True)
+ # get the current version of torch
+ torch_version = importlib.metadata.version("torch")
+ torch_version_trunc = torch_version[: torch_version.rindex(".")]
+ xla_wheel = f"https://storage.googleapis.com/tpu-pytorch/wheels/colab/torch_xla-{torch_version_trunc}-cp37-cp37m-linux_x86_64.whl"
+ xla_install_cmd = ["pip", "install", xla_wheel]
+ subprocess.run(xla_install_cmd, check=True)
+ else:
+ raise RuntimeError("`install_xla` utility works only on google colab.")
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/tqdm.py b/venv/lib/python3.11/site-packages/accelerate/utils/tqdm.py
new file mode 100644
index 0000000000000000000000000000000000000000..2d4873c1573eb2ee7392162f440a76d4f07cd8ce
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/tqdm.py
@@ -0,0 +1,43 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+
+from .imports import is_tqdm_available
+
+
+if is_tqdm_available():
+ from tqdm.auto import tqdm as _tqdm
+
+from ..state import PartialState
+
+
+def tqdm(*args, main_process_only: bool = True, **kwargs):
+ """
+ Wrapper around `tqdm.tqdm` that optionally displays only on the main process.
+
+ Args:
+ main_process_only (`bool`, *optional*):
+ Whether to display the progress bar only on the main process
+ """
+ if not is_tqdm_available():
+ raise ImportError("Accelerate's `tqdm` module requires `tqdm` to be installed. Please run `pip install tqdm`.")
+ if len(args) > 0 and isinstance(args[0], bool):
+ raise ValueError(
+ "Passing `True` or `False` as the first argument to Accelerate's `tqdm` wrapper is unsupported. "
+ "Please use the `main_process_only` keyword argument instead."
+ )
+ disable = kwargs.pop("disable", False)
+ if main_process_only and not disable:
+ disable = PartialState().local_process_index != 0
+ return _tqdm(*args, **kwargs, disable=disable)
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/transformer_engine.py b/venv/lib/python3.11/site-packages/accelerate/utils/transformer_engine.py
new file mode 100644
index 0000000000000000000000000000000000000000..53159a010eb99a0f597b0487eb201ce72593b9e9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/transformer_engine.py
@@ -0,0 +1,160 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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.
+
+from types import MethodType
+
+import torch.nn as nn
+
+from .imports import is_fp8_available, is_hpu_available
+from .operations import GatheredParameters
+
+
+# Do not import `transformer_engine` at package level to avoid potential issues
+
+
+def convert_model(model, to_transformer_engine=True, _convert_linear=True, _convert_ln=True):
+ """
+ Recursively converts the linear and layernorm layers of a model to their `transformers_engine` counterpart.
+ """
+ if not is_fp8_available():
+ raise ImportError("Using `convert_model` requires transformer_engine to be installed.")
+
+ if is_hpu_available():
+ import intel_transformer_engine as te
+ else:
+ import transformer_engine.pytorch as te
+
+ for name, module in model.named_children():
+ if isinstance(module, nn.Linear) and to_transformer_engine and _convert_linear:
+ has_bias = module.bias is not None
+ params_to_gather = [module.weight]
+ if has_bias:
+ params_to_gather.append(module.bias)
+
+ with GatheredParameters(params_to_gather, modifier_rank=0):
+ if any(p % 16 != 0 for p in module.weight.shape):
+ return
+ te_module = te.Linear(
+ module.in_features, module.out_features, bias=has_bias, params_dtype=module.weight.dtype
+ )
+ te_module.weight.copy_(module.weight)
+ if has_bias:
+ te_module.bias.copy_(module.bias)
+
+ setattr(model, name, te_module)
+ # Note: @xrsrke (Phuc) found that te.LayerNorm doesn't have any real memory savings or speedups over nn.LayerNorm
+ elif isinstance(module, nn.LayerNorm) and to_transformer_engine and _convert_ln:
+ with GatheredParameters([module.weight, module.bias], modifier_rank=0):
+ te_module = te.LayerNorm(module.normalized_shape[0], eps=module.eps, params_dtype=module.weight.dtype)
+ te_module.weight.copy_(module.weight)
+ te_module.bias.copy_(module.bias)
+
+ setattr(model, name, te_module)
+ elif isinstance(module, te.Linear) and not to_transformer_engine and _convert_linear:
+ has_bias = module.bias is not None
+ new_module = nn.Linear(
+ module.in_features, module.out_features, bias=has_bias, params_dtype=module.weight.dtype
+ )
+ new_module.weight.copy_(module.weight)
+ if has_bias:
+ new_module.bias.copy_(module.bias)
+
+ setattr(model, name, new_module)
+ elif isinstance(module, te.LayerNorm) and not to_transformer_engine and _convert_ln:
+ new_module = nn.LayerNorm(module.normalized_shape[0], eps=module.eps, params_dtype=module.weight.dtype)
+ new_module.weight.copy_(module.weight)
+ new_module.bias.copy_(module.bias)
+
+ setattr(model, name, new_module)
+ else:
+ convert_model(
+ module,
+ to_transformer_engine=to_transformer_engine,
+ _convert_linear=_convert_linear,
+ _convert_ln=_convert_ln,
+ )
+
+
+def has_transformer_engine_layers(model):
+ """
+ Returns whether a given model has some `transformer_engine` layer or not.
+ """
+ if not is_fp8_available():
+ raise ImportError("Using `has_transformer_engine_layers` requires transformer_engine to be installed.")
+
+ if is_hpu_available():
+ import intel_transformer_engine as te
+
+ module_cls_to_check = te.Linear
+ else:
+ import transformer_engine.pytorch as te
+
+ module_cls_to_check = (te.LayerNorm, te.Linear, te.TransformerLayer)
+
+ for m in model.modules():
+ if isinstance(m, module_cls_to_check):
+ return True
+
+ return False
+
+
+def contextual_fp8_autocast(model_forward, fp8_recipe, use_during_eval=False):
+ """
+ Wrapper for a model's forward method to apply FP8 autocast. Is context aware, meaning that by default it will
+ disable FP8 autocast during eval mode, which is generally better for more accurate metrics.
+ """
+ if not is_fp8_available():
+ raise ImportError("Using `contextual_fp8_autocast` requires transformer_engine to be installed.")
+
+ if is_hpu_available():
+ from intel_transformer_engine import fp8_autocast
+ else:
+ from transformer_engine.pytorch import fp8_autocast
+
+ def forward(self, *args, **kwargs):
+ enabled = use_during_eval or self.training
+ with fp8_autocast(enabled=enabled, fp8_recipe=fp8_recipe):
+ return model_forward(*args, **kwargs)
+
+ # To act like a decorator so that it can be popped when doing `extract_model_from_parallel`
+ forward.__wrapped__ = model_forward
+
+ return forward
+
+
+def apply_fp8_autowrap(model, fp8_recipe_handler):
+ """
+ Applies FP8 context manager to the model's forward method
+ """
+ if not is_fp8_available():
+ raise ImportError("Using `apply_fp8_autowrap` requires transformer_engine to be installed.")
+
+ if is_hpu_available():
+ import intel_transformer_engine.recipe as te_recipe
+ else:
+ import transformer_engine.common.recipe as te_recipe
+
+ kwargs = fp8_recipe_handler.to_kwargs() if fp8_recipe_handler is not None else {}
+ if "fp8_format" in kwargs:
+ kwargs["fp8_format"] = getattr(te_recipe.Format, kwargs["fp8_format"])
+ use_during_eval = kwargs.pop("use_autocast_during_eval", False)
+ fp8_recipe = te_recipe.DelayedScaling(**kwargs)
+ new_forward = contextual_fp8_autocast(model.forward, fp8_recipe, use_during_eval)
+
+ if hasattr(model.forward, "__func__"):
+ model.forward = MethodType(new_forward, model)
+ else:
+ model.forward = new_forward
+
+ return model
diff --git a/venv/lib/python3.11/site-packages/accelerate/utils/versions.py b/venv/lib/python3.11/site-packages/accelerate/utils/versions.py
new file mode 100644
index 0000000000000000000000000000000000000000..985c918f0e057bacc70c372f6906071bb73db577
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/accelerate/utils/versions.py
@@ -0,0 +1,56 @@
+# Copyright 2022 The HuggingFace Team. All rights reserved.
+#
+# 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 importlib.metadata
+from typing import Union
+
+from packaging.version import Version, parse
+
+from .constants import STR_OPERATION_TO_FUNC
+
+
+torch_version = parse(importlib.metadata.version("torch"))
+
+
+def compare_versions(library_or_version: Union[str, Version], operation: str, requirement_version: str):
+ """
+ Compares a library version to some requirement using a given operation.
+
+ Args:
+ library_or_version (`str` or `packaging.version.Version`):
+ A library name or a version to check.
+ operation (`str`):
+ A string representation of an operator, such as `">"` or `"<="`.
+ requirement_version (`str`):
+ The version to compare the library version against
+ """
+ if operation not in STR_OPERATION_TO_FUNC.keys():
+ raise ValueError(f"`operation` must be one of {list(STR_OPERATION_TO_FUNC.keys())}, received {operation}")
+ operation = STR_OPERATION_TO_FUNC[operation]
+ if isinstance(library_or_version, str):
+ library_or_version = parse(importlib.metadata.version(library_or_version))
+ return operation(library_or_version, parse(requirement_version))
+
+
+def is_torch_version(operation: str, version: str):
+ """
+ Compares the current PyTorch version to a given reference with an operation.
+
+ Args:
+ operation (`str`):
+ A string representation of an operator, such as `">"` or `"<="`
+ version (`str`):
+ A string version of PyTorch
+ """
+ return compare_versions(torch_version, operation, version)
diff --git a/venv/lib/python3.11/site-packages/ada92cb5d92a588d1b93__mypyc.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/ada92cb5d92a588d1b93__mypyc.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..6c705cc6708e4726009f99c0b8d703cba6354c45
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/ada92cb5d92a588d1b93__mypyc.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:bff8964f346c6579f65b70f06199dab85af861fd00f30f31b3ff23ee67d36c5c
+size 445952
diff --git a/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/INSTALLER b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/INSTALLER
new file mode 100644
index 0000000000000000000000000000000000000000..a1b589e38a32041e49332e5e81c2d363dc418d68
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/INSTALLER
@@ -0,0 +1 @@
+pip
diff --git a/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/METADATA b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/METADATA
new file mode 100644
index 0000000000000000000000000000000000000000..9bf7a9e800778c5a8c3f1357450ab0849d13d953
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/METADATA
@@ -0,0 +1,145 @@
+Metadata-Version: 2.4
+Name: annotated-doc
+Version: 0.0.4
+Summary: Document parameters, class attributes, return types, and variables inline, with Annotated.
+Author-Email: =?utf-8?q?Sebasti=C3=A1n_Ram=C3=ADrez?=
+License-Expression: MIT
+License-File: LICENSE
+Classifier: Intended Audience :: Information Technology
+Classifier: Intended Audience :: System Administrators
+Classifier: Operating System :: OS Independent
+Classifier: Programming Language :: Python :: 3
+Classifier: Programming Language :: Python
+Classifier: Topic :: Internet
+Classifier: Topic :: Software Development :: Libraries :: Application Frameworks
+Classifier: Topic :: Software Development :: Libraries :: Python Modules
+Classifier: Topic :: Software Development :: Libraries
+Classifier: Topic :: Software Development
+Classifier: Typing :: Typed
+Classifier: Development Status :: 4 - Beta
+Classifier: Intended Audience :: Developers
+Classifier: Programming Language :: Python :: 3 :: Only
+Classifier: Programming Language :: Python :: 3.8
+Classifier: Programming Language :: Python :: 3.9
+Classifier: Programming Language :: Python :: 3.10
+Classifier: Programming Language :: Python :: 3.11
+Classifier: Programming Language :: Python :: 3.12
+Classifier: Programming Language :: Python :: 3.13
+Classifier: Programming Language :: Python :: 3.14
+Project-URL: Homepage, https://github.com/fastapi/annotated-doc
+Project-URL: Documentation, https://github.com/fastapi/annotated-doc
+Project-URL: Repository, https://github.com/fastapi/annotated-doc
+Project-URL: Issues, https://github.com/fastapi/annotated-doc/issues
+Project-URL: Changelog, https://github.com/fastapi/annotated-doc/release-notes.md
+Requires-Python: >=3.8
+Description-Content-Type: text/markdown
+
+# Annotated Doc
+
+Document parameters, class attributes, return types, and variables inline, with `Annotated`.
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+## Installation
+
+```bash
+pip install annotated-doc
+```
+
+Or with `uv`:
+
+```Python
+uv add annotated-doc
+```
+
+## Usage
+
+Import `Doc` and pass a single literal string with the documentation for the specific parameter, class attribute, return type, or variable.
+
+For example, to document a parameter `name` in a function `hi` you could do:
+
+```Python
+from typing import Annotated
+
+from annotated_doc import Doc
+
+def hi(name: Annotated[str, Doc("Who to say hi to")]) -> None:
+ print(f"Hi, {name}!")
+```
+
+You can also use it to document class attributes:
+
+```Python
+from typing import Annotated
+
+from annotated_doc import Doc
+
+class User:
+ name: Annotated[str, Doc("The user's name")]
+ age: Annotated[int, Doc("The user's age")]
+```
+
+The same way, you could document return types and variables, or anything that could have a type annotation with `Annotated`.
+
+## Who Uses This
+
+`annotated-doc` was made for:
+
+* [FastAPI](https://fastapi.tiangolo.com/)
+* [Typer](https://typer.tiangolo.com/)
+* [SQLModel](https://sqlmodel.tiangolo.com/)
+* [Asyncer](https://asyncer.tiangolo.com/)
+
+`annotated-doc` is supported by [griffe-typingdoc](https://github.com/mkdocstrings/griffe-typingdoc), which powers reference documentation like the one in the [FastAPI Reference](https://fastapi.tiangolo.com/reference/).
+
+## Reasons not to use `annotated-doc`
+
+You are already comfortable with one of the existing docstring formats, like:
+
+* Sphinx
+* numpydoc
+* Google
+* Keras
+
+Your team is already comfortable using them.
+
+You prefer having the documentation about parameters all together in a docstring, separated from the code defining them.
+
+You care about a specific set of users, using one specific editor, and that editor already has support for the specific docstring format you use.
+
+## Reasons to use `annotated-doc`
+
+* No micro-syntax to learn for newcomers, it’s **just Python** syntax.
+* **Editing** would be already fully supported by default by any editor (current or future) supporting Python syntax, including syntax errors, syntax highlighting, etc.
+* **Rendering** would be relatively straightforward to implement by static tools (tools that don't need runtime execution), as the information can be extracted from the AST they normally already create.
+* **Deduplication of information**: the name of a parameter would be defined in a single place, not duplicated inside of a docstring.
+* **Elimination** of the possibility of having **inconsistencies** when removing a parameter or class variable and **forgetting to remove** its documentation.
+* **Minimization** of the probability of adding a new parameter or class variable and **forgetting to add its documentation**.
+* **Elimination** of the possibility of having **inconsistencies** between the **name** of a parameter in the **signature** and the name in the docstring when it is renamed.
+* **Access** to the documentation string for each symbol at **runtime**, including existing (older) Python versions.
+* A more formalized way to document other symbols, like type aliases, that could use Annotated.
+* **Support** for apps using FastAPI, Typer and others.
+* **AI Accessibility**: AI tools will have an easier way understanding each parameter as the distance from documentation to parameter is much closer.
+
+## History
+
+I ([@tiangolo](https://github.com/tiangolo)) originally wanted for this to be part of the Python standard library (in [PEP 727](https://peps.python.org/pep-0727/)), but the proposal was withdrawn as there was a fair amount of negative feedback and opposition.
+
+The conclusion was that this was better done as an external effort, in a third-party library.
+
+So, here it is, with a simpler approach, as a third-party library, in a way that can be used by others, starting with FastAPI and friends.
+
+## License
+
+This project is licensed under the terms of the MIT license.
diff --git a/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/RECORD b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/RECORD
new file mode 100644
index 0000000000000000000000000000000000000000..e46a002c68701be55bc8bd318ae983b7940f769d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/RECORD
@@ -0,0 +1,11 @@
+annotated_doc-0.0.4.dist-info/INSTALLER,sha256=zuuue4knoyJ-UwPPXg8fezS7VCrXJQrAP7zeNuwvFQg,4
+annotated_doc-0.0.4.dist-info/METADATA,sha256=Irm5KJua33dY2qKKAjJ-OhKaVBVIfwFGej_dSe3Z1TU,6566
+annotated_doc-0.0.4.dist-info/RECORD,,
+annotated_doc-0.0.4.dist-info/WHEEL,sha256=9P2ygRxDrTJz3gsagc0Z96ukrxjr-LFBGOgv3AuKlCA,90
+annotated_doc-0.0.4.dist-info/entry_points.txt,sha256=6OYgBcLyFCUgeqLgnvMyOJxPCWzgy7se4rLPKtNonMs,34
+annotated_doc-0.0.4.dist-info/licenses/LICENSE,sha256=__Fwd5pqy_ZavbQFwIfxzuF4ZpHkqWpANFF-SlBKDN8,1086
+annotated_doc/__init__.py,sha256=VuyxxUe80kfEyWnOrCx_Bk8hybo3aKo6RYBlkBBYW8k,52
+annotated_doc/__pycache__/__init__.cpython-311.pyc,,
+annotated_doc/__pycache__/main.cpython-311.pyc,,
+annotated_doc/main.py,sha256=5Zfvxv80SwwLqpRW73AZyZyiM4bWma9QWRbp_cgD20s,1075
+annotated_doc/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
diff --git a/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/WHEEL b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/WHEEL
new file mode 100644
index 0000000000000000000000000000000000000000..045c8acdea31cbca5be986e915f784c1aafc720f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/WHEEL
@@ -0,0 +1,4 @@
+Wheel-Version: 1.0
+Generator: pdm-backend (2.4.5)
+Root-Is-Purelib: true
+Tag: py3-none-any
diff --git a/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/entry_points.txt b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/entry_points.txt
new file mode 100644
index 0000000000000000000000000000000000000000..c3ad4726d437022e5c606a4206ffb6007347a008
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/entry_points.txt
@@ -0,0 +1,4 @@
+[console_scripts]
+
+[gui_scripts]
+
diff --git a/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE
new file mode 100644
index 0000000000000000000000000000000000000000..7a254464cc78ccea32b3ded00513c44c4e4da412
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE
@@ -0,0 +1,21 @@
+The MIT License (MIT)
+
+Copyright (c) 2025 Sebastián Ramírez
+
+Permission is hereby granted, free of charge, to any person obtaining a copy
+of this software and associated documentation files (the "Software"), to deal
+in the Software without restriction, including without limitation the rights
+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+copies of the Software, and to permit persons to whom the Software is
+furnished to do so, subject to the following conditions:
+
+The above copyright notice and this permission notice shall be included in
+all copies or substantial portions of the Software.
+
+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+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.
diff --git a/venv/lib/python3.11/site-packages/annotated_doc/__init__.py b/venv/lib/python3.11/site-packages/annotated_doc/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..a0152a7d12abc2db37fb26e764a61e0c894a43f3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc/__init__.py
@@ -0,0 +1,3 @@
+from .main import Doc as Doc
+
+__version__ = "0.0.4"
diff --git a/venv/lib/python3.11/site-packages/annotated_doc/main.py b/venv/lib/python3.11/site-packages/annotated_doc/main.py
new file mode 100644
index 0000000000000000000000000000000000000000..7063c59e4500a1d02bfc9b41887f9e95f8163507
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_doc/main.py
@@ -0,0 +1,36 @@
+class Doc:
+ """Define the documentation of a type annotation using `Annotated`, to be
+ used in class attributes, function and method parameters, return values,
+ and variables.
+
+ The value should be a positional-only string literal to allow static tools
+ like editors and documentation generators to use it.
+
+ This complements docstrings.
+
+ The string value passed is available in the attribute `documentation`.
+
+ Example:
+
+ ```Python
+ from typing import Annotated
+ from annotated_doc import Doc
+
+ def hi(name: Annotated[str, Doc("Who to say hi to")]) -> None:
+ print(f"Hi, {name}!")
+ ```
+ """
+
+ def __init__(self, documentation: str, /) -> None:
+ self.documentation = documentation
+
+ def __repr__(self) -> str:
+ return f"Doc({self.documentation!r})"
+
+ def __hash__(self) -> int:
+ return hash(self.documentation)
+
+ def __eq__(self, other: object) -> bool:
+ if not isinstance(other, Doc):
+ return NotImplemented
+ return self.documentation == other.documentation
diff --git a/venv/lib/python3.11/site-packages/annotated_doc/py.typed b/venv/lib/python3.11/site-packages/annotated_doc/py.typed
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/INSTALLER b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/INSTALLER
new file mode 100644
index 0000000000000000000000000000000000000000..a1b589e38a32041e49332e5e81c2d363dc418d68
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/INSTALLER
@@ -0,0 +1 @@
+pip
diff --git a/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/METADATA b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/METADATA
new file mode 100644
index 0000000000000000000000000000000000000000..eb5a08847f8a91e58ad8325a7059715e9be44ed6
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/METADATA
@@ -0,0 +1,295 @@
+Metadata-Version: 2.4
+Name: annotated-types
+Version: 0.8.0
+Summary: Reusable constraint types to use with typing.Annotated
+Project-URL: Homepage, https://github.com/annotated-types/annotated-types
+Project-URL: Source, https://github.com/annotated-types/annotated-types
+Project-URL: Changelog, https://github.com/annotated-types/annotated-types/releases
+Author-email: Adrian Garcia Badaracco <1755071+adriangb@users.noreply.github.com>, Samuel Colvin , Zac Hatfield-Dodds
+License-Expression: MIT
+License-File: LICENSE
+Classifier: Development Status :: 4 - Beta
+Classifier: Environment :: Console
+Classifier: Environment :: MacOS X
+Classifier: Intended Audience :: Developers
+Classifier: Intended Audience :: Information Technology
+Classifier: License :: OSI Approved :: MIT License
+Classifier: Operating System :: POSIX :: Linux
+Classifier: Operating System :: Unix
+Classifier: Programming Language :: Python :: 3 :: Only
+Classifier: Programming Language :: Python :: 3.10
+Classifier: Programming Language :: Python :: 3.11
+Classifier: Programming Language :: Python :: 3.12
+Classifier: Programming Language :: Python :: 3.13
+Classifier: Programming Language :: Python :: 3.14
+Classifier: Topic :: Software Development :: Libraries :: Python Modules
+Classifier: Typing :: Typed
+Requires-Python: >=3.10
+Description-Content-Type: text/markdown
+
+# annotated-types
+
+[](https://github.com/annotated-types/annotated-types/actions?query=event%3Apush+branch%3Amain+workflow%3ACI)
+[](https://pypi.python.org/pypi/annotated-types)
+[](https://github.com/annotated-types/annotated-types)
+[](https://github.com/annotated-types/annotated-types/blob/main/LICENSE)
+
+[PEP-593](https://peps.python.org/pep-0593/) added `typing.Annotated` as a way of
+adding context-specific metadata to existing types, and specifies that
+`Annotated[T, x]` _should_ be treated as `T` by any tool or library without special
+logic for `x`.
+
+This package provides metadata objects which can be used to represent common
+constraints such as upper and lower bounds on scalar values and collection sizes,
+a `Predicate` marker for runtime checks, and
+descriptions of how we intend these metadata to be interpreted. In some cases,
+we also note alternative representations which do not require this package.
+
+## Install
+
+```bash
+pip install annotated-types
+```
+
+## Examples
+
+```python
+from typing import Annotated
+from annotated_types import Gt, Len, Predicate
+
+class MyClass:
+ age: Annotated[int, Gt(18)] # Valid: 19, 20, ...
+ # Invalid: 17, 18, "19", 19.0, ...
+ factors: list[Annotated[int, Predicate(is_prime)]] # Valid: 2, 3, 5, 7, 11, ...
+ # Invalid: 4, 8, -2, 5.0, "prime", ...
+
+ my_list: Annotated[list[int], Len(0, 10)] # Valid: [], [10, 20, 30, 40, 50]
+ # Invalid: (1, 2), ["abc"], [0] * 20
+```
+
+## Documentation
+
+_While `annotated-types` avoids runtime checks for performance, users should not
+construct invalid combinations such as `MultipleOf("non-numeric")` or `Annotated[int, Len(3)]`.
+Downstream implementors may choose to raise an error, emit a warning, silently ignore
+a metadata item, etc., if the metadata objects described below are used with an
+incompatible type - or for any other reason!_
+
+### Gt, Ge, Lt, Le
+
+Express inclusive and/or exclusive bounds on orderable values - which may be numbers,
+dates, times, strings, sets, etc. Note that the boundary value need not be of the
+same type that was annotated, so long as they can be compared: `Annotated[int, Gt(1.5)]`
+is fine, for example, and implies that the value is an integer x such that `x > 1.5`.
+
+We suggest that implementors may also interpret `functools.partial(operator.le, 1.5)`
+as being equivalent to `Gt(1.5)`, for users who wish to avoid a runtime dependency on
+the `annotated-types` package.
+
+To be explicit, these types have the following meanings:
+
+* `Gt(x)` - value must be "Greater Than" `x` - equivalent to exclusive minimum
+* `Ge(x)` - value must be "Greater than or Equal" to `x` - equivalent to inclusive minimum
+* `Lt(x)` - value must be "Less Than" `x` - equivalent to exclusive maximum
+* `Le(x)` - value must be "Less than or Equal" to `x` - equivalent to inclusive maximum
+
+### Interval
+
+`Interval(gt, ge, lt, le)` allows you to specify an upper and lower bound with a single
+metadata object. `None` attributes should be ignored, and non-`None` attributes
+treated as per the single bounds above.
+
+### MultipleOf
+
+`MultipleOf(multiple_of=x)` might be interpreted in two ways:
+
+1. Python semantics, implying `value % multiple_of == 0`, or
+2. [JSONschema semantics](https://json-schema.org/draft/2020-12/json-schema-validation.html#rfc.section.6.2.1),
+ where `int(value / multiple_of) == value / multiple_of`.
+
+We encourage users to be aware of these two common interpretations and their
+distinct behaviours, especially since very large or non-integer numbers make
+it easy to cause silent data corruption due to floating-point imprecision.
+
+We encourage libraries to carefully document which interpretation they implement.
+
+### MinLen, MaxLen, Len
+
+`Len()` implies that `min_length <= len(value) <= max_length` - lower and upper bounds are inclusive.
+
+As well as `Len()` which can optionally include upper and lower bounds, we also
+provide `MinLen(x)` and `MaxLen(y)` which are equivalent to `Len(min_length=x)`
+and `Len(max_length=y)` respectively.
+
+`Len`, `MinLen`, and `MaxLen` may be used with any type which supports `len(value)`.
+
+Examples of usage:
+
+* `Annotated[list, MaxLen(10)]` (or `Annotated[list, Len(max_length=10)]`) - list must have a length of 10 or less
+* `Annotated[str, MaxLen(10)]` - string must have a length of 10 or less
+* `Annotated[list, MinLen(3)]` (or `Annotated[list, Len(min_length=3)]`) - list must have a length of 3 or more
+* `Annotated[list, Len(4, 6)]` - list must have a length of 4, 5, or 6
+* `Annotated[list, Len(8, 8)]` - list must have a length of exactly 8
+
+#### Changed in v0.4.0
+
+* `min_inclusive` has been renamed to `min_length`, no change in meaning
+* `max_exclusive` has been renamed to `max_length`, upper bound is now **inclusive** instead of **exclusive**
+* The recommendation that slices are interpreted as `Len` has been removed due to ambiguity and different semantic
+ meaning of the upper bound in slices vs. `Len`
+
+See [issue #23](https://github.com/annotated-types/annotated-types/issues/23) for discussion.
+
+### Timezone
+
+`Timezone` can be used with a `datetime` or a `time` to express which timezones
+are allowed. `Annotated[datetime, Timezone(None)]` must be a naive datetime.
+`Timezone[...]` ([literal ellipsis](https://docs.python.org/3/library/constants.html#Ellipsis))
+expresses that any timezone-aware datetime is allowed. You may also pass a specific
+timezone string or [`tzinfo`](https://docs.python.org/3/library/datetime.html#tzinfo-objects)
+object such as `Timezone(timezone.utc)` or `Timezone("Africa/Abidjan")` to express that you only
+allow a specific timezone, though we note that this is often a symptom of fragile design.
+
+#### Changed in v0.x.x
+
+* `Timezone` accepts [`tzinfo`](https://docs.python.org/3/library/datetime.html#tzinfo-objects) objects instead of
+ `timezone`, extending compatibility to [`zoneinfo`](https://docs.python.org/3/library/zoneinfo.html) and third party libraries.
+
+### Unit
+
+`Unit(unit: str)` expresses that the annotated numeric value is the magnitude of
+a quantity with the specified unit. For example, `Annotated[float, Unit("m/s")]`
+would be a float representing a velocity in meters per second.
+
+Please note that `annotated_types` itself makes no attempt to parse or validate
+the unit string in any way. That is left entirely to downstream libraries,
+such as [`pint`](https://pint.readthedocs.io) or
+[`astropy.units`](https://docs.astropy.org/en/stable/units/).
+
+An example of how a library might use this metadata:
+
+```python
+from annotated_types import Unit
+from typing import Annotated, TypeVar, Callable, Any, get_origin, get_args
+
+# given a type annotated with a unit:
+Meters = Annotated[float, Unit("m")]
+
+
+# you can cast the annotation to a specific unit type with any
+# callable that accepts a string and returns the desired type
+T = TypeVar("T")
+def cast_unit(tp: Any, unit_cls: Callable[[str], T]) -> T | None:
+ if get_origin(tp) is Annotated:
+ for arg in get_args(tp):
+ if isinstance(arg, Unit):
+ return unit_cls(arg.unit)
+ return None
+
+
+# using `pint`
+import pint
+pint_unit = cast_unit(Meters, pint.Unit)
+
+
+# using `astropy.units`
+import astropy.units as u
+astropy_unit = cast_unit(Meters, u.Unit)
+```
+
+### Predicate
+
+`Predicate(func: Callable)` expresses that `func(value)` is truthy for valid values.
+Users should prefer the statically inspectable metadata above, but if you need
+the full power and flexibility of arbitrary runtime predicates... here it is.
+
+For some common constraints, we provide generic types:
+
+* `LowerCase = Annotated[T, Predicate(str.islower)]`
+* `UpperCase = Annotated[T, Predicate(str.isupper)]`
+* `IsDigit = Annotated[T, Predicate(str.isdigit)]`
+* `IsFinite = Annotated[T, Predicate(math.isfinite)]`
+* `IsNotFinite = Annotated[T, Predicate(Not(math.isfinite))]`
+* `IsNan = Annotated[T, Predicate(math.isnan)]`
+* `IsNotNan = Annotated[T, Predicate(Not(math.isnan))]`
+* `IsInfinite = Annotated[T, Predicate(math.isinf)]`
+* `IsNotInfinite = Annotated[T, Predicate(Not(math.isinf))]`
+
+so that you can write e.g. `x: IsFinite[float] = 2.0` instead of the longer
+(but exactly equivalent) `x: Annotated[float, Predicate(math.isfinite)] = 2.0`.
+
+Some libraries might have special logic to handle known or understandable predicates,
+for example by checking for `str.isdigit` and using its presence to both call custom
+logic to enforce digit-only strings, and customise some generated external schema.
+Users are therefore encouraged to avoid indirection like `lambda s: s.lower()`, in
+favor of introspectable methods such as `str.lower` or `re.compile("pattern").search`.
+
+To enable basic negation of commonly used predicates like `math.isnan` without introducing introspection that makes it impossible for implementers to introspect the predicate we provide a `Not` wrapper that simply negates the predicate in an introspectable manner. Several of the predicates listed above are created in this manner.
+
+We do not specify what behaviour should be expected for predicates that raise
+an exception. For example `Annotated[int, Predicate(str.isdigit)]` might silently
+skip invalid constraints, or statically raise an error; or it might try calling it
+and then propagate or discard the resulting
+`TypeError: descriptor 'isdigit' for 'str' objects doesn't apply to a 'int' object`
+exception. We encourage libraries to document the behaviour they choose.
+
+### Doc
+
+`doc()` can be used to add documentation information in `Annotated`, for function and method parameters, variables, class attributes, return types, and any place where `Annotated` can be used.
+
+It expects a value that can be statically analyzed, as the main use case is for static analysis, editors, documentation generators, and similar tools.
+
+It returns a `DocInfo` class with a single attribute `documentation` containing the value passed to `doc()`.
+
+This is the early adopter's alternative form of the [`typing-doc` proposal](https://github.com/tiangolo/fastapi/blob/typing-doc/typing_doc.md).
+
+### Integrating downstream types with `GroupedMetadata`
+
+Implementers may choose to provide a convenience wrapper that groups multiple pieces of metadata.
+This can help reduce verbosity and cognitive overhead for users.
+For example, an implementer like Pydantic might provide a `Field` or `Meta` type that accepts keyword arguments and transforms these into low-level metadata:
+
+```python
+from dataclasses import dataclass
+from typing import Iterator
+from annotated_types import GroupedMetadata, Ge
+
+@dataclass
+class Field(GroupedMetadata):
+ ge: int | None = None
+ description: str | None = None
+
+ def __iter__(self) -> Iterator[object]:
+ # Iterating over a GroupedMetadata object should yield annotated-types
+ # constraint metadata objects which describe it as fully as possible,
+ # and may include other unknown objects too.
+ if self.ge is not None:
+ yield Ge(self.ge)
+ if self.description is not None:
+ yield Description(self.description)
+```
+
+Libraries consuming annotated-types constraints should check for `GroupedMetadata` and unpack it by iterating over the object and treating the results as if they had been "unpacked" in the `Annotated` type. The same logic should be applied to the [PEP 646 `Unpack` type](https://peps.python.org/pep-0646/), so that `Annotated[T, Field(...)]`, `Annotated[T, Unpack[Field(...)]]` and `Annotated[T, *Field(...)]` are all treated consistently.
+
+Libraries consuming annotated-types should also ignore any metadata they do not recongize that came from unpacking a `GroupedMetadata`, just like they ignore unrecognized metadata in `Annotated` itself.
+
+Our own `annotated_types.Interval` class is a `GroupedMetadata` which unpacks itself into `Gt`, `Lt`, etc., so this is not an abstract concern. Similarly, `annotated_types.Len` is a `GroupedMetadata` which unpacks itself into `MinLen` (optionally) and `MaxLen`.
+
+### Consuming metadata
+
+We intend to not be prescriptive as to _how_ the metadata and constraints are used, but as an example of how one might parse constraints from types annotations see our [implementation in `test_main.py`](https://github.com/annotated-types/annotated-types/blob/f59cf6d1b5255a0fe359b93896759a180bec30ae/tests/test_main.py#L94-L103).
+
+It is up to the implementer to determine how this metadata is used.
+You could use the metadata for runtime type checking, for generating schemas or to generate example data, amongst other use cases.
+
+## Design & History
+
+This package was designed at the PyCon 2022 sprints by the maintainers of Pydantic
+and Hypothesis, with the goal of making it as easy as possible for end-users to
+provide more informative annotations for use by runtime libraries.
+
+It is deliberately minimal, and following PEP-593 allows considerable downstream
+discretion in what (if anything!) they choose to support. Nonetheless, we expect
+that staying simple and covering _only_ the most common use-cases will give users
+and maintainers the best experience we can. If you'd like more constraints for your
+types - follow our lead, by defining them and documenting them downstream!
diff --git a/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/RECORD b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/RECORD
new file mode 100644
index 0000000000000000000000000000000000000000..744476917caa6b674d4c4a9890674a6d7dd5eb79
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/RECORD
@@ -0,0 +1,10 @@
+annotated_types-0.8.0.dist-info/INSTALLER,sha256=zuuue4knoyJ-UwPPXg8fezS7VCrXJQrAP7zeNuwvFQg,4
+annotated_types-0.8.0.dist-info/METADATA,sha256=YUmFsnj2Abjvhj-CeDuZHMOtpLifWS63HHxIxfLvIpg,15009
+annotated_types-0.8.0.dist-info/RECORD,,
+annotated_types-0.8.0.dist-info/WHEEL,sha256=lCkmxWfQsSc9CfIClYeavTdQeEX2toPqufh9gI35EQA,87
+annotated_types-0.8.0.dist-info/licenses/LICENSE,sha256=_hBJiEsaDZNCkB6I4H8ykl0ksxIdmXK2poBfuYJLCV0,1083
+annotated_types/__init__.py,sha256=pxBKTUObJ6n3T8C-I2ubobeDHmBEAmgCogWrwSmKm8g,13273
+annotated_types/__pycache__/__init__.cpython-311.pyc,,
+annotated_types/__pycache__/test_cases.cpython-311.pyc,,
+annotated_types/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+annotated_types/test_cases.py,sha256=2GdHKstXuBzpf3iXKCjmvFZpyk2zP9Hm9kXGGDhGBo4,6310
diff --git a/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/WHEEL b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/WHEEL
new file mode 100644
index 0000000000000000000000000000000000000000..7401812e19af977cc5088f3b8fb1ef6bc0441c0a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/WHEEL
@@ -0,0 +1,4 @@
+Wheel-Version: 1.0
+Generator: hatchling 1.31.0
+Root-Is-Purelib: true
+Tag: py3-none-any
diff --git a/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/licenses/LICENSE b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/licenses/LICENSE
new file mode 100644
index 0000000000000000000000000000000000000000..d99323a9965f146d5b0888c4ca1bf0727e12b04f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_types-0.8.0.dist-info/licenses/LICENSE
@@ -0,0 +1,21 @@
+The MIT License (MIT)
+
+Copyright (c) 2022 the contributors
+
+Permission is hereby granted, free of charge, to any person obtaining a copy
+of this software and associated documentation files (the "Software"), to deal
+in the Software without restriction, including without limitation the rights
+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+copies of the Software, and to permit persons to whom the Software is
+furnished to do so, subject to the following conditions:
+
+The above copyright notice and this permission notice shall be included in all
+copies or substantial portions of the Software.
+
+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+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.
diff --git a/venv/lib/python3.11/site-packages/annotated_types/__init__.py b/venv/lib/python3.11/site-packages/annotated_types/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..dcb35c54a56b0729f90278c22b402c73127be4d8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_types/__init__.py
@@ -0,0 +1,416 @@
+import math
+import types
+from collections.abc import Callable, Iterator
+from dataclasses import dataclass
+from datetime import tzinfo
+from types import EllipsisType
+from typing import (
+ TYPE_CHECKING,
+ Annotated,
+ Any,
+ Literal,
+ Protocol,
+ SupportsFloat,
+ SupportsIndex,
+ TypeVar,
+ Union,
+ runtime_checkable,
+)
+
+__all__ = (
+ 'BaseMetadata',
+ 'GroupedMetadata',
+ 'Gt',
+ 'Ge',
+ 'Lt',
+ 'Le',
+ 'Interval',
+ 'MultipleOf',
+ 'MinLen',
+ 'MaxLen',
+ 'Len',
+ 'Timezone',
+ 'Predicate',
+ 'LowerCase',
+ 'UpperCase',
+ 'IsDigits',
+ 'IsFinite',
+ 'IsNotFinite',
+ 'IsNan',
+ 'IsNotNan',
+ 'IsInfinite',
+ 'IsNotInfinite',
+ 'doc',
+ 'DocInfo',
+ '__version__',
+)
+
+__version__ = '0.8.0'
+
+
+T = TypeVar('T')
+
+
+# arguments that start with __ are considered
+# positional only
+# see https://peps.python.org/pep-0484/#positional-only-arguments
+
+
+class SupportsGt(Protocol):
+ def __gt__(self: T, __other: T) -> bool:
+ ...
+
+
+class SupportsGe(Protocol):
+ def __ge__(self: T, __other: T) -> bool:
+ ...
+
+
+class SupportsLt(Protocol):
+ def __lt__(self: T, __other: T) -> bool:
+ ...
+
+
+class SupportsLe(Protocol):
+ def __le__(self: T, __other: T) -> bool:
+ ...
+
+
+class SupportsMod(Protocol):
+ def __mod__(self: T, __other: T) -> T:
+ ...
+
+
+class SupportsDiv(Protocol):
+ def __div__(self: T, __other: T) -> T:
+ ...
+
+
+class BaseMetadata:
+ """Base class for all metadata.
+
+ This exists mainly so that implementers
+ can do `isinstance(..., BaseMetadata)` while traversing field annotations.
+ """
+
+ __slots__ = ()
+
+
+@dataclass(frozen=True, slots=True)
+class Gt(BaseMetadata):
+ """Gt(gt=x) implies that the value must be greater than x.
+
+ It can be used with any type that supports the ``>`` operator,
+ including numbers, dates and times, strings, sets, and so on.
+ """
+
+ gt: SupportsGt
+
+
+@dataclass(frozen=True, slots=True)
+class Ge(BaseMetadata):
+ """Ge(ge=x) implies that the value must be greater than or equal to x.
+
+ It can be used with any type that supports the ``>=`` operator,
+ including numbers, dates and times, strings, sets, and so on.
+ """
+
+ ge: SupportsGe
+
+
+@dataclass(frozen=True, slots=True)
+class Lt(BaseMetadata):
+ """Lt(lt=x) implies that the value must be less than x.
+
+ It can be used with any type that supports the ``<`` operator,
+ including numbers, dates and times, strings, sets, and so on.
+ """
+
+ lt: SupportsLt
+
+
+@dataclass(frozen=True, slots=True)
+class Le(BaseMetadata):
+ """Le(le=x) implies that the value must be less than or equal to x.
+
+ It can be used with any type that supports the ``<=`` operator,
+ including numbers, dates and times, strings, sets, and so on.
+ """
+
+ le: SupportsLe
+
+
+@runtime_checkable
+class GroupedMetadata(Protocol):
+ """A grouping of multiple objects, like typing.Unpack.
+
+ `GroupedMetadata` on its own is not metadata and has no meaning.
+ All of the constraints and metadata should be fully expressable
+ in terms of the `BaseMetadata`'s returned by `GroupedMetadata.__iter__()`.
+
+ Concrete implementations should override `GroupedMetadata.__iter__()`
+ to add their own metadata.
+ For example:
+
+ >>> @dataclass
+ >>> class Field(GroupedMetadata):
+ >>> gt: float | None = None
+ >>> description: str | None = None
+ ...
+ >>> def __iter__(self) -> Iterable[object]:
+ >>> if self.gt is not None:
+ >>> yield Gt(self.gt)
+ >>> if self.description is not None:
+ >>> yield Description(self.gt)
+
+ Also see the implementation of `Interval` below for an example.
+
+ Parsers should recognize this and unpack it so that it can be used
+ both with and without unpacking:
+
+ - `Annotated[int, Field(...)]` (parser must unpack Field)
+ - `Annotated[int, *Field(...)]` (PEP-646)
+ """ # noqa: trailing-whitespace
+
+ @property
+ def __is_annotated_types_grouped_metadata__(self) -> Literal[True]:
+ return True
+
+ def __iter__(self) -> Iterator[object]:
+ ...
+
+ if not TYPE_CHECKING:
+ __slots__ = () # allow subclasses to use slots
+
+ def __init_subclass__(cls, *args: Any, **kwargs: Any) -> None:
+ # Basic ABC like functionality without the complexity of an ABC
+ super().__init_subclass__(*args, **kwargs)
+ if cls.__iter__ is GroupedMetadata.__iter__:
+ raise TypeError("Can't subclass GroupedMetadata without implementing __iter__")
+
+ def __iter__(self) -> Iterator[object]: # noqa: F811
+ raise NotImplementedError # more helpful than "None has no attribute..." type errors
+
+
+@dataclass(frozen=True, kw_only=True, slots=True)
+class Interval(GroupedMetadata):
+ """Interval can express inclusive or exclusive bounds with a single object.
+
+ It accepts keyword arguments ``gt``, ``ge``, ``lt``, and/or ``le``, which
+ are interpreted the same way as the single-bound constraints.
+ """
+
+ gt: SupportsGt | None = None
+ ge: SupportsGe | None = None
+ lt: SupportsLt | None = None
+ le: SupportsLe | None = None
+
+ def __iter__(self) -> Iterator[BaseMetadata]:
+ """Unpack an Interval into zero or more single-bounds."""
+ if self.gt is not None:
+ yield Gt(self.gt)
+ if self.ge is not None:
+ yield Ge(self.ge)
+ if self.lt is not None:
+ yield Lt(self.lt)
+ if self.le is not None:
+ yield Le(self.le)
+
+
+@dataclass(frozen=True, slots=True)
+class MultipleOf(BaseMetadata):
+ """MultipleOf(multiple_of=x) might be interpreted in two ways:
+
+ 1. Python semantics, implying ``value % multiple_of == 0``, or
+ 2. JSONschema semantics, where ``int(value / multiple_of) == value / multiple_of``
+
+ We encourage users to be aware of these two common interpretations,
+ and libraries to carefully document which they implement.
+ """
+
+ multiple_of: SupportsDiv | SupportsMod
+
+
+@dataclass(frozen=True, slots=True)
+class MinLen(BaseMetadata):
+ """
+ MinLen() implies minimum inclusive length,
+ e.g. ``len(value) >= min_length``.
+ """
+
+ min_length: Annotated[int, Ge(0)]
+
+
+@dataclass(frozen=True, slots=True)
+class MaxLen(BaseMetadata):
+ """
+ MaxLen() implies maximum inclusive length,
+ e.g. ``len(value) <= max_length``.
+ """
+
+ max_length: Annotated[int, Ge(0)]
+
+
+@dataclass(frozen=True, slots=True)
+class Len(GroupedMetadata):
+ """
+ Len() implies that ``min_length <= len(value) <= max_length``.
+
+ Upper bound may be omitted or ``None`` to indicate no upper length bound.
+ """
+
+ min_length: Annotated[int, Ge(0)] = 0
+ max_length: Annotated[int, Ge(0)] | None = None
+
+ def __iter__(self) -> Iterator[BaseMetadata]:
+ """Unpack a Len into zero or more single-bounds."""
+ if self.min_length > 0:
+ yield MinLen(self.min_length)
+ if self.max_length is not None:
+ yield MaxLen(self.max_length)
+
+
+@dataclass(frozen=True, slots=True)
+class Timezone(BaseMetadata):
+ """Timezone(tz=...) requires a datetime to be aware (or ``tz=None``, naive).
+
+ ``Annotated[datetime, Timezone(None)]`` must be a naive datetime.
+ ``Timezone(...)`` (the ellipsis literal) expresses that the datetime must be
+ tz-aware but any timezone is allowed.
+
+ You may also pass a specific timezone string or tzinfo object such as
+ ``Timezone(timezone.utc)`` or ``Timezone("Africa/Abidjan")`` to express that
+ you only allow a specific timezone, though we note that this is often
+ a symptom of poor design.
+ """
+
+ tz: str | tzinfo | EllipsisType | None
+
+
+@dataclass(frozen=True, slots=True)
+class Unit(BaseMetadata):
+ """Indicates that the value is a physical quantity with the specified unit.
+
+ It is intended for usage with numeric types, where the value represents the
+ magnitude of the quantity. For example, ``distance: Annotated[float, Unit('m')]``
+ or ``speed: Annotated[float, Unit('m/s')]``.
+
+ Interpretation of the unit string is left to the discretion of the consumer.
+ It is suggested to follow conventions established by python libraries that work
+ with physical quantities, such as
+
+ - ``pint`` :
+ - ``astropy.units``:
+
+ For indicating a quantity with a certain dimensionality but without a specific unit
+ it is recommended to use square brackets, e.g. `Annotated[float, Unit('[time]')]`.
+ Note, however, ``annotated_types`` itself makes no use of the unit string.
+ """
+
+ unit: str
+
+
+@dataclass(frozen=True, slots=True)
+class Predicate(BaseMetadata):
+ """``Predicate(func: Callable)`` implies `func(value)` is truthy for valid values.
+
+ Users should prefer statically inspectable metadata, but if you need the full
+ power and flexibility of arbitrary runtime predicates... here it is.
+
+ We provide a few predefined predicates for common string constraints:
+ ``LowerCase = Predicate(str.islower)``, ``UpperCase = Predicate(str.isupper)``, and
+ ``IsDigits = Predicate(str.isdigit)``. Users are encouraged to use methods which
+ can be given special handling, and avoid indirection like ``lambda s: s.lower()``.
+
+ Some libraries might have special logic to handle certain predicates, e.g. by
+ checking for `str.isdigit` and using its presence to both call custom logic to
+ enforce digit-only strings, and customise some generated external schema.
+
+ We do not specify what behaviour should be expected for predicates that raise
+ an exception. For example `Annotated[int, Predicate(str.isdigit)]` might silently
+ skip invalid constraints, or statically raise an error; or it might try calling it
+ and then propagate or discard the resulting exception.
+ """
+
+ func: Callable[[Any], bool]
+
+ def __repr__(self) -> str:
+ if getattr(self.func, "__name__", "") == "":
+ return f"{self.__class__.__name__}({self.func!r})"
+ if isinstance(self.func, (types.MethodType, types.BuiltinMethodType)) and (
+ namespace := getattr(self.func.__self__, "__name__", None)
+ ):
+ return f"{self.__class__.__name__}({namespace}.{self.func.__name__})"
+ if isinstance(self.func, type(str.isascii)): # method descriptor
+ return f"{self.__class__.__name__}({self.func.__qualname__})"
+ return f"{self.__class__.__name__}({self.func.__name__})"
+
+
+@dataclass
+class Not:
+ func: Callable[[Any], bool]
+
+ def __call__(self, __v: Any) -> bool:
+ return not self.func(__v)
+
+
+_StrType = TypeVar("_StrType", bound=str)
+
+LowerCase = Annotated[_StrType, Predicate(str.islower)]
+"""
+Return True if the string is a lowercase string, False otherwise.
+
+A string is lowercase if all cased characters in the string are lowercase and there is at least one cased character in the string.
+""" # noqa: E501
+UpperCase = Annotated[_StrType, Predicate(str.isupper)]
+"""
+Return True if the string is an uppercase string, False otherwise.
+
+A string is uppercase if all cased characters in the string are uppercase and there is at least one cased character in the string.
+""" # noqa: E501
+IsDigit = Annotated[_StrType, Predicate(str.isdigit)]
+IsDigits = IsDigit # type: ignore # plural for backwards compatibility, see #63
+"""
+Return True if the string is a digit string, False otherwise.
+
+A string is a digit string if all characters in the string are digits and there is at least one character in the string.
+""" # noqa: E501
+IsAscii = Annotated[_StrType, Predicate(str.isascii)]
+"""
+Return True if all characters in the string are ASCII, False otherwise.
+
+ASCII characters have code points in the range U+0000-U+007F. Empty string is ASCII too.
+"""
+
+_NumericType = TypeVar('_NumericType', bound=Union[SupportsFloat, SupportsIndex])
+IsFinite = Annotated[_NumericType, Predicate(math.isfinite)]
+"""Return True if x is neither an infinity nor a NaN, and False otherwise."""
+IsNotFinite = Annotated[_NumericType, Predicate(Not(math.isfinite))]
+"""Return True if x is one of infinity or NaN, and False otherwise"""
+IsNan = Annotated[_NumericType, Predicate(math.isnan)]
+"""Return True if x is a NaN (not a number), and False otherwise."""
+IsNotNan = Annotated[_NumericType, Predicate(Not(math.isnan))]
+"""Return True if x is anything but NaN (not a number), and False otherwise."""
+IsInfinite = Annotated[_NumericType, Predicate(math.isinf)]
+"""Return True if x is a positive or negative infinity, and False otherwise."""
+IsNotInfinite = Annotated[_NumericType, Predicate(Not(math.isinf))]
+"""Return True if x is neither a positive or negative infinity, and False otherwise."""
+
+try:
+ # PEP 727 – Documentation in Annotated Metadata
+ from typing_extensions import Doc # type: ignore[attr-defined]
+except ImportError:
+
+ @dataclass(frozen=True, slots=True)
+ class Doc: # type: ignore [no-redef]
+ """ "
+ The return value of doc(), mainly to be used by tools that want to extract the
+ Annotated documentation at runtime.
+ """
+
+ documentation: str
+ """The documentation string passed to doc()."""
+
+
+DocInfo = Doc # backwards compatibility
+doc = Doc
diff --git a/venv/lib/python3.11/site-packages/annotated_types/py.typed b/venv/lib/python3.11/site-packages/annotated_types/py.typed
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/annotated_types/test_cases.py b/venv/lib/python3.11/site-packages/annotated_types/test_cases.py
new file mode 100644
index 0000000000000000000000000000000000000000..fba3c257354ef7064e148f956de3733e709402f6
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/annotated_types/test_cases.py
@@ -0,0 +1,146 @@
+import math
+from collections.abc import Iterable, Iterator
+from datetime import date, datetime, timedelta, timezone
+from decimal import Decimal
+from typing import Annotated, Any, NamedTuple
+
+import annotated_types as at
+
+
+class Case(NamedTuple):
+ """
+ A test case for `annotated_types`.
+ """
+
+ annotation: Any
+ valid_cases: Iterable[Any]
+ invalid_cases: Iterable[Any]
+
+
+def cases() -> Iterable[Case]:
+ # Gt, Ge, Lt, Le
+ yield Case(Annotated[int, at.Gt(4)], (5, 6, 1000), (4, 0, -1))
+ yield Case(Annotated[float, at.Gt(0.5)], (0.6, 0.7, 0.8, 0.9), (0.5, 0.0, -0.1))
+ yield Case(
+ Annotated[datetime, at.Gt(datetime(2000, 1, 1))],
+ [datetime(2000, 1, 2), datetime(2000, 1, 3)],
+ [datetime(2000, 1, 1), datetime(1999, 12, 31)],
+ )
+ yield Case(
+ Annotated[datetime, at.Gt(date(2000, 1, 1))],
+ [date(2000, 1, 2), date(2000, 1, 3)],
+ [date(2000, 1, 1), date(1999, 12, 31)],
+ )
+ yield Case(
+ Annotated[datetime, at.Gt(Decimal('1.123'))],
+ [Decimal('1.1231'), Decimal('123')],
+ [Decimal('1.123'), Decimal('0')],
+ )
+
+ yield Case(Annotated[int, at.Ge(4)], (4, 5, 6, 1000, 4), (0, -1))
+ yield Case(Annotated[float, at.Ge(0.5)], (0.5, 0.6, 0.7, 0.8, 0.9), (0.4, 0.0, -0.1))
+ yield Case(
+ Annotated[datetime, at.Ge(datetime(2000, 1, 1))],
+ [datetime(2000, 1, 2), datetime(2000, 1, 3)],
+ [datetime(1998, 1, 1), datetime(1999, 12, 31)],
+ )
+
+ yield Case(Annotated[int, at.Lt(4)], (0, -1), (4, 5, 6, 1000, 4))
+ yield Case(Annotated[float, at.Lt(0.5)], (0.4, 0.0, -0.1), (0.5, 0.6, 0.7, 0.8, 0.9))
+ yield Case(
+ Annotated[datetime, at.Lt(datetime(2000, 1, 1))],
+ [datetime(1999, 12, 31), datetime(1999, 12, 31)],
+ [datetime(2000, 1, 2), datetime(2000, 1, 3)],
+ )
+
+ yield Case(Annotated[int, at.Le(4)], (4, 0, -1), (5, 6, 1000))
+ yield Case(Annotated[float, at.Le(0.5)], (0.5, 0.0, -0.1), (0.6, 0.7, 0.8, 0.9))
+ yield Case(
+ Annotated[datetime, at.Le(datetime(2000, 1, 1))],
+ [datetime(2000, 1, 1), datetime(1999, 12, 31)],
+ [datetime(2000, 1, 2), datetime(2000, 1, 3)],
+ )
+
+ # Interval
+ yield Case(Annotated[int, at.Interval(gt=4)], (5, 6, 1000), (4, 0, -1))
+ yield Case(Annotated[int, at.Interval(gt=4, lt=10)], (5, 6), (4, 10, 1000, 0, -1))
+ yield Case(Annotated[float, at.Interval(ge=0.5, le=1)], (0.5, 0.9, 1), (0.49, 1.1))
+ yield Case(
+ Annotated[datetime, at.Interval(gt=datetime(2000, 1, 1), le=datetime(2000, 1, 3))],
+ [datetime(2000, 1, 2), datetime(2000, 1, 3)],
+ [datetime(2000, 1, 1), datetime(2000, 1, 4)],
+ )
+
+ yield Case(Annotated[int, at.MultipleOf(multiple_of=3)], (0, 3, 9), (1, 2, 4))
+ yield Case(Annotated[float, at.MultipleOf(multiple_of=0.5)], (0, 0.5, 1, 1.5), (0.4, 1.1))
+
+ # lengths
+
+ yield Case(Annotated[str, at.MinLen(3)], ('123', '1234', 'x' * 10), ('', '1', '12'))
+ yield Case(Annotated[str, at.Len(3)], ('123', '1234', 'x' * 10), ('', '1', '12'))
+ yield Case(Annotated[list[int], at.MinLen(3)], ([1, 2, 3], [1, 2, 3, 4], [1] * 10), ([], [1], [1, 2]))
+ yield Case(Annotated[list[int], at.Len(3)], ([1, 2, 3], [1, 2, 3, 4], [1] * 10), ([], [1], [1, 2]))
+
+ yield Case(Annotated[str, at.MaxLen(4)], ('', '1234'), ('12345', 'x' * 10))
+ yield Case(Annotated[str, at.Len(0, 4)], ('', '1234'), ('12345', 'x' * 10))
+ yield Case(Annotated[list[str], at.MaxLen(4)], ([], ['a', 'bcdef'], ['a', 'b', 'c']), (['a'] * 5, ['b'] * 10))
+ yield Case(Annotated[list[str], at.Len(0, 4)], ([], ['a', 'bcdef'], ['a', 'b', 'c']), (['a'] * 5, ['b'] * 10))
+
+ yield Case(Annotated[str, at.Len(3, 5)], ('123', '12345'), ('', '1', '12', '123456', 'x' * 10))
+ yield Case(Annotated[str, at.Len(3, 3)], ('123',), ('12', '1234'))
+
+ yield Case(Annotated[dict[int, int], at.Len(2, 3)], [{1: 1, 2: 2}], [{}, {1: 1}, {1: 1, 2: 2, 3: 3, 4: 4}])
+ yield Case(Annotated[set[int], at.Len(2, 3)], ({1, 2}, {1, 2, 3}), (set(), {1}, {1, 2, 3, 4}))
+ yield Case(Annotated[tuple[int, ...], at.Len(2, 3)], ((1, 2), (1, 2, 3)), ((), (1,), (1, 2, 3, 4)))
+
+ # Timezone
+
+ yield Case(
+ Annotated[datetime, at.Timezone(None)], [datetime(2000, 1, 1)], [datetime(2000, 1, 1, tzinfo=timezone.utc)]
+ )
+ yield Case(
+ Annotated[datetime, at.Timezone(...)], [datetime(2000, 1, 1, tzinfo=timezone.utc)], [datetime(2000, 1, 1)]
+ )
+ yield Case(
+ Annotated[datetime, at.Timezone(timezone.utc)],
+ [datetime(2000, 1, 1, tzinfo=timezone.utc)],
+ [datetime(2000, 1, 1), datetime(2000, 1, 1, tzinfo=timezone(timedelta(hours=6)))],
+ )
+ yield Case(
+ Annotated[datetime, at.Timezone('Europe/London')],
+ [datetime(2000, 1, 1, tzinfo=timezone(timedelta(0), name='Europe/London'))],
+ [datetime(2000, 1, 1), datetime(2000, 1, 1, tzinfo=timezone(timedelta(hours=6)))],
+ )
+
+ # Quantity
+
+ yield Case(Annotated[float, at.Unit(unit='m')], (5, 4.2), ('5m', '4.2m'))
+
+ # predicate types
+
+ yield Case(at.LowerCase[str], ['abc', 'foobar'], ['', 'A', 'Boom'])
+ yield Case(at.UpperCase[str], ['ABC', 'DEFO'], ['', 'a', 'abc', 'AbC'])
+ yield Case(at.IsDigit[str], ['123'], ['', 'ab', 'a1b2'])
+ yield Case(at.IsAscii[str], ['123', 'foo bar'], ['£100', '😊', 'whatever 👀'])
+
+ yield Case(Annotated[int, at.Predicate(lambda x: x % 2 == 0)], [0, 2, 4], [1, 3, 5])
+
+ yield Case(at.IsFinite[float], [1.23], [math.nan, math.inf, -math.inf])
+ yield Case(at.IsNotFinite[float], [math.nan, math.inf], [1.23])
+ yield Case(at.IsNan[float], [math.nan], [1.23, math.inf])
+ yield Case(at.IsNotNan[float], [1.23, math.inf], [math.nan])
+ yield Case(at.IsInfinite[float], [math.inf], [math.nan, 1.23])
+ yield Case(at.IsNotInfinite[float], [math.nan, 1.23], [math.inf])
+
+ # check stacked predicates
+ yield Case(at.IsInfinite[Annotated[float, at.Predicate(lambda x: x > 0)]], [math.inf], [-math.inf, 1.23, math.nan])
+
+ # doc
+ yield Case(Annotated[int, at.doc("A number")], [1, 2], [])
+
+ # custom GroupedMetadata
+ class MyCustomGroupedMetadata(at.GroupedMetadata):
+ def __iter__(self) -> Iterator[at.Predicate]:
+ yield at.Predicate(lambda x: float(x).is_integer())
+
+ yield Case(Annotated[float, MyCustomGroupedMetadata()], [0, 2.0], [0.01, 1.5])
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/INSTALLER b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/INSTALLER
new file mode 100644
index 0000000000000000000000000000000000000000..a1b589e38a32041e49332e5e81c2d363dc418d68
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/INSTALLER
@@ -0,0 +1 @@
+pip
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/METADATA b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/METADATA
new file mode 100644
index 0000000000000000000000000000000000000000..126a728ba495267b9d9208a6df9ff46221700e3a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/METADATA
@@ -0,0 +1,107 @@
+Metadata-Version: 2.4
+Name: anyio
+Version: 4.14.2
+Summary: High-level concurrency and networking framework on top of asyncio or Trio
+Author-email: Alex Grönholm
+License-Expression: MIT
+Project-URL: Documentation, https://anyio.readthedocs.io/en/latest/
+Project-URL: Changelog, https://anyio.readthedocs.io/en/stable/versionhistory.html
+Project-URL: Source code, https://github.com/agronholm/anyio
+Project-URL: Issue tracker, https://github.com/agronholm/anyio/issues
+Classifier: Development Status :: 5 - Production/Stable
+Classifier: Intended Audience :: Developers
+Classifier: Framework :: AnyIO
+Classifier: Typing :: Typed
+Classifier: Programming Language :: Python
+Classifier: Programming Language :: Python :: 3
+Classifier: Programming Language :: Python :: 3.10
+Classifier: Programming Language :: Python :: 3.11
+Classifier: Programming Language :: Python :: 3.12
+Classifier: Programming Language :: Python :: 3.13
+Classifier: Programming Language :: Python :: 3.14
+Classifier: Programming Language :: Python :: 3.15
+Requires-Python: >=3.10
+Description-Content-Type: text/x-rst
+License-File: LICENSE
+Requires-Dist: exceptiongroup>=1.0.2; python_version < "3.11"
+Requires-Dist: idna>=2.8
+Requires-Dist: typing_extensions>=4.5; python_version < "3.13"
+Provides-Extra: trio
+Requires-Dist: trio>=0.32.0; extra == "trio"
+Dynamic: license-file
+
+.. image:: https://github.com/agronholm/anyio/actions/workflows/test.yml/badge.svg
+ :target: https://github.com/agronholm/anyio/actions/workflows/test.yml
+ :alt: Build Status
+.. image:: https://coveralls.io/repos/github/agronholm/anyio/badge.svg?branch=master
+ :target: https://coveralls.io/github/agronholm/anyio?branch=master
+ :alt: Code Coverage
+.. image:: https://readthedocs.org/projects/anyio/badge/?version=latest
+ :target: https://anyio.readthedocs.io/en/latest/?badge=latest
+ :alt: Documentation
+.. image:: https://badges.gitter.im/gitterHQ/gitter.svg
+ :target: https://gitter.im/python-trio/AnyIO
+ :alt: Gitter chat
+.. image:: https://tidelift.com/badges/package/pypi/anyio
+ :target: https://tidelift.com/subscription/pkg/pypi-anyio
+ :alt: Tidelift
+
+AnyIO is an asynchronous networking and concurrency library that works on top of either asyncio_ or
+Trio_. It implements Trio-like `structured concurrency`_ (SC) on top of asyncio and works in harmony
+with the native SC of Trio itself.
+
+Applications and libraries written against AnyIO's API will run unmodified on either asyncio_ or
+Trio_. AnyIO can also be adopted into a library or application incrementally – bit by bit, no full
+refactoring necessary. It will blend in with the native libraries of your chosen backend.
+
+To find out why you might want to use AnyIO's APIs instead of asyncio's, you can read about it
+`here `_.
+
+Documentation
+-------------
+
+View full documentation at: https://anyio.readthedocs.io/
+
+Features
+--------
+
+AnyIO offers the following functionality:
+
+* Task groups (nurseries_ in trio terminology)
+* High-level networking (TCP, UDP and UNIX sockets)
+
+ * `Happy eyeballs`_ algorithm for TCP connections (more robust than that of asyncio on Python
+ 3.8)
+ * async/await style UDP sockets (unlike asyncio where you still have to use Transports and
+ Protocols)
+
+* A versatile API for byte streams and object streams
+* Inter-task synchronization and communication (locks, conditions, events, semaphores, object
+ streams)
+* Worker threads
+* Subprocesses
+* Subinterpreter support for code parallelization (on Python 3.13 and later)
+* Asynchronous file I/O (using worker threads)
+* Signal handling
+* Asynchronous versions of the functools_ and itertools_ modules
+
+AnyIO also comes with its own pytest_ plugin which also supports asynchronous fixtures.
+It even works with the popular Hypothesis_ library.
+
+.. _asyncio: https://docs.python.org/3/library/asyncio.html
+.. _Trio: https://github.com/python-trio/trio
+.. _structured concurrency: https://en.wikipedia.org/wiki/Structured_concurrency
+.. _nurseries: https://trio.readthedocs.io/en/stable/reference-core.html#nurseries-and-spawning
+.. _Happy eyeballs: https://en.wikipedia.org/wiki/Happy_Eyeballs
+.. _pytest: https://docs.pytest.org/en/latest/
+.. _functools: https://docs.python.org/3/library/functools.html
+.. _itertools: https://docs.python.org/3/library/itertools.html
+.. _Hypothesis: https://hypothesis.works/
+
+Security contact information
+----------------------------
+
+To report a security vulnerability, please use the `Tidelift security contact`_.
+Tidelift will coordinate the fix and disclosure.
+
+.. _Tidelift security contact: https://tidelift.com/security
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/RECORD b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/RECORD
new file mode 100644
index 0000000000000000000000000000000000000000..72f493d98fe5e177e1526a8fec884d4ed852fdac
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/RECORD
@@ -0,0 +1,96 @@
+anyio-4.14.2.dist-info/INSTALLER,sha256=zuuue4knoyJ-UwPPXg8fezS7VCrXJQrAP7zeNuwvFQg,4
+anyio-4.14.2.dist-info/METADATA,sha256=xeb8Tf2DMxmROyh756MNmTXpIIx--RdjBiGFWqrE1Bk,4645
+anyio-4.14.2.dist-info/RECORD,,
+anyio-4.14.2.dist-info/WHEEL,sha256=K260EYznzXsJYBQGqmI8VTxEdiZYNvDZwW9cBh9-_MA,91
+anyio-4.14.2.dist-info/entry_points.txt,sha256=_d6Yu6uiaZmNe0CydowirE9Cmg7zUL2g08tQpoS3Qvc,39
+anyio-4.14.2.dist-info/licenses/LICENSE,sha256=U2GsncWPLvX9LpsJxoKXwX8ElQkJu8gCO9uC6s8iwrA,1081
+anyio-4.14.2.dist-info/scm_file_list.json,sha256=wDSXGv8Ehn5ZW5BhB-RlaAc16zY_OfO27qrlMfMMZy8,3654
+anyio-4.14.2.dist-info/scm_version.json,sha256=KgaUx31SyaqGFQFKpG3FO9kCk-ygOKS7Uf0yk56OzaY,161
+anyio-4.14.2.dist-info/top_level.txt,sha256=QglSMiWX8_5dpoVAEIHdEYzvqFMdSYWmCj6tYw2ITkQ,6
+anyio/__init__.py,sha256=HitUIfzvAojSeaHVmJ9rFn8k_yI63G6s_jUL2QChf4U,6405
+anyio/__pycache__/__init__.cpython-311.pyc,,
+anyio/__pycache__/from_thread.cpython-311.pyc,,
+anyio/__pycache__/functools.cpython-311.pyc,,
+anyio/__pycache__/itertools.cpython-311.pyc,,
+anyio/__pycache__/lowlevel.cpython-311.pyc,,
+anyio/__pycache__/pytest_plugin.cpython-311.pyc,,
+anyio/__pycache__/to_interpreter.cpython-311.pyc,,
+anyio/__pycache__/to_process.cpython-311.pyc,,
+anyio/__pycache__/to_thread.cpython-311.pyc,,
+anyio/_backends/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+anyio/_backends/__pycache__/__init__.cpython-311.pyc,,
+anyio/_backends/__pycache__/_asyncio.cpython-311.pyc,,
+anyio/_backends/__pycache__/_trio.cpython-311.pyc,,
+anyio/_backends/_asyncio.py,sha256=eK0j8PI0O7y3bX8cUBsafoWP8wsXR-szP6ciyP4Mkkk,104400
+anyio/_backends/_trio.py,sha256=hoTXL8zI81v4UE-YT_cr2fzh_zY1mKdmZD3oQHmeR0E,45583
+anyio/_core/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+anyio/_core/__pycache__/__init__.cpython-311.pyc,,
+anyio/_core/__pycache__/_asyncio_selector_thread.cpython-311.pyc,,
+anyio/_core/__pycache__/_contextmanagers.cpython-311.pyc,,
+anyio/_core/__pycache__/_eventloop.cpython-311.pyc,,
+anyio/_core/__pycache__/_exceptions.cpython-311.pyc,,
+anyio/_core/__pycache__/_fileio.cpython-311.pyc,,
+anyio/_core/__pycache__/_resources.cpython-311.pyc,,
+anyio/_core/__pycache__/_signals.cpython-311.pyc,,
+anyio/_core/__pycache__/_sockets.cpython-311.pyc,,
+anyio/_core/__pycache__/_streams.cpython-311.pyc,,
+anyio/_core/__pycache__/_subprocesses.cpython-311.pyc,,
+anyio/_core/__pycache__/_synchronization.cpython-311.pyc,,
+anyio/_core/__pycache__/_tasks.cpython-311.pyc,,
+anyio/_core/__pycache__/_tempfile.cpython-311.pyc,,
+anyio/_core/__pycache__/_testing.cpython-311.pyc,,
+anyio/_core/__pycache__/_typedattr.cpython-311.pyc,,
+anyio/_core/_asyncio_selector_thread.py,sha256=2PdxFM3cs02Kp6BSppbvmRT7q7asreTW5FgBxEsflBo,5626
+anyio/_core/_contextmanagers.py,sha256=YInBCabiEeS-UaP_Jdxa1CaFC71ETPW8HZTHIM8Rsc8,7215
+anyio/_core/_eventloop.py,sha256=ByZUeJD9alMfcyTseRo5IzTO0IltEul_Gyq9iqSjqDk,6658
+anyio/_core/_exceptions.py,sha256=OfzLO4Z3Hog1TnipbIn72YNtkoYxS4lHW9MqKDeGc88,4936
+anyio/_core/_fileio.py,sha256=hHfyV0bXDL-R2ZNnInwse3nmTAd36AIz1cBxgmAwzAQ,31358
+anyio/_core/_resources.py,sha256=NbmU5O5UX3xEyACnkmYX28Fmwdl-f-ny0tHym26e0w0,435
+anyio/_core/_signals.py,sha256=mjTBB2hTKNPRlU0IhnijeQedpWOGERDiMjSlJQsFrug,1016
+anyio/_core/_sockets.py,sha256=HtjiSH-yzehlqh_LpD3PGIafDuUCwd8gueUuk5MFeNk,35288
+anyio/_core/_streams.py,sha256=FczFwIgDpnkK0bODWJXMpsUJYdvAD04kaUaGzJU8DK0,1806
+anyio/_core/_subprocesses.py,sha256=M2GCc4NKXCbB_GtEskJgndiM2b0VS0NK_ohmFxri7O8,7923
+anyio/_core/_synchronization.py,sha256=kgPk88-eVOmY-pDNs-ReRbcEelY1a7YczuVBPsZwc8A,21591
+anyio/_core/_tasks.py,sha256=y99vRi-AFzEv6kyvIwqrzuskE2NMiNItWnuWGk7eOr4,13126
+anyio/_core/_tempfile.py,sha256=jE2w59FRF3yRo4vjkjfZF2YcqsBZvc66VWRwrJGDYGk,19624
+anyio/_core/_testing.py,sha256=u7MPqGXwpTxqI7hclSdNA30z2GH1Nw258uwKvy_RfBg,2340
+anyio/_core/_typedattr.py,sha256=P4ozZikn3-DbpoYcvyghS_FOYAgbmUxeoU8-L_07pZM,2508
+anyio/abc/__init__.py,sha256=6mWhcl_pGXhrgZVHP_TCfMvIXIOp9mroEFM90fYCU_U,2869
+anyio/abc/__pycache__/__init__.cpython-311.pyc,,
+anyio/abc/__pycache__/_eventloop.cpython-311.pyc,,
+anyio/abc/__pycache__/_resources.cpython-311.pyc,,
+anyio/abc/__pycache__/_sockets.cpython-311.pyc,,
+anyio/abc/__pycache__/_streams.cpython-311.pyc,,
+anyio/abc/__pycache__/_subprocesses.cpython-311.pyc,,
+anyio/abc/__pycache__/_tasks.cpython-311.pyc,,
+anyio/abc/__pycache__/_testing.cpython-311.pyc,,
+anyio/abc/_eventloop.py,sha256=OqWYSEj0TmwL_xniCJt3_jHFWsuMk9THk8tCTGsKapI,10681
+anyio/abc/_resources.py,sha256=DrYvkNN1hH6Uvv5_5uKySvDsnknGVDe8FCKfko0VtN8,783
+anyio/abc/_sockets.py,sha256=OmVDrfemVvF9c5K1tpBgQyV6fn5v0XyCExLAqBOGz9o,13124
+anyio/abc/_streams.py,sha256=rnSwRy-6y80TPQtIXit5LjMMiU1CCWS1oMNsmJQHRTg,7582
+anyio/abc/_subprocesses.py,sha256=cumAPJTktOQtw63IqG0lDpyZqu_l1EElvQHMiwJgL08,2067
+anyio/abc/_tasks.py,sha256=m-FtE4phxeNIELSG7A3H7VUz3jA2Ib5J2JIew8-PS6o,6642
+anyio/abc/_testing.py,sha256=9YYM2AXsYFvf4PLjUEr6yRxDiUeB5QbY_gOg0X_C6lY,2034
+anyio/from_thread.py,sha256=JYsbaCaIB_Iit6kNhtXSteJGt4PcQ7ncq0nIpcelIrg,19265
+anyio/functools.py,sha256=T4JS8IXq-x1S0Lbo2owF8l9fza2KypO147QLeyz4cjs,11797
+anyio/itertools.py,sha256=QV-9mnRCr2yBph8g01QFvN-bQ_Yle-8Sl13YSydBlMI,16168
+anyio/lowlevel.py,sha256=-59484Z6K5jF7XJJhHxYHIKnvFH_uAeMhl9a_AfaSHI,6280
+anyio/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+anyio/pytest_plugin.py,sha256=paMpI_VMNQf2bir0LfvgMpXSiYJoHDzWdKUVTyoHmvQ,13609
+anyio/streams/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+anyio/streams/__pycache__/__init__.cpython-311.pyc,,
+anyio/streams/__pycache__/buffered.cpython-311.pyc,,
+anyio/streams/__pycache__/file.cpython-311.pyc,,
+anyio/streams/__pycache__/memory.cpython-311.pyc,,
+anyio/streams/__pycache__/stapled.cpython-311.pyc,,
+anyio/streams/__pycache__/text.cpython-311.pyc,,
+anyio/streams/__pycache__/tls.cpython-311.pyc,,
+anyio/streams/buffered.py,sha256=u7hCD8SNrYHcutG6K5wEiy1F88-pizgjvEFM22Kq2Cw,6746
+anyio/streams/file.py,sha256=6jujI2m-QJITqqKFamrupX_DNsU7y2Fz3omLZxOLuY0,4524
+anyio/streams/memory.py,sha256=ZmKWCLpyItOCXmvCQT-L8IyJHNFaB-OVJbrn88CaMo0,10776
+anyio/streams/stapled.py,sha256=mDNF9Gj4deXfOuKSZmgkEG-QExYAKAGjYBHbrs-rJaQ,4486
+anyio/streams/text.py,sha256=BcVAGJw1VRvtIqnv-o0Rb0pwH7p8vwlvl21xHq522ag,5765
+anyio/streams/tls.py,sha256=Gvs--YOoFxcyn-hakXOiPM8H-aEp8nsU6k9rTrKqJA4,15801
+anyio/to_interpreter.py,sha256=_mLngrMy97TMR6VbW4Y6YzDUk9ZuPcQMPlkuyRh3C9k,7100
+anyio/to_process.py,sha256=jHw7v6XHBNIfYa6DqnqKY0gsRo8ScOCh8fo19fqFvjU,9848
+anyio/to_thread.py,sha256=bYszW0lCDfTmLnXbAYic05HEe_Di_0UWKabgL2vU0T0,2750
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/WHEEL b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/WHEEL
new file mode 100644
index 0000000000000000000000000000000000000000..1d472b6c22838de58d1c3c0dd2c795a5cde9e415
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/WHEEL
@@ -0,0 +1,5 @@
+Wheel-Version: 1.0
+Generator: setuptools (83.0.0)
+Root-Is-Purelib: true
+Tag: py3-none-any
+
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/entry_points.txt b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/entry_points.txt
new file mode 100644
index 0000000000000000000000000000000000000000..44dd9bdc3039122cc98014c1439ca254313fd014
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/entry_points.txt
@@ -0,0 +1,2 @@
+[pytest11]
+anyio = anyio.pytest_plugin
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/licenses/LICENSE b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/licenses/LICENSE
new file mode 100644
index 0000000000000000000000000000000000000000..104eebf5a3002fccdaceef3a4cb936173c1c2035
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/licenses/LICENSE
@@ -0,0 +1,20 @@
+The MIT License (MIT)
+
+Copyright (c) 2018 Alex Grönholm
+
+Permission is hereby granted, free of charge, to any person obtaining a copy of
+this software and associated documentation files (the "Software"), to deal in
+the Software without restriction, including without limitation the rights to
+use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of
+the Software, and to permit persons to whom the Software is furnished to do so,
+subject to the following conditions:
+
+The above copyright notice and this permission notice shall be included in all
+copies or substantial portions of the Software.
+
+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, 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.
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/scm_file_list.json b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/scm_file_list.json
new file mode 100644
index 0000000000000000000000000000000000000000..72a48145f533714ddb8ec6198464e494fa5c0c13
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/scm_file_list.json
@@ -0,0 +1,119 @@
+{
+ "files": [
+ ".pre-commit-config.yaml",
+ "LICENSE",
+ "pyproject.toml",
+ "AGENTS.md",
+ "README.rst",
+ "CLAUDE.md",
+ ".readthedocs.yml",
+ ".gitignore",
+ "docs/tempfile.rst",
+ "docs/signals.rst",
+ "docs/synchronization.rst",
+ "docs/contextmanagers.rst",
+ "docs/testing.rst",
+ "docs/networking.rst",
+ "docs/contributing.rst",
+ "docs/index.rst",
+ "docs/versionhistory.rst",
+ "docs/threads.rst",
+ "docs/api.rst",
+ "docs/typedattrs.rst",
+ "docs/basics.rst",
+ "docs/fileio.rst",
+ "docs/cancellation.rst",
+ "docs/support.rst",
+ "docs/streams.rst",
+ "docs/why.rst",
+ "docs/tasks.rst",
+ "docs/migration.rst",
+ "docs/conf.py",
+ "docs/subprocesses.rst",
+ "docs/faq.rst",
+ "docs/subinterpreters.rst",
+ "src/anyio/functools.py",
+ "src/anyio/py.typed",
+ "src/anyio/__init__.py",
+ "src/anyio/pytest_plugin.py",
+ "src/anyio/itertools.py",
+ "src/anyio/to_interpreter.py",
+ "src/anyio/from_thread.py",
+ "src/anyio/to_process.py",
+ "src/anyio/to_thread.py",
+ "src/anyio/lowlevel.py",
+ "src/anyio/_backends/_trio.py",
+ "src/anyio/_backends/__init__.py",
+ "src/anyio/_backends/_asyncio.py",
+ "src/anyio/streams/memory.py",
+ "src/anyio/streams/__init__.py",
+ "src/anyio/streams/tls.py",
+ "src/anyio/streams/file.py",
+ "src/anyio/streams/text.py",
+ "src/anyio/streams/stapled.py",
+ "src/anyio/streams/buffered.py",
+ "src/anyio/abc/_eventloop.py",
+ "src/anyio/abc/__init__.py",
+ "src/anyio/abc/_sockets.py",
+ "src/anyio/abc/_tasks.py",
+ "src/anyio/abc/_subprocesses.py",
+ "src/anyio/abc/_resources.py",
+ "src/anyio/abc/_streams.py",
+ "src/anyio/abc/_testing.py",
+ "src/anyio/_core/_typedattr.py",
+ "src/anyio/_core/_eventloop.py",
+ "src/anyio/_core/__init__.py",
+ "src/anyio/_core/_tempfile.py",
+ "src/anyio/_core/_sockets.py",
+ "src/anyio/_core/_tasks.py",
+ "src/anyio/_core/_fileio.py",
+ "src/anyio/_core/_synchronization.py",
+ "src/anyio/_core/_subprocesses.py",
+ "src/anyio/_core/_resources.py",
+ "src/anyio/_core/_contextmanagers.py",
+ "src/anyio/_core/_exceptions.py",
+ "src/anyio/_core/_streams.py",
+ "src/anyio/_core/_signals.py",
+ "src/anyio/_core/_asyncio_selector_thread.py",
+ "src/anyio/_core/_testing.py",
+ "tests/test_itertools.py",
+ "tests/test_functools.py",
+ "tests/test_eventloop.py",
+ "tests/__init__.py",
+ "tests/test_to_thread.py",
+ "tests/test_from_thread.py",
+ "tests/test_lowlevel.py",
+ "tests/test_to_interpreter.py",
+ "tests/test_sockets.py",
+ "tests/test_typedattr.py",
+ "tests/test_to_process.py",
+ "tests/test_all_attributes.py",
+ "tests/test_synchronization.py",
+ "tests/test_debugging.py",
+ "tests/test_contextmanagers.py",
+ "tests/test_fileio.py",
+ "tests/conftest.py",
+ "tests/test_signals.py",
+ "tests/test_deprecations.py",
+ "tests/test_tempfile.py",
+ "tests/test_taskgroups.py",
+ "tests/test_pytest_plugin.py",
+ "tests/test_subprocesses.py",
+ "tests/streams/test_text.py",
+ "tests/streams/test_memory.py",
+ "tests/streams/__init__.py",
+ "tests/streams/test_file.py",
+ "tests/streams/test_stapled.py",
+ "tests/streams/test_tls.py",
+ "tests/streams/test_buffered.py",
+ ".github/pull_request_template.md",
+ ".github/dependabot.yml",
+ ".github/FUNDING.yml",
+ ".github/ISSUE_TEMPLATE/features_request.yaml",
+ ".github/ISSUE_TEMPLATE/bug_report.yaml",
+ ".github/ISSUE_TEMPLATE/config.yml",
+ ".github/workflows/test.yml",
+ ".github/workflows/test-downstream.yml",
+ ".github/workflows/publish.yml"
+ ]
+}
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/scm_version.json b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/scm_version.json
new file mode 100644
index 0000000000000000000000000000000000000000..13d71062b9fa3d7ce79a838448f646e6da5a3d2d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/scm_version.json
@@ -0,0 +1,8 @@
+{
+ "tag": "4.14.2",
+ "distance": 0,
+ "node": "gc384f99687c64c59ed8a11c3a0f11a2d57daff71",
+ "dirty": false,
+ "branch": "HEAD",
+ "node_date": "2026-07-12"
+}
diff --git a/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/top_level.txt b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/top_level.txt
new file mode 100644
index 0000000000000000000000000000000000000000..c77c069ecc9b7f8b1f97dbcfec905725db0253a8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio-4.14.2.dist-info/top_level.txt
@@ -0,0 +1 @@
+anyio
diff --git a/venv/lib/python3.11/site-packages/anyio/__init__.py b/venv/lib/python3.11/site-packages/anyio/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..2502c760bcc1d640be2de20f620d52fc3364cb55
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/__init__.py
@@ -0,0 +1,115 @@
+from __future__ import annotations
+
+from ._core._contextmanagers import AsyncContextManagerMixin as AsyncContextManagerMixin
+from ._core._contextmanagers import ContextManagerMixin as ContextManagerMixin
+from ._core._eventloop import current_time as current_time
+from ._core._eventloop import get_all_backends as get_all_backends
+from ._core._eventloop import get_available_backends as get_available_backends
+from ._core._eventloop import get_cancelled_exc_class as get_cancelled_exc_class
+from ._core._eventloop import run as run
+from ._core._eventloop import sleep as sleep
+from ._core._eventloop import sleep_forever as sleep_forever
+from ._core._eventloop import sleep_until as sleep_until
+from ._core._exceptions import BrokenResourceError as BrokenResourceError
+from ._core._exceptions import BrokenWorkerInterpreter as BrokenWorkerInterpreter
+from ._core._exceptions import BrokenWorkerProcess as BrokenWorkerProcess
+from ._core._exceptions import BusyResourceError as BusyResourceError
+from ._core._exceptions import ClosedResourceError as ClosedResourceError
+from ._core._exceptions import ConnectionFailed as ConnectionFailed
+from ._core._exceptions import DelimiterNotFound as DelimiterNotFound
+from ._core._exceptions import EndOfStream as EndOfStream
+from ._core._exceptions import IncompleteRead as IncompleteRead
+from ._core._exceptions import NoEventLoopError as NoEventLoopError
+from ._core._exceptions import RunFinishedError as RunFinishedError
+from ._core._exceptions import TaskCancelled as TaskCancelled
+from ._core._exceptions import TaskFailed as TaskFailed
+from ._core._exceptions import TaskNotFinished as TaskNotFinished
+from ._core._exceptions import TypedAttributeLookupError as TypedAttributeLookupError
+from ._core._exceptions import WouldBlock as WouldBlock
+from ._core._fileio import AsyncFile as AsyncFile
+from ._core._fileio import Path as Path
+from ._core._fileio import open_file as open_file
+from ._core._fileio import wrap_file as wrap_file
+from ._core._resources import aclose_forcefully as aclose_forcefully
+from ._core._signals import open_signal_receiver as open_signal_receiver
+from ._core._sockets import TCPConnectable as TCPConnectable
+from ._core._sockets import UNIXConnectable as UNIXConnectable
+from ._core._sockets import as_connectable as as_connectable
+from ._core._sockets import connect_tcp as connect_tcp
+from ._core._sockets import connect_unix as connect_unix
+from ._core._sockets import create_connected_udp_socket as create_connected_udp_socket
+from ._core._sockets import (
+ create_connected_unix_datagram_socket as create_connected_unix_datagram_socket,
+)
+from ._core._sockets import create_tcp_listener as create_tcp_listener
+from ._core._sockets import create_udp_socket as create_udp_socket
+from ._core._sockets import create_unix_datagram_socket as create_unix_datagram_socket
+from ._core._sockets import create_unix_listener as create_unix_listener
+from ._core._sockets import getaddrinfo as getaddrinfo
+from ._core._sockets import getnameinfo as getnameinfo
+from ._core._sockets import notify_closing as notify_closing
+from ._core._sockets import wait_readable as wait_readable
+from ._core._sockets import wait_socket_readable as wait_socket_readable
+from ._core._sockets import wait_socket_writable as wait_socket_writable
+from ._core._sockets import wait_writable as wait_writable
+from ._core._streams import create_memory_object_stream as create_memory_object_stream
+from ._core._subprocesses import open_process as open_process
+from ._core._subprocesses import run_process as run_process
+from ._core._synchronization import CapacityLimiter as CapacityLimiter
+from ._core._synchronization import (
+ CapacityLimiterStatistics as CapacityLimiterStatistics,
+)
+from ._core._synchronization import Condition as Condition
+from ._core._synchronization import ConditionStatistics as ConditionStatistics
+from ._core._synchronization import Event as Event
+from ._core._synchronization import EventStatistics as EventStatistics
+from ._core._synchronization import Lock as Lock
+from ._core._synchronization import LockStatistics as LockStatistics
+from ._core._synchronization import ResourceGuard as ResourceGuard
+from ._core._synchronization import Semaphore as Semaphore
+from ._core._synchronization import SemaphoreStatistics as SemaphoreStatistics
+from ._core._tasks import TASK_STATUS_IGNORED as TASK_STATUS_IGNORED
+from ._core._tasks import CancelScope as CancelScope
+from ._core._tasks import TaskHandle as TaskHandle
+from ._core._tasks import create_task_group as create_task_group
+from ._core._tasks import current_effective_deadline as current_effective_deadline
+from ._core._tasks import fail_after as fail_after
+from ._core._tasks import move_on_after as move_on_after
+from ._core._tempfile import NamedTemporaryFile as NamedTemporaryFile
+from ._core._tempfile import SpooledTemporaryFile as SpooledTemporaryFile
+from ._core._tempfile import TemporaryDirectory as TemporaryDirectory
+from ._core._tempfile import TemporaryFile as TemporaryFile
+from ._core._tempfile import gettempdir as gettempdir
+from ._core._tempfile import gettempdirb as gettempdirb
+from ._core._tempfile import mkdtemp as mkdtemp
+from ._core._tempfile import mkstemp as mkstemp
+from ._core._testing import TaskInfo as TaskInfo
+from ._core._testing import get_current_task as get_current_task
+from ._core._testing import get_running_tasks as get_running_tasks
+from ._core._testing import wait_all_tasks_blocked as wait_all_tasks_blocked
+from ._core._typedattr import TypedAttributeProvider as TypedAttributeProvider
+from ._core._typedattr import TypedAttributeSet as TypedAttributeSet
+from ._core._typedattr import typed_attribute as typed_attribute
+
+# Re-export imports so they look like they live directly in this package
+for __value in list(locals().values()):
+ if getattr(__value, "__module__", "").startswith("anyio."):
+ __value.__module__ = __name__
+
+
+del __value
+
+
+def __getattr__(attr: str) -> type[BrokenWorkerInterpreter]:
+ """Support deprecated aliases."""
+ if attr == "BrokenWorkerIntepreter":
+ import warnings
+
+ warnings.warn(
+ "The 'BrokenWorkerIntepreter' alias is deprecated, use 'BrokenWorkerInterpreter' instead.",
+ DeprecationWarning,
+ stacklevel=2,
+ )
+ return BrokenWorkerInterpreter
+
+ raise AttributeError(f"module {__name__!r} has no attribute {attr!r}")
diff --git a/venv/lib/python3.11/site-packages/anyio/_backends/__init__.py b/venv/lib/python3.11/site-packages/anyio/_backends/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/anyio/_backends/_asyncio.py b/venv/lib/python3.11/site-packages/anyio/_backends/_asyncio.py
new file mode 100644
index 0000000000000000000000000000000000000000..c00c2cd9be4d773c1ed38498659021bc45099e0c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_backends/_asyncio.py
@@ -0,0 +1,3136 @@
+from __future__ import annotations
+
+import array
+import asyncio
+import concurrent.futures
+import contextvars
+import math
+import os
+import socket
+import sys
+import threading
+import weakref
+from asyncio import (
+ AbstractEventLoop,
+ CancelledError,
+ all_tasks,
+ create_task,
+ current_task,
+ get_running_loop,
+ sleep,
+)
+from asyncio.base_events import _run_until_complete_cb # type: ignore[attr-defined]
+from collections import OrderedDict, deque
+from collections.abc import (
+ AsyncGenerator,
+ AsyncIterator,
+ Awaitable,
+ Callable,
+ Collection,
+ Coroutine,
+ Iterable,
+ Sequence,
+)
+from concurrent.futures import Future
+from contextlib import AbstractContextManager
+from contextvars import Context, copy_context
+from dataclasses import dataclass, field
+from functools import partial, wraps
+from inspect import (
+ CORO_RUNNING,
+ CORO_SUSPENDED,
+ getcoroutinestate,
+)
+from io import IOBase
+from os import PathLike
+from queue import Queue
+from signal import Signals
+from socket import AddressFamily, SocketKind
+from threading import Thread
+from types import CodeType, TracebackType
+from typing import (
+ IO,
+ TYPE_CHECKING,
+ Any,
+ Literal,
+ ParamSpec,
+ TypeVar,
+ cast,
+)
+from weakref import WeakKeyDictionary
+
+from .. import (
+ CapacityLimiterStatistics,
+ EventStatistics,
+ LockStatistics,
+ TaskInfo,
+ abc,
+)
+from .._core._eventloop import (
+ claim_worker_thread,
+ set_current_async_library,
+ threadlocals,
+)
+from .._core._exceptions import (
+ BrokenResourceError,
+ BusyResourceError,
+ ClosedResourceError,
+ EndOfStream,
+ RunFinishedError,
+ WouldBlock,
+)
+from .._core._sockets import convert_ipv6_sockaddr
+from .._core._streams import create_memory_object_stream
+from .._core._synchronization import (
+ CapacityLimiter as BaseCapacityLimiter,
+)
+from .._core._synchronization import Event as BaseEvent
+from .._core._synchronization import Lock as BaseLock
+from .._core._synchronization import (
+ ResourceGuard,
+ SemaphoreStatistics,
+)
+from .._core._synchronization import Semaphore as BaseSemaphore
+from .._core._tasks import CancelScope as BaseCancelScope
+from .._core._tasks import TaskHandle
+from ..abc import (
+ AsyncBackend,
+ IPSockAddrType,
+ SocketListener,
+ UDPPacketType,
+ UNIXDatagramPacketType,
+)
+from ..abc._eventloop import StrOrBytesPath
+from ..abc._tasks import call_for_coroutine, get_callable_name
+from ..lowlevel import RunVar
+from ..streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
+
+if TYPE_CHECKING:
+ from _typeshed import FileDescriptorLike
+else:
+ FileDescriptorLike = object
+
+if sys.version_info >= (3, 11):
+ from asyncio import Runner
+ from typing import TypeVarTuple, Unpack
+else:
+ import contextvars
+ import enum
+ import signal
+ from asyncio import coroutines, events, exceptions, tasks
+
+ from exceptiongroup import BaseExceptionGroup
+ from typing_extensions import TypeVarTuple, Unpack
+
+ class _State(enum.Enum):
+ CREATED = "created"
+ INITIALIZED = "initialized"
+ CLOSED = "closed"
+
+ class Runner:
+ # Copied from CPython 3.11
+ def __init__(
+ self,
+ *,
+ debug: bool | None = None,
+ loop_factory: Callable[[], AbstractEventLoop] | None = None,
+ ):
+ self._state = _State.CREATED
+ self._debug = debug
+ self._loop_factory = loop_factory
+ self._loop: AbstractEventLoop | None = None
+ self._context = None
+ self._interrupt_count = 0
+ self._set_event_loop = False
+
+ def __enter__(self) -> Runner:
+ self._lazy_init()
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self.close()
+
+ def close(self) -> None:
+ """Shutdown and close event loop."""
+ loop = self._loop
+ if self._state is not _State.INITIALIZED or loop is None:
+ return
+ try:
+ _cancel_all_tasks(loop)
+ loop.run_until_complete(loop.shutdown_asyncgens())
+ if hasattr(loop, "shutdown_default_executor"):
+ loop.run_until_complete(loop.shutdown_default_executor())
+ else:
+ loop.run_until_complete(_shutdown_default_executor(loop))
+ finally:
+ if self._set_event_loop:
+ events.set_event_loop(None)
+ loop.close()
+ self._loop = None
+ self._state = _State.CLOSED
+
+ def get_loop(self) -> AbstractEventLoop:
+ """Return embedded event loop."""
+ self._lazy_init()
+ return self._loop
+
+ def run(self, coro: Coroutine[T_Retval], *, context=None) -> T_Retval:
+ """Run a coroutine inside the embedded event loop."""
+ if not coroutines.iscoroutine(coro):
+ raise ValueError(f"a coroutine was expected, got {coro!r}")
+
+ if events._get_running_loop() is not None:
+ # fail fast with short traceback
+ raise RuntimeError(
+ "Runner.run() cannot be called from a running event loop"
+ )
+
+ self._lazy_init()
+
+ if context is None:
+ context = self._context
+ task = context.run(self._loop.create_task, coro)
+
+ if (
+ threading.current_thread() is threading.main_thread()
+ and signal.getsignal(signal.SIGINT) is signal.default_int_handler
+ ):
+ sigint_handler = partial(self._on_sigint, main_task=task)
+ try:
+ signal.signal(signal.SIGINT, sigint_handler)
+ except ValueError:
+ # `signal.signal` may throw if `threading.main_thread` does
+ # not support signals (e.g. embedded interpreter with signals
+ # not registered - see gh-91880)
+ sigint_handler = None
+ else:
+ sigint_handler = None
+
+ self._interrupt_count = 0
+ try:
+ return self._loop.run_until_complete(task)
+ except exceptions.CancelledError:
+ if self._interrupt_count > 0:
+ uncancel = getattr(task, "uncancel", None)
+ if uncancel is not None and uncancel() == 0:
+ raise KeyboardInterrupt # noqa: B904
+ raise # CancelledError
+ finally:
+ if (
+ sigint_handler is not None
+ and signal.getsignal(signal.SIGINT) is sigint_handler
+ ):
+ signal.signal(signal.SIGINT, signal.default_int_handler)
+
+ def _lazy_init(self) -> None:
+ if self._state is _State.CLOSED:
+ raise RuntimeError("Runner is closed")
+ if self._state is _State.INITIALIZED:
+ return
+ if self._loop_factory is None:
+ self._loop = events.new_event_loop()
+ if not self._set_event_loop:
+ # Call set_event_loop only once to avoid calling
+ # attach_loop multiple times on child watchers
+ events.set_event_loop(self._loop)
+ self._set_event_loop = True
+ else:
+ self._loop = self._loop_factory()
+ if self._debug is not None:
+ self._loop.set_debug(self._debug)
+ self._context = contextvars.copy_context()
+ self._state = _State.INITIALIZED
+
+ def _on_sigint(self, signum, frame, main_task: asyncio.Task) -> None:
+ self._interrupt_count += 1
+ if self._interrupt_count == 1 and not main_task.done():
+ main_task.cancel()
+ # wakeup loop if it is blocked by select() with long timeout
+ self._loop.call_soon_threadsafe(lambda: None)
+ return
+ raise KeyboardInterrupt()
+
+ def _cancel_all_tasks(loop: AbstractEventLoop) -> None:
+ to_cancel = tasks.all_tasks(loop)
+ if not to_cancel:
+ return
+
+ for task in to_cancel:
+ task.cancel()
+
+ loop.run_until_complete(tasks.gather(*to_cancel, return_exceptions=True))
+
+ for task in to_cancel:
+ if task.cancelled():
+ continue
+ if task.exception() is not None:
+ loop.call_exception_handler(
+ {
+ "message": "unhandled exception during asyncio.run() shutdown",
+ "exception": task.exception(),
+ "task": task,
+ }
+ )
+
+ async def _shutdown_default_executor(loop: AbstractEventLoop) -> None:
+ """Schedule the shutdown of the default executor."""
+
+ def _do_shutdown(future: asyncio.futures.Future) -> None:
+ try:
+ loop._default_executor.shutdown(wait=True) # type: ignore[attr-defined]
+ loop.call_soon_threadsafe(future.set_result, None)
+ except Exception as ex:
+ loop.call_soon_threadsafe(future.set_exception, ex)
+
+ loop._executor_shutdown_called = True
+ if loop._default_executor is None:
+ return
+ future = loop.create_future()
+ thread = threading.Thread(target=_do_shutdown, args=(future,))
+ thread.start()
+ try:
+ await future
+ finally:
+ thread.join()
+
+
+T_Retval = TypeVar("T_Retval")
+T_co = TypeVar("T_co", covariant=True)
+T_contra = TypeVar("T_contra", contravariant=True)
+PosArgsT = TypeVarTuple("PosArgsT")
+P = ParamSpec("P")
+
+_root_task: RunVar[asyncio.Task | None] = RunVar("_root_task")
+
+
+def find_root_task() -> asyncio.Task:
+ root_task = _root_task.get(None)
+ if root_task is not None and not root_task.done():
+ return root_task
+
+ # Look for a task that has been started via run_until_complete()
+ for task in all_tasks():
+ if task._callbacks and not task.done():
+ callbacks = [cb for cb, context in task._callbacks]
+ for cb in callbacks:
+ if (
+ cb is _run_until_complete_cb
+ or getattr(cb, "__module__", None) == "uvloop.loop"
+ ):
+ _root_task.set(task)
+ return task
+
+ # Look up the topmost task in the AnyIO task tree, if possible
+ task = cast(asyncio.Task, current_task())
+ state = _task_states.get(task)
+ if state:
+ cancel_scope = state.cancel_scope
+ while cancel_scope and cancel_scope._parent_scope is not None:
+ cancel_scope = cancel_scope._parent_scope
+
+ if cancel_scope is not None:
+ return cast(asyncio.Task, cancel_scope._host_task)
+
+ return task
+
+
+#
+# Event loop
+#
+
+_run_vars: WeakKeyDictionary[asyncio.AbstractEventLoop, Any] = WeakKeyDictionary()
+
+
+def _task_started(task: asyncio.Task) -> bool:
+ """Return ``True`` if the task has been started and has not finished."""
+ # The task coro should never be None here, as we never add finished tasks to the
+ # task list
+ coro = task.get_coro()
+ assert coro is not None
+ return getcoroutinestate(coro) in (CORO_RUNNING, CORO_SUSPENDED)
+
+
+#
+# Timeouts and cancellation
+#
+
+
+def is_anyio_cancellation(exc: CancelledError) -> bool:
+ # Sometimes third party frameworks catch a CancelledError and raise a new one, so as
+ # a workaround we have to look at the previous ones in __context__ too for a
+ # matching cancel message
+ while True:
+ if (
+ exc.args
+ and isinstance(exc.args[0], str)
+ and exc.args[0].startswith("Cancelled via cancel scope ")
+ ):
+ return True
+
+ if isinstance(exc.__context__, CancelledError):
+ exc = exc.__context__
+ continue
+
+ return False
+
+
+class CancelScope(BaseCancelScope):
+ __slots__ = (
+ "_active",
+ "_cancel_called",
+ "_cancel_handle",
+ "_cancel_reason",
+ "_cancelled_caught",
+ "_child_scopes",
+ "_deadline",
+ "_host_task",
+ "_parent_scope",
+ "_pending_uncancellations",
+ "_shield",
+ "_tasks",
+ "_timeout_handle",
+ )
+
+ def __new__(
+ cls, *, deadline: float = math.inf, shield: bool = False
+ ) -> CancelScope:
+ return object.__new__(cls)
+
+ def __init__(self, deadline: float = math.inf, shield: bool = False):
+ self._deadline = deadline
+ self._shield = shield
+ self._parent_scope: CancelScope | None = None
+ self._child_scopes: set[CancelScope] = set()
+ self._cancel_called = False
+ self._cancel_reason: str | None = None
+ self._cancelled_caught = False
+ self._active = False
+ self._timeout_handle: asyncio.TimerHandle | None = None
+ self._cancel_handle: asyncio.Handle | None = None
+ self._tasks: set[asyncio.Task] = set()
+ self._host_task: asyncio.Task | None = None
+ if sys.version_info >= (3, 11):
+ self._pending_uncancellations: int | None = 0
+ else:
+ self._pending_uncancellations = None
+
+ def __enter__(self) -> CancelScope:
+ if self._active:
+ raise RuntimeError(
+ "Each CancelScope may only be used for a single 'with' block"
+ )
+
+ self._host_task = host_task = cast(asyncio.Task, current_task())
+ self._tasks.add(host_task)
+ try:
+ task_state = _task_states[host_task]
+ except KeyError:
+ task_state = TaskState(None, self)
+ _task_states[host_task] = task_state
+ else:
+ self._parent_scope = task_state.cancel_scope
+ task_state.cancel_scope = self
+ if self._parent_scope is not None:
+ # If using an eager task factory, the parent scope may not even contain
+ # the host task
+ self._parent_scope._child_scopes.add(self)
+ self._parent_scope._tasks.discard(host_task)
+
+ self._timeout()
+ self._active = True
+
+ # Start cancelling the host task if the scope was cancelled before entering
+ if self._cancel_called:
+ self._deliver_cancellation(self)
+
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool:
+ del exc_tb
+
+ if not self._active:
+ raise RuntimeError("This cancel scope is not active")
+ if current_task() is not self._host_task:
+ raise RuntimeError(
+ "Attempted to exit cancel scope in a different task than it was "
+ "entered in"
+ )
+
+ assert self._host_task is not None
+ host_task_state = _task_states.get(self._host_task)
+ if host_task_state is None or host_task_state.cancel_scope is not self:
+ raise RuntimeError(
+ "Attempted to exit a cancel scope that isn't the current tasks's "
+ "current cancel scope"
+ )
+
+ try:
+ self._active = False
+ if self._timeout_handle:
+ self._timeout_handle.cancel()
+ self._timeout_handle = None
+
+ self._tasks.remove(self._host_task)
+ if self._parent_scope is not None:
+ self._parent_scope._child_scopes.remove(self)
+ self._parent_scope._tasks.add(self._host_task)
+
+ host_task_state.cancel_scope = self._parent_scope
+
+ # Restart the cancellation effort in the closest visible, cancelled parent
+ # scope if necessary
+ self._restart_cancellation_in_parent()
+
+ # We only swallow the exception iff it was an AnyIO CancelledError, either
+ # directly as exc_val or inside an exception group and there are no cancelled
+ # parent cancel scopes visible to us here
+ if self._cancel_called and not self._parent_cancellation_is_visible_to_us:
+ # For each level-cancel() call made on the host task, call uncancel()
+ while self._pending_uncancellations:
+ self._host_task.uncancel()
+ self._pending_uncancellations -= 1
+
+ # Update cancelled_caught and check for exceptions we must not swallow
+ if isinstance(exc_val, BaseExceptionGroup):
+ cancelleds_caught, remaining = exc_val.split(
+ lambda exc: (
+ isinstance(exc, CancelledError)
+ and is_anyio_cancellation(exc)
+ )
+ )
+
+ if cancelleds_caught is None:
+ return False
+
+ self._cancelled_caught = True
+
+ if remaining is None:
+ return True
+
+ context = remaining.__context__
+ try:
+ # Preserve __cause__ and __suppress_context__ by avoiding `raise
+ # ... from ...`
+ raise remaining
+ finally:
+ # Preserve __context__
+ remaining.__context__ = context
+ del context
+ else:
+ if isinstance(exc_val, CancelledError) and is_anyio_cancellation(
+ exc_val
+ ):
+ self._cancelled_caught = True
+ return True
+ else:
+ return False
+ else:
+ if self._pending_uncancellations:
+ assert self._parent_scope is not None
+ assert self._parent_scope._pending_uncancellations is not None
+ self._parent_scope._pending_uncancellations += (
+ self._pending_uncancellations
+ )
+ self._pending_uncancellations = 0
+
+ return False
+ finally:
+ self._host_task = None
+ del exc_val
+
+ @property
+ def _effectively_cancelled(self) -> bool:
+ cancel_scope: CancelScope | None = self
+ while cancel_scope is not None:
+ if cancel_scope._cancel_called:
+ return True
+
+ if cancel_scope.shield:
+ return False
+
+ cancel_scope = cancel_scope._parent_scope
+
+ return False
+
+ @property
+ def _parent_cancellation_is_visible_to_us(self) -> bool:
+ return (
+ self._parent_scope is not None
+ and not self.shield
+ and self._parent_scope._effectively_cancelled
+ )
+
+ def _timeout(self) -> None:
+ if self._deadline != math.inf:
+ loop = get_running_loop()
+ if loop.time() >= self._deadline:
+ self.cancel("deadline exceeded")
+ else:
+ self._timeout_handle = loop.call_at(self._deadline, self._timeout)
+
+ def _deliver_cancellation(self, origin: CancelScope) -> bool:
+ """
+ Deliver cancellation to directly contained tasks and nested cancel scopes.
+
+ Schedule another run at the end if we still have tasks eligible for
+ cancellation.
+
+ :param origin: the cancel scope that originated the cancellation
+ :return: ``True`` if the delivery needs to be retried on the next cycle
+
+ """
+ should_retry = False
+ current = current_task()
+ for task in self._tasks:
+ # Always skip tasks that are already done (see issue #1111)
+ if task.done():
+ continue
+
+ should_retry = True
+ if task._must_cancel: # type: ignore[attr-defined]
+ continue
+
+ # The task is eligible for cancellation if it has started
+ if task is not current and (task is self._host_task or _task_started(task)):
+ waiter = task._fut_waiter # type: ignore[attr-defined]
+ if not isinstance(waiter, asyncio.Future) or not waiter.done():
+ task.cancel(origin._cancel_reason)
+ if (
+ task is origin._host_task
+ and origin._pending_uncancellations is not None
+ ):
+ origin._pending_uncancellations += 1
+
+ # Deliver cancellation to child scopes that aren't shielded or running their own
+ # cancellation callbacks
+ for scope in self._child_scopes:
+ if not scope._shield and not scope.cancel_called:
+ should_retry = scope._deliver_cancellation(origin) or should_retry
+
+ # Schedule another callback if there are still tasks left
+ if origin is self:
+ if should_retry:
+ self._cancel_handle = get_running_loop().call_soon(
+ self._deliver_cancellation, origin
+ )
+ else:
+ self._cancel_handle = None
+
+ return should_retry
+
+ def _restart_cancellation_in_parent(self) -> None:
+ """
+ Restart the cancellation effort in the closest directly cancelled parent scope.
+
+ """
+ scope = self._parent_scope
+ while scope is not None:
+ if scope._cancel_called:
+ if scope._cancel_handle is None:
+ scope._deliver_cancellation(scope)
+
+ break
+
+ # No point in looking beyond any shielded scope
+ if scope._shield:
+ break
+
+ scope = scope._parent_scope
+
+ def cancel(self, reason: str | None = None) -> None:
+ if not self._cancel_called:
+ if self._timeout_handle:
+ self._timeout_handle.cancel()
+ self._timeout_handle = None
+
+ self._cancel_called = True
+ self._cancel_reason = f"Cancelled via cancel scope {id(self):x}"
+ if task := current_task():
+ self._cancel_reason += f" by {task}"
+
+ if reason:
+ self._cancel_reason += f"; reason: {reason}"
+
+ if self._host_task is not None:
+ self._deliver_cancellation(self)
+
+ @property
+ def deadline(self) -> float:
+ return self._deadline
+
+ @deadline.setter
+ def deadline(self, value: float) -> None:
+ self._deadline = float(value)
+ if self._timeout_handle is not None:
+ self._timeout_handle.cancel()
+ self._timeout_handle = None
+
+ if self._active and not self._cancel_called:
+ self._timeout()
+
+ @property
+ def cancel_called(self) -> bool:
+ return self._cancel_called
+
+ @property
+ def cancelled_caught(self) -> bool:
+ return self._cancelled_caught
+
+ @property
+ def shield(self) -> bool:
+ return self._shield
+
+ @shield.setter
+ def shield(self, value: bool) -> None:
+ if self._shield != value:
+ self._shield = value
+ if not value:
+ self._restart_cancellation_in_parent()
+
+
+#
+# Task states
+#
+
+
+class TaskState:
+ """
+ Encapsulates auxiliary task information that cannot be added to the Task instance
+ itself because there are no guarantees about its implementation.
+ """
+
+ __slots__ = "parent_id", "cancel_scope", "__weakref__"
+
+ def __init__(self, parent_id: int | None, cancel_scope: CancelScope | None):
+ self.parent_id = parent_id
+ self.cancel_scope = cancel_scope
+
+
+_task_states: WeakKeyDictionary[asyncio.Task, TaskState] = WeakKeyDictionary()
+
+
+#
+# Task groups
+#
+
+
+class _AsyncioTaskStatus(abc.TaskStatus):
+ def __init__(self, future: asyncio.Future, parent_id: int):
+ self._future = future
+ self._parent_id = parent_id
+
+ def started(self, value: T_contra | None = None) -> None:
+ try:
+ self._future.set_result(value)
+ except asyncio.InvalidStateError:
+ if not self._future.cancelled():
+ raise RuntimeError(
+ "called 'started' twice on the same task status"
+ ) from None
+
+ task = cast(asyncio.Task, current_task())
+ _task_states[task].parent_id = self._parent_id
+
+
+if sys.version_info >= (3, 12):
+ _eager_task_factory_code: CodeType | None = asyncio.eager_task_factory.__code__
+else:
+ _eager_task_factory_code = None
+
+
+class TaskGroup(abc.TaskGroup):
+ def __init__(self) -> None:
+ self.cancel_scope: CancelScope = CancelScope()
+ self._entered = False
+ self._exceptions: list[BaseException] = []
+ self._tasks: set[asyncio.Task] = set()
+ self._on_completed_fut: asyncio.Future[None] | None = None
+
+ async def __aenter__(self) -> TaskGroup:
+ if self._entered:
+ raise RuntimeError("TaskGroup cannot be entered more than once")
+
+ self._entered = True
+
+ self.cancel_scope.__enter__()
+ return self
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool:
+ try:
+ if exc_val is not None:
+ self.cancel_scope.cancel()
+ if not isinstance(exc_val, CancelledError):
+ self._exceptions.append(exc_val)
+
+ loop = get_running_loop()
+ try:
+ if self._tasks:
+ with CancelScope() as wait_scope:
+ while self._tasks:
+ self._on_completed_fut = loop.create_future()
+
+ try:
+ await self._on_completed_fut
+ except CancelledError as exc:
+ # Shield the scope against further cancellation attempts,
+ # as they're not productive (#695)
+ wait_scope.shield = True
+ self.cancel_scope.cancel()
+
+ # Set exc_val from the cancellation exception if it was
+ # previously unset. However, we should not replace a native
+ # cancellation exception with one raise by a cancel scope.
+ if exc_val is None or (
+ isinstance(exc_val, CancelledError)
+ and not is_anyio_cancellation(exc)
+ ):
+ exc_val = exc
+
+ self._on_completed_fut = None
+ else:
+ # If there are no child tasks to wait on, run at least one checkpoint
+ # anyway
+ await AsyncIOBackend.cancel_shielded_checkpoint()
+
+ if self._exceptions:
+ # The exception that got us here should already have been
+ # added to self._exceptions so it's ok to break exception
+ # chaining and avoid adding a "During handling of above..."
+ # for each nesting level.
+ raise BaseExceptionGroup(
+ "unhandled errors in a TaskGroup", self._exceptions
+ ) from None
+ elif exc_val:
+ raise exc_val
+ except BaseException as exc:
+ if self.cancel_scope.__exit__(type(exc), exc, exc.__traceback__):
+ return True
+
+ raise
+
+ return self.cancel_scope.__exit__(exc_type, exc_val, exc_tb)
+ finally:
+ del exc_val, exc_tb, self._exceptions
+
+ def _spawn(
+ self,
+ coro: Coroutine[Any, Any, T_co],
+ name: object,
+ task_status_future: asyncio.Future | None = None,
+ ) -> TaskHandle[T_co]:
+ def task_done(_task: asyncio.Task) -> None:
+ if sys.version_info >= (3, 14) and self.cancel_scope._host_task is not None:
+ asyncio.future_discard_from_awaited_by(
+ _task, self.cancel_scope._host_task
+ )
+
+ task_state = _task_states[_task]
+ assert task_state.cancel_scope is not None
+ assert _task in task_state.cancel_scope._tasks
+ task_state.cancel_scope._tasks.remove(_task)
+ self._tasks.remove(task)
+ del _task_states[_task]
+
+ if self._on_completed_fut is not None and not self._tasks:
+ try:
+ self._on_completed_fut.set_result(None)
+ except asyncio.InvalidStateError:
+ pass
+
+ try:
+ exc = _task.exception()
+ except CancelledError as e:
+ while isinstance(e.__context__, CancelledError):
+ e = e.__context__
+
+ exc = e
+
+ if exc is not None:
+ # The future can only be in the cancelled state if the host task was
+ # cancelled, so return immediately instead of adding one more
+ # CancelledError to the exceptions list
+ if task_status_future is not None and task_status_future.cancelled():
+ return
+
+ if task_status_future is None or task_status_future.done():
+ if not isinstance(exc, CancelledError):
+ self._exceptions.append(exc)
+
+ if not self.cancel_scope._effectively_cancelled:
+ self.cancel_scope.cancel()
+ else:
+ task_status_future.set_exception(exc)
+ elif task_status_future is not None and not task_status_future.done():
+ task_status_future.set_exception(
+ RuntimeError("Child exited without calling task_status.started()")
+ )
+
+ if task_status_future:
+ parent_id = id(current_task())
+ else:
+ parent_id = id(self.cancel_scope._host_task)
+
+ handle = TaskHandle(coro, name)
+ loop = asyncio.get_running_loop()
+ wrapper_coro = handle._run_coro()
+ if (
+ (factory := loop.get_task_factory())
+ and getattr(factory, "__code__", None) is _eager_task_factory_code
+ and (closure := getattr(factory, "__closure__", None))
+ ):
+ custom_task_constructor = closure[0].cell_contents
+ task = custom_task_constructor(wrapper_coro, loop=loop, name=handle.name)
+ else:
+ task = loop.create_task(wrapper_coro, name=handle.name)
+
+ # Make the spawned task inherit the task group's cancel scope
+ _task_states[task] = TaskState(
+ parent_id=parent_id, cancel_scope=self.cancel_scope
+ )
+ self.cancel_scope._tasks.add(task)
+ self._tasks.add(task)
+ if sys.version_info >= (3, 14) and self.cancel_scope._host_task is not None:
+ asyncio.future_add_to_awaited_by(task, self.cancel_scope._host_task)
+
+ task.add_done_callback(task_done)
+ return handle
+
+ def create_task(
+ self,
+ coro: Coroutine[Any, Any, T_co],
+ *,
+ name: object = None,
+ context: Context | None = None,
+ ) -> TaskHandle[T_co]:
+ if not isinstance(coro, Coroutine):
+ raise TypeError(f"expected a coroutine, got {coro.__class__.__qualname__}")
+
+ if not self._entered or not self.cancel_scope._active:
+ coro.close()
+ raise RuntimeError(
+ "This task group is not active; no new tasks can be started."
+ )
+
+ if context is not None:
+ return context.run(self._spawn, coro, name=name)
+ else:
+ return self._spawn(coro, name=name)
+
+ async def start(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ *args: Unpack[PosArgsT],
+ name: object = None,
+ return_handle: Literal[False] | Literal[True] = False,
+ ) -> Any:
+ if not self._entered or not self.cancel_scope._active:
+ raise RuntimeError(
+ "This task group is not active; no new tasks can be started."
+ )
+
+ future: asyncio.Future = asyncio.Future()
+ final_name = get_callable_name(func, name)
+ task_status = _AsyncioTaskStatus(future, id(self.cancel_scope._host_task))
+ coro = call_for_coroutine(func, args, task_status=task_status)
+ handle = self._spawn(coro, final_name, future)
+
+ # If the task raises an exception after sending a start value without a switch
+ # point between, the task group is cancelled and this method never proceeds to
+ # process the completed future. That's why we have to have a shielded cancel
+ # scope here.
+ try:
+ await future
+ except BaseException:
+ if handle.status is TaskHandle.Status.PENDING:
+ # Cancel the task and wait for it to exit before returning
+ handle.cancel()
+ with CancelScope(shield=True):
+ await handle.wait()
+
+ raise
+
+ if return_handle:
+ handle._start_value = future.result()
+ return handle
+ else:
+ return future.result()
+
+
+#
+# Threads
+#
+
+_Retval_Queue_Type = tuple[T_Retval | None, BaseException | None]
+
+
+class WorkerThread(Thread):
+ MAX_IDLE_TIME = 10 # seconds
+
+ def __init__(
+ self,
+ root_task: asyncio.Task,
+ workers: set[WorkerThread],
+ idle_workers: deque[WorkerThread],
+ ):
+ super().__init__(name="AnyIO worker thread")
+ self.root_task = root_task
+ self.workers = workers
+ self.idle_workers = idle_workers
+ self.loop = root_task._loop
+ self.queue: Queue[
+ tuple[Context, Callable, tuple, asyncio.Future, CancelScope] | None
+ ] = Queue(2)
+ self.idle_since = AsyncIOBackend.current_time()
+ self.stopping = False
+
+ def _report_result(
+ self, future: asyncio.Future, result: Any, exc: BaseException | None
+ ) -> None:
+ self.idle_since = AsyncIOBackend.current_time()
+ if not self.stopping:
+ self.idle_workers.append(self)
+
+ if not future.cancelled():
+ if exc is not None:
+ if isinstance(exc, StopIteration):
+ new_exc = RuntimeError("coroutine raised StopIteration")
+ new_exc.__cause__ = exc
+ exc = new_exc
+
+ future.set_exception(exc)
+ else:
+ future.set_result(result)
+
+ def run(self) -> None:
+ with claim_worker_thread(AsyncIOBackend, self.loop):
+ while True:
+ item = self.queue.get()
+ if item is None:
+ # Shutdown command received
+ return
+
+ context, func, args, future, cancel_scope = item
+ if not future.cancelled():
+ result = None
+ exception: BaseException | None = None
+ threadlocals.current_cancel_scope = cancel_scope
+ try:
+ result = context.run(func, *args)
+ except BaseException as exc:
+ exception = exc
+ finally:
+ del threadlocals.current_cancel_scope
+
+ if not self.loop.is_closed():
+ self.loop.call_soon_threadsafe(
+ self._report_result, future, result, exception
+ )
+
+ del result, exception
+
+ self.queue.task_done()
+ del item, context, func, args, future, cancel_scope
+
+ def stop(self, f: asyncio.Task | None = None) -> None:
+ self.stopping = True
+ self.queue.put_nowait(None)
+ self.workers.discard(self)
+ try:
+ self.idle_workers.remove(self)
+ except ValueError:
+ pass
+
+
+_threadpool_idle_workers: RunVar[deque[WorkerThread]] = RunVar(
+ "_threadpool_idle_workers"
+)
+_threadpool_workers: RunVar[set[WorkerThread]] = RunVar("_threadpool_workers")
+
+
+#
+# Subprocesses
+#
+
+
+@dataclass(eq=False)
+class StreamReaderWrapper(abc.ByteReceiveStream):
+ _stream: asyncio.StreamReader
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ data = await self._stream.read(max_bytes)
+ if data:
+ return data
+ else:
+ raise EndOfStream
+
+ async def aclose(self) -> None:
+ self._stream.set_exception(ClosedResourceError())
+ await AsyncIOBackend.checkpoint()
+
+
+@dataclass(eq=False)
+class StreamWriterWrapper(abc.ByteSendStream):
+ _stream: asyncio.StreamWriter
+ _closed: bool = field(init=False, default=False)
+
+ async def send(self, item: bytes) -> None:
+ await AsyncIOBackend.checkpoint_if_cancelled()
+ stream_paused = self._stream._protocol._paused # type: ignore[attr-defined]
+ try:
+ self._stream.write(item)
+ await self._stream.drain()
+ except (ConnectionResetError, BrokenPipeError, RuntimeError) as exc:
+ # If closed by us and/or the peer:
+ # * on stdlib, drain() raises ConnectionResetError or BrokenPipeError
+ # * on uvloop and Winloop, write() eventually starts raising RuntimeError
+ if self._closed:
+ raise ClosedResourceError from exc
+ elif self._stream.is_closing():
+ raise BrokenResourceError from exc
+
+ raise
+
+ if not stream_paused:
+ await AsyncIOBackend.cancel_shielded_checkpoint()
+
+ async def aclose(self) -> None:
+ self._closed = True
+ self._stream.close()
+ await AsyncIOBackend.checkpoint()
+
+
+@dataclass(eq=False)
+class Process(abc.Process):
+ _process: asyncio.subprocess.Process
+ _stdin: StreamWriterWrapper | None
+ _stdout: StreamReaderWrapper | None
+ _stderr: StreamReaderWrapper | None
+ _exited: asyncio.Event
+ _transport: asyncio.SubprocessTransport
+
+ async def aclose(self) -> None:
+ with CancelScope(shield=True) as scope:
+ # We need to close the underlying pipe_transports as well to allow a
+ # process blocking on full buffers to receive SIGPIPE and exit.
+ if self._stdin:
+ await self._stdin.aclose()
+ if pipe := self._transport.get_pipe_transport(0):
+ pipe.close()
+ if self._stdout:
+ await self._stdout.aclose()
+ if pipe := self._transport.get_pipe_transport(1):
+ pipe.close()
+ if self._stderr:
+ await self._stderr.aclose()
+ if pipe := self._transport.get_pipe_transport(2):
+ pipe.close()
+
+ scope.shield = False
+ try:
+ await self.wait()
+ except BaseException:
+ scope.shield = True
+ # Closing the transport on asyncio also handles sending kill
+ self._transport.close()
+ await self.wait()
+ raise
+
+ async def wait(self) -> int:
+ await self._exited.wait()
+ assert self._process.returncode is not None
+ return self._process.returncode
+
+ def terminate(self) -> None:
+ self._process.terminate()
+
+ def kill(self) -> None:
+ self._process.kill()
+
+ def send_signal(self, signal: int) -> None:
+ self._process.send_signal(signal)
+
+ @property
+ def pid(self) -> int:
+ return self._process.pid
+
+ @property
+ def returncode(self) -> int | None:
+ return self._process.returncode
+
+ @property
+ def stdin(self) -> abc.ByteSendStream | None:
+ return self._stdin
+
+ @property
+ def stdout(self) -> abc.ByteReceiveStream | None:
+ return self._stdout
+
+ @property
+ def stderr(self) -> abc.ByteReceiveStream | None:
+ return self._stderr
+
+
+def _forcibly_shutdown_process_pool_on_exit(
+ workers: set[Process], _task: object
+) -> None:
+ """
+ Forcibly shuts down worker processes belonging to this event loop."""
+ child_watcher: asyncio.AbstractChildWatcher | None = None # type: ignore[name-defined]
+ if sys.version_info < (3, 12):
+ try:
+ child_watcher = asyncio.get_event_loop_policy().get_child_watcher()
+ except NotImplementedError:
+ pass
+
+ # Close as much as possible (w/o async/await) to avoid warnings
+ for process in workers.copy():
+ if process.returncode is not None:
+ continue
+
+ process._stdin._stream._transport.close() # type: ignore[union-attr]
+ process._stdout._stream._transport.close() # type: ignore[union-attr]
+ process._stderr._stream._transport.close() # type: ignore[union-attr]
+ process.kill()
+ if child_watcher:
+ child_watcher.remove_child_handler(process.pid)
+
+
+async def _shutdown_process_pool_on_exit(workers: set[abc.Process]) -> None:
+ """
+ Shuts down worker processes belonging to this event loop.
+
+ NOTE: this only works when the event loop was started using asyncio.run() or
+ anyio.run().
+
+ """
+ process: abc.Process
+ try:
+ await sleep(math.inf)
+ except asyncio.CancelledError:
+ workers = workers.copy()
+ for process in workers:
+ if process.returncode is None:
+ process.kill()
+
+ for process in workers:
+ await process.aclose()
+
+
+#
+# Sockets and networking
+#
+
+
+class StreamProtocol(asyncio.Protocol):
+ read_queue: deque[bytes]
+ read_event: asyncio.Event
+ write_event: asyncio.Event
+ exception: Exception | None = None
+ is_at_eof: bool = False
+
+ def connection_made(self, transport: asyncio.BaseTransport) -> None:
+ self.read_queue = deque()
+ self.read_event = asyncio.Event()
+ self.write_event = asyncio.Event()
+ self.write_event.set()
+ cast(asyncio.Transport, transport).set_write_buffer_limits(0)
+
+ def connection_lost(self, exc: Exception | None) -> None:
+ if exc:
+ self.exception = exc
+
+ self.read_event.set()
+ self.write_event.set()
+
+ def data_received(self, data: bytes) -> None:
+ # ProactorEventloop sometimes sends bytearray instead of bytes
+ self.read_queue.append(bytes(data))
+ self.read_event.set()
+
+ def eof_received(self) -> bool | None:
+ self.is_at_eof = True
+ self.read_event.set()
+ return True
+
+ def pause_writing(self) -> None:
+ self.write_event = asyncio.Event()
+
+ def resume_writing(self) -> None:
+ self.write_event.set()
+
+
+class DatagramProtocol(asyncio.DatagramProtocol):
+ read_queue: deque[tuple[bytes, IPSockAddrType]]
+ read_event: asyncio.Event
+ write_event: asyncio.Event
+ closed_event: asyncio.Event
+ exception: Exception | None = None
+
+ def connection_made(self, transport: asyncio.BaseTransport) -> None:
+ self.read_queue = deque(maxlen=100) # arbitrary value
+ self.read_event = asyncio.Event()
+ self.write_event = asyncio.Event()
+ self.closed_event = asyncio.Event()
+ self.write_event.set()
+
+ def connection_lost(self, exc: Exception | None) -> None:
+ self.read_event.set()
+ self.write_event.set()
+ self.closed_event.set()
+
+ def datagram_received(self, data: bytes, addr: IPSockAddrType) -> None:
+ addr = convert_ipv6_sockaddr(addr)
+ self.read_queue.append((data, addr))
+ self.read_event.set()
+
+ def error_received(self, exc: Exception) -> None:
+ self.exception = exc
+
+ def pause_writing(self) -> None:
+ self.write_event.clear()
+
+ def resume_writing(self) -> None:
+ self.write_event.set()
+
+
+class SocketStream(abc.SocketStream):
+ def __init__(self, transport: asyncio.Transport, protocol: StreamProtocol):
+ self._transport = transport
+ self._protocol = protocol
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+ self._closed = False
+
+ @property
+ def _raw_socket(self) -> socket.socket:
+ return self._transport.get_extra_info("socket")
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ with self._receive_guard:
+ if (
+ not self._protocol.read_event.is_set()
+ and not self._transport.is_closing()
+ and not self._protocol.is_at_eof
+ ):
+ self._transport.resume_reading()
+ await self._protocol.read_event.wait()
+ self._transport.pause_reading()
+ else:
+ await AsyncIOBackend.checkpoint()
+
+ try:
+ chunk = self._protocol.read_queue.popleft()
+ except IndexError:
+ if self._closed:
+ raise ClosedResourceError from None
+ elif self._protocol.exception:
+ raise BrokenResourceError from self._protocol.exception
+ else:
+ raise EndOfStream from None
+
+ if len(chunk) > max_bytes:
+ # Split the oversized chunk
+ chunk, leftover = chunk[:max_bytes], chunk[max_bytes:]
+ self._protocol.read_queue.appendleft(leftover)
+
+ # If the read queue is empty, clear the flag so that the next call will
+ # block until data is available
+ if not self._protocol.read_queue:
+ self._protocol.read_event.clear()
+
+ return chunk
+
+ async def send(self, item: bytes) -> None:
+ with self._send_guard:
+ await AsyncIOBackend.checkpoint()
+
+ if self._closed:
+ raise ClosedResourceError
+ elif self._protocol.exception is not None:
+ raise BrokenResourceError from self._protocol.exception
+
+ try:
+ self._transport.write(item)
+ except RuntimeError as exc:
+ if self._transport.is_closing():
+ raise BrokenResourceError from exc
+ else:
+ raise
+
+ await self._protocol.write_event.wait()
+
+ async def send_eof(self) -> None:
+ try:
+ self._transport.write_eof()
+ except OSError:
+ pass
+
+ async def aclose(self) -> None:
+ self._closed = True
+ if not self._transport.is_closing():
+ try:
+ self._transport.write_eof()
+ except OSError:
+ pass
+
+ self._transport.close()
+ await sleep(0)
+ self._transport.abort()
+
+
+class _RawSocketMixin:
+ _receive_future: asyncio.Future | None = None
+ _send_future: asyncio.Future | None = None
+ _closing = False
+
+ def __init__(self, raw_socket: socket.socket):
+ self.__raw_socket = raw_socket
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+
+ @property
+ def _raw_socket(self) -> socket.socket:
+ return self.__raw_socket
+
+ def _wait_until_readable(self, loop: asyncio.AbstractEventLoop) -> asyncio.Future:
+ def callback(f: object) -> None:
+ del self._receive_future
+ loop.remove_reader(self.__raw_socket)
+
+ f = self._receive_future = asyncio.Future()
+ loop.add_reader(self.__raw_socket, f.set_result, None)
+ f.add_done_callback(callback)
+ return f
+
+ def _wait_until_writable(self, loop: asyncio.AbstractEventLoop) -> asyncio.Future:
+ def callback(f: object) -> None:
+ del self._send_future
+ loop.remove_writer(self.__raw_socket)
+
+ f = self._send_future = asyncio.Future()
+ loop.add_writer(self.__raw_socket, f.set_result, None)
+ f.add_done_callback(callback)
+ return f
+
+ async def aclose(self) -> None:
+ if not self._closing:
+ self._closing = True
+ if self.__raw_socket.fileno() != -1:
+ self.__raw_socket.close()
+
+ if self._receive_future:
+ self._receive_future.set_result(None)
+ if self._send_future:
+ self._send_future.set_result(None)
+
+
+class UNIXSocketStream(_RawSocketMixin, abc.UNIXSocketStream):
+ async def send_eof(self) -> None:
+ with self._send_guard:
+ self._raw_socket.shutdown(socket.SHUT_WR)
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ loop = get_running_loop()
+ await AsyncIOBackend.checkpoint()
+ with self._receive_guard:
+ while True:
+ try:
+ data = self._raw_socket.recv(max_bytes)
+ except BlockingIOError:
+ await self._wait_until_readable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+ else:
+ if not data:
+ raise EndOfStream
+
+ return data
+
+ async def send(self, item: bytes) -> None:
+ loop = get_running_loop()
+ await AsyncIOBackend.checkpoint()
+ with self._send_guard:
+ view = memoryview(item)
+ while view:
+ try:
+ bytes_sent = self._raw_socket.send(view)
+ except BlockingIOError:
+ await self._wait_until_writable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+ else:
+ view = view[bytes_sent:]
+
+ async def receive_fds(self, msglen: int, maxfds: int) -> tuple[bytes, list[int]]:
+ if not isinstance(msglen, int) or msglen < 0:
+ raise ValueError("msglen must be a non-negative integer")
+ if not isinstance(maxfds, int) or maxfds < 1:
+ raise ValueError("maxfds must be a positive integer")
+
+ loop = get_running_loop()
+ fds = array.array("i")
+ await AsyncIOBackend.checkpoint()
+ with self._receive_guard:
+ while True:
+ try:
+ message, ancdata, flags, addr = self._raw_socket.recvmsg(
+ msglen, socket.CMSG_LEN(maxfds * fds.itemsize)
+ )
+ except BlockingIOError:
+ await self._wait_until_readable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+ else:
+ if not message and not ancdata:
+ raise EndOfStream
+
+ break
+
+ for cmsg_level, cmsg_type, cmsg_data in ancdata:
+ if cmsg_level != socket.SOL_SOCKET or cmsg_type != socket.SCM_RIGHTS:
+ raise RuntimeError(
+ f"Received unexpected ancillary data; message = {message!r}, "
+ f"cmsg_level = {cmsg_level}, cmsg_type = {cmsg_type}"
+ )
+
+ fds.frombytes(cmsg_data[: len(cmsg_data) - (len(cmsg_data) % fds.itemsize)])
+
+ return message, list(fds)
+
+ async def send_fds(self, message: bytes, fds: Collection[int | IOBase]) -> None:
+ if not message:
+ raise ValueError("message must not be empty")
+ if not fds:
+ raise ValueError("fds must not be empty")
+
+ loop = get_running_loop()
+ filenos: list[int] = []
+ for fd in fds:
+ if isinstance(fd, int):
+ filenos.append(fd)
+ elif isinstance(fd, IOBase):
+ filenos.append(fd.fileno())
+
+ fdarray = array.array("i", filenos)
+ await AsyncIOBackend.checkpoint()
+ with self._send_guard:
+ while True:
+ try:
+ # The ignore can be removed after mypy picks up
+ # https://github.com/python/typeshed/pull/5545
+ self._raw_socket.sendmsg(
+ [message], [(socket.SOL_SOCKET, socket.SCM_RIGHTS, fdarray)]
+ )
+ break
+ except BlockingIOError:
+ await self._wait_until_writable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+
+
+class TCPSocketListener(abc.SocketListener):
+ _accept_scope: CancelScope | None = None
+ _closed = False
+
+ def __init__(self, raw_socket: socket.socket):
+ self.__raw_socket = raw_socket
+ self._loop = cast(asyncio.BaseEventLoop, get_running_loop())
+ self._accept_guard = ResourceGuard("accepting connections from")
+
+ @property
+ def _raw_socket(self) -> socket.socket:
+ return self.__raw_socket
+
+ async def accept(self) -> abc.SocketStream:
+ if self._closed:
+ raise ClosedResourceError
+
+ with self._accept_guard:
+ await AsyncIOBackend.checkpoint()
+ with CancelScope() as self._accept_scope:
+ try:
+ client_sock, _addr = await self._loop.sock_accept(self._raw_socket)
+ except asyncio.CancelledError:
+ # Workaround for https://bugs.python.org/issue41317
+ try:
+ self._loop.remove_reader(self._raw_socket)
+ except (ValueError, NotImplementedError):
+ pass
+
+ if self._closed:
+ raise ClosedResourceError from None
+
+ raise
+ finally:
+ self._accept_scope = None
+
+ client_sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
+ transport, protocol = await self._loop.connect_accepted_socket(
+ StreamProtocol, client_sock
+ )
+ return SocketStream(transport, protocol)
+
+ async def aclose(self) -> None:
+ if self._closed:
+ return
+
+ self._closed = True
+ if self._accept_scope:
+ # Workaround for https://bugs.python.org/issue41317
+ try:
+ self._loop.remove_reader(self._raw_socket)
+ except (ValueError, NotImplementedError):
+ pass
+
+ self._accept_scope.cancel()
+ await sleep(0)
+
+ self._raw_socket.close()
+
+
+class UNIXSocketListener(abc.SocketListener):
+ def __init__(self, raw_socket: socket.socket):
+ self.__raw_socket = raw_socket
+ self._loop = get_running_loop()
+ self._accept_guard = ResourceGuard("accepting connections from")
+ self._closed = False
+
+ async def accept(self) -> abc.SocketStream:
+ await AsyncIOBackend.checkpoint()
+ with self._accept_guard:
+ while True:
+ try:
+ client_sock, _ = self.__raw_socket.accept()
+ client_sock.setblocking(False)
+ return UNIXSocketStream(client_sock)
+ except BlockingIOError:
+ f: asyncio.Future = asyncio.Future()
+ self._loop.add_reader(self.__raw_socket, f.set_result, None)
+ f.add_done_callback(
+ lambda _: self._loop.remove_reader(self.__raw_socket)
+ )
+ await f
+ except OSError as exc:
+ if self._closed:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+
+ async def aclose(self) -> None:
+ self._closed = True
+ self.__raw_socket.close()
+
+ @property
+ def _raw_socket(self) -> socket.socket:
+ return self.__raw_socket
+
+
+class UDPSocket(abc.UDPSocket):
+ def __init__(
+ self, transport: asyncio.DatagramTransport, protocol: DatagramProtocol
+ ):
+ self._transport = transport
+ self._protocol = protocol
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+ self._closed = False
+
+ @property
+ def _raw_socket(self) -> socket.socket:
+ return self._transport.get_extra_info("socket")
+
+ async def aclose(self) -> None:
+ self._closed = True
+ if not self._transport.is_closing():
+ self._transport.close()
+
+ await self._protocol.closed_event.wait()
+
+ async def receive(self) -> tuple[bytes, IPSockAddrType]:
+ with self._receive_guard:
+ await AsyncIOBackend.checkpoint()
+
+ # If the buffer is empty, ask for more data
+ if not self._protocol.read_queue and not self._transport.is_closing():
+ self._protocol.read_event.clear()
+ await self._protocol.read_event.wait()
+
+ try:
+ return self._protocol.read_queue.popleft()
+ except IndexError:
+ if self._closed:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from None
+
+ async def send(self, item: UDPPacketType) -> None:
+ with self._send_guard:
+ await AsyncIOBackend.checkpoint()
+ await self._protocol.write_event.wait()
+ if self._closed:
+ raise ClosedResourceError
+ elif self._transport.is_closing():
+ raise BrokenResourceError
+ else:
+ self._transport.sendto(*item)
+
+
+class ConnectedUDPSocket(abc.ConnectedUDPSocket):
+ def __init__(
+ self, transport: asyncio.DatagramTransport, protocol: DatagramProtocol
+ ):
+ self._transport = transport
+ self._protocol = protocol
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+ self._closed = False
+
+ @property
+ def _raw_socket(self) -> socket.socket:
+ return self._transport.get_extra_info("socket")
+
+ async def aclose(self) -> None:
+ self._closed = True
+ if not self._transport.is_closing():
+ self._transport.close()
+
+ await self._protocol.closed_event.wait()
+
+ async def receive(self) -> bytes:
+ with self._receive_guard:
+ await AsyncIOBackend.checkpoint()
+
+ # If the buffer is empty, ask for more data
+ if not self._protocol.read_queue and not self._transport.is_closing():
+ self._protocol.read_event.clear()
+ await self._protocol.read_event.wait()
+
+ try:
+ packet = self._protocol.read_queue.popleft()
+ except IndexError:
+ if self._closed:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from None
+
+ return packet[0]
+
+ async def send(self, item: bytes) -> None:
+ with self._send_guard:
+ await AsyncIOBackend.checkpoint()
+ await self._protocol.write_event.wait()
+ if self._closed:
+ raise ClosedResourceError
+ elif self._transport.is_closing():
+ raise BrokenResourceError
+ else:
+ self._transport.sendto(item)
+
+
+class UNIXDatagramSocket(_RawSocketMixin, abc.UNIXDatagramSocket):
+ async def receive(self) -> UNIXDatagramPacketType:
+ loop = get_running_loop()
+ await AsyncIOBackend.checkpoint()
+ with self._receive_guard:
+ while True:
+ try:
+ data = self._raw_socket.recvfrom(65536)
+ except BlockingIOError:
+ await self._wait_until_readable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+ else:
+ return data
+
+ async def send(self, item: UNIXDatagramPacketType) -> None:
+ loop = get_running_loop()
+ await AsyncIOBackend.checkpoint()
+ with self._send_guard:
+ while True:
+ try:
+ self._raw_socket.sendto(*item)
+ except BlockingIOError:
+ await self._wait_until_writable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+ else:
+ return
+
+
+class ConnectedUNIXDatagramSocket(_RawSocketMixin, abc.ConnectedUNIXDatagramSocket):
+ async def receive(self) -> bytes:
+ loop = get_running_loop()
+ await AsyncIOBackend.checkpoint()
+ with self._receive_guard:
+ while True:
+ try:
+ data = self._raw_socket.recv(65536)
+ except BlockingIOError:
+ await self._wait_until_readable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+ else:
+ return data
+
+ async def send(self, item: bytes) -> None:
+ loop = get_running_loop()
+ await AsyncIOBackend.checkpoint()
+ with self._send_guard:
+ while True:
+ try:
+ self._raw_socket.send(item)
+ except BlockingIOError:
+ await self._wait_until_writable(loop)
+ except OSError as exc:
+ if self._closing:
+ raise ClosedResourceError from None
+ else:
+ raise BrokenResourceError from exc
+ else:
+ return
+
+
+_read_events: RunVar[dict[int, asyncio.Future[bool]]] = RunVar("read_events")
+_write_events: RunVar[dict[int, asyncio.Future[bool]]] = RunVar("write_events")
+
+
+#
+# Synchronization
+#
+
+
+class Event(BaseEvent):
+ __slots__ = ("_event",)
+
+ def __new__(cls) -> Event:
+ return object.__new__(cls)
+
+ def __init__(self) -> None:
+ self._event = asyncio.Event()
+
+ def set(self) -> None:
+ self._event.set()
+
+ def is_set(self) -> bool:
+ return self._event.is_set()
+
+ async def wait(self) -> None:
+ if self.is_set():
+ await AsyncIOBackend.checkpoint()
+ else:
+ await self._event.wait()
+
+ def statistics(self) -> EventStatistics:
+ return EventStatistics(len(self._event._waiters))
+
+
+class Lock(BaseLock):
+ __slots__ = "_fast_acquire", "_owner_task", "_waiters"
+
+ def __new__(cls, *, fast_acquire: bool = False) -> Lock:
+ return object.__new__(cls)
+
+ def __init__(self, *, fast_acquire: bool = False) -> None:
+ self._fast_acquire = fast_acquire
+ self._owner_task: asyncio.Task | None = None
+ self._waiters: deque[tuple[asyncio.Task, asyncio.Future]] = deque()
+
+ async def acquire(self) -> None:
+ task = cast(asyncio.Task, current_task())
+ if self._owner_task is None and not self._waiters:
+ await AsyncIOBackend.checkpoint_if_cancelled()
+ self._owner_task = task
+
+ # Unless on the "fast path", yield control of the event loop so that other
+ # tasks can run too
+ if not self._fast_acquire:
+ try:
+ await AsyncIOBackend.cancel_shielded_checkpoint()
+ except CancelledError:
+ self.release()
+ raise
+
+ return
+
+ if self._owner_task == task:
+ raise RuntimeError("Attempted to acquire an already held Lock")
+
+ fut: asyncio.Future[None] = asyncio.Future()
+ item = task, fut
+ self._waiters.append(item)
+ try:
+ await fut
+ except CancelledError:
+ if fut.cancelled():
+ try:
+ self._waiters.remove(item)
+ except ValueError:
+ pass
+ else:
+ self.release()
+
+ raise
+
+ def acquire_nowait(self) -> None:
+ task = cast(asyncio.Task, current_task())
+ if self._owner_task is None and not self._waiters:
+ self._owner_task = task
+ return
+
+ if self._owner_task is task:
+ raise RuntimeError("Attempted to acquire an already held Lock")
+
+ raise WouldBlock
+
+ def locked(self) -> bool:
+ return self._owner_task is not None
+
+ def release(self) -> None:
+ if self._owner_task != current_task():
+ raise RuntimeError("The current task is not holding this lock")
+
+ # A cancelled waiter that already received ownership removes itself from
+ # _waiters before calling release(); any cancelled waiter still queued here
+ # was cancelled before being woken, so drop it.
+ while self._waiters:
+ task, fut = self._waiters.popleft()
+ if fut.cancelled():
+ continue
+
+ self._owner_task = task
+ fut.set_result(None)
+ return
+
+ self._owner_task = None
+
+ def statistics(self) -> LockStatistics:
+ task_info = AsyncIOTaskInfo(self._owner_task) if self._owner_task else None
+ return LockStatistics(self.locked(), task_info, len(self._waiters))
+
+
+class Semaphore(BaseSemaphore):
+ __slots__ = "_value", "_max_value", "_fast_acquire", "_waiters"
+
+ def __new__(
+ cls,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> Semaphore:
+ return object.__new__(cls)
+
+ def __init__(
+ self,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ):
+ super().__init__(initial_value, max_value=max_value)
+ self._value = initial_value
+ self._max_value = max_value
+ self._fast_acquire = fast_acquire
+ self._waiters: deque[asyncio.Future[None]] = deque()
+
+ async def acquire(self) -> None:
+ if self._value > 0 and not self._waiters:
+ await AsyncIOBackend.checkpoint_if_cancelled()
+ self._value -= 1
+
+ # Unless on the "fast path", yield control of the event loop so that other
+ # tasks can run too
+ if not self._fast_acquire:
+ try:
+ await AsyncIOBackend.cancel_shielded_checkpoint()
+ except CancelledError:
+ self.release()
+ raise
+
+ return
+
+ fut: asyncio.Future[None] = asyncio.Future()
+ self._waiters.append(fut)
+ try:
+ await fut
+ except CancelledError:
+ if fut.cancelled():
+ try:
+ self._waiters.remove(fut)
+ except ValueError:
+ pass
+ else:
+ self.release()
+
+ raise
+
+ def acquire_nowait(self) -> None:
+ if self._value == 0:
+ raise WouldBlock
+
+ self._value -= 1
+
+ def release(self) -> None:
+ if self._max_value is not None and self._value == self._max_value:
+ raise ValueError("semaphore released too many times")
+
+ while self._waiters:
+ fut = self._waiters.popleft()
+ if fut.cancelled():
+ continue
+
+ fut.set_result(None)
+ return
+
+ self._value += 1
+
+ @property
+ def value(self) -> int:
+ return self._value
+
+ @property
+ def max_value(self) -> int | None:
+ return self._max_value
+
+ def statistics(self) -> SemaphoreStatistics:
+ return SemaphoreStatistics(len(self._waiters))
+
+
+class CapacityLimiter(BaseCapacityLimiter):
+ __slots__ = "_total_tokens", "_borrowers", "_wait_queue"
+
+ def __new__(cls, total_tokens: float) -> CapacityLimiter:
+ return object.__new__(cls)
+
+ def __init__(self, total_tokens: float):
+ self._total_tokens: float = 0
+ self._borrowers: set[Any] = set()
+ self._wait_queue: OrderedDict[Any, asyncio.Event] = OrderedDict()
+ self.total_tokens = total_tokens
+
+ async def __aenter__(self) -> None:
+ await self.acquire()
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self.release()
+
+ @property
+ def total_tokens(self) -> float:
+ return self._total_tokens
+
+ @total_tokens.setter
+ def total_tokens(self, value: float) -> None:
+ if not isinstance(value, int) and not math.isinf(value):
+ raise TypeError("total_tokens must be an int or math.inf")
+
+ if value < 0:
+ raise ValueError("total_tokens must be >= 0")
+
+ waiters_to_notify = max(value - self._total_tokens, 0)
+ self._total_tokens = value
+
+ # Notify waiting tasks that they have acquired the limiter
+ while self._wait_queue and waiters_to_notify:
+ borrower, event = self._wait_queue.popitem(last=False)
+ self._borrowers.add(borrower)
+ event.set()
+ waiters_to_notify -= 1
+
+ @property
+ def borrowed_tokens(self) -> int:
+ return len(self._borrowers)
+
+ @property
+ def available_tokens(self) -> float:
+ return self._total_tokens - len(self._borrowers)
+
+ def _notify_next_waiter(self) -> None:
+ """Hand a free token to the next task in line, if any."""
+ if self._wait_queue and len(self._borrowers) < self._total_tokens:
+ borrower, event = self._wait_queue.popitem(last=False)
+ self._borrowers.add(borrower)
+ event.set()
+
+ def acquire_nowait(self) -> None:
+ self.acquire_on_behalf_of_nowait(current_task())
+
+ def acquire_on_behalf_of_nowait(self, borrower: object) -> None:
+ if borrower in self._borrowers:
+ raise RuntimeError(
+ "this borrower is already holding one of this CapacityLimiter's tokens"
+ )
+
+ if self._wait_queue or len(self._borrowers) >= self._total_tokens:
+ raise WouldBlock
+
+ self._borrowers.add(borrower)
+
+ async def acquire(self) -> None:
+ return await self.acquire_on_behalf_of(current_task())
+
+ async def acquire_on_behalf_of(self, borrower: object) -> None:
+ await AsyncIOBackend.checkpoint_if_cancelled()
+ try:
+ self.acquire_on_behalf_of_nowait(borrower)
+ except WouldBlock:
+ event = asyncio.Event()
+ self._wait_queue[borrower] = event
+ try:
+ await event.wait()
+ except BaseException:
+ self._wait_queue.pop(borrower, None)
+ if event.is_set():
+ self._borrowers.discard(borrower)
+ self._notify_next_waiter()
+
+ raise
+ else:
+ try:
+ await AsyncIOBackend.cancel_shielded_checkpoint()
+ except BaseException:
+ self.release()
+ raise
+
+ def release(self) -> None:
+ self.release_on_behalf_of(current_task())
+
+ def release_on_behalf_of(self, borrower: object) -> None:
+ try:
+ self._borrowers.remove(borrower)
+ except KeyError:
+ raise RuntimeError(
+ "this borrower isn't holding any of this CapacityLimiter's tokens"
+ ) from None
+
+ self._notify_next_waiter()
+
+ def statistics(self) -> CapacityLimiterStatistics:
+ return CapacityLimiterStatistics(
+ self.borrowed_tokens,
+ self.total_tokens,
+ tuple(self._borrowers),
+ len(self._wait_queue),
+ )
+
+
+_default_thread_limiter: RunVar[CapacityLimiter] = RunVar("_default_thread_limiter")
+
+
+#
+# Operating system signals
+#
+
+
+class _SignalReceiver:
+ def __init__(self, signals: tuple[Signals, ...]):
+ self._signals = signals
+ self._loop = get_running_loop()
+ self._signal_queue: deque[Signals] = deque()
+ self._future: asyncio.Future = asyncio.Future()
+ self._handled_signals: set[Signals] = set()
+
+ def _deliver(self, signum: Signals) -> None:
+ self._signal_queue.append(signum)
+ if not self._future.done():
+ self._future.set_result(None)
+
+ def __enter__(self) -> _SignalReceiver:
+ for sig in set(self._signals):
+ self._loop.add_signal_handler(sig, self._deliver, sig)
+ self._handled_signals.add(sig)
+
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ for sig in self._handled_signals:
+ self._loop.remove_signal_handler(sig)
+
+ def __aiter__(self) -> _SignalReceiver:
+ return self
+
+ async def __anext__(self) -> Signals:
+ await AsyncIOBackend.checkpoint()
+ if not self._signal_queue:
+ self._future = asyncio.Future()
+ await self._future
+
+ return self._signal_queue.popleft()
+
+
+#
+# Testing and debugging
+#
+
+
+class AsyncIOTaskInfo(TaskInfo):
+ def __init__(self, task: asyncio.Task):
+ task_state = _task_states.get(task)
+ if task_state is None:
+ parent_id = None
+ else:
+ parent_id = task_state.parent_id
+
+ coro = task.get_coro()
+ assert coro is not None, "created TaskInfo from a completed Task"
+ super().__init__(id(task), parent_id, task.get_name(), coro)
+ self._task = weakref.ref(task)
+
+ def has_pending_cancellation(self) -> bool:
+ if not (task := self._task()):
+ # If the task isn't around anymore, it won't have a pending cancellation
+ return False
+
+ if task._must_cancel: # type: ignore[attr-defined]
+ return True
+ elif (
+ isinstance(task._fut_waiter, asyncio.Future) # type: ignore[attr-defined]
+ and task._fut_waiter.cancelled() # type: ignore[attr-defined]
+ ):
+ return True
+
+ if task_state := _task_states.get(task):
+ if cancel_scope := task_state.cancel_scope:
+ return cancel_scope._effectively_cancelled
+
+ return False
+
+
+class TestRunner(abc.TestRunner):
+ _send_stream: MemoryObjectSendStream[tuple[Awaitable[Any], asyncio.Future[Any]]]
+
+ def __init__(
+ self,
+ *,
+ debug: bool | None = None,
+ use_uvloop: bool = False,
+ loop_factory: Callable[[], AbstractEventLoop] | None = None,
+ ) -> None:
+ if use_uvloop and loop_factory is None:
+ if sys.platform != "win32":
+ import uvloop
+
+ loop_factory = uvloop.new_event_loop
+ else:
+ import winloop
+
+ loop_factory = winloop.new_event_loop
+
+ self._runner = Runner(debug=debug, loop_factory=loop_factory)
+ self._exceptions: list[BaseException] = []
+ self._runner_task: asyncio.Task | None = None
+
+ def __enter__(self) -> TestRunner:
+ self._runner.__enter__()
+ self.get_loop().set_exception_handler(self._exception_handler)
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self._runner.__exit__(exc_type, exc_val, exc_tb)
+
+ def get_loop(self) -> AbstractEventLoop:
+ return self._runner.get_loop()
+
+ def is_running(self) -> bool:
+ try:
+ asyncio.get_running_loop()
+ return True
+ except RuntimeError:
+ return False
+
+ def _exception_handler(
+ self, loop: asyncio.AbstractEventLoop, context: dict[str, Any]
+ ) -> None:
+ if isinstance(context.get("exception"), Exception):
+ self._exceptions.append(context["exception"])
+ else:
+ loop.default_exception_handler(context)
+
+ def _raise_async_exceptions(self) -> None:
+ # Re-raise any exceptions raised in asynchronous callbacks
+ if self._exceptions:
+ exceptions, self._exceptions = self._exceptions, []
+ if len(exceptions) == 1:
+ raise exceptions[0]
+ elif exceptions:
+ raise BaseExceptionGroup(
+ "Multiple exceptions occurred in asynchronous callbacks", exceptions
+ )
+
+ async def _run_tests_and_fixtures(
+ self,
+ receive_stream: MemoryObjectReceiveStream[
+ tuple[Awaitable[T_Retval], asyncio.Future[T_Retval]]
+ ],
+ ) -> None:
+ from _pytest.outcomes import OutcomeException
+
+ with receive_stream, self._send_stream:
+ async for coro, future in receive_stream:
+ try:
+ retval = await coro
+ except CancelledError as exc:
+ if not future.cancelled():
+ future.cancel(*exc.args)
+
+ raise
+ except BaseException as exc:
+ if not future.cancelled():
+ future.set_exception(exc)
+
+ if not isinstance(exc, (Exception, OutcomeException)):
+ raise
+ else:
+ if not future.cancelled():
+ future.set_result(retval)
+
+ async def _call_in_runner_task(
+ self,
+ func: Callable[P, Awaitable[T_Retval]],
+ /,
+ *args: P.args,
+ **kwargs: P.kwargs,
+ ) -> T_Retval:
+ if not self._runner_task:
+ self._send_stream, receive_stream = create_memory_object_stream[
+ tuple[Awaitable[Any], asyncio.Future]
+ ](1)
+ self._runner_task = self.get_loop().create_task(
+ self._run_tests_and_fixtures(receive_stream)
+ )
+
+ coro = func(*args, **kwargs)
+ future: asyncio.Future[T_Retval] = self.get_loop().create_future()
+ self._send_stream.send_nowait((coro, future))
+ return await future
+
+ def run_asyncgen_fixture(
+ self,
+ fixture_func: Callable[..., AsyncGenerator[T_Retval, Any]],
+ kwargs: dict[str, Any],
+ ) -> Iterable[T_Retval]:
+ asyncgen = fixture_func(**kwargs)
+ fixturevalue: T_Retval = self.get_loop().run_until_complete(
+ self._call_in_runner_task(asyncgen.asend, None)
+ )
+ self._raise_async_exceptions()
+
+ yield fixturevalue
+
+ try:
+ self.get_loop().run_until_complete(
+ self._call_in_runner_task(asyncgen.asend, None)
+ )
+ except StopAsyncIteration:
+ self._raise_async_exceptions()
+ else:
+ self.get_loop().run_until_complete(asyncgen.aclose())
+ raise RuntimeError("Async generator fixture did not stop")
+
+ def run_fixture(
+ self,
+ fixture_func: Callable[..., Coroutine[Any, Any, T_Retval]],
+ kwargs: dict[str, Any],
+ ) -> T_Retval:
+ retval = self.get_loop().run_until_complete(
+ self._call_in_runner_task(fixture_func, **kwargs)
+ )
+ self._raise_async_exceptions()
+ return retval
+
+ def run_test(
+ self, test_func: Callable[..., Coroutine[Any, Any, Any]], kwargs: dict[str, Any]
+ ) -> None:
+ from _pytest.outcomes import OutcomeException
+
+ try:
+ self.get_loop().run_until_complete(
+ self._call_in_runner_task(test_func, **kwargs)
+ )
+ except Exception as exc:
+ self._exceptions.append(exc)
+ except OutcomeException:
+ raise
+ except BaseException:
+ # A BaseException (e.g. KeyboardInterrupt, SystemExit) interrupted the event loop before
+ # the test completed. Cancel _runner_task so it does not resume when the event
+ # loop is re-entered during async generator fixture teardown.
+ if self._runner_task is not None and not self._runner_task.done():
+ self._runner_task.cancel()
+ self._send_stream.close()
+ try:
+ self.get_loop().run_until_complete(self._runner_task)
+ except CancelledError:
+ pass
+ finally:
+ self._runner_task = None
+ raise
+ self._raise_async_exceptions()
+
+
+class _ProcessStreamProtocol(asyncio.subprocess.SubprocessStreamProtocol):
+ """
+ A subprocess protocol that allows us to be notified of ``process_exited``
+
+ asyncio's own ``Process.wait()`` only resolves once every pipe transport has
+ disconnected so to get same semantics as on trio and uvloop we need this.
+ """
+
+ def __init__(self) -> None:
+ # Match the standard factory for asyncio.create_process
+ super().__init__(limit=2**16, loop=asyncio.get_running_loop())
+ self.exited = asyncio.Event()
+
+ def process_exited(self) -> None:
+ super().process_exited()
+ self.exited.set()
+
+
+class AsyncIOBackend(AsyncBackend):
+ @classmethod
+ def run(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
+ args: tuple[Unpack[PosArgsT]],
+ kwargs: dict[str, Any],
+ options: dict[str, Any],
+ ) -> T_Retval:
+ @wraps(func)
+ async def wrapper() -> T_Retval:
+ task = cast(asyncio.Task, current_task())
+ task.set_name(get_callable_name(func))
+ _task_states[task] = TaskState(None, None)
+
+ try:
+ return await func(*args)
+ finally:
+ del _task_states[task]
+
+ debug = options.get("debug", None)
+ loop_factory = options.get("loop_factory", None)
+ if loop_factory is None and options.get("use_uvloop", False):
+ if sys.platform != "win32":
+ import uvloop
+
+ loop_factory = uvloop.new_event_loop
+ else:
+ import winloop
+
+ loop_factory = winloop.new_event_loop
+
+ with Runner(debug=debug, loop_factory=loop_factory) as runner:
+ return runner.run(wrapper())
+
+ @classmethod
+ def current_token(cls) -> object:
+ return get_running_loop()
+
+ @classmethod
+ def current_time(cls) -> float:
+ return get_running_loop().time()
+
+ @classmethod
+ def cancelled_exception_class(cls) -> type[BaseException]:
+ return CancelledError
+
+ @classmethod
+ async def checkpoint(cls) -> None:
+ await sleep(0)
+
+ @classmethod
+ async def checkpoint_if_cancelled(cls) -> None:
+ task = current_task()
+ if task is None:
+ return
+
+ try:
+ cancel_scope = _task_states[task].cancel_scope
+ except KeyError:
+ return
+
+ while cancel_scope:
+ if cancel_scope.cancel_called:
+ await sleep(0)
+ elif cancel_scope.shield:
+ break
+ else:
+ cancel_scope = cancel_scope._parent_scope
+
+ @classmethod
+ async def cancel_shielded_checkpoint(cls) -> None:
+ with CancelScope(shield=True):
+ await sleep(0)
+
+ @classmethod
+ async def sleep(cls, delay: float) -> None:
+ await sleep(delay)
+
+ @classmethod
+ def create_cancel_scope(
+ cls, *, deadline: float = math.inf, shield: bool = False
+ ) -> CancelScope:
+ return CancelScope(deadline=deadline, shield=shield)
+
+ @classmethod
+ def current_effective_deadline(cls) -> float:
+ if (task := current_task()) is None:
+ return math.inf
+
+ try:
+ cancel_scope = _task_states[task].cancel_scope
+ except KeyError:
+ return math.inf
+
+ deadline = math.inf
+ while cancel_scope:
+ deadline = min(deadline, cancel_scope.deadline)
+ if cancel_scope._cancel_called:
+ deadline = -math.inf
+ break
+ elif cancel_scope.shield:
+ break
+ else:
+ cancel_scope = cancel_scope._parent_scope
+
+ return deadline
+
+ @classmethod
+ def create_task_group(cls) -> abc.TaskGroup:
+ return TaskGroup()
+
+ @classmethod
+ def create_event(cls) -> abc.Event:
+ return Event()
+
+ @classmethod
+ def create_lock(cls, *, fast_acquire: bool) -> abc.Lock:
+ return Lock(fast_acquire=fast_acquire)
+
+ @classmethod
+ def create_semaphore(
+ cls,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> abc.Semaphore:
+ return Semaphore(initial_value, max_value=max_value, fast_acquire=fast_acquire)
+
+ @classmethod
+ def create_capacity_limiter(cls, total_tokens: float) -> abc.CapacityLimiter:
+ return CapacityLimiter(total_tokens)
+
+ @classmethod
+ async def run_sync_in_worker_thread( # type: ignore[return]
+ cls,
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ abandon_on_cancel: bool = False,
+ limiter: abc.CapacityLimiter | None = None,
+ ) -> T_Retval:
+ await cls.checkpoint()
+
+ # If this is the first run in this event loop thread, set up the necessary
+ # variables
+ try:
+ idle_workers = _threadpool_idle_workers.get()
+ workers = _threadpool_workers.get()
+ except LookupError:
+ idle_workers = deque()
+ workers = set()
+ _threadpool_idle_workers.set(idle_workers)
+ _threadpool_workers.set(workers)
+
+ async with limiter or cls.current_default_thread_limiter():
+ with CancelScope(shield=not abandon_on_cancel) as scope:
+ future = asyncio.Future[T_Retval]()
+ root_task = find_root_task()
+ if not idle_workers:
+ worker = WorkerThread(root_task, workers, idle_workers)
+ worker.start()
+ workers.add(worker)
+ root_task.add_done_callback(
+ worker.stop, context=contextvars.Context()
+ )
+ else:
+ worker = idle_workers.pop()
+
+ # Prune any other workers that have been idle for MAX_IDLE_TIME
+ # seconds or longer
+ now = cls.current_time()
+ while idle_workers:
+ if (
+ now - idle_workers[0].idle_since
+ < WorkerThread.MAX_IDLE_TIME
+ ):
+ break
+
+ expired_worker = idle_workers.popleft()
+ expired_worker.root_task.remove_done_callback(
+ expired_worker.stop
+ )
+ expired_worker.stop()
+
+ context = copy_context()
+ context.run(set_current_async_library, None)
+ if abandon_on_cancel or scope._parent_scope is None:
+ worker_scope = scope
+ else:
+ worker_scope = scope._parent_scope
+
+ worker.queue.put_nowait((context, func, args, future, worker_scope))
+ return await future
+
+ @classmethod
+ def check_cancelled(cls) -> None:
+ scope: CancelScope | None = threadlocals.current_cancel_scope
+ while scope is not None:
+ if scope.cancel_called:
+ raise CancelledError(f"Cancelled via cancel scope {id(scope):x}")
+
+ if scope.shield:
+ return
+
+ scope = scope._parent_scope
+
+ @classmethod
+ def run_async_from_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ args: tuple[Unpack[PosArgsT]],
+ token: object,
+ ) -> T_co:
+ async def task_wrapper() -> T_co:
+ __tracebackhide__ = True
+ if scope is not None:
+ task = cast(asyncio.Task, current_task())
+ _task_states[task] = TaskState(None, scope)
+ scope._tasks.add(task)
+ try:
+ return await func(*args)
+ except CancelledError as exc:
+ raise concurrent.futures.CancelledError(str(exc)) from None
+ finally:
+ if scope is not None:
+ scope._tasks.discard(task)
+
+ loop = cast(
+ "AbstractEventLoop", token or threadlocals.current_token.native_token
+ )
+ if loop.is_closed():
+ raise RunFinishedError
+
+ context = copy_context()
+ context.run(set_current_async_library, "asyncio")
+ scope = getattr(threadlocals, "current_cancel_scope", None)
+ f: concurrent.futures.Future[T_co] = context.run(
+ asyncio.run_coroutine_threadsafe, task_wrapper(), loop=loop
+ )
+ return f.result()
+
+ @classmethod
+ def run_sync_from_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ token: object,
+ ) -> T_Retval:
+ @wraps(func)
+ def wrapper() -> None:
+ try:
+ set_current_async_library("asyncio")
+ f.set_result(func(*args))
+ except BaseException as exc:
+ f.set_exception(exc)
+ if not isinstance(exc, Exception):
+ raise
+
+ loop = cast(
+ "AbstractEventLoop", token or threadlocals.current_token.native_token
+ )
+ if loop.is_closed():
+ raise RunFinishedError
+
+ f: concurrent.futures.Future[T_Retval] = Future()
+ loop.call_soon_threadsafe(wrapper)
+ return f.result()
+
+ @classmethod
+ async def open_process(
+ cls,
+ command: StrOrBytesPath | Sequence[StrOrBytesPath],
+ *,
+ stdin: int | IO[Any] | None,
+ stdout: int | IO[Any] | None,
+ stderr: int | IO[Any] | None,
+ **kwargs: Any,
+ ) -> Process:
+ await cls.checkpoint()
+ if isinstance(command, PathLike):
+ command = os.fspath(command)
+
+ # Use loop.subprocess_shell()/subprocess_exec() rather than their
+ # asyncio.create_subprocess_*() counterparts to get access to
+ # transport/protocol.
+ loop = asyncio.get_running_loop()
+ if isinstance(command, (str, bytes)):
+ transport, protocol = await loop.subprocess_shell(
+ _ProcessStreamProtocol,
+ command,
+ stdin=stdin,
+ stdout=stdout,
+ stderr=stderr,
+ **kwargs,
+ )
+ else:
+ transport, protocol = await loop.subprocess_exec(
+ _ProcessStreamProtocol,
+ *command,
+ stdin=stdin,
+ stdout=stdout,
+ stderr=stderr,
+ **kwargs,
+ )
+
+ process = asyncio.subprocess.Process(transport, protocol, loop)
+ stdin_stream = StreamWriterWrapper(process.stdin) if process.stdin else None
+ stdout_stream = StreamReaderWrapper(process.stdout) if process.stdout else None
+ stderr_stream = StreamReaderWrapper(process.stderr) if process.stderr else None
+ return Process(
+ process,
+ stdin_stream,
+ stdout_stream,
+ stderr_stream,
+ protocol.exited,
+ transport,
+ )
+
+ @classmethod
+ def setup_process_pool_exit_at_shutdown(cls, workers: set[abc.Process]) -> None:
+ create_task(
+ _shutdown_process_pool_on_exit(workers),
+ name="AnyIO process pool shutdown task",
+ )
+ find_root_task().add_done_callback(
+ partial(_forcibly_shutdown_process_pool_on_exit, workers) # type:ignore[arg-type]
+ )
+
+ @classmethod
+ async def connect_tcp(
+ cls, host: str, port: int, local_address: IPSockAddrType | None = None
+ ) -> abc.SocketStream:
+ transport, protocol = cast(
+ tuple[asyncio.Transport, StreamProtocol],
+ await get_running_loop().create_connection(
+ StreamProtocol, host, port, local_addr=local_address
+ ),
+ )
+ transport.pause_reading()
+ return SocketStream(transport, protocol)
+
+ @classmethod
+ async def connect_unix(cls, path: str | bytes) -> abc.UNIXSocketStream:
+ await cls.checkpoint()
+ loop = get_running_loop()
+ raw_socket = socket.socket(socket.AF_UNIX)
+ raw_socket.setblocking(False)
+ while True:
+ try:
+ raw_socket.connect(path)
+ except BlockingIOError:
+ f: asyncio.Future = asyncio.Future()
+ loop.add_writer(raw_socket, f.set_result, None)
+ f.add_done_callback(lambda _: loop.remove_writer(raw_socket))
+ await f
+ except BaseException:
+ raw_socket.close()
+ raise
+ else:
+ return UNIXSocketStream(raw_socket)
+
+ @classmethod
+ def create_tcp_listener(cls, sock: socket.socket) -> SocketListener:
+ return TCPSocketListener(sock)
+
+ @classmethod
+ def create_unix_listener(cls, sock: socket.socket) -> SocketListener:
+ return UNIXSocketListener(sock)
+
+ @classmethod
+ async def create_udp_socket(
+ cls,
+ family: AddressFamily,
+ local_address: IPSockAddrType | None,
+ remote_address: IPSockAddrType | None,
+ reuse_port: bool,
+ ) -> UDPSocket | ConnectedUDPSocket:
+ transport, protocol = await get_running_loop().create_datagram_endpoint(
+ DatagramProtocol,
+ local_addr=local_address,
+ remote_addr=remote_address,
+ family=family,
+ reuse_port=reuse_port,
+ )
+ if protocol.exception:
+ transport.close()
+ raise protocol.exception
+
+ if not remote_address:
+ return UDPSocket(transport, protocol)
+ else:
+ return ConnectedUDPSocket(transport, protocol)
+
+ @classmethod
+ async def create_unix_datagram_socket( # type: ignore[override]
+ cls, raw_socket: socket.socket, remote_path: str | bytes | None
+ ) -> abc.UNIXDatagramSocket | abc.ConnectedUNIXDatagramSocket:
+ await cls.checkpoint()
+ loop = get_running_loop()
+
+ if remote_path:
+ while True:
+ try:
+ raw_socket.connect(remote_path)
+ except BlockingIOError:
+ f: asyncio.Future = asyncio.Future()
+ loop.add_writer(raw_socket, f.set_result, None)
+ f.add_done_callback(lambda _: loop.remove_writer(raw_socket))
+ await f
+ except BaseException:
+ raw_socket.close()
+ raise
+ else:
+ return ConnectedUNIXDatagramSocket(raw_socket)
+ else:
+ return UNIXDatagramSocket(raw_socket)
+
+ @classmethod
+ async def getaddrinfo(
+ cls,
+ host: bytes | str | None,
+ port: str | int | None,
+ *,
+ family: int | AddressFamily = 0,
+ type: int | SocketKind = 0,
+ proto: int = 0,
+ flags: int = 0,
+ ) -> Sequence[
+ tuple[
+ AddressFamily,
+ SocketKind,
+ int,
+ str,
+ tuple[str, int] | tuple[str, int, int, int] | tuple[int, bytes],
+ ]
+ ]:
+ return await get_running_loop().getaddrinfo(
+ host, port, family=family, type=type, proto=proto, flags=flags
+ )
+
+ @classmethod
+ async def getnameinfo(
+ cls, sockaddr: IPSockAddrType, flags: int = 0
+ ) -> tuple[str, str]:
+ return await get_running_loop().getnameinfo(sockaddr, flags)
+
+ @classmethod
+ async def wait_readable(cls, obj: FileDescriptorLike) -> None:
+ try:
+ read_events = _read_events.get()
+ except LookupError:
+ read_events = {}
+ _read_events.set(read_events)
+
+ fd = obj if isinstance(obj, int) else obj.fileno()
+ if read_events.get(fd):
+ raise BusyResourceError("reading from")
+
+ loop = get_running_loop()
+ fut: asyncio.Future[bool] = loop.create_future()
+
+ def cb() -> None:
+ try:
+ del read_events[fd]
+ except KeyError:
+ pass
+ else:
+ remove_reader(fd)
+
+ try:
+ fut.set_result(True)
+ except asyncio.InvalidStateError:
+ pass
+
+ try:
+ loop.add_reader(fd, cb)
+ except NotImplementedError:
+ from anyio._core._asyncio_selector_thread import get_selector
+
+ selector = get_selector()
+ selector.add_reader(fd, cb)
+ remove_reader = selector.remove_reader
+ else:
+ remove_reader = loop.remove_reader
+
+ read_events[fd] = fut
+ try:
+ success = await fut
+ finally:
+ try:
+ del read_events[fd]
+ except KeyError:
+ pass
+ else:
+ remove_reader(fd)
+
+ if not success:
+ raise ClosedResourceError
+
+ @classmethod
+ async def wait_writable(cls, obj: FileDescriptorLike) -> None:
+ try:
+ write_events = _write_events.get()
+ except LookupError:
+ write_events = {}
+ _write_events.set(write_events)
+
+ fd = obj if isinstance(obj, int) else obj.fileno()
+ if write_events.get(fd):
+ raise BusyResourceError("writing to")
+
+ loop = get_running_loop()
+ fut: asyncio.Future[bool] = loop.create_future()
+
+ def cb() -> None:
+ try:
+ del write_events[fd]
+ except KeyError:
+ pass
+ else:
+ remove_writer(fd)
+
+ try:
+ fut.set_result(True)
+ except asyncio.InvalidStateError:
+ pass
+
+ try:
+ loop.add_writer(fd, cb)
+ except NotImplementedError:
+ from anyio._core._asyncio_selector_thread import get_selector
+
+ selector = get_selector()
+ selector.add_writer(fd, cb)
+ remove_writer = selector.remove_writer
+ else:
+ remove_writer = loop.remove_writer
+
+ write_events[fd] = fut
+ try:
+ success = await fut
+ finally:
+ try:
+ del write_events[fd]
+ except KeyError:
+ pass
+ else:
+ remove_writer(fd)
+
+ if not success:
+ raise ClosedResourceError
+
+ @classmethod
+ def notify_closing(cls, obj: FileDescriptorLike) -> None:
+ fd = obj if isinstance(obj, int) else obj.fileno()
+ loop = get_running_loop()
+
+ try:
+ write_events = _write_events.get()
+ except LookupError:
+ pass
+ else:
+ try:
+ fut = write_events.pop(fd)
+ except KeyError:
+ pass
+ else:
+ try:
+ fut.set_result(False)
+ except asyncio.InvalidStateError:
+ pass
+
+ try:
+ loop.remove_writer(fd)
+ except NotImplementedError:
+ from anyio._core._asyncio_selector_thread import get_selector
+
+ get_selector().remove_writer(fd)
+
+ try:
+ read_events = _read_events.get()
+ except LookupError:
+ pass
+ else:
+ try:
+ fut = read_events.pop(fd)
+ except KeyError:
+ pass
+ else:
+ try:
+ fut.set_result(False)
+ except asyncio.InvalidStateError:
+ pass
+
+ try:
+ loop.remove_reader(fd)
+ except NotImplementedError:
+ from anyio._core._asyncio_selector_thread import get_selector
+
+ get_selector().remove_reader(fd)
+
+ @classmethod
+ async def wrap_listener_socket(cls, sock: socket.socket) -> SocketListener:
+ if hasattr(socket, "AF_UNIX") and sock.family == socket.AF_UNIX:
+ return UNIXSocketListener(sock)
+
+ return TCPSocketListener(sock)
+
+ @classmethod
+ async def wrap_stream_socket(cls, sock: socket.socket) -> SocketStream:
+ transport, protocol = await get_running_loop().create_connection(
+ StreamProtocol, sock=sock
+ )
+ return SocketStream(transport, protocol)
+
+ @classmethod
+ async def wrap_unix_stream_socket(cls, sock: socket.socket) -> UNIXSocketStream:
+ return UNIXSocketStream(sock)
+
+ @classmethod
+ async def wrap_udp_socket(cls, sock: socket.socket) -> UDPSocket:
+ transport, protocol = await get_running_loop().create_datagram_endpoint(
+ DatagramProtocol, sock=sock
+ )
+ return UDPSocket(transport, protocol)
+
+ @classmethod
+ async def wrap_connected_udp_socket(cls, sock: socket.socket) -> ConnectedUDPSocket:
+ transport, protocol = await get_running_loop().create_datagram_endpoint(
+ DatagramProtocol, sock=sock
+ )
+ return ConnectedUDPSocket(transport, protocol)
+
+ @classmethod
+ async def wrap_unix_datagram_socket(cls, sock: socket.socket) -> UNIXDatagramSocket:
+ return UNIXDatagramSocket(sock)
+
+ @classmethod
+ async def wrap_connected_unix_datagram_socket(
+ cls, sock: socket.socket
+ ) -> ConnectedUNIXDatagramSocket:
+ return ConnectedUNIXDatagramSocket(sock)
+
+ @classmethod
+ def current_default_thread_limiter(cls) -> CapacityLimiter:
+ try:
+ return _default_thread_limiter.get()
+ except LookupError:
+ limiter = CapacityLimiter(40)
+ _default_thread_limiter.set(limiter)
+ return limiter
+
+ @classmethod
+ def open_signal_receiver(
+ cls, *signals: Signals
+ ) -> AbstractContextManager[AsyncIterator[Signals]]:
+ return _SignalReceiver(signals)
+
+ @classmethod
+ def get_current_task(cls) -> TaskInfo:
+ return AsyncIOTaskInfo(current_task()) # type: ignore[arg-type]
+
+ @classmethod
+ def get_running_tasks(cls) -> Sequence[TaskInfo]:
+ return [AsyncIOTaskInfo(task) for task in all_tasks() if not task.done()]
+
+ @classmethod
+ async def wait_all_tasks_blocked(cls) -> None:
+ await cls.checkpoint()
+ this_task = current_task()
+ while True:
+ for task in all_tasks():
+ if task is this_task:
+ continue
+
+ waiter = task._fut_waiter # type: ignore[attr-defined]
+ if waiter is None or waiter.done():
+ await sleep(0.1)
+ break
+ else:
+ return
+
+ @classmethod
+ def create_test_runner(cls, options: dict[str, Any]) -> TestRunner:
+ return TestRunner(**options)
+
+
+backend_class = AsyncIOBackend
diff --git a/venv/lib/python3.11/site-packages/anyio/_backends/_trio.py b/venv/lib/python3.11/site-packages/anyio/_backends/_trio.py
new file mode 100644
index 0000000000000000000000000000000000000000..43d24d23315322eef2df539d5d959b66fc1e2dbb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_backends/_trio.py
@@ -0,0 +1,1468 @@
+from __future__ import annotations
+
+import array
+import math
+import os
+import socket
+import sys
+import types
+import weakref
+from collections.abc import (
+ AsyncGenerator,
+ AsyncIterator,
+ Awaitable,
+ Callable,
+ Collection,
+ Coroutine,
+ Iterable,
+ Sequence,
+)
+from contextlib import AbstractContextManager
+from contextvars import Context
+from dataclasses import dataclass
+from functools import partial, wraps
+from io import IOBase
+from os import PathLike
+from signal import Signals
+from socket import AddressFamily, SocketKind
+from types import TracebackType
+from typing import (
+ IO,
+ TYPE_CHECKING,
+ Any,
+ Generic,
+ Literal,
+ NoReturn,
+ ParamSpec,
+ TypeVar,
+ cast,
+ overload,
+)
+
+import trio.from_thread
+import trio.lowlevel
+from outcome import Error, Outcome, Value
+from trio.lowlevel import (
+ current_root_task,
+ current_task,
+ notify_closing,
+ wait_readable,
+ wait_writable,
+)
+from trio.socket import SocketType as TrioSocketType
+from trio.to_thread import run_sync
+
+from .. import (
+ CapacityLimiterStatistics,
+ EventStatistics,
+ LockStatistics,
+ RunFinishedError,
+ TaskInfo,
+ WouldBlock,
+ abc,
+)
+from .._core._eventloop import claim_worker_thread
+from .._core._exceptions import (
+ BrokenResourceError,
+ BusyResourceError,
+ ClosedResourceError,
+ EndOfStream,
+)
+from .._core._sockets import convert_ipv6_sockaddr
+from .._core._streams import create_memory_object_stream
+from .._core._synchronization import (
+ CapacityLimiter as BaseCapacityLimiter,
+)
+from .._core._synchronization import Event as BaseEvent
+from .._core._synchronization import Lock as BaseLock
+from .._core._synchronization import (
+ ResourceGuard,
+ SemaphoreStatistics,
+)
+from .._core._synchronization import Semaphore as BaseSemaphore
+from .._core._tasks import CancelScope as BaseCancelScope
+from .._core._tasks import TaskHandle
+from ..abc import IPSockAddrType, UDPPacketType, UNIXDatagramPacketType
+from ..abc._eventloop import AsyncBackend, StrOrBytesPath
+from ..abc._tasks import T_contra, call_for_coroutine, get_callable_name
+from ..streams.memory import MemoryObjectSendStream
+
+if TYPE_CHECKING:
+ from _typeshed import FileDescriptorLike
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from exceptiongroup import BaseExceptionGroup
+ from typing_extensions import TypeVarTuple, Unpack
+
+T = TypeVar("T")
+T_Retval = TypeVar("T_Retval")
+T_co = TypeVar("T_co", covariant=True)
+T_SockAddr = TypeVar("T_SockAddr", str, IPSockAddrType)
+PosArgsT = TypeVarTuple("PosArgsT")
+P = ParamSpec("P")
+
+
+def ensure_returns_coro(
+ func: Callable[P, Awaitable[T_Retval]],
+) -> Callable[P, Coroutine[Any, Any, T_Retval]]:
+ @wraps(func)
+ def wrapper(*args: P.args, **kwargs: P.kwargs) -> Coroutine[Any, Any, T_Retval]:
+ awaitable = func(*args, **kwargs)
+ # Check the common case first.
+ if isinstance(awaitable, Coroutine):
+ return awaitable
+ elif not isinstance(awaitable, Awaitable):
+ # The user violated the type annotations. Still, we should pass this on to
+ # Trio so it can raise with an appropriate message.
+ return awaitable
+ else:
+
+ @wraps(func)
+ async def inner_wrapper() -> T_Retval:
+ return await awaitable
+
+ return inner_wrapper()
+
+ return wrapper
+
+
+#
+# Event loop
+#
+
+RunVar = trio.lowlevel.RunVar
+
+
+#
+# Timeouts and cancellation
+#
+
+
+class CancelScope(BaseCancelScope):
+ __slots__ = ("__original",)
+
+ def __new__(
+ cls, original: trio.CancelScope | None = None, **kwargs: object
+ ) -> CancelScope:
+ return object.__new__(cls)
+
+ def __init__(self, original: trio.CancelScope | None = None, **kwargs: Any) -> None:
+ self.__original = original or trio.CancelScope(**kwargs)
+
+ def __enter__(self) -> CancelScope:
+ self.__original.__enter__()
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool:
+ return self.__original.__exit__(exc_type, exc_val, exc_tb)
+
+ def cancel(self, reason: str | None = None) -> None:
+ self.__original.cancel(reason)
+
+ @property
+ def deadline(self) -> float:
+ return self.__original.deadline
+
+ @deadline.setter
+ def deadline(self, value: float) -> None:
+ self.__original.deadline = value
+
+ @property
+ def cancel_called(self) -> bool:
+ return self.__original.cancel_called
+
+ @property
+ def cancelled_caught(self) -> bool:
+ return self.__original.cancelled_caught
+
+ @property
+ def shield(self) -> bool:
+ return self.__original.shield
+
+ @shield.setter
+ def shield(self, value: bool) -> None:
+ self.__original.shield = value
+
+
+#
+# Task groups
+#
+
+empty_start_value = object()
+
+
+class _TrioTaskStatus(Generic[T_contra], abc.TaskStatus[T_contra]):
+ early_start_value: T_contra | object = empty_start_value
+ real_task_status: trio.TaskStatus[T_contra | None] | None = None
+
+ def started(self, value: T_contra | None = None) -> None:
+ if self.real_task_status is None:
+ if self.early_start_value is not empty_start_value:
+ raise RuntimeError("called 'started' twice on the same task status")
+
+ self.early_start_value = value
+ else:
+ self.real_task_status.started(value)
+
+
+class TaskGroup(abc.TaskGroup):
+ def __init__(self) -> None:
+ self._entered = False
+ self._active = False
+ self._nursery_manager = trio.open_nursery(strict_exception_groups=True)
+ self.cancel_scope = None # type: ignore[assignment]
+
+ async def __aenter__(self) -> TaskGroup:
+ if self._entered:
+ raise RuntimeError("TaskGroup cannot be entered more than once")
+
+ self._entered = True
+ self._active = True
+ self._nursery = await self._nursery_manager.__aenter__()
+ self.cancel_scope = CancelScope(self._nursery.cancel_scope)
+ return self
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool:
+ try:
+ # trio.Nursery.__exit__ returns bool; .open_nursery has wrong type
+ return await self._nursery_manager.__aexit__(exc_type, exc_val, exc_tb) # type: ignore[return-value]
+ except BaseExceptionGroup as exc:
+ if not exc.split(trio.Cancelled)[1]:
+ raise trio.Cancelled._create() from exc
+
+ raise
+ finally:
+ del exc_val, exc_tb
+ self._active = False
+
+ def _check_active(self, coro: Coroutine | None = None) -> None:
+ if not self._active:
+ if coro is not None:
+ coro.close()
+
+ raise RuntimeError(
+ "This task group is not active; no new tasks can be started."
+ )
+
+ def create_task(
+ self,
+ coro: Coroutine[Any, Any, T_co],
+ *,
+ name: object = None,
+ context: Context | None = None,
+ ) -> TaskHandle[T_co]:
+ if not isinstance(coro, Coroutine):
+ raise TypeError(f"expected a coroutine, got {coro.__class__.__qualname__}")
+
+ self._check_active(coro)
+ handle = TaskHandle(coro, name)
+ if context is not None:
+ context.run(
+ partial(self._nursery.start_soon, handle._run_coro, name=handle.name)
+ )
+ else:
+ self._nursery.start_soon(handle._run_coro, name=handle.name)
+
+ return handle
+
+ async def start(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ *args: Unpack[PosArgsT],
+ name: object = None,
+ return_handle: Literal[False] | Literal[True] = False,
+ ) -> Any:
+ handle: TaskHandle[T_co]
+
+ async def run_coro_with_task_status(
+ *, task_status: trio.TaskStatus[Any]
+ ) -> None:
+ nonlocal handle
+ wrapper_task_status = _TrioTaskStatus()
+ coro = call_for_coroutine(func, args, task_status=wrapper_task_status)
+ if wrapper_task_status.early_start_value is not empty_start_value:
+ task_status.started(wrapper_task_status.early_start_value)
+ else:
+ wrapper_task_status.real_task_status = task_status
+
+ handle = TaskHandle(coro, name)
+ await handle._run_coro()
+
+ self._check_active()
+ final_name = get_callable_name(func, name)
+ start_value = await self._nursery.start(
+ run_coro_with_task_status, name=final_name
+ )
+ if return_handle:
+ handle._start_value = start_value
+ return handle
+ else:
+ return start_value
+
+
+#
+# Subprocesses
+#
+
+
+@dataclass(eq=False)
+class ReceiveStreamWrapper(abc.ByteReceiveStream):
+ _stream: trio.abc.ReceiveStream
+
+ async def receive(self, max_bytes: int | None = None) -> bytes:
+ if max_bytes is not None and max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ try:
+ data = await self._stream.receive_some(max_bytes)
+ except trio.ClosedResourceError as exc:
+ raise ClosedResourceError from exc.__cause__
+ except trio.BrokenResourceError as exc:
+ raise BrokenResourceError from exc.__cause__
+
+ if data:
+ return bytes(data)
+ else:
+ raise EndOfStream
+
+ async def aclose(self) -> None:
+ await self._stream.aclose()
+
+
+@dataclass(eq=False)
+class SendStreamWrapper(abc.ByteSendStream):
+ _stream: trio.abc.SendStream
+
+ async def send(self, item: bytes) -> None:
+ try:
+ await self._stream.send_all(item)
+ except trio.ClosedResourceError as exc:
+ raise ClosedResourceError from exc.__cause__
+ except trio.BrokenResourceError as exc:
+ raise BrokenResourceError from exc.__cause__
+
+ async def aclose(self) -> None:
+ await self._stream.aclose()
+
+
+@dataclass(eq=False)
+class Process(abc.Process):
+ _process: trio.Process
+ _stdin: abc.ByteSendStream | None
+ _stdout: abc.ByteReceiveStream | None
+ _stderr: abc.ByteReceiveStream | None
+
+ async def aclose(self) -> None:
+ with CancelScope(shield=True):
+ if self._stdin:
+ await self._stdin.aclose()
+ if self._stdout:
+ await self._stdout.aclose()
+ if self._stderr:
+ await self._stderr.aclose()
+
+ try:
+ await self.wait()
+ except BaseException:
+ self.kill()
+ with CancelScope(shield=True):
+ await self.wait()
+ raise
+
+ async def wait(self) -> int:
+ return await self._process.wait()
+
+ def terminate(self) -> None:
+ self._process.terminate()
+
+ def kill(self) -> None:
+ self._process.kill()
+
+ def send_signal(self, signal: Signals) -> None:
+ self._process.send_signal(signal)
+
+ @property
+ def pid(self) -> int:
+ return self._process.pid
+
+ @property
+ def returncode(self) -> int | None:
+ return self._process.returncode
+
+ @property
+ def stdin(self) -> abc.ByteSendStream | None:
+ return self._stdin
+
+ @property
+ def stdout(self) -> abc.ByteReceiveStream | None:
+ return self._stdout
+
+ @property
+ def stderr(self) -> abc.ByteReceiveStream | None:
+ return self._stderr
+
+
+class _ProcessPoolShutdownInstrument(trio.abc.Instrument):
+ def after_run(self) -> None:
+ super().after_run()
+
+
+current_default_worker_process_limiter: trio.lowlevel.RunVar = RunVar(
+ "current_default_worker_process_limiter"
+)
+
+
+async def _shutdown_process_pool(workers: set[abc.Process]) -> None:
+ try:
+ await trio.sleep(math.inf)
+ except trio.Cancelled:
+ for process in workers:
+ if process.returncode is None:
+ process.kill()
+
+ with CancelScope(shield=True):
+ for process in workers:
+ await process.aclose()
+
+
+#
+# Sockets and networking
+#
+
+
+class _TrioSocketMixin(Generic[T_SockAddr]):
+ def __init__(self, trio_socket: TrioSocketType) -> None:
+ self._trio_socket = trio_socket
+ self._closed = False
+
+ def _check_closed(self) -> None:
+ if self._closed:
+ raise ClosedResourceError
+ if self._trio_socket.fileno() < 0:
+ raise BrokenResourceError
+
+ @property
+ def _raw_socket(self) -> socket.socket:
+ return self._trio_socket._sock # type: ignore[attr-defined]
+
+ async def aclose(self) -> None:
+ if self._trio_socket.fileno() >= 0:
+ self._closed = True
+ self._trio_socket.close()
+
+ def _convert_socket_error(self, exc: BaseException) -> NoReturn:
+ if isinstance(exc, trio.ClosedResourceError):
+ raise ClosedResourceError from exc
+ elif self._trio_socket.fileno() < 0 and self._closed:
+ raise ClosedResourceError from None
+ elif isinstance(exc, OSError):
+ raise BrokenResourceError from exc
+ else:
+ raise exc
+
+
+class SocketStream(_TrioSocketMixin, abc.SocketStream):
+ def __init__(self, trio_socket: TrioSocketType) -> None:
+ super().__init__(trio_socket)
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ with self._receive_guard:
+ try:
+ data = await self._trio_socket.recv(max_bytes)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ if data:
+ return data
+ else:
+ raise EndOfStream
+
+ async def send(self, item: bytes) -> None:
+ with self._send_guard:
+ view = memoryview(item)
+ while view:
+ try:
+ bytes_sent = await self._trio_socket.send(view)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ view = view[bytes_sent:]
+
+ async def send_eof(self) -> None:
+ self._trio_socket.shutdown(socket.SHUT_WR)
+
+
+class UNIXSocketStream(SocketStream, abc.UNIXSocketStream):
+ async def receive_fds(self, msglen: int, maxfds: int) -> tuple[bytes, list[int]]:
+ if not isinstance(msglen, int) or msglen < 0:
+ raise ValueError("msglen must be a non-negative integer")
+ if not isinstance(maxfds, int) or maxfds < 1:
+ raise ValueError("maxfds must be a positive integer")
+
+ fds = array.array("i")
+ await trio.lowlevel.checkpoint()
+ with self._receive_guard:
+ while True:
+ try:
+ message, ancdata, flags, addr = await self._trio_socket.recvmsg(
+ msglen, socket.CMSG_LEN(maxfds * fds.itemsize)
+ )
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+ else:
+ if not message and not ancdata:
+ raise EndOfStream
+
+ break
+
+ for cmsg_level, cmsg_type, cmsg_data in ancdata:
+ if cmsg_level != socket.SOL_SOCKET or cmsg_type != socket.SCM_RIGHTS:
+ raise RuntimeError(
+ f"Received unexpected ancillary data; message = {message!r}, "
+ f"cmsg_level = {cmsg_level}, cmsg_type = {cmsg_type}"
+ )
+
+ fds.frombytes(cmsg_data[: len(cmsg_data) - (len(cmsg_data) % fds.itemsize)])
+
+ return message, list(fds)
+
+ async def send_fds(self, message: bytes, fds: Collection[int | IOBase]) -> None:
+ if not message:
+ raise ValueError("message must not be empty")
+ if not fds:
+ raise ValueError("fds must not be empty")
+
+ filenos: list[int] = []
+ for fd in fds:
+ if isinstance(fd, int):
+ filenos.append(fd)
+ elif isinstance(fd, IOBase):
+ filenos.append(fd.fileno())
+
+ fdarray = array.array("i", filenos)
+ await trio.lowlevel.checkpoint()
+ with self._send_guard:
+ while True:
+ try:
+ await self._trio_socket.sendmsg(
+ [message],
+ [
+ (
+ socket.SOL_SOCKET,
+ socket.SCM_RIGHTS,
+ fdarray,
+ )
+ ],
+ )
+ break
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+
+class TCPSocketListener(_TrioSocketMixin, abc.SocketListener):
+ def __init__(self, raw_socket: socket.socket):
+ super().__init__(trio.socket.from_stdlib_socket(raw_socket))
+ self._accept_guard = ResourceGuard("accepting connections from")
+
+ async def accept(self) -> SocketStream:
+ with self._accept_guard:
+ try:
+ trio_socket, _addr = await self._trio_socket.accept()
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ trio_socket.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
+ return SocketStream(trio_socket)
+
+
+class UNIXSocketListener(_TrioSocketMixin, abc.SocketListener):
+ def __init__(self, raw_socket: socket.socket):
+ super().__init__(trio.socket.from_stdlib_socket(raw_socket))
+ self._accept_guard = ResourceGuard("accepting connections from")
+
+ async def accept(self) -> UNIXSocketStream:
+ with self._accept_guard:
+ try:
+ trio_socket, _addr = await self._trio_socket.accept()
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ return UNIXSocketStream(trio_socket)
+
+
+class UDPSocket(_TrioSocketMixin[IPSockAddrType], abc.UDPSocket):
+ def __init__(self, trio_socket: TrioSocketType) -> None:
+ super().__init__(trio_socket)
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+
+ async def receive(self) -> tuple[bytes, IPSockAddrType]:
+ with self._receive_guard:
+ try:
+ data, addr = await self._trio_socket.recvfrom(65536)
+ return data, convert_ipv6_sockaddr(addr)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ async def send(self, item: UDPPacketType) -> None:
+ with self._send_guard:
+ try:
+ await self._trio_socket.sendto(*item)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+
+class ConnectedUDPSocket(_TrioSocketMixin[IPSockAddrType], abc.ConnectedUDPSocket):
+ def __init__(self, trio_socket: TrioSocketType) -> None:
+ super().__init__(trio_socket)
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+
+ async def receive(self) -> bytes:
+ with self._receive_guard:
+ try:
+ return await self._trio_socket.recv(65536)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ async def send(self, item: bytes) -> None:
+ with self._send_guard:
+ try:
+ await self._trio_socket.send(item)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+
+class UNIXDatagramSocket(_TrioSocketMixin[str], abc.UNIXDatagramSocket):
+ def __init__(self, trio_socket: TrioSocketType) -> None:
+ super().__init__(trio_socket)
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+
+ async def receive(self) -> UNIXDatagramPacketType:
+ with self._receive_guard:
+ try:
+ data, addr = await self._trio_socket.recvfrom(65536)
+ return data, addr
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ async def send(self, item: UNIXDatagramPacketType) -> None:
+ with self._send_guard:
+ try:
+ await self._trio_socket.sendto(*item)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+
+class ConnectedUNIXDatagramSocket(
+ _TrioSocketMixin[str], abc.ConnectedUNIXDatagramSocket
+):
+ def __init__(self, trio_socket: TrioSocketType) -> None:
+ super().__init__(trio_socket)
+ self._receive_guard = ResourceGuard("reading from")
+ self._send_guard = ResourceGuard("writing to")
+
+ async def receive(self) -> bytes:
+ with self._receive_guard:
+ try:
+ return await self._trio_socket.recv(65536)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+ async def send(self, item: bytes) -> None:
+ with self._send_guard:
+ try:
+ await self._trio_socket.send(item)
+ except BaseException as exc:
+ self._convert_socket_error(exc)
+
+
+#
+# Synchronization
+#
+
+
+class Event(BaseEvent):
+ __slots__ = ("__original",)
+
+ def __new__(cls) -> Event:
+ return object.__new__(cls)
+
+ def __init__(self) -> None:
+ self.__original = trio.Event()
+
+ def is_set(self) -> bool:
+ return self.__original.is_set()
+
+ async def wait(self) -> None:
+ return await self.__original.wait()
+
+ def statistics(self) -> EventStatistics:
+ orig_statistics = self.__original.statistics()
+ return EventStatistics(tasks_waiting=orig_statistics.tasks_waiting)
+
+ def set(self) -> None:
+ self.__original.set()
+
+
+class Lock(BaseLock):
+ __slots__ = "_fast_acquire", "__original"
+
+ def __new__(cls, *, fast_acquire: bool = False) -> Lock:
+ return object.__new__(cls)
+
+ def __init__(self, *, fast_acquire: bool = False) -> None:
+ self._fast_acquire = fast_acquire
+ self.__original = trio.Lock()
+
+ @staticmethod
+ def _convert_runtime_error_msg(exc: RuntimeError) -> None:
+ if exc.args == ("attempt to re-acquire an already held Lock",):
+ exc.args = ("Attempted to acquire an already held Lock",)
+
+ async def acquire(self) -> None:
+ if not self._fast_acquire:
+ try:
+ await self.__original.acquire()
+ except RuntimeError as exc:
+ self._convert_runtime_error_msg(exc)
+ raise
+
+ return
+
+ # This is the "fast path" where we don't let other tasks run
+ await trio.lowlevel.checkpoint_if_cancelled()
+ try:
+ self.__original.acquire_nowait()
+ except trio.WouldBlock:
+ await self.__original._lot.park()
+ except RuntimeError as exc:
+ self._convert_runtime_error_msg(exc)
+ raise
+
+ def acquire_nowait(self) -> None:
+ try:
+ self.__original.acquire_nowait()
+ except trio.WouldBlock:
+ raise WouldBlock from None
+ except RuntimeError as exc:
+ self._convert_runtime_error_msg(exc)
+ raise
+
+ def locked(self) -> bool:
+ return self.__original.locked()
+
+ def release(self) -> None:
+ self.__original.release()
+
+ def statistics(self) -> LockStatistics:
+ orig_statistics = self.__original.statistics()
+ owner = TrioTaskInfo(orig_statistics.owner) if orig_statistics.owner else None
+ return LockStatistics(
+ orig_statistics.locked, owner, orig_statistics.tasks_waiting
+ )
+
+
+class Semaphore(BaseSemaphore):
+ __slots__ = ("__original",)
+
+ def __new__(
+ cls,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> Semaphore:
+ return object.__new__(cls)
+
+ def __init__(
+ self,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> None:
+ super().__init__(initial_value, max_value=max_value, fast_acquire=fast_acquire)
+ self.__original = trio.Semaphore(initial_value, max_value=max_value)
+
+ async def acquire(self) -> None:
+ if not self._fast_acquire:
+ await self.__original.acquire()
+ return
+
+ # This is the "fast path" where we don't let other tasks run
+ await trio.lowlevel.checkpoint_if_cancelled()
+ try:
+ self.__original.acquire_nowait()
+ except trio.WouldBlock:
+ await self.__original._lot.park()
+
+ def acquire_nowait(self) -> None:
+ try:
+ self.__original.acquire_nowait()
+ except trio.WouldBlock:
+ raise WouldBlock from None
+
+ @property
+ def max_value(self) -> int | None:
+ return self.__original.max_value
+
+ @property
+ def value(self) -> int:
+ return self.__original.value
+
+ def release(self) -> None:
+ self.__original.release()
+
+ def statistics(self) -> SemaphoreStatistics:
+ orig_statistics = self.__original.statistics()
+ return SemaphoreStatistics(orig_statistics.tasks_waiting)
+
+
+class CapacityLimiter(BaseCapacityLimiter):
+ __slots__ = ("__original",)
+
+ def __new__(
+ cls,
+ total_tokens: float | None = None,
+ *,
+ original: trio.CapacityLimiter | None = None,
+ ) -> CapacityLimiter:
+ return object.__new__(cls)
+
+ def __init__(
+ self,
+ total_tokens: float | None = None,
+ *,
+ original: trio.CapacityLimiter | None = None,
+ ) -> None:
+ if original is not None:
+ self.__original = original
+ else:
+ assert total_tokens is not None
+ self.__original = trio.CapacityLimiter(total_tokens)
+
+ async def __aenter__(self) -> None:
+ return await self.__original.__aenter__()
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ await self.__original.__aexit__(exc_type, exc_val, exc_tb)
+
+ @property
+ def total_tokens(self) -> float:
+ return self.__original.total_tokens
+
+ @total_tokens.setter
+ def total_tokens(self, value: float) -> None:
+ self.__original.total_tokens = value
+
+ @property
+ def borrowed_tokens(self) -> int:
+ return self.__original.borrowed_tokens
+
+ @property
+ def available_tokens(self) -> float:
+ return self.__original.available_tokens
+
+ def acquire_nowait(self) -> None:
+ try:
+ self.__original.acquire_nowait()
+ except trio.WouldBlock:
+ raise WouldBlock from None
+
+ def acquire_on_behalf_of_nowait(self, borrower: object) -> None:
+ try:
+ self.__original.acquire_on_behalf_of_nowait(borrower)
+ except trio.WouldBlock:
+ raise WouldBlock from None
+
+ async def acquire(self) -> None:
+ await self.__original.acquire()
+
+ async def acquire_on_behalf_of(self, borrower: object) -> None:
+ await self.__original.acquire_on_behalf_of(borrower)
+
+ def release(self) -> None:
+ return self.__original.release()
+
+ def release_on_behalf_of(self, borrower: object) -> None:
+ return self.__original.release_on_behalf_of(borrower)
+
+ def statistics(self) -> CapacityLimiterStatistics:
+ orig = self.__original.statistics()
+ return CapacityLimiterStatistics(
+ borrowed_tokens=orig.borrowed_tokens,
+ total_tokens=orig.total_tokens,
+ borrowers=tuple(orig.borrowers),
+ tasks_waiting=orig.tasks_waiting,
+ )
+
+
+_capacity_limiter_wrapper: trio.lowlevel.RunVar = RunVar("_capacity_limiter_wrapper")
+
+
+#
+# Signal handling
+#
+
+
+class _SignalReceiver:
+ _iterator: AsyncIterator[int]
+
+ def __init__(self, signals: tuple[Signals, ...]):
+ self._signals = signals
+
+ def __enter__(self) -> _SignalReceiver:
+ self._cm = trio.open_signal_receiver(*self._signals)
+ self._iterator = self._cm.__enter__()
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool | None:
+ return self._cm.__exit__(exc_type, exc_val, exc_tb)
+
+ def __aiter__(self) -> _SignalReceiver:
+ return self
+
+ async def __anext__(self) -> Signals:
+ signum = await self._iterator.__anext__()
+ return Signals(signum)
+
+
+#
+# Testing and debugging
+#
+
+
+class TestRunner(abc.TestRunner):
+ def __init__(self, **options: Any) -> None:
+ from queue import Queue
+
+ self._call_queue: Queue[Callable[[], object]] = Queue()
+ self._send_stream: (
+ MemoryObjectSendStream[tuple[Awaitable[Any], list[Outcome]]] | None
+ ) = None
+ self._options = options
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: types.TracebackType | None,
+ ) -> None:
+ if self._send_stream:
+ self._send_stream.close()
+ while self._send_stream is not None:
+ self._call_queue.get()()
+
+ def is_running(self) -> bool:
+ return trio.lowlevel.in_trio_task()
+
+ async def _run_tests_and_fixtures(self) -> None:
+ self._send_stream, receive_stream = create_memory_object_stream[
+ tuple[Awaitable[Any], list[Outcome]]
+ ](1)
+ with receive_stream:
+ async for awaitable, outcome_holder in receive_stream:
+ try:
+ retval = await awaitable
+ except BaseException as exc:
+ outcome_holder.append(Error(exc))
+ else:
+ outcome_holder.append(Value(retval))
+
+ def _main_task_finished(self, outcome: object) -> None:
+ self._send_stream = None
+
+ def _call_in_runner_task(
+ self,
+ func: Callable[P, Awaitable[T_Retval]],
+ /,
+ *args: P.args,
+ **kwargs: P.kwargs,
+ ) -> T_Retval:
+ if self._send_stream is None:
+ trio.lowlevel.start_guest_run(
+ self._run_tests_and_fixtures,
+ run_sync_soon_threadsafe=self._call_queue.put,
+ done_callback=self._main_task_finished,
+ **self._options,
+ )
+ while self._send_stream is None:
+ self._call_queue.get()()
+
+ outcome_holder: list[Outcome] = []
+ self._send_stream.send_nowait((func(*args, **kwargs), outcome_holder))
+ while not outcome_holder:
+ self._call_queue.get()()
+
+ return outcome_holder[0].unwrap()
+
+ def run_asyncgen_fixture(
+ self,
+ fixture_func: Callable[..., AsyncGenerator[T_Retval, Any]],
+ kwargs: dict[str, Any],
+ ) -> Iterable[T_Retval]:
+ asyncgen = fixture_func(**kwargs)
+ fixturevalue: T_Retval = self._call_in_runner_task(asyncgen.asend, None)
+
+ yield fixturevalue
+
+ try:
+ self._call_in_runner_task(asyncgen.asend, None)
+ except StopAsyncIteration:
+ pass
+ else:
+ self._call_in_runner_task(asyncgen.aclose)
+ raise RuntimeError("Async generator fixture did not stop")
+
+ def run_fixture(
+ self,
+ fixture_func: Callable[..., Coroutine[Any, Any, T_Retval]],
+ kwargs: dict[str, Any],
+ ) -> T_Retval:
+ return self._call_in_runner_task(fixture_func, **kwargs)
+
+ def run_test(
+ self, test_func: Callable[..., Coroutine[Any, Any, Any]], kwargs: dict[str, Any]
+ ) -> None:
+ self._call_in_runner_task(test_func, **kwargs)
+
+
+class TrioTaskInfo(TaskInfo):
+ def __init__(self, task: trio.lowlevel.Task):
+ parent_id = None
+ if task.parent_nursery and task.parent_nursery.parent_task:
+ parent_id = id(task.parent_nursery.parent_task)
+
+ super().__init__(id(task), parent_id, task.name, task.coro)
+ self._task = weakref.proxy(task)
+
+ def has_pending_cancellation(self) -> bool:
+ try:
+ return self._task._cancel_status.effectively_cancelled
+ except ReferenceError:
+ # If the task is no longer around, it surely doesn't have a cancellation
+ # pending
+ return False
+
+
+class TrioBackend(AsyncBackend):
+ @classmethod
+ def run(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
+ args: tuple[Unpack[PosArgsT]],
+ kwargs: dict[str, Any],
+ options: dict[str, Any],
+ ) -> T_Retval:
+ assert not kwargs, "unreachable, and not supported by Trio"
+ return trio.run(ensure_returns_coro(func), *args, **options)
+
+ @classmethod
+ def current_token(cls) -> object:
+ return trio.lowlevel.current_trio_token()
+
+ @classmethod
+ def current_time(cls) -> float:
+ return trio.current_time()
+
+ @classmethod
+ def cancelled_exception_class(cls) -> type[BaseException]:
+ return trio.Cancelled
+
+ @classmethod
+ async def checkpoint(cls) -> None:
+ await trio.lowlevel.checkpoint()
+
+ @classmethod
+ async def checkpoint_if_cancelled(cls) -> None:
+ await trio.lowlevel.checkpoint_if_cancelled()
+
+ @classmethod
+ async def cancel_shielded_checkpoint(cls) -> None:
+ await trio.lowlevel.cancel_shielded_checkpoint()
+
+ @classmethod
+ async def sleep(cls, delay: float) -> None:
+ await trio.sleep(delay)
+
+ @classmethod
+ def create_cancel_scope(
+ cls, *, deadline: float = math.inf, shield: bool = False
+ ) -> abc.CancelScope:
+ return CancelScope(deadline=deadline, shield=shield)
+
+ @classmethod
+ def current_effective_deadline(cls) -> float:
+ return trio.current_effective_deadline()
+
+ @classmethod
+ def create_task_group(cls) -> abc.TaskGroup:
+ return TaskGroup()
+
+ @classmethod
+ def create_event(cls) -> abc.Event:
+ return Event()
+
+ @classmethod
+ def create_lock(cls, *, fast_acquire: bool) -> Lock:
+ return Lock(fast_acquire=fast_acquire)
+
+ @classmethod
+ def create_semaphore(
+ cls,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> abc.Semaphore:
+ return Semaphore(initial_value, max_value=max_value, fast_acquire=fast_acquire)
+
+ @classmethod
+ def create_capacity_limiter(cls, total_tokens: float) -> CapacityLimiter:
+ return CapacityLimiter(total_tokens)
+
+ @classmethod
+ async def run_sync_in_worker_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ abandon_on_cancel: bool = False,
+ limiter: abc.CapacityLimiter | None = None,
+ ) -> T_Retval:
+ def wrapper() -> T_Retval:
+ with claim_worker_thread(TrioBackend, token):
+ return func(*args)
+
+ token = TrioBackend.current_token()
+ return await run_sync(
+ wrapper,
+ abandon_on_cancel=abandon_on_cancel,
+ limiter=cast(trio.CapacityLimiter, limiter),
+ )
+
+ @classmethod
+ def check_cancelled(cls) -> None:
+ trio.from_thread.check_cancelled()
+
+ @classmethod
+ def run_async_from_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ args: tuple[Unpack[PosArgsT]],
+ token: object,
+ ) -> T_co:
+ trio_token = cast("trio.lowlevel.TrioToken | None", token)
+ try:
+ return trio.from_thread.run(func, *args, trio_token=trio_token)
+ except trio.RunFinishedError:
+ raise RunFinishedError from None
+
+ @classmethod
+ def run_sync_from_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ token: object,
+ ) -> T_Retval:
+ trio_token = cast("trio.lowlevel.TrioToken | None", token)
+ try:
+ return trio.from_thread.run_sync(func, *args, trio_token=trio_token)
+ except trio.RunFinishedError:
+ raise RunFinishedError from None
+
+ @classmethod
+ async def open_process(
+ cls,
+ command: StrOrBytesPath | Sequence[StrOrBytesPath],
+ *,
+ stdin: int | IO[Any] | None,
+ stdout: int | IO[Any] | None,
+ stderr: int | IO[Any] | None,
+ **kwargs: Any,
+ ) -> Process:
+ def convert_item(item: StrOrBytesPath) -> str:
+ str_or_bytes = os.fspath(item)
+ if isinstance(str_or_bytes, str):
+ return str_or_bytes
+ else:
+ return os.fsdecode(str_or_bytes)
+
+ if isinstance(command, (str, bytes, PathLike)):
+ process = await trio.lowlevel.open_process(
+ convert_item(command),
+ stdin=stdin,
+ stdout=stdout,
+ stderr=stderr,
+ shell=True,
+ **kwargs,
+ )
+ else:
+ process = await trio.lowlevel.open_process(
+ [convert_item(item) for item in command],
+ stdin=stdin,
+ stdout=stdout,
+ stderr=stderr,
+ shell=False,
+ **kwargs,
+ )
+
+ stdin_stream = SendStreamWrapper(process.stdin) if process.stdin else None
+ stdout_stream = ReceiveStreamWrapper(process.stdout) if process.stdout else None
+ stderr_stream = ReceiveStreamWrapper(process.stderr) if process.stderr else None
+ return Process(process, stdin_stream, stdout_stream, stderr_stream)
+
+ @classmethod
+ def setup_process_pool_exit_at_shutdown(cls, workers: set[abc.Process]) -> None:
+ trio.lowlevel.spawn_system_task(_shutdown_process_pool, workers)
+
+ @classmethod
+ async def connect_tcp(
+ cls, host: str, port: int, local_address: IPSockAddrType | None = None
+ ) -> SocketStream:
+ family = socket.AF_INET6 if ":" in host else socket.AF_INET
+ trio_socket = trio.socket.socket(family)
+ trio_socket.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
+ if local_address:
+ await trio_socket.bind(local_address)
+
+ try:
+ await trio_socket.connect((host, port))
+ except BaseException:
+ trio_socket.close()
+ raise
+
+ return SocketStream(trio_socket)
+
+ @classmethod
+ async def connect_unix(cls, path: str | bytes) -> abc.UNIXSocketStream:
+ trio_socket = trio.socket.socket(socket.AF_UNIX)
+ try:
+ await trio_socket.connect(path)
+ except BaseException:
+ trio_socket.close()
+ raise
+
+ return UNIXSocketStream(trio_socket)
+
+ @classmethod
+ def create_tcp_listener(cls, sock: socket.socket) -> abc.SocketListener:
+ return TCPSocketListener(sock)
+
+ @classmethod
+ def create_unix_listener(cls, sock: socket.socket) -> abc.SocketListener:
+ return UNIXSocketListener(sock)
+
+ @classmethod
+ async def create_udp_socket(
+ cls,
+ family: socket.AddressFamily,
+ local_address: IPSockAddrType | None,
+ remote_address: IPSockAddrType | None,
+ reuse_port: bool,
+ ) -> UDPSocket | ConnectedUDPSocket:
+ trio_socket = trio.socket.socket(family=family, type=socket.SOCK_DGRAM)
+
+ if reuse_port:
+ trio_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1)
+
+ if local_address:
+ await trio_socket.bind(local_address)
+
+ if remote_address:
+ await trio_socket.connect(remote_address)
+ return ConnectedUDPSocket(trio_socket)
+ else:
+ return UDPSocket(trio_socket)
+
+ @classmethod
+ @overload
+ async def create_unix_datagram_socket(
+ cls, raw_socket: socket.socket, remote_path: None
+ ) -> abc.UNIXDatagramSocket: ...
+
+ @classmethod
+ @overload
+ async def create_unix_datagram_socket(
+ cls, raw_socket: socket.socket, remote_path: str | bytes
+ ) -> abc.ConnectedUNIXDatagramSocket: ...
+
+ @classmethod
+ async def create_unix_datagram_socket(
+ cls, raw_socket: socket.socket, remote_path: str | bytes | None
+ ) -> abc.UNIXDatagramSocket | abc.ConnectedUNIXDatagramSocket:
+ trio_socket = trio.socket.from_stdlib_socket(raw_socket)
+
+ if remote_path:
+ await trio_socket.connect(remote_path)
+ return ConnectedUNIXDatagramSocket(trio_socket)
+ else:
+ return UNIXDatagramSocket(trio_socket)
+
+ @classmethod
+ async def getaddrinfo(
+ cls,
+ host: bytes | str | None,
+ port: str | int | None,
+ *,
+ family: int | AddressFamily = 0,
+ type: int | SocketKind = 0,
+ proto: int = 0,
+ flags: int = 0,
+ ) -> Sequence[
+ tuple[
+ AddressFamily,
+ SocketKind,
+ int,
+ str,
+ tuple[str, int] | tuple[str, int, int, int] | tuple[int, bytes],
+ ]
+ ]:
+ return await trio.socket.getaddrinfo(host, port, family, type, proto, flags)
+
+ @classmethod
+ async def getnameinfo(
+ cls, sockaddr: IPSockAddrType, flags: int = 0
+ ) -> tuple[str, str]:
+ return await trio.socket.getnameinfo(sockaddr, flags)
+
+ @classmethod
+ async def wait_readable(cls, obj: FileDescriptorLike) -> None:
+ try:
+ await wait_readable(obj)
+ except trio.ClosedResourceError as exc:
+ raise ClosedResourceError().with_traceback(exc.__traceback__) from None
+ except trio.BusyResourceError:
+ raise BusyResourceError("reading from") from None
+
+ @classmethod
+ async def wait_writable(cls, obj: FileDescriptorLike) -> None:
+ try:
+ await wait_writable(obj)
+ except trio.ClosedResourceError as exc:
+ raise ClosedResourceError().with_traceback(exc.__traceback__) from None
+ except trio.BusyResourceError:
+ raise BusyResourceError("writing to") from None
+
+ @classmethod
+ def notify_closing(cls, obj: FileDescriptorLike) -> None:
+ notify_closing(obj)
+
+ @classmethod
+ async def wrap_listener_socket(cls, sock: socket.socket) -> abc.SocketListener:
+ if hasattr(socket, "AF_UNIX") and sock.family == socket.AF_UNIX:
+ return UNIXSocketListener(sock)
+
+ return TCPSocketListener(sock)
+
+ @classmethod
+ async def wrap_stream_socket(cls, sock: socket.socket) -> SocketStream:
+ trio_sock = trio.socket.from_stdlib_socket(sock)
+ return SocketStream(trio_sock)
+
+ @classmethod
+ async def wrap_unix_stream_socket(cls, sock: socket.socket) -> UNIXSocketStream:
+ trio_sock = trio.socket.from_stdlib_socket(sock)
+ return UNIXSocketStream(trio_sock)
+
+ @classmethod
+ async def wrap_udp_socket(cls, sock: socket.socket) -> UDPSocket:
+ trio_sock = trio.socket.from_stdlib_socket(sock)
+ return UDPSocket(trio_sock)
+
+ @classmethod
+ async def wrap_connected_udp_socket(cls, sock: socket.socket) -> ConnectedUDPSocket:
+ trio_sock = trio.socket.from_stdlib_socket(sock)
+ return ConnectedUDPSocket(trio_sock)
+
+ @classmethod
+ async def wrap_unix_datagram_socket(cls, sock: socket.socket) -> UNIXDatagramSocket:
+ trio_sock = trio.socket.from_stdlib_socket(sock)
+ return UNIXDatagramSocket(trio_sock)
+
+ @classmethod
+ async def wrap_connected_unix_datagram_socket(
+ cls, sock: socket.socket
+ ) -> ConnectedUNIXDatagramSocket:
+ trio_sock = trio.socket.from_stdlib_socket(sock)
+ return ConnectedUNIXDatagramSocket(trio_sock)
+
+ @classmethod
+ def current_default_thread_limiter(cls) -> CapacityLimiter:
+ try:
+ return _capacity_limiter_wrapper.get()
+ except LookupError:
+ limiter = CapacityLimiter(
+ original=trio.to_thread.current_default_thread_limiter()
+ )
+ _capacity_limiter_wrapper.set(limiter)
+ return limiter
+
+ @classmethod
+ def open_signal_receiver(
+ cls, *signals: Signals
+ ) -> AbstractContextManager[AsyncIterator[Signals]]:
+ return _SignalReceiver(signals)
+
+ @classmethod
+ def get_current_task(cls) -> TaskInfo:
+ task = current_task()
+ return TrioTaskInfo(task)
+
+ @classmethod
+ def get_running_tasks(cls) -> Sequence[TaskInfo]:
+ root_task = current_root_task()
+ assert root_task
+ task_infos = [TrioTaskInfo(root_task)]
+ nurseries = root_task.child_nurseries
+ while nurseries:
+ new_nurseries: list[trio.Nursery] = []
+ for nursery in nurseries:
+ for task in nursery.child_tasks:
+ task_infos.append(TrioTaskInfo(task))
+ new_nurseries.extend(task.child_nurseries)
+
+ nurseries = new_nurseries
+
+ return task_infos
+
+ @classmethod
+ async def wait_all_tasks_blocked(cls) -> None:
+ from trio.testing import wait_all_tasks_blocked
+
+ await wait_all_tasks_blocked()
+
+ @classmethod
+ def create_test_runner(cls, options: dict[str, Any]) -> TestRunner:
+ return TestRunner(**options)
+
+
+backend_class = TrioBackend
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/__init__.py b/venv/lib/python3.11/site-packages/anyio/_core/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_asyncio_selector_thread.py b/venv/lib/python3.11/site-packages/anyio/_core/_asyncio_selector_thread.py
new file mode 100644
index 0000000000000000000000000000000000000000..9f35bae568e33e6a9e1219761c83cc8350fa0532
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_asyncio_selector_thread.py
@@ -0,0 +1,167 @@
+from __future__ import annotations
+
+import asyncio
+import socket
+import threading
+from collections.abc import Callable
+from selectors import EVENT_READ, EVENT_WRITE, DefaultSelector
+from typing import TYPE_CHECKING, Any
+
+if TYPE_CHECKING:
+ from _typeshed import FileDescriptorLike
+
+_selector_lock = threading.Lock()
+_selector: Selector | None = None
+
+
+class Selector:
+ def __init__(self) -> None:
+ self._thread = threading.Thread(target=self.run, name="AnyIO socket selector")
+ self._selector = DefaultSelector()
+ self._send, self._receive = socket.socketpair()
+ self._send.setblocking(False)
+ self._receive.setblocking(False)
+ # This somewhat reduces the amount of memory wasted queueing up data
+ # for wakeups. With these settings, maximum number of 1-byte sends
+ # before getting BlockingIOError:
+ # Linux 4.8: 6
+ # macOS (darwin 15.5): 1
+ # Windows 10: 525347
+ # Windows you're weird. (And on Windows setting SNDBUF to 0 makes send
+ # blocking, even on non-blocking sockets, so don't do that.)
+ self._receive.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, 1)
+ self._send.setsockopt(socket.SOL_SOCKET, socket.SO_SNDBUF, 1)
+ # On Windows this is a TCP socket so this might matter. On other
+ # platforms this fails b/c AF_UNIX sockets aren't actually TCP.
+ try:
+ self._send.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
+ except OSError:
+ pass
+
+ self._selector.register(self._receive, EVENT_READ)
+ self._closed = False
+
+ def start(self) -> None:
+ self._thread.start()
+ threading._register_atexit(self._stop) # type: ignore[attr-defined]
+
+ def _stop(self) -> None:
+ global _selector
+ self._closed = True
+ self._notify_self()
+ self._send.close()
+ self._thread.join()
+ self._selector.unregister(self._receive)
+ self._receive.close()
+ self._selector.close()
+ _selector = None
+ assert not self._selector.get_map(), (
+ "selector still has registered file descriptors after shutdown"
+ )
+
+ def _notify_self(self) -> None:
+ try:
+ self._send.send(b"\x00")
+ except BlockingIOError:
+ pass
+
+ def add_reader(self, fd: FileDescriptorLike, callback: Callable[[], Any]) -> None:
+ loop = asyncio.get_running_loop()
+ try:
+ key = self._selector.get_key(fd)
+ except KeyError:
+ self._selector.register(fd, EVENT_READ, {EVENT_READ: (loop, callback)})
+ else:
+ if EVENT_READ in key.data:
+ raise ValueError(
+ "this file descriptor is already registered for reading"
+ )
+
+ key.data[EVENT_READ] = loop, callback
+ self._selector.modify(fd, key.events | EVENT_READ, key.data)
+
+ self._notify_self()
+
+ def add_writer(self, fd: FileDescriptorLike, callback: Callable[[], Any]) -> None:
+ loop = asyncio.get_running_loop()
+ try:
+ key = self._selector.get_key(fd)
+ except KeyError:
+ self._selector.register(fd, EVENT_WRITE, {EVENT_WRITE: (loop, callback)})
+ else:
+ if EVENT_WRITE in key.data:
+ raise ValueError(
+ "this file descriptor is already registered for writing"
+ )
+
+ key.data[EVENT_WRITE] = loop, callback
+ self._selector.modify(fd, key.events | EVENT_WRITE, key.data)
+
+ self._notify_self()
+
+ def remove_reader(self, fd: FileDescriptorLike) -> bool:
+ try:
+ key = self._selector.get_key(fd)
+ except KeyError:
+ return False
+
+ if new_events := key.events ^ EVENT_READ:
+ del key.data[EVENT_READ]
+ self._selector.modify(fd, new_events, key.data)
+ else:
+ self._selector.unregister(fd)
+
+ return True
+
+ def remove_writer(self, fd: FileDescriptorLike) -> bool:
+ try:
+ key = self._selector.get_key(fd)
+ except KeyError:
+ return False
+
+ if new_events := key.events ^ EVENT_WRITE:
+ del key.data[EVENT_WRITE]
+ self._selector.modify(fd, new_events, key.data)
+ else:
+ self._selector.unregister(fd)
+
+ return True
+
+ def run(self) -> None:
+ while not self._closed:
+ for key, events in self._selector.select():
+ if key.fileobj is self._receive:
+ try:
+ while self._receive.recv(4096):
+ pass
+ except BlockingIOError:
+ pass
+
+ continue
+
+ if events & EVENT_READ:
+ loop, callback = key.data[EVENT_READ]
+ self.remove_reader(key.fd)
+ try:
+ loop.call_soon_threadsafe(callback)
+ except RuntimeError:
+ pass # the loop was already closed
+
+ if events & EVENT_WRITE:
+ loop, callback = key.data[EVENT_WRITE]
+ self.remove_writer(key.fd)
+ try:
+ loop.call_soon_threadsafe(callback)
+ except RuntimeError:
+ pass # the loop was already closed
+
+
+def get_selector() -> Selector:
+ global _selector
+
+ with _selector_lock:
+ if _selector is None:
+ _selector = Selector()
+ _selector.start()
+
+ return _selector
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_contextmanagers.py b/venv/lib/python3.11/site-packages/anyio/_core/_contextmanagers.py
new file mode 100644
index 0000000000000000000000000000000000000000..302f32b0c78a7071605b195c55054cfdb0b55f37
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_contextmanagers.py
@@ -0,0 +1,200 @@
+from __future__ import annotations
+
+from abc import abstractmethod
+from contextlib import AbstractAsyncContextManager, AbstractContextManager
+from inspect import isasyncgen, iscoroutine, isgenerator
+from types import TracebackType
+from typing import Protocol, TypeVar, cast, final
+
+_T_co = TypeVar("_T_co", covariant=True)
+_ExitT_co = TypeVar("_ExitT_co", covariant=True, bound="bool | None")
+
+
+class _SupportsCtxMgr(Protocol[_T_co, _ExitT_co]):
+ def __contextmanager__(self) -> AbstractContextManager[_T_co, _ExitT_co]: ...
+
+
+class _SupportsAsyncCtxMgr(Protocol[_T_co, _ExitT_co]):
+ def __asynccontextmanager__(
+ self,
+ ) -> AbstractAsyncContextManager[_T_co, _ExitT_co]: ...
+
+
+class ContextManagerMixin:
+ """
+ Mixin class providing context manager functionality via a generator-based
+ implementation.
+
+ This class allows you to implement a context manager via :meth:`__contextmanager__`
+ which should return a generator. The mechanics are meant to mirror those of
+ :func:`@contextmanager `.
+
+ .. note:: Classes using this mix-in are not reentrant as context managers, meaning
+ that once you enter it, you can't re-enter before first exiting it.
+
+ .. seealso:: :doc:`contextmanagers`
+ """
+
+ __cm: AbstractContextManager[object, bool | None] | None = None
+
+ @final
+ def __enter__(self: _SupportsCtxMgr[_T_co, bool | None]) -> _T_co:
+ # Needed for mypy to assume self still has the __cm member
+ assert isinstance(self, ContextManagerMixin)
+ if self.__cm is not None:
+ raise RuntimeError(
+ f"this {self.__class__.__qualname__} has already been entered"
+ )
+
+ cm = self.__contextmanager__()
+ if not isinstance(cm, AbstractContextManager):
+ if isgenerator(cm):
+ raise TypeError(
+ "__contextmanager__() returned a generator object instead of "
+ "a context manager. Did you forget to add the @contextmanager "
+ "decorator?"
+ )
+
+ raise TypeError(
+ f"__contextmanager__() did not return a context manager object, "
+ f"but {cm.__class__!r}"
+ )
+
+ if cm is self:
+ raise TypeError(
+ f"{self.__class__.__qualname__}.__contextmanager__() returned "
+ f"self. Did you forget to add the @contextmanager decorator and a "
+ f"'yield' statement?"
+ )
+
+ value = cm.__enter__()
+ self.__cm = cm
+ return value
+
+ @final
+ def __exit__(
+ self: _SupportsCtxMgr[object, _ExitT_co],
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> _ExitT_co:
+ # Needed for mypy to assume self still has the __cm member
+ assert isinstance(self, ContextManagerMixin)
+ if self.__cm is None:
+ raise RuntimeError(
+ f"this {self.__class__.__qualname__} has not been entered yet"
+ )
+
+ # Prevent circular references
+ cm = self.__cm
+ del self.__cm
+
+ return cast(_ExitT_co, cm.__exit__(exc_type, exc_val, exc_tb))
+
+ @abstractmethod
+ def __contextmanager__(self) -> AbstractContextManager[object, bool | None]:
+ """
+ Implement your context manager logic here.
+
+ This method **must** be decorated with
+ :func:`@contextmanager `.
+
+ .. note:: Remember that the ``yield`` will raise any exception raised in the
+ enclosed context block, so use a ``finally:`` block to clean up resources!
+
+ :return: a context manager object
+ """
+
+
+class AsyncContextManagerMixin:
+ """
+ Mixin class providing async context manager functionality via a generator-based
+ implementation.
+
+ This class allows you to implement a context manager via
+ :meth:`__asynccontextmanager__`. The mechanics are meant to mirror those of
+ :func:`@asynccontextmanager `.
+
+ .. note:: Classes using this mix-in are not reentrant as context managers, meaning
+ that once you enter it, you can't re-enter before first exiting it.
+
+ .. seealso:: :doc:`contextmanagers`
+ """
+
+ __cm: AbstractAsyncContextManager[object, bool | None] | None = None
+
+ @final
+ async def __aenter__(self: _SupportsAsyncCtxMgr[_T_co, bool | None]) -> _T_co:
+ # Needed for mypy to assume self still has the __cm member
+ assert isinstance(self, AsyncContextManagerMixin)
+ if self.__cm is not None:
+ raise RuntimeError(
+ f"this {self.__class__.__qualname__} has already been entered"
+ )
+
+ cm = self.__asynccontextmanager__()
+ if not isinstance(cm, AbstractAsyncContextManager):
+ if isasyncgen(cm):
+ raise TypeError(
+ "__asynccontextmanager__() returned an async generator instead of "
+ "an async context manager. Did you forget to add the "
+ "@asynccontextmanager decorator?"
+ )
+ elif iscoroutine(cm):
+ cm.close()
+ raise TypeError(
+ "__asynccontextmanager__() returned a coroutine object instead of "
+ "an async context manager. Did you forget to add the "
+ "@asynccontextmanager decorator and a 'yield' statement?"
+ )
+
+ raise TypeError(
+ f"__asynccontextmanager__() did not return an async context manager, "
+ f"but {cm.__class__!r}"
+ )
+
+ if cm is self:
+ raise TypeError(
+ f"{self.__class__.__qualname__}.__asynccontextmanager__() returned "
+ f"self. Did you forget to add the @asynccontextmanager decorator and a "
+ f"'yield' statement?"
+ )
+
+ value = await cm.__aenter__()
+ self.__cm = cm
+ return value
+
+ @final
+ async def __aexit__(
+ self: _SupportsAsyncCtxMgr[object, _ExitT_co],
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> _ExitT_co:
+ assert isinstance(self, AsyncContextManagerMixin)
+ if self.__cm is None:
+ raise RuntimeError(
+ f"this {self.__class__.__qualname__} has not been entered yet"
+ )
+
+ # Prevent circular references
+ cm = self.__cm
+ del self.__cm
+
+ return cast(_ExitT_co, await cm.__aexit__(exc_type, exc_val, exc_tb))
+
+ @abstractmethod
+ def __asynccontextmanager__(
+ self,
+ ) -> AbstractAsyncContextManager[object, bool | None]:
+ """
+ Implement your async context manager logic here.
+
+ This method **must** be decorated with
+ :func:`@asynccontextmanager `.
+
+ .. note:: Remember that the ``yield`` will raise any exception raised in the
+ enclosed context block, so use a ``finally:`` block to clean up resources!
+
+ :return: an async context manager object
+ """
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_eventloop.py b/venv/lib/python3.11/site-packages/anyio/_core/_eventloop.py
new file mode 100644
index 0000000000000000000000000000000000000000..a3e2ab1cc172f2fbd94f244ab110cb16aa7eb8eb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_eventloop.py
@@ -0,0 +1,240 @@
+from __future__ import annotations
+
+import math
+import sys
+import threading
+from collections.abc import Awaitable, Callable, Generator
+from contextlib import contextmanager
+from contextvars import Token
+from importlib import import_module
+from typing import TYPE_CHECKING, Any, TypeVar
+
+from ._exceptions import NoEventLoopError
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+sniffio: Any
+try:
+ import sniffio
+except ModuleNotFoundError:
+ sniffio = None
+
+if TYPE_CHECKING:
+ from ..abc import AsyncBackend
+
+# This must be updated when new backends are introduced
+BACKENDS = "asyncio", "trio"
+
+T_Retval = TypeVar("T_Retval")
+PosArgsT = TypeVarTuple("PosArgsT")
+
+threadlocals = threading.local()
+loaded_backends: dict[str, type[AsyncBackend]] = {}
+
+
+def run(
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
+ *args: Unpack[PosArgsT],
+ backend: str = "asyncio",
+ backend_options: dict[str, Any] | None = None,
+) -> T_Retval:
+ """
+ Run the given coroutine function in an asynchronous event loop.
+
+ The current thread must not be already running an event loop.
+
+ :param func: a coroutine function
+ :param args: positional arguments to ``func``
+ :param backend: name of the asynchronous event loop implementation – currently
+ either ``asyncio`` or ``trio``
+ :param backend_options: keyword arguments to call the backend ``run()``
+ implementation with (documented :ref:`here `)
+ :return: the return value of the coroutine function
+ :raises RuntimeError: if an asynchronous event loop is already running in this
+ thread
+ :raises LookupError: if the named backend is not found
+
+ """
+ if asynclib_name := current_async_library():
+ raise RuntimeError(f"Already running {asynclib_name} in this thread")
+
+ try:
+ async_backend = get_async_backend(backend)
+ except ImportError as exc:
+ if backend in BACKENDS:
+ raise LookupError(
+ f"Backend {backend!r} is not available. "
+ f"Install it with: pip install anyio[{backend}]"
+ ) from exc
+
+ raise LookupError(f"No such backend: {backend}") from exc
+
+ token = None
+ if asynclib_name is None:
+ # Since we're in control of the event loop, we can cache the name of the async
+ # library
+ token = set_current_async_library(backend)
+
+ try:
+ backend_options = backend_options or {}
+ return async_backend.run(func, args, {}, backend_options)
+ finally:
+ reset_current_async_library(token)
+
+
+async def sleep(delay: float) -> None:
+ """
+ Pause the current task for the specified duration.
+
+ :param delay: the duration, in seconds
+
+ """
+ return await get_async_backend().sleep(delay)
+
+
+async def sleep_forever() -> None:
+ """
+ Pause the current task until it's cancelled.
+
+ This is a shortcut for ``sleep(math.inf)``.
+
+ .. versionadded:: 3.1
+
+ """
+ await sleep(math.inf)
+
+
+async def sleep_until(deadline: float) -> None:
+ """
+ Pause the current task until the given time.
+
+ :param deadline: the absolute time to wake up at (according to the internal
+ monotonic clock of the event loop)
+
+ .. versionadded:: 3.1
+
+ """
+ now = current_time()
+ await sleep(max(deadline - now, 0))
+
+
+def current_time() -> float:
+ """
+ Return the current value of the event loop's internal clock.
+
+ :return: the clock value (seconds)
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().current_time()
+
+
+def get_all_backends() -> tuple[str, ...]:
+ """Return a tuple of the names of all built-in backends."""
+ return BACKENDS
+
+
+def get_available_backends() -> tuple[str, ...]:
+ """
+ Test for the availability of built-in backends.
+
+ :return a tuple of the built-in backend names that were successfully imported
+
+ .. versionadded:: 4.12
+
+ """
+ available_backends: list[str] = []
+ for backend_name in get_all_backends():
+ try:
+ get_async_backend(backend_name)
+ except ImportError:
+ continue
+
+ available_backends.append(backend_name)
+
+ return tuple(available_backends)
+
+
+def get_cancelled_exc_class() -> type[BaseException]:
+ """
+ Return the current async library's cancellation exception class.
+
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().cancelled_exception_class()
+
+
+#
+# Private API
+#
+
+
+@contextmanager
+def claim_worker_thread(
+ backend_class: type[AsyncBackend], token: object
+) -> Generator[Any, None, None]:
+ from ..lowlevel import EventLoopToken
+
+ threadlocals.current_token = EventLoopToken(backend_class, token)
+ try:
+ yield
+ finally:
+ del threadlocals.current_token
+
+
+def get_async_backend(asynclib_name: str | None = None) -> type[AsyncBackend]:
+ if asynclib_name is None:
+ asynclib_name = current_async_library()
+ if not asynclib_name:
+ raise NoEventLoopError(
+ f"Not currently running on any asynchronous event loop. "
+ f"Available async backends: {', '.join(get_all_backends())}"
+ )
+
+ # We use our own dict instead of sys.modules to get the already imported back-end
+ # class because the appropriate modules in sys.modules could potentially be only
+ # partially initialized
+ try:
+ return loaded_backends[asynclib_name]
+ except KeyError:
+ module = import_module(f"anyio._backends._{asynclib_name}")
+ loaded_backends[asynclib_name] = module.backend_class
+ return module.backend_class
+
+
+def current_async_library() -> str | None:
+ if sniffio is None:
+ # If sniffio is not installed, we assume we're either running asyncio or nothing
+ import asyncio
+
+ try:
+ asyncio.get_running_loop()
+ return "asyncio"
+ except RuntimeError:
+ pass
+ else:
+ try:
+ return sniffio.current_async_library()
+ except sniffio.AsyncLibraryNotFoundError:
+ pass
+
+ return None
+
+
+def set_current_async_library(asynclib_name: str | None) -> Token | None:
+ # no-op if sniffio is not installed
+ if sniffio is None:
+ return None
+
+ return sniffio.current_async_library_cvar.set(asynclib_name)
+
+
+def reset_current_async_library(token: Token | None) -> None:
+ if token is not None:
+ sniffio.current_async_library_cvar.reset(token)
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_exceptions.py b/venv/lib/python3.11/site-packages/anyio/_core/_exceptions.py
new file mode 100644
index 0000000000000000000000000000000000000000..cd6eb9b5ca82c70014fc610670a34cd076b872fe
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_exceptions.py
@@ -0,0 +1,177 @@
+from __future__ import annotations
+
+import sys
+from collections.abc import Generator
+from textwrap import dedent
+from typing import Any
+
+if sys.version_info < (3, 11):
+ from exceptiongroup import BaseExceptionGroup
+
+
+class BrokenResourceError(Exception):
+ """
+ Raised when trying to use a resource that has been rendered unusable due to external
+ causes (e.g. a send stream whose peer has disconnected).
+ """
+
+
+class BrokenWorkerProcess(Exception):
+ """
+ Raised by :meth:`~anyio.to_process.run_sync` if the worker process terminates abruptly or
+ otherwise misbehaves.
+ """
+
+
+class BrokenWorkerInterpreter(Exception):
+ """
+ Raised by :meth:`~anyio.to_interpreter.run_sync` if an unexpected exception is
+ raised in the subinterpreter.
+ """
+
+ def __init__(self, excinfo: Any):
+ # This was adapted from concurrent.futures.interpreter.ExecutionFailed
+ msg = excinfo.formatted
+ if not msg:
+ if excinfo.type and excinfo.msg:
+ msg = f"{excinfo.type.__name__}: {excinfo.msg}"
+ else:
+ msg = excinfo.type.__name__ or excinfo.msg
+
+ super().__init__(msg)
+ self.excinfo = excinfo
+
+ def __str__(self) -> str:
+ try:
+ formatted = self.excinfo.errdisplay
+ except Exception:
+ return super().__str__()
+ else:
+ return dedent(
+ f"""
+ {super().__str__()}
+
+ Uncaught in the interpreter:
+
+ {formatted}
+ """.strip()
+ )
+
+
+class BusyResourceError(Exception):
+ """
+ Raised when two tasks are trying to read from or write to the same resource
+ concurrently.
+ """
+
+ def __init__(self, action: str):
+ super().__init__(f"Another task is already {action} this resource")
+
+
+class ClosedResourceError(Exception):
+ """Raised when trying to use a resource that has been closed."""
+
+
+class ConnectionFailed(OSError):
+ """
+ Raised when a connection attempt fails.
+
+ .. note:: This class inherits from :exc:`OSError` for backwards compatibility.
+ """
+
+
+def iterate_exceptions(
+ exception: BaseException,
+) -> Generator[BaseException, None, None]:
+ if isinstance(exception, BaseExceptionGroup):
+ for exc in exception.exceptions:
+ yield from iterate_exceptions(exc)
+ else:
+ yield exception
+
+
+class DelimiterNotFound(Exception):
+ """
+ Raised during
+ :meth:`~anyio.streams.buffered.BufferedByteReceiveStream.receive_until` if the
+ maximum number of bytes has been read without the delimiter being found.
+ """
+
+ def __init__(self, max_bytes: int) -> None:
+ super().__init__(
+ f"The delimiter was not found among the first {max_bytes} bytes"
+ )
+
+
+class EndOfStream(Exception):
+ """
+ Raised when trying to read from a stream that has been closed from the other end.
+ """
+
+
+class IncompleteRead(Exception):
+ """
+ Raised during
+ :meth:`~anyio.streams.buffered.BufferedByteReceiveStream.receive_exactly` or
+ :meth:`~anyio.streams.buffered.BufferedByteReceiveStream.receive_until` if the
+ connection is closed before the requested amount of bytes has been read.
+ """
+
+ def __init__(self) -> None:
+ super().__init__(
+ "The stream was closed before the read operation could be completed"
+ )
+
+
+class TypedAttributeLookupError(LookupError):
+ """
+ Raised by :meth:`~anyio.TypedAttributeProvider.extra` when the given typed attribute
+ is not found and no default value has been given.
+ """
+
+
+class WouldBlock(Exception):
+ """Raised by ``X_nowait`` functions if ``X()`` would block."""
+
+
+class NoEventLoopError(RuntimeError):
+ """
+ Raised by several functions that require an event loop to be running in the current
+ thread when there is no running event loop.
+
+ This is also raised by :func:`.from_thread.run` and :func:`.from_thread.run_sync`
+ if not calling from an AnyIO worker thread, and no ``token`` was passed.
+ """
+
+
+class RunFinishedError(RuntimeError):
+ """
+ Raised by :func:`.from_thread.run` and :func:`.from_thread.run_sync` if the event
+ loop associated with the explicitly passed token has already finished.
+ """
+
+ def __init__(self) -> None:
+ super().__init__(
+ "The event loop associated with the given token has already finished"
+ )
+
+
+class TaskFailed(Exception):
+ """
+ Raised when awaiting on, or attempting to access the return value of, a
+ :class:`.TaskHandle` that raised an exception.
+ """
+
+
+class TaskCancelled(TaskFailed):
+ """
+ Raised when awaiting on, or attempting to access the return value of, a
+ :class:`.TaskHandle` that was cancelled.
+ """
+
+
+class TaskNotFinished(Exception):
+ """
+ Raised when attempting to access the return value or exception of a
+ :class:`.TaskHandle` that is still pending completion.
+ """
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_fileio.py b/venv/lib/python3.11/site-packages/anyio/_core/_fileio.py
new file mode 100644
index 0000000000000000000000000000000000000000..692c754b3fd5b8a2e795d7c981b2785423161216
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_fileio.py
@@ -0,0 +1,960 @@
+from __future__ import annotations
+
+import os
+import pathlib
+import sys
+from collections.abc import (
+ AsyncIterator,
+ Callable,
+ Iterable,
+ Iterator,
+ Sequence,
+)
+from dataclasses import dataclass
+from functools import partial
+from os import PathLike
+from typing import (
+ IO,
+ TYPE_CHECKING,
+ Any,
+ AnyStr,
+ ClassVar,
+ Final,
+ Generic,
+ TypeVar,
+ overload,
+)
+
+from .. import to_thread
+from ..abc import AsyncResource
+from ._synchronization import CapacityLimiter
+
+if sys.version_info >= (3, 11):
+ from typing import Self
+else:
+ from typing_extensions import Self
+
+if sys.version_info >= (3, 14):
+ from pathlib.types import PathInfo
+
+if TYPE_CHECKING:
+ from types import ModuleType
+
+ from _typeshed import OpenBinaryMode, OpenTextMode, ReadableBuffer, WriteableBuffer
+else:
+ ReadableBuffer = OpenBinaryMode = OpenTextMode = WriteableBuffer = object
+
+
+T = TypeVar("T", bound="Path")
+
+
+class AsyncFile(AsyncResource, Generic[AnyStr]):
+ """
+ An asynchronous file object.
+
+ This class wraps a standard file object and provides async friendly versions of the
+ following blocking methods (where available on the original file object):
+
+ * read
+ * read1
+ * readline
+ * readlines
+ * readinto
+ * readinto1
+ * write
+ * writelines
+ * truncate
+ * seek
+ * tell
+ * flush
+
+ All other methods are directly passed through.
+
+ This class supports the asynchronous context manager protocol which closes the
+ underlying file at the end of the context block.
+
+ This class also supports asynchronous iteration::
+
+ async with await open_file(...) as f:
+ async for line in f:
+ print(line)
+ """
+
+ def __init__(
+ self, fp: IO[AnyStr], *, limiter: CapacityLimiter | None = None
+ ) -> None:
+ if limiter is not None and not isinstance(limiter, CapacityLimiter):
+ raise TypeError(
+ f"limiter must be a CapacityLimiter or None, not "
+ f"{limiter.__class__.__name__}"
+ )
+
+ self._fp: Any = fp
+ self._limiter = limiter
+
+ def __getattr__(self, name: str) -> object:
+ return getattr(self._fp, name)
+
+ @property
+ def limiter(self) -> CapacityLimiter | None:
+ """The capacity limiter used by this file object, if not the global limiter."""
+ return self._limiter
+
+ @property
+ def wrapped(self) -> IO[AnyStr]:
+ """The wrapped file object."""
+ return self._fp
+
+ async def __aiter__(self) -> AsyncIterator[AnyStr]:
+ while True:
+ line = await self.readline()
+ if line:
+ yield line
+ else:
+ break
+
+ async def aclose(self) -> None:
+ return await to_thread.run_sync(self._fp.close, limiter=self._limiter)
+
+ async def read(self, size: int = -1) -> AnyStr:
+ return await to_thread.run_sync(self._fp.read, size, limiter=self._limiter)
+
+ async def read1(self: AsyncFile[bytes], size: int = -1) -> bytes:
+ return await to_thread.run_sync(self._fp.read1, size, limiter=self._limiter)
+
+ async def readline(self) -> AnyStr:
+ return await to_thread.run_sync(self._fp.readline, limiter=self._limiter)
+
+ async def readlines(self) -> list[AnyStr]:
+ return await to_thread.run_sync(self._fp.readlines, limiter=self._limiter)
+
+ async def readinto(self: AsyncFile[bytes], b: WriteableBuffer) -> int:
+ return await to_thread.run_sync(self._fp.readinto, b, limiter=self._limiter)
+
+ async def readinto1(self: AsyncFile[bytes], b: WriteableBuffer) -> int:
+ return await to_thread.run_sync(self._fp.readinto1, b, limiter=self._limiter)
+
+ @overload
+ async def write(self: AsyncFile[bytes], b: ReadableBuffer) -> int: ...
+
+ @overload
+ async def write(self: AsyncFile[str], b: str) -> int: ...
+
+ async def write(self, b: ReadableBuffer | str) -> int:
+ return await to_thread.run_sync(self._fp.write, b, limiter=self._limiter)
+
+ @overload
+ async def writelines(
+ self: AsyncFile[bytes], lines: Iterable[ReadableBuffer]
+ ) -> None: ...
+
+ @overload
+ async def writelines(self: AsyncFile[str], lines: Iterable[str]) -> None: ...
+
+ async def writelines(self, lines: Iterable[ReadableBuffer] | Iterable[str]) -> None:
+ return await to_thread.run_sync(
+ self._fp.writelines, lines, limiter=self._limiter
+ )
+
+ async def truncate(self, size: int | None = None) -> int:
+ return await to_thread.run_sync(self._fp.truncate, size, limiter=self._limiter)
+
+ async def seek(self, offset: int, whence: int | None = os.SEEK_SET) -> int:
+ return await to_thread.run_sync(
+ self._fp.seek, offset, whence, limiter=self._limiter
+ )
+
+ async def tell(self) -> int:
+ return await to_thread.run_sync(self._fp.tell, limiter=self._limiter)
+
+ async def flush(self) -> None:
+ return await to_thread.run_sync(self._fp.flush, limiter=self._limiter)
+
+
+@overload
+async def open_file(
+ file: str | PathLike[str] | int,
+ mode: OpenBinaryMode,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ errors: str | None = ...,
+ newline: str | None = ...,
+ closefd: bool = ...,
+ opener: Callable[[str, int], int] | None = ...,
+ *,
+ limiter: CapacityLimiter | None = ...,
+) -> AsyncFile[bytes]: ...
+
+
+@overload
+async def open_file(
+ file: str | PathLike[str] | int,
+ mode: OpenTextMode = ...,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ errors: str | None = ...,
+ newline: str | None = ...,
+ closefd: bool = ...,
+ opener: Callable[[str, int], int] | None = ...,
+ *,
+ limiter: CapacityLimiter | None = ...,
+) -> AsyncFile[str]: ...
+
+
+async def open_file(
+ file: str | PathLike[str] | int,
+ mode: str = "r",
+ buffering: int = -1,
+ encoding: str | None = None,
+ errors: str | None = None,
+ newline: str | None = None,
+ closefd: bool = True,
+ opener: Callable[[str, int], int] | None = None,
+ *,
+ limiter: CapacityLimiter | None = None,
+) -> AsyncFile[Any]:
+ """
+ Open a file asynchronously.
+
+ Except for ``limiter``, the arguments are exactly the same as for the builtin :func:`open`.
+
+ :param limiter: an optional capacity limiter to use with the file
+ instead of the default one
+ :return: an asynchronous file object
+
+ .. versionchanged:: 4.14.0
+ Added the ``limiter`` keyword argument.
+
+ """
+ fp = await to_thread.run_sync(
+ open,
+ file,
+ mode,
+ buffering,
+ encoding,
+ errors,
+ newline,
+ closefd,
+ opener,
+ limiter=limiter,
+ )
+ return AsyncFile(fp, limiter=limiter)
+
+
+def wrap_file(
+ file: IO[AnyStr], *, limiter: CapacityLimiter | None = None
+) -> AsyncFile[AnyStr]:
+ """
+ Wrap an existing file as an asynchronous file.
+
+ :param file: an existing file-like object
+ :param limiter: an optional capacity limiter to use with the file
+ instead of the default one
+ :return: an asynchronous file object
+
+ .. versionchanged:: 4.14.0
+ Added the ``limiter`` keyword argument.
+
+ """
+ return AsyncFile(file, limiter=limiter)
+
+
+@dataclass(eq=False)
+class _PathIterator(AsyncIterator[T]):
+ iterator: Iterator[PathLike[str]]
+ limiter: CapacityLimiter | None
+ # This was added to ensure that iterating over a subclass of Path yields instances
+ # of that subclass rather than the base Path class.
+ path_cls: type[T]
+
+ async def __anext__(self) -> T:
+ nextval = await to_thread.run_sync(
+ next, self.iterator, None, abandon_on_cancel=True, limiter=self.limiter
+ )
+ if nextval is None:
+ raise StopAsyncIteration from None
+
+ return self.path_cls(nextval, limiter=self.limiter)
+
+
+class Path:
+ """
+ An asynchronous version of :class:`pathlib.Path`.
+
+ This class cannot be substituted for :class:`pathlib.Path` or
+ :class:`pathlib.PurePath`, but it is compatible with the :class:`os.PathLike`
+ interface.
+
+ It implements the Python 3.10 version of :class:`pathlib.Path` interface, except for
+ the deprecated :meth:`~pathlib.Path.link_to` method.
+
+ Some methods may be unavailable or have limited functionality, based on the Python
+ version:
+
+ * :meth:`~pathlib.Path.copy` (available on Python 3.14 or later)
+ * :meth:`~pathlib.Path.copy_into` (available on Python 3.14 or later)
+ * :meth:`~pathlib.Path.from_uri` (available on Python 3.13 or later)
+ * :meth:`~pathlib.PurePath.full_match` (available on Python 3.13 or later)
+ * :attr:`~pathlib.Path.info` (available on Python 3.14 or later)
+ * :meth:`~pathlib.Path.is_junction` (available on Python 3.12 or later)
+ * :meth:`~pathlib.PurePath.match` (the ``case_sensitive`` parameter is only
+ available on Python 3.13 or later)
+ * :meth:`~pathlib.Path.move` (available on Python 3.14 or later)
+ * :meth:`~pathlib.Path.move_into` (available on Python 3.14 or later)
+ * :meth:`~pathlib.PurePath.relative_to` (the ``walk_up`` parameter is only available
+ on Python 3.12 or later)
+ * :meth:`~pathlib.Path.walk` (available on Python 3.12 or later)
+
+ Any methods that do disk I/O need to be awaited on. These methods are:
+
+ * :meth:`~pathlib.Path.absolute`
+ * :meth:`~pathlib.Path.chmod`
+ * :meth:`~pathlib.Path.cwd`
+ * :meth:`~pathlib.Path.exists`
+ * :meth:`~pathlib.Path.expanduser`
+ * :meth:`~pathlib.Path.group`
+ * :meth:`~pathlib.Path.hardlink_to`
+ * :meth:`~pathlib.Path.home`
+ * :meth:`~pathlib.Path.is_block_device`
+ * :meth:`~pathlib.Path.is_char_device`
+ * :meth:`~pathlib.Path.is_dir`
+ * :meth:`~pathlib.Path.is_fifo`
+ * :meth:`~pathlib.Path.is_file`
+ * :meth:`~pathlib.Path.is_junction`
+ * :meth:`~pathlib.Path.is_mount`
+ * :meth:`~pathlib.Path.is_socket`
+ * :meth:`~pathlib.Path.is_symlink`
+ * :meth:`~pathlib.Path.lchmod`
+ * :meth:`~pathlib.Path.lstat`
+ * :meth:`~pathlib.Path.mkdir`
+ * :meth:`~pathlib.Path.open`
+ * :meth:`~pathlib.Path.owner`
+ * :meth:`~pathlib.Path.read_bytes`
+ * :meth:`~pathlib.Path.read_text`
+ * :meth:`~pathlib.Path.readlink`
+ * :meth:`~pathlib.Path.rename`
+ * :meth:`~pathlib.Path.replace`
+ * :meth:`~pathlib.Path.resolve`
+ * :meth:`~pathlib.Path.rmdir`
+ * :meth:`~pathlib.Path.samefile`
+ * :meth:`~pathlib.Path.stat`
+ * :meth:`~pathlib.Path.symlink_to`
+ * :meth:`~pathlib.Path.touch`
+ * :meth:`~pathlib.Path.unlink`
+ * :meth:`~pathlib.Path.walk`
+ * :meth:`~pathlib.Path.write_bytes`
+ * :meth:`~pathlib.Path.write_text`
+
+ Additionally, the following methods return an async iterator yielding
+ :class:`~.Path` objects:
+
+ * :meth:`~pathlib.Path.glob`
+ * :meth:`~pathlib.Path.iterdir`
+ * :meth:`~pathlib.Path.rglob`
+
+ .. versionchanged:: 4.14.0
+ Added the ``limiter`` keyword argument.
+ """
+
+ __slots__ = "_path", "_limiter", "__weakref__"
+
+ __weakref__: Any
+
+ def __init__(
+ self, *args: str | PathLike[str], limiter: CapacityLimiter | None = None
+ ) -> None:
+ if limiter is not None and not isinstance(limiter, CapacityLimiter):
+ raise TypeError(
+ f"limiter must be a CapacityLimiter or None, not "
+ f"{limiter.__class__.__name__}"
+ )
+
+ self._path: Final[pathlib.Path] = pathlib.Path(*args)
+ self._limiter = limiter
+
+ def __fspath__(self) -> str:
+ return self._path.__fspath__()
+
+ if sys.version_info >= (3, 15):
+
+ def __vfspath__(self) -> str:
+ return self._path.__vfspath__()
+
+ def __str__(self) -> str:
+ return self._path.__str__()
+
+ def __repr__(self) -> str:
+ return f"{self.__class__.__name__}({self.as_posix()!r})"
+
+ def __bytes__(self) -> bytes:
+ return self._path.__bytes__()
+
+ def __hash__(self) -> int:
+ return self._path.__hash__()
+
+ def __eq__(self, other: object) -> bool:
+ target = other._path if isinstance(other, Path) else other
+ return self._path.__eq__(target)
+
+ def __lt__(self, other: pathlib.PurePath | Path) -> bool:
+ target = other._path if isinstance(other, Path) else other
+ return self._path.__lt__(target)
+
+ def __le__(self, other: pathlib.PurePath | Path) -> bool:
+ target = other._path if isinstance(other, Path) else other
+ return self._path.__le__(target)
+
+ def __gt__(self, other: pathlib.PurePath | Path) -> bool:
+ target = other._path if isinstance(other, Path) else other
+ return self._path.__gt__(target)
+
+ def __ge__(self, other: pathlib.PurePath | Path) -> bool:
+ target = other._path if isinstance(other, Path) else other
+ return self._path.__ge__(target)
+
+ def __truediv__(self, other: str | PathLike[str]) -> Self:
+ return type(self)(self._path / other, limiter=self._limiter)
+
+ def __rtruediv__(self, other: str | PathLike[str]) -> Self:
+ return type(self)(other, limiter=self._limiter) / self
+
+ @property
+ def limiter(self) -> CapacityLimiter | None:
+ """The capacity limiter used by this path, if not the global limiter."""
+ return self._limiter
+
+ @property
+ def parts(self) -> tuple[str, ...]:
+ return self._path.parts
+
+ @property
+ def drive(self) -> str:
+ return self._path.drive
+
+ @property
+ def root(self) -> str:
+ return self._path.root
+
+ @property
+ def anchor(self) -> str:
+ return self._path.anchor
+
+ @property
+ def parents(self) -> Sequence[Self]:
+ return tuple(type(self)(p, limiter=self._limiter) for p in self._path.parents)
+
+ @property
+ def parent(self) -> Self:
+ return type(self)(self._path.parent, limiter=self._limiter)
+
+ @property
+ def name(self) -> str:
+ return self._path.name
+
+ @property
+ def suffix(self) -> str:
+ return self._path.suffix
+
+ @property
+ def suffixes(self) -> list[str]:
+ return self._path.suffixes
+
+ @property
+ def stem(self) -> str:
+ return self._path.stem
+
+ async def absolute(self) -> Self:
+ path = await to_thread.run_sync(self._path.absolute, limiter=self._limiter)
+ return type(self)(path, limiter=self._limiter)
+
+ def as_posix(self) -> str:
+ return self._path.as_posix()
+
+ def as_uri(self) -> str:
+ return self._path.as_uri()
+
+ if sys.version_info >= (3, 13):
+ parser: ClassVar[ModuleType] = pathlib.Path.parser
+
+ @classmethod
+ def from_uri(cls, uri: str, *, limiter: CapacityLimiter | None = None) -> Self:
+ return cls(pathlib.Path.from_uri(uri), limiter=limiter)
+
+ def full_match(
+ self, path_pattern: str, *, case_sensitive: bool | None = None
+ ) -> bool:
+ return self._path.full_match(path_pattern, case_sensitive=case_sensitive)
+
+ def match(
+ self, path_pattern: str, *, case_sensitive: bool | None = None
+ ) -> bool:
+ return self._path.match(path_pattern, case_sensitive=case_sensitive)
+ else:
+
+ def match(self, path_pattern: str) -> bool:
+ return self._path.match(path_pattern)
+
+ if sys.version_info >= (3, 14):
+
+ @property
+ def info(self) -> PathInfo:
+ return self._path.info
+
+ async def copy(
+ self,
+ target: str | os.PathLike[str],
+ *,
+ follow_symlinks: bool = True,
+ preserve_metadata: bool = False,
+ ) -> Self:
+ func = partial(
+ self._path.copy,
+ follow_symlinks=follow_symlinks,
+ preserve_metadata=preserve_metadata,
+ )
+ return type(self)(
+ await to_thread.run_sync(
+ func, pathlib.Path(target), limiter=self._limiter
+ ),
+ limiter=self._limiter,
+ )
+
+ async def copy_into(
+ self,
+ target_dir: str | os.PathLike[str],
+ *,
+ follow_symlinks: bool = True,
+ preserve_metadata: bool = False,
+ ) -> Self:
+ func = partial(
+ self._path.copy_into,
+ follow_symlinks=follow_symlinks,
+ preserve_metadata=preserve_metadata,
+ )
+ return type(self)(
+ await to_thread.run_sync(
+ func, pathlib.Path(target_dir), limiter=self._limiter
+ ),
+ limiter=self._limiter,
+ )
+
+ async def move(self, target: str | os.PathLike[str]) -> Self:
+ # Upstream does not handle anyio.Path properly as a PathLike
+ target = pathlib.Path(target)
+ return type(self)(
+ await to_thread.run_sync(
+ self._path.move, target, limiter=self._limiter
+ ),
+ limiter=self._limiter,
+ )
+
+ async def move_into(
+ self,
+ target_dir: str | os.PathLike[str],
+ ) -> Self:
+ return type(self)(
+ await to_thread.run_sync(
+ self._path.move_into, target_dir, limiter=self._limiter
+ ),
+ limiter=self._limiter,
+ )
+
+ def is_relative_to(self, other: str | PathLike[str]) -> bool:
+ try:
+ self.relative_to(other)
+ return True
+ except ValueError:
+ return False
+
+ async def chmod(self, mode: int, *, follow_symlinks: bool = True) -> None:
+ func = partial(os.chmod, follow_symlinks=follow_symlinks)
+ return await to_thread.run_sync(func, self._path, mode, limiter=self._limiter)
+
+ @classmethod
+ async def cwd(cls, *, limiter: CapacityLimiter | None = None) -> Self:
+ path = await to_thread.run_sync(pathlib.Path.cwd, limiter=limiter)
+ return cls(path, limiter=limiter)
+
+ async def exists(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.exists, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def expanduser(self) -> Self:
+ return type(self)(
+ await to_thread.run_sync(
+ self._path.expanduser, abandon_on_cancel=True, limiter=self._limiter
+ ),
+ limiter=self._limiter,
+ )
+
+ if sys.version_info < (3, 12):
+ # Python 3.11 and earlier
+ def glob(self, pattern: str) -> AsyncIterator[Self]:
+ gen = self._path.glob(pattern)
+ return _PathIterator(gen, self._limiter, type(self))
+ elif (3, 12) <= sys.version_info < (3, 13):
+ # changed in Python 3.12:
+ # - The case_sensitive parameter was added.
+ def glob(
+ self,
+ pattern: str,
+ *,
+ case_sensitive: bool | None = None,
+ ) -> AsyncIterator[Self]:
+ gen = self._path.glob(pattern, case_sensitive=case_sensitive)
+ return _PathIterator(gen, self._limiter, type(self))
+ elif sys.version_info >= (3, 13):
+ # Changed in Python 3.13:
+ # - The recurse_symlinks parameter was added.
+ # - The pattern parameter accepts a path-like object.
+ def glob( # type: ignore[misc] # mypy doesn't allow for differing signatures in a conditional block
+ self,
+ pattern: str | PathLike[str],
+ *,
+ case_sensitive: bool | None = None,
+ recurse_symlinks: bool = False,
+ ) -> AsyncIterator[Self]:
+ gen = self._path.glob(
+ pattern, # type: ignore[arg-type]
+ case_sensitive=case_sensitive,
+ recurse_symlinks=recurse_symlinks,
+ )
+ return _PathIterator(gen, self._limiter, type(self))
+
+ async def group(self) -> str:
+ return await to_thread.run_sync(
+ self._path.group, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def hardlink_to(
+ self, target: str | bytes | PathLike[str] | PathLike[bytes]
+ ) -> None:
+ if isinstance(target, Path):
+ target = target._path
+
+ await to_thread.run_sync(os.link, target, self, limiter=self._limiter)
+
+ @classmethod
+ async def home(cls, *, limiter: CapacityLimiter | None = None) -> Self:
+ home_path = await to_thread.run_sync(pathlib.Path.home, limiter=limiter)
+ return cls(home_path, limiter=limiter)
+
+ def is_absolute(self) -> bool:
+ return self._path.is_absolute()
+
+ async def is_block_device(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_block_device, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def is_char_device(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_char_device, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def is_dir(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_dir, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def is_fifo(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_fifo, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def is_file(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_file, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ if sys.version_info >= (3, 12):
+
+ async def is_junction(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_junction, limiter=self._limiter
+ )
+
+ async def is_mount(self) -> bool:
+ return await to_thread.run_sync(
+ os.path.ismount, self._path, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ if sys.version_info < (3, 15):
+
+ def is_reserved(self) -> bool:
+ return self._path.is_reserved()
+
+ async def is_socket(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_socket, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def is_symlink(self) -> bool:
+ return await to_thread.run_sync(
+ self._path.is_symlink, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def iterdir(self) -> AsyncIterator[Self]:
+ gen = (
+ self._path.iterdir()
+ if sys.version_info < (3, 13)
+ else await to_thread.run_sync(
+ self._path.iterdir, abandon_on_cancel=True, limiter=self._limiter
+ )
+ )
+ async for path in _PathIterator(gen, self._limiter, type(self)):
+ yield path
+
+ def joinpath(self, *args: str | PathLike[str]) -> Self:
+ return type(self)(self._path.joinpath(*args), limiter=self._limiter)
+
+ async def lchmod(self, mode: int) -> None:
+ await to_thread.run_sync(self._path.lchmod, mode, limiter=self._limiter)
+
+ async def lstat(self) -> os.stat_result:
+ return await to_thread.run_sync(
+ self._path.lstat, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def mkdir(
+ self, mode: int = 0o777, parents: bool = False, exist_ok: bool = False
+ ) -> None:
+ await to_thread.run_sync(
+ self._path.mkdir, mode, parents, exist_ok, limiter=self._limiter
+ )
+
+ @overload
+ async def open(
+ self,
+ mode: OpenBinaryMode,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ errors: str | None = ...,
+ newline: str | None = ...,
+ ) -> AsyncFile[bytes]: ...
+
+ @overload
+ async def open(
+ self,
+ mode: OpenTextMode = ...,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ errors: str | None = ...,
+ newline: str | None = ...,
+ ) -> AsyncFile[str]: ...
+
+ async def open(
+ self,
+ mode: str = "r",
+ buffering: int = -1,
+ encoding: str | None = None,
+ errors: str | None = None,
+ newline: str | None = None,
+ ) -> AsyncFile[Any]:
+ fp = await to_thread.run_sync(
+ self._path.open,
+ mode,
+ buffering,
+ encoding,
+ errors,
+ newline,
+ limiter=self._limiter,
+ )
+ return AsyncFile(fp, limiter=self._limiter)
+
+ async def owner(self) -> str:
+ return await to_thread.run_sync(
+ self._path.owner, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def read_bytes(self) -> bytes:
+ return await to_thread.run_sync(self._path.read_bytes, limiter=self._limiter)
+
+ async def read_text(
+ self, encoding: str | None = None, errors: str | None = None
+ ) -> str:
+ return await to_thread.run_sync(
+ self._path.read_text, encoding, errors, limiter=self._limiter
+ )
+
+ if sys.version_info >= (3, 12):
+
+ def relative_to(
+ self, *other: str | PathLike[str], walk_up: bool = False
+ ) -> Self:
+ # relative_to() should work with any PathLike but it doesn't
+ others = [pathlib.Path(other) for other in other]
+ return type(self)(
+ self._path.relative_to(*others, walk_up=walk_up), limiter=self._limiter
+ )
+
+ else:
+
+ def relative_to(self, *other: str | PathLike[str]) -> Self:
+ return type(self)(self._path.relative_to(*other), limiter=self._limiter)
+
+ async def readlink(self) -> Self:
+ target = await to_thread.run_sync(
+ os.readlink, self._path, limiter=self._limiter
+ )
+ return type(self)(target, limiter=self._limiter)
+
+ async def rename(self, target: str | pathlib.PurePath | Path) -> Self:
+ if isinstance(target, Path):
+ target = target._path
+
+ await to_thread.run_sync(self._path.rename, target, limiter=self._limiter)
+ return type(self)(target, limiter=self._limiter)
+
+ async def replace(self, target: str | pathlib.PurePath | Path) -> Self:
+ if isinstance(target, Path):
+ target = target._path
+
+ await to_thread.run_sync(self._path.replace, target, limiter=self._limiter)
+ return type(self)(target, limiter=self._limiter)
+
+ async def resolve(self, strict: bool = False) -> Self:
+ func = partial(self._path.resolve, strict=strict)
+ return type(self)(
+ await to_thread.run_sync(
+ func, abandon_on_cancel=True, limiter=self._limiter
+ ),
+ limiter=self._limiter,
+ )
+
+ if sys.version_info < (3, 12):
+ # Pre Python 3.12
+ def rglob(self, pattern: str) -> AsyncIterator[Self]:
+ gen = self._path.rglob(pattern)
+ return _PathIterator(gen, self._limiter, type(self))
+ elif (3, 12) <= sys.version_info < (3, 13):
+ # Changed in Python 3.12:
+ # - The case_sensitive parameter was added.
+ def rglob(
+ self, pattern: str, *, case_sensitive: bool | None = None
+ ) -> AsyncIterator[Self]:
+ gen = self._path.rglob(pattern, case_sensitive=case_sensitive)
+ return _PathIterator(gen, self._limiter, type(self))
+ elif sys.version_info >= (3, 13):
+ # Changed in Python 3.13:
+ # - The recurse_symlinks parameter was added.
+ # - The pattern parameter accepts a path-like object.
+ def rglob( # type: ignore[misc] # mypy doesn't allow for differing signatures in a conditional block
+ self,
+ pattern: str | PathLike[str],
+ *,
+ case_sensitive: bool | None = None,
+ recurse_symlinks: bool = False,
+ ) -> AsyncIterator[Self]:
+ gen = self._path.rglob(
+ pattern, # type: ignore[arg-type]
+ case_sensitive=case_sensitive,
+ recurse_symlinks=recurse_symlinks,
+ )
+ return _PathIterator(gen, self._limiter, type(self))
+
+ async def rmdir(self) -> None:
+ await to_thread.run_sync(self._path.rmdir, limiter=self._limiter)
+
+ async def samefile(self, other_path: str | PathLike[str]) -> bool:
+ if isinstance(other_path, Path):
+ other_path = other_path._path
+
+ return await to_thread.run_sync(
+ self._path.samefile,
+ other_path,
+ abandon_on_cancel=True,
+ limiter=self._limiter,
+ )
+
+ async def stat(self, *, follow_symlinks: bool = True) -> os.stat_result:
+ func = partial(os.stat, follow_symlinks=follow_symlinks)
+ return await to_thread.run_sync(
+ func, self._path, abandon_on_cancel=True, limiter=self._limiter
+ )
+
+ async def symlink_to(
+ self,
+ target: str | bytes | PathLike[str] | PathLike[bytes],
+ target_is_directory: bool = False,
+ ) -> None:
+ if isinstance(target, Path):
+ target = target._path
+
+ await to_thread.run_sync(
+ self._path.symlink_to, target, target_is_directory, limiter=self._limiter
+ )
+
+ async def touch(self, mode: int = 0o666, exist_ok: bool = True) -> None:
+ await to_thread.run_sync(
+ self._path.touch, mode, exist_ok, limiter=self._limiter
+ )
+
+ async def unlink(self, missing_ok: bool = False) -> None:
+ try:
+ await to_thread.run_sync(self._path.unlink, limiter=self._limiter)
+ except FileNotFoundError:
+ if not missing_ok:
+ raise
+
+ if sys.version_info >= (3, 12):
+
+ async def walk(
+ self,
+ top_down: bool = True,
+ on_error: Callable[[OSError], object] | None = None,
+ follow_symlinks: bool = False,
+ ) -> AsyncIterator[tuple[Self, list[str], list[str]]]:
+ def get_next_value() -> tuple[pathlib.Path, list[str], list[str]] | None:
+ try:
+ return next(gen)
+ except StopIteration:
+ return None
+
+ gen = self._path.walk(top_down, on_error, follow_symlinks)
+ while True:
+ value = await to_thread.run_sync(get_next_value, limiter=self._limiter)
+ if value is None:
+ return
+
+ root, dirs, paths = value
+ yield type(self)(root, limiter=self._limiter), dirs, paths
+
+ def with_name(self, name: str) -> Self:
+ return type(self)(self._path.with_name(name), limiter=self._limiter)
+
+ def with_stem(self, stem: str) -> Self:
+ return type(self)(
+ self._path.with_name(stem + self._path.suffix), limiter=self._limiter
+ )
+
+ def with_suffix(self, suffix: str) -> Self:
+ return type(self)(self._path.with_suffix(suffix), limiter=self._limiter)
+
+ def with_segments(self, *pathsegments: str | PathLike[str]) -> Self:
+ return type(self)(*pathsegments, limiter=self._limiter)
+
+ async def write_bytes(self, data: ReadableBuffer) -> int:
+ return await to_thread.run_sync(
+ self._path.write_bytes, data, limiter=self._limiter
+ )
+
+ async def write_text(
+ self,
+ data: str,
+ encoding: str | None = None,
+ errors: str | None = None,
+ newline: str | None = None,
+ ) -> int:
+ return await to_thread.run_sync(
+ self._path.write_text,
+ data,
+ encoding,
+ errors,
+ newline,
+ limiter=self._limiter,
+ )
+
+
+PathLike.register(Path)
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_resources.py b/venv/lib/python3.11/site-packages/anyio/_core/_resources.py
new file mode 100644
index 0000000000000000000000000000000000000000..b9a5344aef2962670f9b305a02cd0b11f2087d2f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_resources.py
@@ -0,0 +1,18 @@
+from __future__ import annotations
+
+from ..abc import AsyncResource
+from ._tasks import CancelScope
+
+
+async def aclose_forcefully(resource: AsyncResource) -> None:
+ """
+ Close an asynchronous resource in a cancelled scope.
+
+ Doing this closes the resource without waiting on anything.
+
+ :param resource: the resource to close
+
+ """
+ with CancelScope() as scope:
+ scope.cancel()
+ await resource.aclose()
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_signals.py b/venv/lib/python3.11/site-packages/anyio/_core/_signals.py
new file mode 100644
index 0000000000000000000000000000000000000000..e24c79e10d4b76775679f7dd0dbe3f5860150451
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_signals.py
@@ -0,0 +1,29 @@
+from __future__ import annotations
+
+from collections.abc import AsyncIterator
+from contextlib import AbstractContextManager
+from signal import Signals
+
+from ._eventloop import get_async_backend
+
+
+def open_signal_receiver(
+ *signals: Signals,
+) -> AbstractContextManager[AsyncIterator[Signals]]:
+ """
+ Start receiving operating system signals.
+
+ :param signals: signals to receive (e.g. ``signal.SIGINT``)
+ :return: an asynchronous context manager for an asynchronous iterator which yields
+ signal numbers
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ .. warning:: Windows does not support signals natively so it is best to avoid
+ relying on this in cross-platform applications.
+
+ .. warning:: On asyncio, this permanently replaces any previous signal handler for
+ the given signals, as set via :meth:`~asyncio.loop.add_signal_handler`.
+
+ """
+ return get_async_backend().open_signal_receiver(*signals)
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_sockets.py b/venv/lib/python3.11/site-packages/anyio/_core/_sockets.py
new file mode 100644
index 0000000000000000000000000000000000000000..c75791b85b5549e7b56ac501da2fbe33bc58f99e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_sockets.py
@@ -0,0 +1,1011 @@
+from __future__ import annotations
+
+import errno
+import os
+import socket
+import ssl
+import stat
+import sys
+from collections.abc import Awaitable
+from dataclasses import dataclass
+from ipaddress import IPv4Address, IPv6Address, ip_address
+from os import PathLike, chmod
+from socket import AddressFamily, SocketKind
+from typing import TYPE_CHECKING, Any, Literal, cast, overload
+
+from .. import ConnectionFailed, to_thread
+from ..abc import (
+ ByteStreamConnectable,
+ ConnectedUDPSocket,
+ ConnectedUNIXDatagramSocket,
+ IPAddressType,
+ IPSockAddrType,
+ SocketListener,
+ SocketStream,
+ UDPSocket,
+ UNIXDatagramSocket,
+ UNIXSocketStream,
+)
+from ..streams.stapled import MultiListener
+from ..streams.tls import TLSConnectable, TLSStream
+from ._eventloop import get_async_backend
+from ._resources import aclose_forcefully
+from ._synchronization import Event
+from ._tasks import create_task_group, move_on_after
+
+if TYPE_CHECKING:
+ from _typeshed import FileDescriptorLike
+else:
+ FileDescriptorLike = object
+
+if sys.version_info < (3, 11):
+ from exceptiongroup import ExceptionGroup
+
+if sys.version_info >= (3, 12):
+ from typing import override
+else:
+ from typing_extensions import override
+
+if sys.version_info < (3, 13):
+ from typing_extensions import deprecated
+else:
+ from warnings import deprecated
+
+IPPROTO_IPV6 = getattr(socket, "IPPROTO_IPV6", 41) # https://bugs.python.org/issue29515
+
+AnyIPAddressFamily = Literal[
+ AddressFamily.AF_UNSPEC, AddressFamily.AF_INET, AddressFamily.AF_INET6
+]
+IPAddressFamily = Literal[AddressFamily.AF_INET, AddressFamily.AF_INET6]
+
+
+def idna2008_resolve(host: str) -> bytes:
+ try:
+ return host.encode("ascii")
+ except UnicodeEncodeError:
+ import idna
+
+ return idna.encode(host, uts46=True)
+
+
+# tls_hostname given
+@overload
+async def connect_tcp(
+ remote_host: IPAddressType,
+ remote_port: int,
+ *,
+ local_host: IPAddressType | None = ...,
+ local_port: int | None = ...,
+ ssl_context: ssl.SSLContext | None = ...,
+ tls_standard_compatible: bool = ...,
+ tls_hostname: str,
+ happy_eyeballs_delay: float = ...,
+) -> TLSStream: ...
+
+
+# ssl_context given
+@overload
+async def connect_tcp(
+ remote_host: IPAddressType,
+ remote_port: int,
+ *,
+ local_host: IPAddressType | None = ...,
+ local_port: int | None = ...,
+ ssl_context: ssl.SSLContext,
+ tls_standard_compatible: bool = ...,
+ tls_hostname: str | None = ...,
+ happy_eyeballs_delay: float = ...,
+) -> TLSStream: ...
+
+
+# tls=True
+@overload
+async def connect_tcp(
+ remote_host: IPAddressType,
+ remote_port: int,
+ *,
+ local_host: IPAddressType | None = ...,
+ local_port: int | None = ...,
+ tls: Literal[True],
+ ssl_context: ssl.SSLContext | None = ...,
+ tls_standard_compatible: bool = ...,
+ tls_hostname: str | None = ...,
+ happy_eyeballs_delay: float = ...,
+) -> TLSStream: ...
+
+
+# tls=False
+@overload
+async def connect_tcp(
+ remote_host: IPAddressType,
+ remote_port: int,
+ *,
+ local_host: IPAddressType | None = ...,
+ local_port: int | None = ...,
+ tls: Literal[False],
+ ssl_context: ssl.SSLContext | None = ...,
+ tls_standard_compatible: bool = ...,
+ tls_hostname: str | None = ...,
+ happy_eyeballs_delay: float = ...,
+) -> SocketStream: ...
+
+
+# No TLS arguments
+@overload
+async def connect_tcp(
+ remote_host: IPAddressType,
+ remote_port: int,
+ *,
+ local_host: IPAddressType | None = ...,
+ local_port: int | None = ...,
+ happy_eyeballs_delay: float = ...,
+) -> SocketStream: ...
+
+
+async def connect_tcp(
+ remote_host: IPAddressType,
+ remote_port: int,
+ *,
+ local_host: IPAddressType | None = None,
+ local_port: int | None = None,
+ tls: bool = False,
+ ssl_context: ssl.SSLContext | None = None,
+ tls_standard_compatible: bool = True,
+ tls_hostname: str | None = None,
+ happy_eyeballs_delay: float = 0.25,
+) -> SocketStream | TLSStream:
+ """
+ Connect to a host using the TCP protocol.
+
+ This function implements the stateless version of the Happy Eyeballs algorithm (RFC
+ 6555). If ``remote_host`` is a host name that resolves to multiple IP addresses,
+ each one is tried until one connection attempt succeeds. If the first attempt does
+ not connected within 250 milliseconds, a second attempt is started using the next
+ address in the list, and so on. On IPv6 enabled systems, an IPv6 address (if
+ available) is tried first.
+
+ When the connection has been established, a TLS handshake will be done if either
+ ``ssl_context`` or ``tls_hostname`` is not ``None``, or if ``tls`` is ``True``.
+
+ :param remote_host: the IP address or host name to connect to
+ :param remote_port: port on the target host to connect to
+ :param local_host: the interface address or name to bind the socket to before
+ connecting
+ :param local_port: the local port to bind to (requires ``local_host`` to also be
+ set)
+ :param tls: ``True`` to do a TLS handshake with the connected stream and return a
+ :class:`~anyio.streams.tls.TLSStream` instead
+ :param ssl_context: the SSL context object to use (if omitted, a default context is
+ created)
+ :param tls_standard_compatible: If ``True``, performs the TLS shutdown handshake
+ before closing the stream and requires that the server does this as well.
+ Otherwise, :exc:`~ssl.SSLEOFError` may be raised during reads from the stream.
+ Some protocols, such as HTTP, require this option to be ``False``.
+ See :meth:`~ssl.SSLContext.wrap_socket` for details.
+ :param tls_hostname: host name to check the server certificate against (defaults to
+ the value of ``remote_host``)
+ :param happy_eyeballs_delay: delay (in seconds) before starting the next connection
+ attempt
+ :return: a socket stream object if no TLS handshake was done, otherwise a TLS stream
+ :raises ConnectionFailed: if the connection fails
+
+ """
+ # Placed here due to https://github.com/python/mypy/issues/7057
+ connected_stream: SocketStream | None = None
+
+ async def try_connect(remote_host: str, event: Event) -> None:
+ nonlocal connected_stream
+ try:
+ stream = await asynclib.connect_tcp(remote_host, remote_port, local_address)
+ except OSError as exc:
+ oserrors.append(exc)
+ return
+ else:
+ if connected_stream is None:
+ connected_stream = stream
+ tg.cancel_scope.cancel()
+ else:
+ await stream.aclose()
+ finally:
+ event.set()
+
+ asynclib = get_async_backend()
+ local_address: IPSockAddrType | None = None
+ family = socket.AF_UNSPEC
+ if local_host:
+ gai_res = await getaddrinfo(str(local_host), local_port)
+ family, *_, local_address = gai_res[0]
+
+ target_host = str(remote_host)
+ try:
+ addr_obj = ip_address(remote_host)
+ except ValueError:
+ addr_obj = None
+
+ if addr_obj is not None:
+ if isinstance(addr_obj, IPv6Address):
+ target_addrs = [(socket.AF_INET6, addr_obj.compressed)]
+ else:
+ target_addrs = [(socket.AF_INET, addr_obj.compressed)]
+ else:
+ # getaddrinfo() will raise an exception if name resolution fails
+ gai_res = await getaddrinfo(
+ target_host, remote_port, family=family, type=socket.SOCK_STREAM
+ )
+
+ # Organize the list so that the first address is an IPv6 address (if available)
+ # and the second one is an IPv4 addresses. The rest can be in whatever order.
+ v6_found = v4_found = False
+ target_addrs = []
+ for af, *_, sa in gai_res:
+ if af == socket.AF_INET6 and not v6_found:
+ v6_found = True
+ target_addrs.insert(0, (af, sa[0]))
+ elif af == socket.AF_INET and not v4_found and v6_found:
+ v4_found = True
+ target_addrs.insert(1, (af, sa[0]))
+ else:
+ target_addrs.append((af, sa[0]))
+
+ oserrors: list[OSError] = []
+ try:
+ async with create_task_group() as tg:
+ for _af, addr in target_addrs:
+ event = Event()
+ tg.start_soon(try_connect, addr, event)
+ with move_on_after(happy_eyeballs_delay):
+ await event.wait()
+
+ if connected_stream is None:
+ cause = (
+ oserrors[0]
+ if len(oserrors) == 1
+ else ExceptionGroup("multiple connection attempts failed", oserrors)
+ )
+ raise OSError("All connection attempts failed") from cause
+ finally:
+ oserrors.clear()
+
+ if tls or tls_hostname or ssl_context:
+ try:
+ return await TLSStream.wrap(
+ connected_stream,
+ server_side=False,
+ hostname=tls_hostname or str(remote_host),
+ ssl_context=ssl_context,
+ standard_compatible=tls_standard_compatible,
+ )
+ except BaseException:
+ await aclose_forcefully(connected_stream)
+ raise
+
+ return connected_stream
+
+
+async def connect_unix(path: str | bytes | PathLike[Any]) -> UNIXSocketStream:
+ """
+ Connect to the given UNIX socket.
+
+ Not available on Windows.
+
+ :param path: path to the socket
+ :return: a socket stream object
+ :raises ConnectionFailed: if the connection fails
+
+ """
+ path = os.fspath(path)
+ return await get_async_backend().connect_unix(path)
+
+
+async def create_tcp_listener(
+ *,
+ local_host: IPAddressType | None = None,
+ local_port: int = 0,
+ family: AnyIPAddressFamily = socket.AddressFamily.AF_UNSPEC,
+ backlog: int = 65536,
+ reuse_port: bool = False,
+) -> MultiListener[SocketStream]:
+ """
+ Create a TCP socket listener.
+
+ :param local_port: port number to listen on
+ :param local_host: IP address of the interface to listen on. If omitted, listen on
+ all IPv4 and IPv6 interfaces. To listen on all interfaces on a specific address
+ family, use ``0.0.0.0`` for IPv4 or ``::`` for IPv6.
+ :param family: address family (used if ``local_host`` was omitted)
+ :param backlog: maximum number of queued incoming connections (up to a maximum of
+ 2**16, or 65536)
+ :param reuse_port: ``True`` to allow multiple sockets to bind to the same
+ address/port (not supported on Windows)
+ :return: a multi-listener object containing one or more socket listeners
+ :raises OSError: if there's an error creating a socket, or binding to one or more
+ interfaces failed
+
+ """
+ asynclib = get_async_backend()
+ backlog = min(backlog, 65536)
+ local_host = str(local_host) if local_host is not None else None
+
+ def setup_raw_socket(
+ fam: AddressFamily,
+ bind_addr: tuple[str, int] | tuple[str, int, int, int],
+ *,
+ v6only: bool = True,
+ ) -> socket.socket:
+ sock = socket.socket(fam)
+ try:
+ sock.setblocking(False)
+
+ if fam == AddressFamily.AF_INET6:
+ sock.setsockopt(IPPROTO_IPV6, socket.IPV6_V6ONLY, v6only)
+
+ # For Windows, enable exclusive address use. For others, enable address
+ # reuse.
+ if sys.platform == "win32":
+ sock.setsockopt(socket.SOL_SOCKET, socket.SO_EXCLUSIVEADDRUSE, 1)
+ else:
+ sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
+
+ if reuse_port:
+ sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1)
+
+ # Workaround for #554
+ if fam == socket.AF_INET6 and "%" in bind_addr[0]:
+ addr, scope_id = bind_addr[0].split("%", 1)
+ bind_addr = (addr, bind_addr[1], 0, int(scope_id))
+
+ sock.bind(bind_addr)
+ sock.listen(backlog)
+ except BaseException:
+ sock.close()
+ raise
+
+ return sock
+
+ # We passing type=0 on non-Windows platforms as a workaround for a uvloop bug
+ # where we don't get the correct scope ID for IPv6 link-local addresses when passing
+ # type=socket.SOCK_STREAM to getaddrinfo():
+ # https://github.com/MagicStack/uvloop/issues/539
+ gai_res = await getaddrinfo(
+ local_host,
+ local_port,
+ family=family,
+ type=socket.SOCK_STREAM if sys.platform == "win32" else 0,
+ flags=socket.AI_PASSIVE | socket.AI_ADDRCONFIG,
+ )
+
+ # The set comprehension is here to work around a glibc bug:
+ # https://sourceware.org/bugzilla/show_bug.cgi?id=14969
+ sockaddrs = sorted({res for res in gai_res if res[1] == SocketKind.SOCK_STREAM})
+
+ # Special case for dual-stack binding on the "any" interface
+ if (
+ local_host is None
+ and family == AddressFamily.AF_UNSPEC
+ and socket.has_dualstack_ipv6()
+ and any(fam == AddressFamily.AF_INET6 for fam, *_ in gai_res)
+ ):
+ raw_socket = setup_raw_socket(
+ AddressFamily.AF_INET6, ("::", local_port), v6only=False
+ )
+ listener = asynclib.create_tcp_listener(raw_socket)
+ return MultiListener([listener])
+
+ errors: list[OSError] = []
+ try:
+ for _ in range(len(sockaddrs)):
+ listeners: list[SocketListener] = []
+ bound_ephemeral_port = local_port
+ try:
+ for fam, *_, sockaddr in sockaddrs:
+ sockaddr = sockaddr[0], bound_ephemeral_port, *sockaddr[2:]
+ raw_socket = setup_raw_socket(fam, sockaddr)
+
+ # Store the assigned port if an ephemeral port was requested, so
+ # we'll bind to the same port on all interfaces
+ if local_port == 0 and len(gai_res) > 1:
+ bound_ephemeral_port = raw_socket.getsockname()[1]
+
+ listeners.append(asynclib.create_tcp_listener(raw_socket))
+ except BaseException as exc:
+ for listener in listeners:
+ await listener.aclose()
+
+ # If an ephemeral port was requested but binding the assigned port
+ # failed for another interface, rotate the address list and try again
+ if (
+ isinstance(exc, OSError)
+ and exc.errno == errno.EADDRINUSE
+ and local_port == 0
+ and bound_ephemeral_port
+ ):
+ errors.append(exc)
+ sockaddrs.append(sockaddrs.pop(0))
+ continue
+
+ raise
+
+ return MultiListener(listeners)
+
+ raise OSError(
+ f"Could not create {len(sockaddrs)} listeners with a consistent port"
+ ) from ExceptionGroup("Several bind attempts failed", errors)
+ finally:
+ del errors # Prevent reference cycles
+
+
+async def create_unix_listener(
+ path: str | bytes | PathLike[Any],
+ *,
+ mode: int | None = None,
+ backlog: int = 65536,
+) -> SocketListener:
+ """
+ Create a UNIX socket listener.
+
+ Not available on Windows.
+
+ :param path: path of the socket
+ :param mode: permissions to set on the socket
+ :param backlog: maximum number of queued incoming connections (up to a maximum of
+ 2**16, or 65536)
+ :return: a listener object
+
+ .. versionchanged:: 3.0
+ If a socket already exists on the file system in the given path, it will be
+ removed first.
+
+ """
+ backlog = min(backlog, 65536)
+ raw_socket = await setup_unix_local_socket(path, mode, socket.SOCK_STREAM)
+ try:
+ raw_socket.listen(backlog)
+ return get_async_backend().create_unix_listener(raw_socket)
+ except BaseException:
+ raw_socket.close()
+ raise
+
+
+async def create_udp_socket(
+ family: AnyIPAddressFamily = AddressFamily.AF_UNSPEC,
+ *,
+ local_host: IPAddressType | None = None,
+ local_port: int = 0,
+ reuse_port: bool = False,
+) -> UDPSocket:
+ """
+ Create a UDP socket.
+
+ If ``port`` has been given, the socket will be bound to this port on the local
+ machine, making this socket suitable for providing UDP based services.
+
+ :param family: address family (``AF_INET`` or ``AF_INET6``) – automatically
+ determined from ``local_host`` if omitted
+ :param local_host: IP address or host name of the local interface to bind to
+ :param local_port: local port to bind to
+ :param reuse_port: ``True`` to allow multiple sockets to bind to the same
+ address/port (not supported on Windows)
+ :return: a UDP socket
+
+ """
+ if family is AddressFamily.AF_UNSPEC and not local_host:
+ raise ValueError('Either "family" or "local_host" must be given')
+
+ if local_host:
+ gai_res = await getaddrinfo(
+ str(local_host),
+ local_port,
+ family=family,
+ type=socket.SOCK_DGRAM,
+ flags=socket.AI_PASSIVE | socket.AI_ADDRCONFIG,
+ )
+ family = cast(AnyIPAddressFamily, gai_res[0][0])
+ local_address = gai_res[0][-1]
+ elif family is AddressFamily.AF_INET6:
+ local_address = ("::", 0)
+ else:
+ local_address = ("0.0.0.0", 0)
+
+ sock = await get_async_backend().create_udp_socket(
+ family, local_address, None, reuse_port
+ )
+ return cast(UDPSocket, sock)
+
+
+async def create_connected_udp_socket(
+ remote_host: IPAddressType,
+ remote_port: int,
+ *,
+ family: AnyIPAddressFamily = AddressFamily.AF_UNSPEC,
+ local_host: IPAddressType | None = None,
+ local_port: int = 0,
+ reuse_port: bool = False,
+) -> ConnectedUDPSocket:
+ """
+ Create a connected UDP socket.
+
+ Connected UDP sockets can only communicate with the specified remote host/port, an
+ any packets sent from other sources are dropped.
+
+ :param remote_host: remote host to set as the default target
+ :param remote_port: port on the remote host to set as the default target
+ :param family: address family (``AF_INET`` or ``AF_INET6``) – automatically
+ determined from ``local_host`` or ``remote_host`` if omitted
+ :param local_host: IP address or host name of the local interface to bind to
+ :param local_port: local port to bind to
+ :param reuse_port: ``True`` to allow multiple sockets to bind to the same
+ address/port (not supported on Windows)
+ :return: a connected UDP socket
+
+ """
+ local_address = None
+ if local_host:
+ gai_res = await getaddrinfo(
+ str(local_host),
+ local_port,
+ family=family,
+ type=socket.SOCK_DGRAM,
+ flags=socket.AI_PASSIVE | socket.AI_ADDRCONFIG,
+ )
+ family = cast(AnyIPAddressFamily, gai_res[0][0])
+ local_address = gai_res[0][-1]
+
+ gai_res = await getaddrinfo(
+ str(remote_host), remote_port, family=family, type=socket.SOCK_DGRAM
+ )
+ family = cast(AnyIPAddressFamily, gai_res[0][0])
+ remote_address = gai_res[0][-1]
+
+ sock = await get_async_backend().create_udp_socket(
+ family, local_address, remote_address, reuse_port
+ )
+ return cast(ConnectedUDPSocket, sock)
+
+
+async def create_unix_datagram_socket(
+ *,
+ local_path: None | str | bytes | PathLike[Any] = None,
+ local_mode: int | None = None,
+) -> UNIXDatagramSocket:
+ """
+ Create a UNIX datagram socket.
+
+ Not available on Windows.
+
+ If ``local_path`` has been given, the socket will be bound to this path, making this
+ socket suitable for receiving datagrams from other processes. Other processes can
+ send datagrams to this socket only if ``local_path`` is set.
+
+ If a socket already exists on the file system in the ``local_path``, it will be
+ removed first.
+
+ :param local_path: the path on which to bind to
+ :param local_mode: permissions to set on the local socket
+ :return: a UNIX datagram socket
+
+ """
+ raw_socket = await setup_unix_local_socket(
+ local_path, local_mode, socket.SOCK_DGRAM
+ )
+ return await get_async_backend().create_unix_datagram_socket(raw_socket, None)
+
+
+async def create_connected_unix_datagram_socket(
+ remote_path: str | bytes | PathLike[Any],
+ *,
+ local_path: None | str | bytes | PathLike[Any] = None,
+ local_mode: int | None = None,
+) -> ConnectedUNIXDatagramSocket:
+ """
+ Create a connected UNIX datagram socket.
+
+ Connected datagram sockets can only communicate with the specified remote path.
+
+ If ``local_path`` has been given, the socket will be bound to this path, making
+ this socket suitable for receiving datagrams from other processes. Other processes
+ can send datagrams to this socket only if ``local_path`` is set.
+
+ If a socket already exists on the file system in the ``local_path``, it will be
+ removed first.
+
+ :param remote_path: the path to set as the default target
+ :param local_path: the path on which to bind to
+ :param local_mode: permissions to set on the local socket
+ :return: a connected UNIX datagram socket
+
+ """
+ remote_path = os.fspath(remote_path)
+ raw_socket = await setup_unix_local_socket(
+ local_path, local_mode, socket.SOCK_DGRAM
+ )
+ return await get_async_backend().create_unix_datagram_socket(
+ raw_socket, remote_path
+ )
+
+
+async def getaddrinfo(
+ host: bytes | str | None,
+ port: str | int | None,
+ *,
+ family: int | AddressFamily = 0,
+ type: int | SocketKind = 0,
+ proto: int = 0,
+ flags: int = 0,
+) -> list[tuple[AddressFamily, SocketKind, int, str, tuple[str, int]]]:
+ """
+ Look up a numeric IP address given a host name.
+
+ Internationalized domain names are translated according to the (non-transitional)
+ IDNA 2008 standard.
+
+ .. note:: 4-tuple IPv6 socket addresses are automatically converted to 2-tuples of
+ (host, port), unlike what :func:`socket.getaddrinfo` does.
+
+ :param host: host name
+ :param port: port number
+ :param family: socket family (`'AF_INET``, ...)
+ :param type: socket type (``SOCK_STREAM``, ...)
+ :param proto: protocol number
+ :param flags: flags to pass to upstream ``getaddrinfo()``
+ :return: list of tuples containing (family, type, proto, canonname, sockaddr)
+
+ .. seealso:: :func:`socket.getaddrinfo`
+
+ """
+ # Handle unicode hostnames
+ encoded_host = idna2008_resolve(host) if isinstance(host, str) else host
+ gai_res = await get_async_backend().getaddrinfo(
+ encoded_host, port, family=family, type=type, proto=proto, flags=flags
+ )
+ return [
+ (family, type, proto, canonname, convert_ipv6_sockaddr(sockaddr))
+ for family, type, proto, canonname, sockaddr in gai_res
+ # filter out IPv6 results when IPv6 is disabled
+ if not isinstance(sockaddr[0], int)
+ ]
+
+
+def getnameinfo(sockaddr: IPSockAddrType, flags: int = 0) -> Awaitable[tuple[str, str]]:
+ """
+ Look up the host name of an IP address.
+
+ :param sockaddr: socket address (e.g. (ipaddress, port) for IPv4)
+ :param flags: flags to pass to upstream ``getnameinfo()``
+ :return: a tuple of (host name, service name)
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ .. seealso:: :func:`socket.getnameinfo`
+
+ """
+ return get_async_backend().getnameinfo(sockaddr, flags)
+
+
+@deprecated("This function is deprecated; use `wait_readable` instead")
+def wait_socket_readable(sock: socket.socket) -> Awaitable[None]:
+ """
+ .. deprecated:: 4.7.0
+ Use :func:`wait_readable` instead.
+
+ Wait until the given socket has data to be read.
+
+ .. warning:: Only use this on raw sockets that have not been wrapped by any higher
+ level constructs like socket streams!
+
+ :param sock: a socket object
+ :raises ~anyio.ClosedResourceError: if the socket was closed while waiting for the
+ socket to become readable
+ :raises ~anyio.BusyResourceError: if another task is already waiting for the socket
+ to become readable
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().wait_readable(sock.fileno())
+
+
+@deprecated("This function is deprecated; use `wait_writable` instead")
+def wait_socket_writable(sock: socket.socket) -> Awaitable[None]:
+ """
+ .. deprecated:: 4.7.0
+ Use :func:`wait_writable` instead.
+
+ Wait until the given socket can be written to.
+
+ This does **NOT** work on Windows when using the asyncio backend with a proactor
+ event loop (default on py3.8+).
+
+ .. warning:: Only use this on raw sockets that have not been wrapped by any higher
+ level constructs like socket streams!
+
+ :param sock: a socket object
+ :raises ~anyio.ClosedResourceError: if the socket was closed while waiting for the
+ socket to become writable
+ :raises ~anyio.BusyResourceError: if another task is already waiting for the socket
+ to become writable
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().wait_writable(sock.fileno())
+
+
+def wait_readable(obj: FileDescriptorLike) -> Awaitable[None]:
+ """
+ Wait until the given object has data to be read.
+
+ On Unix systems, ``obj`` must either be an integer file descriptor, or else an
+ object with a ``.fileno()`` method which returns an integer file descriptor. Any
+ kind of file descriptor can be passed, though the exact semantics will depend on
+ your kernel. For example, this probably won't do anything useful for on-disk files.
+
+ On Windows systems, ``obj`` must either be an integer ``SOCKET`` handle, or else an
+ object with a ``.fileno()`` method which returns an integer ``SOCKET`` handle. File
+ descriptors aren't supported, and neither are handles that refer to anything besides
+ a ``SOCKET``.
+
+ On backends where this functionality is not natively provided (asyncio
+ ``ProactorEventLoop`` on Windows), it is provided using a separate selector thread
+ which is set to shut down when the interpreter shuts down.
+
+ .. warning:: Don't use this on raw sockets that have been wrapped by any higher
+ level constructs like socket streams!
+
+ :param obj: an object with a ``.fileno()`` method or an integer handle
+ :raises ~anyio.ClosedResourceError: if the object was closed while waiting for the
+ object to become readable
+ :raises ~anyio.BusyResourceError: if another task is already waiting for the object
+ to become readable
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().wait_readable(obj)
+
+
+def wait_writable(obj: FileDescriptorLike) -> Awaitable[None]:
+ """
+ Wait until the given object can be written to.
+
+ :param obj: an object with a ``.fileno()`` method or an integer handle
+ :raises ~anyio.ClosedResourceError: if the object was closed while waiting for the
+ object to become writable
+ :raises ~anyio.BusyResourceError: if another task is already waiting for the object
+ to become writable
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ .. seealso:: See the documentation of :func:`wait_readable` for the definition of
+ ``obj`` and notes on backend compatibility.
+
+ .. warning:: Don't use this on raw sockets that have been wrapped by any higher
+ level constructs like socket streams!
+
+ """
+ return get_async_backend().wait_writable(obj)
+
+
+def notify_closing(obj: FileDescriptorLike) -> None:
+ """
+ Call this before closing a file descriptor (on Unix) or socket (on
+ Windows). This will cause any `wait_readable` or `wait_writable`
+ calls on the given object to immediately wake up and raise
+ `~anyio.ClosedResourceError`.
+
+ This doesn't actually close the object – you still have to do that
+ yourself afterwards. Also, you want to be careful to make sure no
+ new tasks start waiting on the object in between when you call this
+ and when it's actually closed. So to close something properly, you
+ usually want to do these steps in order:
+
+ 1. Explicitly mark the object as closed, so that any new attempts
+ to use it will abort before they start.
+ 2. Call `notify_closing` to wake up any already-existing users.
+ 3. Actually close the object.
+
+ It's also possible to do them in a different order if that's more
+ convenient, *but only if* you make sure not to have any checkpoints in
+ between the steps. This way they all happen in a single atomic
+ step, so other tasks won't be able to tell what order they happened
+ in anyway.
+
+ :param obj: an object with a ``.fileno()`` method or an integer handle
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ get_async_backend().notify_closing(obj)
+
+
+#
+# Private API
+#
+
+
+def convert_ipv6_sockaddr(
+ sockaddr: tuple[str, int, int, int] | tuple[str, int],
+) -> tuple[str, int]:
+ """
+ Convert a 4-tuple IPv6 socket address to a 2-tuple (address, port) format.
+
+ If the scope ID is nonzero, it is added to the address, separated with ``%``.
+ Otherwise the flow id and scope id are simply cut off from the tuple.
+ Any other kinds of socket addresses are returned as-is.
+
+ :param sockaddr: the result of :meth:`~socket.socket.getsockname`
+ :return: the converted socket address
+
+ """
+ # This is more complicated than it should be because of MyPy
+ if isinstance(sockaddr, tuple) and len(sockaddr) == 4:
+ host, port, flowinfo, scope_id = sockaddr
+ if scope_id:
+ # PyPy (as of v7.3.11) leaves the interface name in the result, so
+ # we discard it and only get the scope ID from the end
+ # (https://foss.heptapod.net/pypy/pypy/-/issues/3938)
+ host = host.split("%")[0]
+
+ # Add scope_id to the address
+ return f"{host}%{scope_id}", port
+ else:
+ return host, port
+ else:
+ return sockaddr
+
+
+async def setup_unix_local_socket(
+ path: None | str | bytes | PathLike[Any],
+ mode: int | None,
+ socktype: int,
+) -> socket.socket:
+ """
+ Create a UNIX local socket object, deleting the socket at the given path if it
+ exists.
+
+ Not available on Windows.
+
+ :param path: path of the socket
+ :param mode: permissions to set on the socket
+ :param socktype: socket.SOCK_STREAM or socket.SOCK_DGRAM
+
+ """
+ path_str: str | None
+ if path is not None:
+ path_str = os.fsdecode(path)
+
+ # Linux abstract namespace sockets aren't backed by a concrete file so skip stat call
+ if not path_str.startswith("\0"):
+ # Copied from pathlib...
+ try:
+ stat_result = os.stat(path)
+ except OSError as e:
+ if e.errno not in (
+ errno.ENOENT,
+ errno.ENOTDIR,
+ errno.EBADF,
+ errno.ELOOP,
+ ):
+ raise
+ else:
+ if stat.S_ISSOCK(stat_result.st_mode):
+ os.unlink(path)
+ else:
+ path_str = None
+
+ raw_socket = socket.socket(socket.AF_UNIX, socktype)
+ raw_socket.setblocking(False)
+
+ if path_str is not None:
+ try:
+ await to_thread.run_sync(raw_socket.bind, path_str, abandon_on_cancel=True)
+ if mode is not None:
+ await to_thread.run_sync(chmod, path_str, mode, abandon_on_cancel=True)
+ except BaseException:
+ raw_socket.close()
+ raise
+
+ return raw_socket
+
+
+@dataclass
+class TCPConnectable(ByteStreamConnectable):
+ """
+ Connects to a TCP server at the given host and port.
+
+ :param host: host name or IP address of the server
+ :param port: TCP port number of the server
+ """
+
+ host: str | IPv4Address | IPv6Address
+ port: int
+
+ def __post_init__(self) -> None:
+ if self.port < 1 or self.port > 65535:
+ raise ValueError("TCP port number out of range")
+
+ @override
+ async def connect(self) -> SocketStream:
+ try:
+ return await connect_tcp(self.host, self.port)
+ except OSError as exc:
+ raise ConnectionFailed(
+ f"error connecting to {self.host}:{self.port}: {exc}"
+ ) from exc
+
+
+@dataclass
+class UNIXConnectable(ByteStreamConnectable):
+ """
+ Connects to a UNIX domain socket at the given path.
+
+ :param path: the file system path of the socket
+ """
+
+ path: str | bytes | PathLike[str] | PathLike[bytes]
+
+ @override
+ async def connect(self) -> UNIXSocketStream:
+ try:
+ return await connect_unix(self.path)
+ except OSError as exc:
+ raise ConnectionFailed(f"error connecting to {self.path!r}: {exc}") from exc
+
+
+def as_connectable(
+ remote: ByteStreamConnectable
+ | tuple[str | IPv4Address | IPv6Address, int]
+ | str
+ | bytes
+ | PathLike[str],
+ /,
+ *,
+ tls: bool = False,
+ ssl_context: ssl.SSLContext | None = None,
+ tls_hostname: str | None = None,
+ tls_standard_compatible: bool = True,
+) -> ByteStreamConnectable:
+ """
+ Return a byte stream connectable from the given object.
+
+ If a bytestream connectable is given, it is returned unchanged.
+ If a tuple of (host, port) is given, a TCP connectable is returned.
+ If a string or bytes path is given, a UNIX connectable is returned.
+
+ If ``tls=True``, the connectable will be wrapped in a
+ :class:`~.streams.tls.TLSConnectable`.
+
+ :param remote: a connectable, a tuple of (host, port) or a path to a UNIX socket
+ :param tls: if ``True``, wrap the plaintext connectable in a
+ :class:`~.streams.tls.TLSConnectable`, using the provided TLS settings)
+ :param ssl_context: if ``tls=True``, the SSLContext object to use (if not provided,
+ a secure default will be created)
+ :param tls_hostname: if ``tls=True``, host name of the server to use for checking
+ the server certificate (defaults to the host portion of the address for TCP
+ connectables)
+ :param tls_standard_compatible: if ``False`` and ``tls=True``, makes the TLS stream
+ skip the closing handshake when closing the connection, so it won't raise an
+ exception if the server does the same
+
+ """
+ connectable: TCPConnectable | UNIXConnectable | TLSConnectable
+ if isinstance(remote, ByteStreamConnectable):
+ return remote
+ elif isinstance(remote, tuple) and len(remote) == 2:
+ connectable = TCPConnectable(*remote)
+ elif isinstance(remote, (str, bytes, PathLike)):
+ connectable = UNIXConnectable(remote)
+ else:
+ raise TypeError(f"cannot convert {remote!r} to a connectable")
+
+ if tls:
+ if not tls_hostname and isinstance(connectable, TCPConnectable):
+ tls_hostname = str(connectable.host)
+
+ connectable = TLSConnectable(
+ connectable,
+ ssl_context=ssl_context,
+ hostname=tls_hostname,
+ standard_compatible=tls_standard_compatible,
+ )
+
+ return connectable
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_streams.py b/venv/lib/python3.11/site-packages/anyio/_core/_streams.py
new file mode 100644
index 0000000000000000000000000000000000000000..2b9c7df200f9520357503c754bcdea1c047bdda3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_streams.py
@@ -0,0 +1,52 @@
+from __future__ import annotations
+
+import math
+from typing import TypeVar
+from warnings import warn
+
+from ..streams.memory import (
+ MemoryObjectReceiveStream,
+ MemoryObjectSendStream,
+ _MemoryObjectStreamState,
+)
+
+T_Item = TypeVar("T_Item")
+
+
+class create_memory_object_stream(
+ tuple[MemoryObjectSendStream[T_Item], MemoryObjectReceiveStream[T_Item]],
+):
+ """
+ Create a memory object stream.
+
+ The stream's item type can be annotated like
+ :func:`create_memory_object_stream[T_Item]`.
+
+ :param max_buffer_size: number of items held in the buffer until ``send()`` starts
+ blocking
+ :param item_type: old way of marking the streams with the right generic type for
+ static typing (does nothing on AnyIO 4)
+
+ .. deprecated:: 4.0
+ Use ``create_memory_object_stream[YourItemType](...)`` instead.
+ :return: a tuple of (send stream, receive stream)
+
+ """
+
+ def __new__( # type: ignore[misc]
+ cls, max_buffer_size: float = 0, item_type: object = None
+ ) -> tuple[MemoryObjectSendStream[T_Item], MemoryObjectReceiveStream[T_Item]]:
+ if max_buffer_size != math.inf and not isinstance(max_buffer_size, int):
+ raise ValueError("max_buffer_size must be either an integer or math.inf")
+ if max_buffer_size < 0:
+ raise ValueError("max_buffer_size cannot be negative")
+ if item_type is not None:
+ warn(
+ "The item_type argument has been deprecated in AnyIO 4.0. "
+ "Use create_memory_object_stream[YourItemType](...) instead.",
+ DeprecationWarning,
+ stacklevel=2,
+ )
+
+ state = _MemoryObjectStreamState[T_Item](max_buffer_size)
+ return (MemoryObjectSendStream(state), MemoryObjectReceiveStream(state))
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_subprocesses.py b/venv/lib/python3.11/site-packages/anyio/_core/_subprocesses.py
new file mode 100644
index 0000000000000000000000000000000000000000..a6590ca623f5dd3a1c0fa7a2a155b2f7637fd82b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_subprocesses.py
@@ -0,0 +1,196 @@
+from __future__ import annotations
+
+from collections.abc import AsyncIterable, Iterable, Mapping, Sequence
+from io import BytesIO
+from os import PathLike
+from subprocess import PIPE, CalledProcessError, CompletedProcess
+from typing import IO, Any, TypeAlias, cast
+
+from ..abc import Process
+from ._eventloop import get_async_backend
+from ._tasks import create_task_group
+
+StrOrBytesPath: TypeAlias = str | bytes | PathLike[str] | PathLike[bytes]
+
+
+async def run_process(
+ command: StrOrBytesPath | Sequence[StrOrBytesPath],
+ *,
+ input: bytes | None = None,
+ stdin: int | IO[Any] | None = None,
+ stdout: int | IO[Any] | None = PIPE,
+ stderr: int | IO[Any] | None = PIPE,
+ check: bool = True,
+ cwd: StrOrBytesPath | None = None,
+ env: Mapping[str, str] | None = None,
+ startupinfo: Any = None,
+ creationflags: int = 0,
+ start_new_session: bool = False,
+ pass_fds: Sequence[int] = (),
+ user: str | int | None = None,
+ group: str | int | None = None,
+ extra_groups: Iterable[str | int] | None = None,
+ umask: int = -1,
+) -> CompletedProcess[bytes]:
+ """
+ Run an external command in a subprocess and wait until it completes.
+
+ .. seealso:: :func:`subprocess.run`
+
+ :param command: either a string to pass to the shell, or an iterable of strings
+ containing the executable name or path and its arguments
+ :param input: bytes passed to the standard input of the subprocess
+ :param stdin: one of :data:`subprocess.PIPE`, :data:`subprocess.DEVNULL`,
+ a file-like object, or `None`; ``input`` overrides this
+ :param stdout: one of :data:`subprocess.PIPE`, :data:`subprocess.DEVNULL`,
+ a file-like object, or `None`
+ :param stderr: one of :data:`subprocess.PIPE`, :data:`subprocess.DEVNULL`,
+ :data:`subprocess.STDOUT`, a file-like object, or `None`
+ :param check: if ``True``, raise :exc:`~subprocess.CalledProcessError` if the
+ process terminates with a return code other than 0
+ :param cwd: If not ``None``, change the working directory to this before running the
+ command
+ :param env: if not ``None``, this mapping replaces the inherited environment
+ variables from the parent process
+ :param startupinfo: an instance of :class:`subprocess.STARTUPINFO` that can be used
+ to specify process startup parameters (Windows only)
+ :param creationflags: flags that can be used to control the creation of the
+ subprocess (see :class:`subprocess.Popen` for the specifics)
+ :param start_new_session: if ``true`` the setsid() system call will be made in the
+ child process prior to the execution of the subprocess. (POSIX only)
+ :param pass_fds: sequence of file descriptors to keep open between the parent and
+ child processes. (POSIX only)
+ :param user: effective user to run the process as (Python >= 3.9, POSIX only)
+ :param group: effective group to run the process as (Python >= 3.9, POSIX only)
+ :param extra_groups: supplementary groups to set in the subprocess (Python >= 3.9,
+ POSIX only)
+ :param umask: if not negative, this umask is applied in the child process before
+ running the given command (Python >= 3.9, POSIX only)
+ :return: an object representing the completed process
+ :raises ~subprocess.CalledProcessError: if ``check`` is ``True`` and the process
+ exits with a nonzero return code
+
+ """
+
+ async def drain_stream(stream: AsyncIterable[bytes], index: int) -> None:
+ buffer = BytesIO()
+ async for chunk in stream:
+ buffer.write(chunk)
+
+ stream_contents[index] = buffer.getvalue()
+
+ if stdin is not None and input is not None:
+ raise ValueError("only one of stdin and input is allowed")
+
+ async with await open_process(
+ command,
+ stdin=PIPE if input else stdin,
+ stdout=stdout,
+ stderr=stderr,
+ cwd=cwd,
+ env=env,
+ startupinfo=startupinfo,
+ creationflags=creationflags,
+ start_new_session=start_new_session,
+ pass_fds=pass_fds,
+ user=user,
+ group=group,
+ extra_groups=extra_groups,
+ umask=umask,
+ ) as process:
+ stream_contents: list[bytes | None] = [None, None]
+ async with create_task_group() as tg:
+ if process.stdout:
+ tg.start_soon(drain_stream, process.stdout, 0)
+
+ if process.stderr:
+ tg.start_soon(drain_stream, process.stderr, 1)
+
+ if process.stdin and input:
+ await process.stdin.send(input)
+ await process.stdin.aclose()
+
+ await process.wait()
+
+ output, errors = stream_contents
+ if check and process.returncode != 0:
+ raise CalledProcessError(cast(int, process.returncode), command, output, errors)
+
+ return CompletedProcess(command, cast(int, process.returncode), output, errors)
+
+
+async def open_process(
+ command: StrOrBytesPath | Sequence[StrOrBytesPath],
+ *,
+ stdin: int | IO[Any] | None = PIPE,
+ stdout: int | IO[Any] | None = PIPE,
+ stderr: int | IO[Any] | None = PIPE,
+ cwd: StrOrBytesPath | None = None,
+ env: Mapping[str, str] | None = None,
+ startupinfo: Any = None,
+ creationflags: int = 0,
+ start_new_session: bool = False,
+ pass_fds: Sequence[int] = (),
+ user: str | int | None = None,
+ group: str | int | None = None,
+ extra_groups: Iterable[str | int] | None = None,
+ umask: int = -1,
+) -> Process:
+ """
+ Start an external command in a subprocess.
+
+ .. seealso:: :class:`subprocess.Popen`
+
+ :param command: either a string to pass to the shell, or an iterable of strings
+ containing the executable name or path and its arguments
+ :param stdin: one of :data:`subprocess.PIPE`, :data:`subprocess.DEVNULL`, a
+ file-like object, or ``None``
+ :param stdout: one of :data:`subprocess.PIPE`, :data:`subprocess.DEVNULL`,
+ a file-like object, or ``None``
+ :param stderr: one of :data:`subprocess.PIPE`, :data:`subprocess.DEVNULL`,
+ :data:`subprocess.STDOUT`, a file-like object, or ``None``
+ :param cwd: If not ``None``, the working directory is changed before executing
+ :param env: If env is not ``None``, it must be a mapping that defines the
+ environment variables for the new process
+ :param creationflags: flags that can be used to control the creation of the
+ subprocess (see :class:`subprocess.Popen` for the specifics)
+ :param startupinfo: an instance of :class:`subprocess.STARTUPINFO` that can be used
+ to specify process startup parameters (Windows only)
+ :param start_new_session: if ``true`` the setsid() system call will be made in the
+ child process prior to the execution of the subprocess. (POSIX only)
+ :param pass_fds: sequence of file descriptors to keep open between the parent and
+ child processes. (POSIX only)
+ :param user: effective user to run the process as (POSIX only)
+ :param group: effective group to run the process as (POSIX only)
+ :param extra_groups: supplementary groups to set in the subprocess (POSIX only)
+ :param umask: if not negative, this umask is applied in the child process before
+ running the given command (POSIX only)
+ :return: an asynchronous process object
+
+ """
+ kwargs: dict[str, Any] = {}
+ if user is not None:
+ kwargs["user"] = user
+
+ if group is not None:
+ kwargs["group"] = group
+
+ if extra_groups is not None:
+ kwargs["extra_groups"] = extra_groups
+
+ if umask >= 0:
+ kwargs["umask"] = umask
+
+ return await get_async_backend().open_process(
+ command,
+ stdin=stdin,
+ stdout=stdout,
+ stderr=stderr,
+ cwd=cwd,
+ env=env,
+ startupinfo=startupinfo,
+ creationflags=creationflags,
+ start_new_session=start_new_session,
+ pass_fds=pass_fds,
+ **kwargs,
+ )
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_synchronization.py b/venv/lib/python3.11/site-packages/anyio/_core/_synchronization.py
new file mode 100644
index 0000000000000000000000000000000000000000..f1990a5003a6b99cbcb3c72ae8fb01e6e299a9aa
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_synchronization.py
@@ -0,0 +1,772 @@
+from __future__ import annotations
+
+import math
+from collections import deque
+from collections.abc import Callable
+from dataclasses import dataclass
+from types import TracebackType
+from typing import TypeVar
+
+from ..lowlevel import checkpoint_if_cancelled
+from ._eventloop import get_async_backend
+from ._exceptions import BusyResourceError, NoEventLoopError
+from ._tasks import CancelScope
+from ._testing import TaskInfo, get_current_task
+
+T = TypeVar("T")
+
+
+@dataclass(frozen=True)
+class EventStatistics:
+ """
+ :ivar int tasks_waiting: number of tasks waiting on :meth:`~.Event.wait`
+ """
+
+ tasks_waiting: int
+
+
+@dataclass(frozen=True)
+class CapacityLimiterStatistics:
+ """
+ :ivar int borrowed_tokens: number of tokens currently borrowed by tasks
+ :ivar float total_tokens: total number of available tokens
+ :ivar tuple borrowers: tasks or other objects currently holding tokens borrowed from
+ this limiter
+ :ivar int tasks_waiting: number of tasks waiting on
+ :meth:`~.CapacityLimiter.acquire` or
+ :meth:`~.CapacityLimiter.acquire_on_behalf_of`
+ """
+
+ borrowed_tokens: int
+ total_tokens: float
+ borrowers: tuple[object, ...]
+ tasks_waiting: int
+
+
+@dataclass(frozen=True)
+class LockStatistics:
+ """
+ :ivar bool locked: flag indicating if this lock is locked or not
+ :ivar ~anyio.TaskInfo owner: task currently holding the lock (or ``None`` if the
+ lock is not held by any task)
+ :ivar int tasks_waiting: number of tasks waiting on :meth:`~.Lock.acquire`
+ """
+
+ locked: bool
+ owner: TaskInfo | None
+ tasks_waiting: int
+
+
+@dataclass(frozen=True)
+class ConditionStatistics:
+ """
+ :ivar int tasks_waiting: number of tasks blocked on :meth:`~.Condition.wait`
+ :ivar ~anyio.LockStatistics lock_statistics: statistics of the underlying
+ :class:`~.Lock`
+ """
+
+ tasks_waiting: int
+ lock_statistics: LockStatistics
+
+
+@dataclass(frozen=True)
+class SemaphoreStatistics:
+ """
+ :ivar int tasks_waiting: number of tasks waiting on :meth:`~.Semaphore.acquire`
+
+ """
+
+ tasks_waiting: int
+
+
+class Event:
+ __slots__ = ("__weakref__",)
+
+ def __new__(cls) -> Event:
+ try:
+ return get_async_backend().create_event()
+ except NoEventLoopError:
+ return EventAdapter()
+
+ def set(self) -> None:
+ """Set the flag, notifying all listeners."""
+ raise NotImplementedError
+
+ def is_set(self) -> bool:
+ """Return ``True`` if the flag is set, ``False`` if not."""
+ raise NotImplementedError
+
+ async def wait(self) -> None:
+ """
+ Wait until the flag has been set.
+
+ If the flag has already been set when this method is called, it returns
+ immediately.
+
+ """
+ raise NotImplementedError
+
+ def statistics(self) -> EventStatistics:
+ """Return statistics about the current state of this event."""
+ raise NotImplementedError
+
+
+class EventAdapter(Event):
+ __slots__ = "_internal_event", "_is_set"
+
+ def __new__(cls) -> EventAdapter:
+ return object.__new__(cls)
+
+ def __init__(self) -> None:
+ self._internal_event: Event | None = None
+ self._is_set = False
+
+ @property
+ def _event(self) -> Event:
+ if self._internal_event is None:
+ self._internal_event = get_async_backend().create_event()
+ if self._is_set:
+ self._internal_event.set()
+
+ return self._internal_event
+
+ def set(self) -> None:
+ if self._internal_event is None:
+ self._is_set = True
+ else:
+ self._event.set()
+
+ def is_set(self) -> bool:
+ if self._internal_event is None:
+ return self._is_set
+
+ return self._internal_event.is_set()
+
+ async def wait(self) -> None:
+ await self._event.wait()
+
+ def statistics(self) -> EventStatistics:
+ if self._internal_event is None:
+ return EventStatistics(tasks_waiting=0)
+
+ return self._internal_event.statistics()
+
+
+class Lock:
+ __slots__ = ("__weakref__",)
+
+ def __new__(cls, *, fast_acquire: bool = False) -> Lock:
+ try:
+ return get_async_backend().create_lock(fast_acquire=fast_acquire)
+ except NoEventLoopError:
+ return LockAdapter(fast_acquire=fast_acquire)
+
+ async def __aenter__(self) -> None:
+ await self.acquire()
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self.release()
+
+ async def acquire(self) -> None:
+ """Acquire the lock."""
+ raise NotImplementedError
+
+ def acquire_nowait(self) -> None:
+ """
+ Acquire the lock, without blocking.
+
+ :raises ~anyio.WouldBlock: if the operation would block
+
+ """
+ raise NotImplementedError
+
+ def release(self) -> None:
+ """Release the lock."""
+ raise NotImplementedError
+
+ def locked(self) -> bool:
+ """Return True if the lock is currently held."""
+ raise NotImplementedError
+
+ def statistics(self) -> LockStatistics:
+ """
+ Return statistics about the current state of this lock.
+
+ .. versionadded:: 3.0
+ """
+ raise NotImplementedError
+
+
+class LockAdapter(Lock):
+ __slots__ = "_internal_lock", "_fast_acquire"
+
+ def __new__(cls, *, fast_acquire: bool = False) -> LockAdapter:
+ return object.__new__(cls)
+
+ def __init__(self, *, fast_acquire: bool = False):
+ self._internal_lock: Lock | None = None
+ self._fast_acquire = fast_acquire
+
+ @property
+ def _lock(self) -> Lock:
+ if self._internal_lock is None:
+ self._internal_lock = get_async_backend().create_lock(
+ fast_acquire=self._fast_acquire
+ )
+
+ return self._internal_lock
+
+ async def __aenter__(self) -> None:
+ await self._lock.acquire()
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ if self._internal_lock is not None:
+ self._internal_lock.release()
+
+ async def acquire(self) -> None:
+ """Acquire the lock."""
+ await self._lock.acquire()
+
+ def acquire_nowait(self) -> None:
+ """
+ Acquire the lock, without blocking.
+
+ :raises ~anyio.WouldBlock: if the operation would block
+
+ """
+ self._lock.acquire_nowait()
+
+ def release(self) -> None:
+ """Release the lock."""
+ self._lock.release()
+
+ def locked(self) -> bool:
+ """Return True if the lock is currently held."""
+ return self._lock.locked()
+
+ def statistics(self) -> LockStatistics:
+ """
+ Return statistics about the current state of this lock.
+
+ .. versionadded:: 3.0
+
+ """
+ if self._internal_lock is None:
+ return LockStatistics(False, None, 0)
+
+ return self._internal_lock.statistics()
+
+
+class Condition:
+ __slots__ = "__weakref__", "_owner_task", "_lock", "_waiters"
+
+ def __init__(self, lock: Lock | None = None):
+ self._owner_task: TaskInfo | None = None
+ self._lock = lock or Lock()
+ self._waiters: deque[Event] = deque()
+
+ async def __aenter__(self) -> None:
+ await self.acquire()
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self.release()
+
+ def _check_acquired(self) -> None:
+ if self._owner_task != get_current_task():
+ raise RuntimeError("The current task is not holding the underlying lock")
+
+ async def acquire(self) -> None:
+ """Acquire the underlying lock."""
+ await self._lock.acquire()
+ self._owner_task = get_current_task()
+
+ def acquire_nowait(self) -> None:
+ """
+ Acquire the underlying lock, without blocking.
+
+ :raises ~anyio.WouldBlock: if the operation would block
+
+ """
+ self._lock.acquire_nowait()
+ self._owner_task = get_current_task()
+
+ def release(self) -> None:
+ """Release the underlying lock."""
+ self._lock.release()
+
+ def locked(self) -> bool:
+ """Return True if the lock is set."""
+ return self._lock.locked()
+
+ def notify(self, n: int = 1) -> None:
+ """Notify exactly n listeners."""
+ self._check_acquired()
+ for _ in range(n):
+ try:
+ event = self._waiters.popleft()
+ except IndexError:
+ break
+
+ event.set()
+
+ def notify_all(self) -> None:
+ """Notify all the listeners."""
+ self._check_acquired()
+ for event in self._waiters:
+ event.set()
+
+ self._waiters.clear()
+
+ async def wait(self) -> None:
+ """Wait for a notification."""
+ await checkpoint_if_cancelled()
+ self._check_acquired()
+ event = Event()
+ self._waiters.append(event)
+ self.release()
+ try:
+ await event.wait()
+ except BaseException:
+ if not event.is_set():
+ self._waiters.remove(event)
+ elif self._waiters:
+ # This task was notified by could not act on it, so pass
+ # it on to the next task
+ self._waiters.popleft().set()
+
+ raise
+ finally:
+ with CancelScope(shield=True):
+ await self.acquire()
+
+ async def wait_for(self, predicate: Callable[[], T]) -> T:
+ """
+ Wait until a predicate becomes true.
+
+ :param predicate: a callable that returns a truthy value when the condition is
+ met
+ :return: the result of the predicate
+
+ .. versionadded:: 4.11.0
+
+ """
+ while not (result := predicate()):
+ await self.wait()
+
+ return result
+
+ def statistics(self) -> ConditionStatistics:
+ """
+ Return statistics about the current state of this condition.
+
+ .. versionadded:: 3.0
+ """
+ return ConditionStatistics(len(self._waiters), self._lock.statistics())
+
+
+class Semaphore:
+ __slots__ = "__weakref__", "_fast_acquire"
+
+ def __new__(
+ cls,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> Semaphore:
+ try:
+ return get_async_backend().create_semaphore(
+ initial_value, max_value=max_value, fast_acquire=fast_acquire
+ )
+ except NoEventLoopError:
+ return SemaphoreAdapter(initial_value, max_value=max_value)
+
+ def __init__(
+ self,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ):
+ if not isinstance(initial_value, int):
+ raise TypeError("initial_value must be an integer")
+ if initial_value < 0:
+ raise ValueError("initial_value must be >= 0")
+ if max_value is not None:
+ if not isinstance(max_value, int):
+ raise TypeError("max_value must be an integer or None")
+ if max_value < initial_value:
+ raise ValueError(
+ "max_value must be equal to or higher than initial_value"
+ )
+
+ self._fast_acquire = fast_acquire
+
+ async def __aenter__(self) -> Semaphore:
+ await self.acquire()
+ return self
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self.release()
+
+ async def acquire(self) -> None:
+ """Decrement the semaphore value, blocking if necessary."""
+ raise NotImplementedError
+
+ def acquire_nowait(self) -> None:
+ """
+ Acquire the underlying lock, without blocking.
+
+ :raises ~anyio.WouldBlock: if the operation would block
+
+ """
+ raise NotImplementedError
+
+ def release(self) -> None:
+ """Increment the semaphore value."""
+ raise NotImplementedError
+
+ @property
+ def value(self) -> int:
+ """The current value of the semaphore."""
+ raise NotImplementedError
+
+ @property
+ def max_value(self) -> int | None:
+ """The maximum value of the semaphore."""
+ raise NotImplementedError
+
+ def statistics(self) -> SemaphoreStatistics:
+ """
+ Return statistics about the current state of this semaphore.
+
+ .. versionadded:: 3.0
+ """
+ raise NotImplementedError
+
+
+class SemaphoreAdapter(Semaphore):
+ __slots__ = "_internal_semaphore", "_initial_value", "_max_value"
+
+ def __new__(
+ cls,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> SemaphoreAdapter:
+ return object.__new__(cls)
+
+ def __init__(
+ self,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> None:
+ super().__init__(initial_value, max_value=max_value, fast_acquire=fast_acquire)
+ self._internal_semaphore: Semaphore | None = None
+ self._initial_value = initial_value
+ self._max_value = max_value
+
+ @property
+ def _semaphore(self) -> Semaphore:
+ if self._internal_semaphore is None:
+ self._internal_semaphore = get_async_backend().create_semaphore(
+ self._initial_value, max_value=self._max_value
+ )
+
+ return self._internal_semaphore
+
+ async def acquire(self) -> None:
+ await self._semaphore.acquire()
+
+ def acquire_nowait(self) -> None:
+ self._semaphore.acquire_nowait()
+
+ def release(self) -> None:
+ self._semaphore.release()
+
+ @property
+ def value(self) -> int:
+ if self._internal_semaphore is None:
+ return self._initial_value
+
+ return self._semaphore.value
+
+ @property
+ def max_value(self) -> int | None:
+ return self._max_value
+
+ def statistics(self) -> SemaphoreStatistics:
+ if self._internal_semaphore is None:
+ return SemaphoreStatistics(tasks_waiting=0)
+
+ return self._semaphore.statistics()
+
+
+class CapacityLimiter:
+ __slots__ = ("__weakref__",)
+
+ def __new__(cls, total_tokens: float) -> CapacityLimiter:
+ try:
+ return get_async_backend().create_capacity_limiter(total_tokens)
+ except NoEventLoopError:
+ return CapacityLimiterAdapter(total_tokens)
+
+ async def __aenter__(self) -> None:
+ raise NotImplementedError
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ raise NotImplementedError
+
+ @property
+ def total_tokens(self) -> float:
+ """
+ The total number of tokens available for borrowing.
+
+ This is a read-write property. If the total number of tokens is increased, the
+ proportionate number of tasks waiting on this limiter will be granted their
+ tokens.
+
+ .. versionchanged:: 3.0
+ The property is now writable.
+ .. versionchanged:: 4.12
+ The value can now be set to 0.
+
+ """
+ raise NotImplementedError
+
+ @total_tokens.setter
+ def total_tokens(self, value: float) -> None:
+ raise NotImplementedError
+
+ @property
+ def borrowed_tokens(self) -> int:
+ """The number of tokens that have currently been borrowed."""
+ raise NotImplementedError
+
+ @property
+ def available_tokens(self) -> float:
+ """The number of tokens currently available to be borrowed"""
+ raise NotImplementedError
+
+ def acquire_nowait(self) -> None:
+ """
+ Acquire a token for the current task without waiting for one to become
+ available.
+
+ :raises ~anyio.WouldBlock: if there are no tokens available for borrowing
+
+ """
+ raise NotImplementedError
+
+ def acquire_on_behalf_of_nowait(self, borrower: object) -> None:
+ """
+ Acquire a token without waiting for one to become available.
+
+ :param borrower: the entity borrowing a token
+ :raises ~anyio.WouldBlock: if there are no tokens available for borrowing
+
+ """
+ raise NotImplementedError
+
+ async def acquire(self) -> None:
+ """
+ Acquire a token for the current task, waiting if necessary for one to become
+ available.
+
+ """
+ raise NotImplementedError
+
+ async def acquire_on_behalf_of(self, borrower: object) -> None:
+ """
+ Acquire a token, waiting if necessary for one to become available.
+
+ :param borrower: the entity borrowing a token
+
+ """
+ raise NotImplementedError
+
+ def release(self) -> None:
+ """
+ Release the token held by the current task.
+
+ :raises RuntimeError: if the current task has not borrowed a token from this
+ limiter.
+
+ """
+ raise NotImplementedError
+
+ def release_on_behalf_of(self, borrower: object) -> None:
+ """
+ Release the token held by the given borrower.
+
+ :raises RuntimeError: if the borrower has not borrowed a token from this
+ limiter.
+
+ """
+ raise NotImplementedError
+
+ def statistics(self) -> CapacityLimiterStatistics:
+ """
+ Return statistics about the current state of this limiter.
+
+ .. versionadded:: 3.0
+
+ """
+ raise NotImplementedError
+
+
+class CapacityLimiterAdapter(CapacityLimiter):
+ __slots__ = "_internal_limiter", "_total_tokens"
+
+ def __new__(cls, total_tokens: float) -> CapacityLimiterAdapter:
+ return object.__new__(cls)
+
+ def __init__(self, total_tokens: float) -> None:
+ self._internal_limiter: CapacityLimiter | None = None
+ self.total_tokens = total_tokens
+
+ @property
+ def _limiter(self) -> CapacityLimiter:
+ if self._internal_limiter is None:
+ self._internal_limiter = get_async_backend().create_capacity_limiter(
+ self._total_tokens
+ )
+
+ return self._internal_limiter
+
+ async def __aenter__(self) -> None:
+ await self._limiter.__aenter__()
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ return await self._limiter.__aexit__(exc_type, exc_val, exc_tb)
+
+ @property
+ def total_tokens(self) -> float:
+ if self._internal_limiter is None:
+ return self._total_tokens
+
+ return self._internal_limiter.total_tokens
+
+ @total_tokens.setter
+ def total_tokens(self, value: float) -> None:
+ if not isinstance(value, int) and not math.isinf(value):
+ raise TypeError("total_tokens must be an int or math.inf")
+ elif value < 0:
+ raise ValueError("total_tokens must be >= 0")
+
+ if self._internal_limiter is None:
+ self._total_tokens = value
+ return
+
+ self._limiter.total_tokens = value
+
+ @property
+ def borrowed_tokens(self) -> int:
+ if self._internal_limiter is None:
+ return 0
+
+ return self._internal_limiter.borrowed_tokens
+
+ @property
+ def available_tokens(self) -> float:
+ if self._internal_limiter is None:
+ return self._total_tokens
+
+ return self._internal_limiter.available_tokens
+
+ def acquire_nowait(self) -> None:
+ self._limiter.acquire_nowait()
+
+ def acquire_on_behalf_of_nowait(self, borrower: object) -> None:
+ self._limiter.acquire_on_behalf_of_nowait(borrower)
+
+ async def acquire(self) -> None:
+ await self._limiter.acquire()
+
+ async def acquire_on_behalf_of(self, borrower: object) -> None:
+ await self._limiter.acquire_on_behalf_of(borrower)
+
+ def release(self) -> None:
+ self._limiter.release()
+
+ def release_on_behalf_of(self, borrower: object) -> None:
+ self._limiter.release_on_behalf_of(borrower)
+
+ def statistics(self) -> CapacityLimiterStatistics:
+ if self._internal_limiter is None:
+ return CapacityLimiterStatistics(
+ borrowed_tokens=0,
+ total_tokens=self.total_tokens,
+ borrowers=(),
+ tasks_waiting=0,
+ )
+
+ return self._internal_limiter.statistics()
+
+
+class ResourceGuard:
+ """
+ A context manager for ensuring that a resource is only used by a single task at a
+ time.
+
+ Entering this context manager while the previous has not exited it yet will trigger
+ :exc:`BusyResourceError`.
+
+ :param action: the action to guard against (visible in the :exc:`BusyResourceError`
+ when triggered, e.g. "Another task is already {action} this resource")
+
+ .. versionadded:: 4.1
+ """
+
+ __slots__ = "__weakref__", "action", "_guarded"
+
+ def __init__(self, action: str = "using"):
+ self.action: str = action
+ self._guarded = False
+
+ def __enter__(self) -> None:
+ if self._guarded:
+ raise BusyResourceError(self.action)
+
+ self._guarded = True
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self._guarded = False
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_tasks.py b/venv/lib/python3.11/site-packages/anyio/_core/_tasks.py
new file mode 100644
index 0000000000000000000000000000000000000000..ced54b2e20f7dea2f238600a1a5833fe0d64d70d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_tasks.py
@@ -0,0 +1,412 @@
+from __future__ import annotations
+
+import math
+import sys
+from collections.abc import (
+ Coroutine,
+ Generator,
+)
+from contextlib import (
+ contextmanager,
+)
+from enum import Enum, auto
+from inspect import iscoroutine
+from types import TracebackType
+from typing import Any, Generic, final
+
+from ..abc import TaskGroup, TaskStatus
+from ._eventloop import get_async_backend, get_cancelled_exc_class
+from ._exceptions import TaskCancelled, TaskFailed, TaskNotFinished
+
+if sys.version_info >= (3, 13):
+ from typing import TypeVar
+else:
+ from typing_extensions import TypeVar
+
+if sys.version_info >= (3, 11):
+ from typing import Never, TypeVarTuple
+else:
+ from typing_extensions import Never, TypeVarTuple
+
+T = TypeVar("T")
+T_co = TypeVar("T_co", covariant=True)
+T_startval = TypeVar("T_startval", covariant=True, default=Never)
+PosArgsT = TypeVarTuple("PosArgsT")
+
+
+class _IgnoredTaskStatus(TaskStatus[object]):
+ def started(self, value: object = None) -> None:
+ pass
+
+
+TASK_STATUS_IGNORED = _IgnoredTaskStatus()
+
+
+class CancelScope:
+ """
+ Wraps a unit of work that can be made separately cancellable.
+
+ :param deadline: The time (clock value) when this scope is cancelled automatically
+ :param shield: ``True`` to shield the cancel scope from external cancellation
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+ """
+
+ __slots__ = ("__weakref__",)
+
+ def __new__(
+ cls, *, deadline: float = math.inf, shield: bool = False
+ ) -> CancelScope:
+ return get_async_backend().create_cancel_scope(shield=shield, deadline=deadline)
+
+ def cancel(self, reason: str | None = None) -> None:
+ """
+ Cancel this scope immediately.
+
+ :param reason: a message describing the reason for the cancellation
+
+ """
+ raise NotImplementedError
+
+ @property
+ def deadline(self) -> float:
+ """
+ The time (clock value) when this scope is cancelled automatically.
+
+ Will be ``float('inf')`` if no timeout has been set.
+
+ """
+ raise NotImplementedError
+
+ @deadline.setter
+ def deadline(self, value: float) -> None:
+ raise NotImplementedError
+
+ @property
+ def cancel_called(self) -> bool:
+ """``True`` if :meth:`cancel` has been called."""
+ raise NotImplementedError
+
+ @property
+ def cancelled_caught(self) -> bool:
+ """
+ ``True`` if this scope suppressed a cancellation exception it itself raised.
+
+ This is typically used to check if any work was interrupted, or to see if the
+ scope was cancelled due to its deadline being reached. The value will, however,
+ only be ``True`` if the cancellation was triggered by the scope itself (and not
+ an outer scope).
+
+ """
+ raise NotImplementedError
+
+ @property
+ def shield(self) -> bool:
+ """
+ ``True`` if this scope is shielded from external cancellation.
+
+ While a scope is shielded, it will not receive cancellations from outside.
+
+ """
+ raise NotImplementedError
+
+ @shield.setter
+ def shield(self, value: bool) -> None:
+ raise NotImplementedError
+
+ def __enter__(self) -> CancelScope:
+ raise NotImplementedError
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool:
+ raise NotImplementedError
+
+
+@contextmanager
+def fail_after(
+ delay: float | None, shield: bool = False
+) -> Generator[CancelScope, None, None]:
+ """
+ Create a context manager which raises a :class:`TimeoutError` if does not finish in
+ time.
+
+ :param delay: maximum allowed time (in seconds) before raising the exception, or
+ ``None`` to disable the timeout
+ :param shield: ``True`` to shield the cancel scope from external cancellation
+ :return: a context manager that yields a cancel scope
+ :rtype: :class:`~typing.ContextManager`\\[:class:`~anyio.CancelScope`\\]
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ current_time = get_async_backend().current_time
+ deadline = (current_time() + delay) if delay is not None else math.inf
+ with get_async_backend().create_cancel_scope(
+ deadline=deadline, shield=shield
+ ) as cancel_scope:
+ yield cancel_scope
+
+ if cancel_scope.cancelled_caught and current_time() >= cancel_scope.deadline:
+ raise TimeoutError
+
+
+def move_on_after(delay: float | None, shield: bool = False) -> CancelScope:
+ """
+ Create a cancel scope with a deadline that expires after the given delay.
+
+ :param delay: maximum allowed time (in seconds) before exiting the context block, or
+ ``None`` to disable the timeout
+ :param shield: ``True`` to shield the cancel scope from external cancellation
+ :return: a cancel scope
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ deadline = (
+ (get_async_backend().current_time() + delay) if delay is not None else math.inf
+ )
+ return get_async_backend().create_cancel_scope(deadline=deadline, shield=shield)
+
+
+def current_effective_deadline() -> float:
+ """
+ Return the nearest deadline among all the cancel scopes effective for the current
+ task.
+
+ :return: a clock value from the event loop's internal clock (or ``float('inf')`` if
+ there is no deadline in effect, or ``float('-inf')`` if the current scope has
+ been cancelled)
+ :rtype: float
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().current_effective_deadline()
+
+
+def create_task_group() -> TaskGroup:
+ """
+ Create a task group.
+
+ :return: a task group
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().create_task_group()
+
+
+@final
+class TaskHandle(Generic[T_co, T_startval]):
+ """
+ Returned from the task-spawning methods of :class:`TaskGroup`. Can be awaited on to
+ get the return value of the task (or the raised exception). If the task was
+ terminated by a :exc:`BaseException`, :exc:`TaskFailed` will be raised (or its
+ subclass :exc:`TaskCancelled` if the task was cancelled).
+
+ .. versionadded:: 4.14.0
+ """
+
+ class Status(Enum):
+ """
+ The status of a task handle.
+
+ .. attribute:: PENDING
+
+ The task has not finished yet.
+ .. attribute:: FINISHED
+
+ The task has finished with a return value.
+ .. attribute:: CANCELLING
+
+ The task has been cancelled but has not finished yet.
+ .. attribute:: CANCELLED
+
+ The task was cancelled and has finished since.
+ .. attribute:: FAILED
+
+ The task raised an exception.
+ """
+
+ PENDING = auto()
+ FINISHED = auto()
+ CANCELLING = auto()
+ CANCELLED = auto()
+ FAILED = auto()
+
+ __slots__ = (
+ "__weakref__",
+ "_coro",
+ "_name",
+ "_cancel_scope",
+ "_finished_event",
+ "_return_value",
+ "_start_value",
+ "_exception",
+ )
+
+ _return_value: T_co
+ _start_value: T_startval
+
+ def __init__(self, coro: Coroutine[Any, Any, T_co], name: object) -> None:
+ from ._synchronization import Event
+
+ self._coro = coro
+ self._cancel_scope = CancelScope()
+ self._finished_event = Event()
+ self._exception: BaseException | None = None
+
+ if name is not None:
+ self._name = str(name)
+ elif iscoroutine(coro):
+ self._name = coro.__qualname__
+ else:
+ self._name = str(coro) # coroutine-like object (e.g. asend() objects)
+
+ async def _run_coro(self) -> None:
+ __tracebackhide__ = True
+
+ with self._cancel_scope:
+ try:
+ retval = await self._coro
+ except BaseException as exc:
+ self._exception = exc
+ raise
+ else:
+ self._return_value = retval
+ finally:
+ self._finished_event.set()
+ del self # Break the reference cycle
+
+ def cancel(self) -> None:
+ """
+ Set the task to a cancelled state.
+
+ This will interrupt any interruptible asynchronous operation, and will cause
+ any further awaits on this task to get immediately cancelled, unless done in
+ a shielded cancel scope.
+
+ If the task has already finished, this method has no effect.
+ """
+ if not self._finished_event.is_set():
+ self._cancel_scope.cancel()
+
+ @property
+ def coro(self) -> Coroutine[Any, Any, T_co]:
+ """
+ The coroutine object that was passed to one of the task-spawning methods in
+ :class:`TaskGroup`.
+ """
+ return self._coro
+
+ @property
+ def status(self) -> TaskHandle.Status:
+ """
+ The current status of the task.
+
+ Every task starts in the :attr:`~TaskHandle.Status.PENDING` state.
+ If a task is cancelled while in this state, it will transition to the
+ :attr:`~TaskHandle.Status.CANCELLING` state. When the task finishes, it will
+ transition to one of the three final states (
+ :attr:`~TaskHandle.Status.FINISHED`, :attr:`~TaskHandle.Status.FAILED`, or
+ :attr:`~TaskHandle.Status.CANCELLING`) depending on the exception the task
+ raised, if any. No other status transitions will happen.
+ """
+ if not self._finished_event.is_set():
+ if self._cancel_scope.cancel_called:
+ return TaskHandle.Status.CANCELLING
+ else:
+ return TaskHandle.Status.PENDING
+ elif self._exception is not None:
+ if isinstance(self._exception, get_cancelled_exc_class()):
+ return TaskHandle.Status.CANCELLED
+ else:
+ return TaskHandle.Status.FAILED
+ else:
+ return TaskHandle.Status.FINISHED
+
+ @property
+ def name(self) -> str:
+ """The name of the task."""
+ return self._name
+
+ @property
+ def exception(self) -> BaseException | None:
+ """
+ The exception raised by the task, or ``None`` if it finished without raising.
+
+ :raises TaskNotFinished: if the task has not finished yet
+ :raises TaskCancelled: if the task was cancelled
+
+ """
+ match self.status:
+ case TaskHandle.Status.PENDING:
+ raise TaskNotFinished("the task has not finished yet")
+ case TaskHandle.Status.FINISHED:
+ return None
+ case TaskHandle.Status.CANCELLING:
+ raise TaskCancelled("the task was cancelled")
+ case TaskHandle.Status.CANCELLED:
+ raise TaskCancelled("the task was cancelled") from self._exception
+ case TaskHandle.Status.FAILED:
+ return self._exception
+
+ @property
+ def return_value(self) -> T_co:
+ """
+ The return value of the task.
+
+ :raises TaskNotFinished: if the task has not finished yet
+ :raises TaskCancelled: if the task was cancelled
+ :raises TaskFailed: if the task raised an exception
+
+ """
+ match self.status:
+ case TaskHandle.Status.PENDING:
+ raise TaskNotFinished("the task has not finished yet")
+ case TaskHandle.Status.FINISHED:
+ return self._return_value
+ case TaskHandle.Status.CANCELLING:
+ raise TaskCancelled("the task was cancelled")
+ case TaskHandle.Status.CANCELLED:
+ raise TaskCancelled("the task was cancelled") from self._exception
+ case TaskHandle.Status.FAILED:
+ raise TaskFailed("the task raised an exception") from self._exception
+
+ @property
+ def start_value(self) -> T_startval:
+ """
+ The value passed to :meth:`task_status.started() <.abc.TaskStatus.started>`,
+
+ :raises RuntimeError: if the task was not started with :meth:`TaskGroup.start()
+ <.abc.TaskGroup.start>`
+ """
+ try:
+ return self._start_value
+ except AttributeError:
+ raise RuntimeError(
+ "the task was not started with TaskGroup.start()"
+ ) from None
+
+ async def wait(self) -> None:
+ """
+ Wait for the task to finish.
+
+ This method will return as soon as the task has finished, no matter how it
+ happened.
+ """
+ await self._finished_event.wait()
+
+ def __await__(self) -> Generator[Any, Any, T_co]:
+ yield from self._finished_event.wait().__await__()
+ return self.return_value
+
+ def __repr__(self) -> str:
+ return (
+ f"<{self.__class__.__name__} {self.status.name.lower()} "
+ f"name={self._name!r} coro={self._coro!r}>"
+ )
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_tempfile.py b/venv/lib/python3.11/site-packages/anyio/_core/_tempfile.py
new file mode 100644
index 0000000000000000000000000000000000000000..75a09f793744b8e60375ce2efab98307d077bc21
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_tempfile.py
@@ -0,0 +1,613 @@
+from __future__ import annotations
+
+import os
+import sys
+import tempfile
+from collections.abc import Iterable
+from io import BytesIO, TextIOWrapper
+from types import TracebackType
+from typing import (
+ TYPE_CHECKING,
+ Any,
+ AnyStr,
+ Generic,
+ overload,
+)
+
+from .. import to_thread
+from .._core._fileio import AsyncFile
+from ..lowlevel import checkpoint_if_cancelled
+
+if TYPE_CHECKING:
+ from _typeshed import OpenBinaryMode, OpenTextMode, ReadableBuffer, WriteableBuffer
+
+
+class TemporaryFile(Generic[AnyStr]):
+ """
+ An asynchronous temporary file that is automatically created and cleaned up.
+
+ This class provides an asynchronous context manager interface to a temporary file.
+ The file is created using Python's standard `tempfile.TemporaryFile` function in a
+ background thread, and is wrapped as an asynchronous file using `AsyncFile`.
+
+ :param mode: The mode in which the file is opened. Defaults to "w+b".
+ :param buffering: The buffering policy (-1 means the default buffering).
+ :param encoding: The encoding used to decode or encode the file. Only applicable in
+ text mode.
+ :param newline: Controls how universal newlines mode works (only applicable in text
+ mode).
+ :param suffix: The suffix for the temporary file name.
+ :param prefix: The prefix for the temporary file name.
+ :param dir: The directory in which the temporary file is created.
+ :param errors: The error handling scheme used for encoding/decoding errors.
+ """
+
+ _async_file: AsyncFile[AnyStr]
+
+ @overload
+ def __init__(
+ self: TemporaryFile[bytes],
+ mode: OpenBinaryMode = ...,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ newline: str | None = ...,
+ suffix: str | None = ...,
+ prefix: str | None = ...,
+ dir: str | None = ...,
+ *,
+ errors: str | None = ...,
+ ): ...
+ @overload
+ def __init__(
+ self: TemporaryFile[str],
+ mode: OpenTextMode,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ newline: str | None = ...,
+ suffix: str | None = ...,
+ prefix: str | None = ...,
+ dir: str | None = ...,
+ *,
+ errors: str | None = ...,
+ ): ...
+
+ def __init__(
+ self,
+ mode: OpenTextMode | OpenBinaryMode = "w+b",
+ buffering: int = -1,
+ encoding: str | None = None,
+ newline: str | None = None,
+ suffix: str | None = None,
+ prefix: str | None = None,
+ dir: str | None = None,
+ *,
+ errors: str | None = None,
+ ) -> None:
+ self.mode = mode
+ self.buffering = buffering
+ self.encoding = encoding
+ self.newline = newline
+ self.suffix: str | None = suffix
+ self.prefix: str | None = prefix
+ self.dir: str | None = dir
+ self.errors = errors
+
+ async def __aenter__(self) -> AsyncFile[AnyStr]:
+ fp = await to_thread.run_sync(
+ lambda: tempfile.TemporaryFile(
+ self.mode,
+ self.buffering,
+ self.encoding,
+ self.newline,
+ self.suffix,
+ self.prefix,
+ self.dir,
+ errors=self.errors,
+ )
+ )
+ self._async_file = AsyncFile(fp)
+ return self._async_file
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_value: BaseException | None,
+ traceback: TracebackType | None,
+ ) -> None:
+ await self._async_file.aclose()
+
+
+class NamedTemporaryFile(Generic[AnyStr]):
+ """
+ An asynchronous named temporary file that is automatically created and cleaned up.
+
+ This class provides an asynchronous context manager for a temporary file with a
+ visible name in the file system. It uses Python's standard
+ :func:`~tempfile.NamedTemporaryFile` function and wraps the file object with
+ :class:`AsyncFile` for asynchronous operations.
+
+ :param mode: The mode in which the file is opened. Defaults to "w+b".
+ :param buffering: The buffering policy (-1 means the default buffering).
+ :param encoding: The encoding used to decode or encode the file. Only applicable in
+ text mode.
+ :param newline: Controls how universal newlines mode works (only applicable in text
+ mode).
+ :param suffix: The suffix for the temporary file name.
+ :param prefix: The prefix for the temporary file name.
+ :param dir: The directory in which the temporary file is created.
+ :param delete: Whether to delete the file when it is closed.
+ :param errors: The error handling scheme used for encoding/decoding errors.
+ :param delete_on_close: (Python 3.12+) Whether to delete the file on close.
+ """
+
+ _async_file: AsyncFile[AnyStr]
+
+ @overload
+ def __init__(
+ self: NamedTemporaryFile[bytes],
+ mode: OpenBinaryMode = ...,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ newline: str | None = ...,
+ suffix: str | None = ...,
+ prefix: str | None = ...,
+ dir: str | None = ...,
+ delete: bool = ...,
+ *,
+ errors: str | None = ...,
+ delete_on_close: bool = ...,
+ ): ...
+ @overload
+ def __init__(
+ self: NamedTemporaryFile[str],
+ mode: OpenTextMode,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ newline: str | None = ...,
+ suffix: str | None = ...,
+ prefix: str | None = ...,
+ dir: str | None = ...,
+ delete: bool = ...,
+ *,
+ errors: str | None = ...,
+ delete_on_close: bool = ...,
+ ): ...
+
+ def __init__(
+ self,
+ mode: OpenBinaryMode | OpenTextMode = "w+b",
+ buffering: int = -1,
+ encoding: str | None = None,
+ newline: str | None = None,
+ suffix: str | None = None,
+ prefix: str | None = None,
+ dir: str | None = None,
+ delete: bool = True,
+ *,
+ errors: str | None = None,
+ delete_on_close: bool = True,
+ ) -> None:
+ self._params: dict[str, Any] = {
+ "mode": mode,
+ "buffering": buffering,
+ "encoding": encoding,
+ "newline": newline,
+ "suffix": suffix,
+ "prefix": prefix,
+ "dir": dir,
+ "delete": delete,
+ "errors": errors,
+ }
+ if sys.version_info >= (3, 12):
+ self._params["delete_on_close"] = delete_on_close
+
+ async def __aenter__(self) -> AsyncFile[AnyStr]:
+ fp = await to_thread.run_sync(
+ lambda: tempfile.NamedTemporaryFile(**self._params)
+ )
+ self._async_file = AsyncFile(fp)
+ return self._async_file
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_value: BaseException | None,
+ traceback: TracebackType | None,
+ ) -> None:
+ await self._async_file.aclose()
+
+
+class SpooledTemporaryFile(AsyncFile[AnyStr]):
+ """
+ An asynchronous spooled temporary file that starts in memory and is spooled to disk.
+
+ This class provides an asynchronous interface to a spooled temporary file, much like
+ Python's standard :class:`~tempfile.SpooledTemporaryFile`. It supports asynchronous
+ write operations and provides a method to force a rollover to disk.
+
+ :param max_size: Maximum size in bytes before the file is rolled over to disk.
+ :param mode: The mode in which the file is opened. Defaults to "w+b".
+ :param buffering: The buffering policy (-1 means the default buffering).
+ :param encoding: The encoding used to decode or encode the file (text mode only).
+ :param newline: Controls how universal newlines mode works (text mode only).
+ :param suffix: The suffix for the temporary file name.
+ :param prefix: The prefix for the temporary file name.
+ :param dir: The directory in which the temporary file is created.
+ :param errors: The error handling scheme used for encoding/decoding errors.
+ """
+
+ _rolled: bool = False
+
+ @overload
+ def __init__(
+ self: SpooledTemporaryFile[bytes],
+ max_size: int = ...,
+ mode: OpenBinaryMode = ...,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ newline: str | None = ...,
+ suffix: str | None = ...,
+ prefix: str | None = ...,
+ dir: str | None = ...,
+ *,
+ errors: str | None = ...,
+ ): ...
+ @overload
+ def __init__(
+ self: SpooledTemporaryFile[str],
+ max_size: int = ...,
+ mode: OpenTextMode = ...,
+ buffering: int = ...,
+ encoding: str | None = ...,
+ newline: str | None = ...,
+ suffix: str | None = ...,
+ prefix: str | None = ...,
+ dir: str | None = ...,
+ *,
+ errors: str | None = ...,
+ ): ...
+
+ def __init__(
+ self,
+ max_size: int = 0,
+ mode: OpenBinaryMode | OpenTextMode = "w+b",
+ buffering: int = -1,
+ encoding: str | None = None,
+ newline: str | None = None,
+ suffix: str | None = None,
+ prefix: str | None = None,
+ dir: str | None = None,
+ *,
+ errors: str | None = None,
+ ) -> None:
+ self._tempfile_params: dict[str, Any] = {
+ "mode": mode,
+ "buffering": buffering,
+ "encoding": encoding,
+ "newline": newline,
+ "suffix": suffix,
+ "prefix": prefix,
+ "dir": dir,
+ "errors": errors,
+ }
+ self._max_size = max_size
+ if "b" in mode:
+ super().__init__(BytesIO()) # type: ignore[arg-type]
+ else:
+ super().__init__(
+ TextIOWrapper( # type: ignore[arg-type]
+ BytesIO(),
+ encoding=encoding,
+ errors=errors,
+ newline=newline,
+ write_through=True,
+ )
+ )
+
+ async def aclose(self) -> None:
+ if not self._rolled:
+ self._fp.close()
+ return
+
+ await super().aclose()
+
+ async def _check(self) -> None:
+ if self._rolled or self._fp.tell() <= self._max_size:
+ return
+
+ await self.rollover()
+
+ async def rollover(self) -> None:
+ if self._rolled:
+ return
+
+ self._rolled = True
+ buffer = self._fp
+ buffer.seek(0)
+ self._fp = await to_thread.run_sync(
+ lambda: tempfile.TemporaryFile(**self._tempfile_params)
+ )
+ await self.write(buffer.read())
+ buffer.close()
+
+ @property
+ def closed(self) -> bool:
+ return self._fp.closed
+
+ async def read(self, size: int = -1) -> AnyStr:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ return self._fp.read(size)
+
+ return await super().read(size) # type: ignore[return-value]
+
+ async def read1(self: SpooledTemporaryFile[bytes], size: int = -1) -> bytes:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ return self._fp.read1(size)
+
+ return await super().read1(size)
+
+ async def readline(self) -> AnyStr:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ return self._fp.readline()
+
+ return await super().readline() # type: ignore[return-value]
+
+ async def readlines(self) -> list[AnyStr]:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ return self._fp.readlines()
+
+ return await super().readlines() # type: ignore[return-value]
+
+ async def readinto(self: SpooledTemporaryFile[bytes], b: WriteableBuffer) -> int:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ self._fp.readinto(b)
+
+ return await super().readinto(b)
+
+ async def readinto1(self: SpooledTemporaryFile[bytes], b: WriteableBuffer) -> int:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ self._fp.readinto(b)
+
+ return await super().readinto1(b)
+
+ async def seek(self, offset: int, whence: int | None = os.SEEK_SET) -> int:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ return self._fp.seek(offset, whence)
+
+ return await super().seek(offset, whence)
+
+ async def tell(self) -> int:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ return self._fp.tell()
+
+ return await super().tell()
+
+ async def truncate(self, size: int | None = None) -> int:
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ return self._fp.truncate(size)
+
+ return await super().truncate(size)
+
+ @overload
+ async def write(self: SpooledTemporaryFile[bytes], b: ReadableBuffer) -> int: ...
+ @overload
+ async def write(self: SpooledTemporaryFile[str], b: str) -> int: ...
+
+ async def write(self, b: ReadableBuffer | str) -> int:
+ """
+ Asynchronously write data to the spooled temporary file.
+
+ If the file has not yet been rolled over, the data is written synchronously,
+ and a rollover is triggered if the size exceeds the maximum size.
+
+ :param s: The data to write.
+ :return: The number of bytes written.
+ :raises RuntimeError: If the underlying file is not initialized.
+
+ """
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ result = self._fp.write(b)
+ await self._check()
+ return result
+
+ return await super().write(b) # type: ignore[misc]
+
+ @overload
+ async def writelines(
+ self: SpooledTemporaryFile[bytes], lines: Iterable[ReadableBuffer]
+ ) -> None: ...
+ @overload
+ async def writelines(
+ self: SpooledTemporaryFile[str], lines: Iterable[str]
+ ) -> None: ...
+
+ async def writelines(self, lines: Iterable[str] | Iterable[ReadableBuffer]) -> None:
+ """
+ Asynchronously write a list of lines to the spooled temporary file.
+
+ If the file has not yet been rolled over, the lines are written synchronously,
+ and a rollover is triggered if the size exceeds the maximum size.
+
+ :param lines: An iterable of lines to write.
+ :raises RuntimeError: If the underlying file is not initialized.
+
+ """
+ if not self._rolled:
+ await checkpoint_if_cancelled()
+ result = self._fp.writelines(lines)
+ await self._check()
+ return result
+
+ return await super().writelines(lines) # type: ignore[misc]
+
+
+class TemporaryDirectory(Generic[AnyStr]):
+ """
+ An asynchronous temporary directory that is created and cleaned up automatically.
+
+ This class provides an asynchronous context manager for creating a temporary
+ directory. It wraps Python's standard :class:`~tempfile.TemporaryDirectory` to
+ perform directory creation and cleanup operations in a background thread.
+
+ :param suffix: Suffix to be added to the temporary directory name.
+ :param prefix: Prefix to be added to the temporary directory name.
+ :param dir: The parent directory where the temporary directory is created.
+ :param ignore_cleanup_errors: Whether to ignore errors during cleanup
+ :param delete: Whether to delete the directory upon closing (Python 3.12+).
+ """
+
+ def __init__(
+ self,
+ suffix: AnyStr | None = None,
+ prefix: AnyStr | None = None,
+ dir: AnyStr | None = None,
+ *,
+ ignore_cleanup_errors: bool = False,
+ delete: bool = True,
+ ) -> None:
+ self.suffix: AnyStr | None = suffix
+ self.prefix: AnyStr | None = prefix
+ self.dir: AnyStr | None = dir
+ self.ignore_cleanup_errors = ignore_cleanup_errors
+ self.delete = delete
+
+ self._tempdir: tempfile.TemporaryDirectory | None = None
+
+ async def __aenter__(self) -> str:
+ params: dict[str, Any] = {
+ "suffix": self.suffix,
+ "prefix": self.prefix,
+ "dir": self.dir,
+ "ignore_cleanup_errors": self.ignore_cleanup_errors,
+ }
+ if sys.version_info >= (3, 12):
+ params["delete"] = self.delete
+
+ self._tempdir = await to_thread.run_sync(
+ lambda: tempfile.TemporaryDirectory(**params)
+ )
+ return await to_thread.run_sync(self._tempdir.__enter__)
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_value: BaseException | None,
+ traceback: TracebackType | None,
+ ) -> None:
+ if self._tempdir is not None:
+ await to_thread.run_sync(
+ self._tempdir.__exit__, exc_type, exc_value, traceback
+ )
+
+ async def cleanup(self) -> None:
+ if self._tempdir is not None:
+ await to_thread.run_sync(self._tempdir.cleanup)
+
+
+@overload
+async def mkstemp(
+ suffix: str | None = None,
+ prefix: str | None = None,
+ dir: str | None = None,
+ text: bool = False,
+) -> tuple[int, str]: ...
+
+
+@overload
+async def mkstemp(
+ suffix: bytes | None = None,
+ prefix: bytes | None = None,
+ dir: bytes | None = None,
+ text: bool = False,
+) -> tuple[int, bytes]: ...
+
+
+async def mkstemp(
+ suffix: AnyStr | None = None,
+ prefix: AnyStr | None = None,
+ dir: AnyStr | None = None,
+ text: bool = False,
+) -> tuple[int, str | bytes]:
+ """
+ Asynchronously create a temporary file and return an OS-level handle and the file
+ name.
+
+ This function wraps `tempfile.mkstemp` and executes it in a background thread.
+
+ :param suffix: Suffix to be added to the file name.
+ :param prefix: Prefix to be added to the file name.
+ :param dir: Directory in which the temporary file is created.
+ :param text: Whether the file is opened in text mode.
+ :return: A tuple containing the file descriptor and the file name.
+
+ """
+ return await to_thread.run_sync(tempfile.mkstemp, suffix, prefix, dir, text)
+
+
+@overload
+async def mkdtemp(
+ suffix: str | None = None,
+ prefix: str | None = None,
+ dir: str | None = None,
+) -> str: ...
+
+
+@overload
+async def mkdtemp(
+ suffix: bytes | None = None,
+ prefix: bytes | None = None,
+ dir: bytes | None = None,
+) -> bytes: ...
+
+
+async def mkdtemp(
+ suffix: AnyStr | None = None,
+ prefix: AnyStr | None = None,
+ dir: AnyStr | None = None,
+) -> str | bytes:
+ """
+ Asynchronously create a temporary directory and return its path.
+
+ This function wraps `tempfile.mkdtemp` and executes it in a background thread.
+
+ :param suffix: Suffix to be added to the directory name.
+ :param prefix: Prefix to be added to the directory name.
+ :param dir: Parent directory where the temporary directory is created.
+ :return: The path of the created temporary directory.
+
+ """
+ return await to_thread.run_sync(tempfile.mkdtemp, suffix, prefix, dir)
+
+
+async def gettempdir() -> str:
+ """
+ Asynchronously return the name of the directory used for temporary files.
+
+ This function wraps `tempfile.gettempdir` and executes it in a background thread.
+
+ :return: The path of the temporary directory as a string.
+
+ """
+ return await to_thread.run_sync(tempfile.gettempdir)
+
+
+async def gettempdirb() -> bytes:
+ """
+ Asynchronously return the name of the directory used for temporary files in bytes.
+
+ This function wraps `tempfile.gettempdirb` and executes it in a background thread.
+
+ :return: The path of the temporary directory as bytes.
+
+ """
+ return await to_thread.run_sync(tempfile.gettempdirb)
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_testing.py b/venv/lib/python3.11/site-packages/anyio/_core/_testing.py
new file mode 100644
index 0000000000000000000000000000000000000000..369e65c068a426e99b7e8571209e80ce35b71f47
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_testing.py
@@ -0,0 +1,82 @@
+from __future__ import annotations
+
+from collections.abc import Awaitable, Generator
+from typing import Any, cast
+
+from ._eventloop import get_async_backend
+
+
+class TaskInfo:
+ """
+ Represents an asynchronous task.
+
+ :ivar int id: the unique identifier of the task
+ :ivar parent_id: the identifier of the parent task, if any
+ :vartype parent_id: Optional[int]
+ :ivar str name: the description of the task (if any)
+ :ivar ~collections.abc.Coroutine coro: the coroutine object of the task
+ """
+
+ __slots__ = "_name", "id", "parent_id", "name", "coro"
+
+ def __init__(
+ self,
+ id: int,
+ parent_id: int | None,
+ name: str | None,
+ coro: Generator[Any, Any, Any] | Awaitable[Any],
+ ):
+ func = get_current_task
+ self._name = f"{func.__module__}.{func.__qualname__}"
+ self.id: int = id
+ self.parent_id: int | None = parent_id
+ self.name: str | None = name
+ self.coro: Generator[Any, Any, Any] | Awaitable[Any] = coro
+
+ def __eq__(self, other: object) -> bool:
+ if isinstance(other, TaskInfo):
+ return self.id == other.id
+
+ return NotImplemented
+
+ def __hash__(self) -> int:
+ return hash(self.id)
+
+ def __repr__(self) -> str:
+ return f"{self.__class__.__name__}(id={self.id!r}, name={self.name!r})"
+
+ def has_pending_cancellation(self) -> bool:
+ """
+ Return ``True`` if the task has a cancellation pending, ``False`` otherwise.
+
+ """
+ return False
+
+
+def get_current_task() -> TaskInfo:
+ """
+ Return the current task.
+
+ :return: a representation of the current task
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().get_current_task()
+
+
+def get_running_tasks() -> list[TaskInfo]:
+ """
+ Return a list of running tasks in the current event loop.
+
+ :return: a list of task info objects
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return cast("list[TaskInfo]", get_async_backend().get_running_tasks())
+
+
+async def wait_all_tasks_blocked() -> None:
+ """Wait until all other tasks are waiting for something."""
+ await get_async_backend().wait_all_tasks_blocked()
diff --git a/venv/lib/python3.11/site-packages/anyio/_core/_typedattr.py b/venv/lib/python3.11/site-packages/anyio/_core/_typedattr.py
new file mode 100644
index 0000000000000000000000000000000000000000..f358a448cb12739fd4eda4f4859d3a24ddd1de63
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/_core/_typedattr.py
@@ -0,0 +1,81 @@
+from __future__ import annotations
+
+from collections.abc import Callable, Mapping
+from typing import Any, TypeVar, final, overload
+
+from ._exceptions import TypedAttributeLookupError
+
+T_Attr = TypeVar("T_Attr")
+T_Default = TypeVar("T_Default")
+undefined = object()
+
+
+def typed_attribute() -> Any:
+ """Return a unique object, used to mark typed attributes."""
+ return object()
+
+
+class TypedAttributeSet:
+ """
+ Superclass for typed attribute collections.
+
+ Checks that every public attribute of every subclass has a type annotation.
+ """
+
+ def __init_subclass__(cls) -> None:
+ annotations: dict[str, Any] = getattr(cls, "__annotations__", {})
+ for attrname in dir(cls):
+ if not attrname.startswith("_") and attrname not in annotations:
+ raise TypeError(
+ f"Attribute {attrname!r} is missing its type annotation"
+ )
+
+ super().__init_subclass__()
+
+
+class TypedAttributeProvider:
+ """Base class for classes that wish to provide typed extra attributes."""
+
+ @property
+ def extra_attributes(self) -> Mapping[T_Attr, Callable[[], T_Attr]]:
+ """
+ A mapping of the extra attributes to callables that return the corresponding
+ values.
+
+ If the provider wraps another provider, the attributes from that wrapper should
+ also be included in the returned mapping (but the wrapper may override the
+ callables from the wrapped instance).
+
+ """
+ return {}
+
+ @overload
+ def extra(self, attribute: T_Attr) -> T_Attr: ...
+
+ @overload
+ def extra(self, attribute: T_Attr, default: T_Default) -> T_Attr | T_Default: ...
+
+ @final
+ def extra(self, attribute: Any, default: object = undefined) -> object:
+ """
+ extra(attribute, default=undefined)
+
+ Return the value of the given typed extra attribute.
+
+ :param attribute: the attribute (member of a :class:`~TypedAttributeSet`) to
+ look for
+ :param default: the value that should be returned if no value is found for the
+ attribute
+ :raises ~anyio.TypedAttributeLookupError: if the search failed and no default
+ value was given
+
+ """
+ try:
+ getter = self.extra_attributes[attribute]
+ except KeyError:
+ if default is undefined:
+ raise TypedAttributeLookupError("Attribute not found") from None
+ else:
+ return default
+
+ return getter()
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/__init__.py b/venv/lib/python3.11/site-packages/anyio/abc/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..d560ce3f1fa45a7ee4a3bc8958aa59702caa9d0c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/__init__.py
@@ -0,0 +1,58 @@
+from __future__ import annotations
+
+from ._eventloop import AsyncBackend as AsyncBackend
+from ._resources import AsyncResource as AsyncResource
+from ._sockets import ConnectedUDPSocket as ConnectedUDPSocket
+from ._sockets import ConnectedUNIXDatagramSocket as ConnectedUNIXDatagramSocket
+from ._sockets import IPAddressType as IPAddressType
+from ._sockets import IPSockAddrType as IPSockAddrType
+from ._sockets import SocketAttribute as SocketAttribute
+from ._sockets import SocketListener as SocketListener
+from ._sockets import SocketStream as SocketStream
+from ._sockets import UDPPacketType as UDPPacketType
+from ._sockets import UDPSocket as UDPSocket
+from ._sockets import UNIXDatagramPacketType as UNIXDatagramPacketType
+from ._sockets import UNIXDatagramSocket as UNIXDatagramSocket
+from ._sockets import UNIXSocketStream as UNIXSocketStream
+from ._streams import AnyByteReceiveStream as AnyByteReceiveStream
+from ._streams import AnyByteSendStream as AnyByteSendStream
+from ._streams import AnyByteStream as AnyByteStream
+from ._streams import AnyByteStreamConnectable as AnyByteStreamConnectable
+from ._streams import AnyUnreliableByteReceiveStream as AnyUnreliableByteReceiveStream
+from ._streams import AnyUnreliableByteSendStream as AnyUnreliableByteSendStream
+from ._streams import AnyUnreliableByteStream as AnyUnreliableByteStream
+from ._streams import ByteReceiveStream as ByteReceiveStream
+from ._streams import ByteSendStream as ByteSendStream
+from ._streams import ByteStream as ByteStream
+from ._streams import ByteStreamConnectable as ByteStreamConnectable
+from ._streams import Listener as Listener
+from ._streams import ObjectReceiveStream as ObjectReceiveStream
+from ._streams import ObjectSendStream as ObjectSendStream
+from ._streams import ObjectStream as ObjectStream
+from ._streams import ObjectStreamConnectable as ObjectStreamConnectable
+from ._streams import UnreliableObjectReceiveStream as UnreliableObjectReceiveStream
+from ._streams import UnreliableObjectSendStream as UnreliableObjectSendStream
+from ._streams import UnreliableObjectStream as UnreliableObjectStream
+from ._subprocesses import Process as Process
+from ._tasks import TaskGroup as TaskGroup
+from ._tasks import TaskStatus as TaskStatus
+from ._testing import TestRunner as TestRunner
+
+# Re-exported here, for backwards compatibility
+# isort: off
+from .._core._synchronization import (
+ CapacityLimiter as CapacityLimiter,
+ Condition as Condition,
+ Event as Event,
+ Lock as Lock,
+ Semaphore as Semaphore,
+)
+from .._core._tasks import CancelScope as CancelScope
+from ..from_thread import BlockingPortal as BlockingPortal
+
+# Re-export imports so they look like they live directly in this package
+for __value in list(locals().values()):
+ if getattr(__value, "__module__", "").startswith("anyio.abc."):
+ __value.__module__ = __name__
+
+del __value
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/_eventloop.py b/venv/lib/python3.11/site-packages/anyio/abc/_eventloop.py
new file mode 100644
index 0000000000000000000000000000000000000000..cad3fa76370ff573469c1c52fd82d2eff0bc83eb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/_eventloop.py
@@ -0,0 +1,410 @@
+from __future__ import annotations
+
+import math
+import sys
+from abc import ABCMeta, abstractmethod
+from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine, Sequence
+from contextlib import AbstractContextManager
+from os import PathLike
+from signal import Signals
+from socket import AddressFamily, SocketKind, socket
+from typing import (
+ IO,
+ TYPE_CHECKING,
+ Any,
+ TypeAlias,
+ TypeVar,
+ overload,
+)
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+if TYPE_CHECKING:
+ from _typeshed import FileDescriptorLike
+
+ from .._core._synchronization import CapacityLimiter, Event, Lock, Semaphore
+ from .._core._tasks import CancelScope
+ from .._core._testing import TaskInfo
+ from ._sockets import (
+ ConnectedUDPSocket,
+ ConnectedUNIXDatagramSocket,
+ IPSockAddrType,
+ SocketListener,
+ SocketStream,
+ UDPSocket,
+ UNIXDatagramSocket,
+ UNIXSocketStream,
+ )
+ from ._subprocesses import Process
+ from ._tasks import TaskGroup
+ from ._testing import TestRunner
+
+T_Retval = TypeVar("T_Retval")
+T_co = TypeVar("T_co", covariant=True)
+PosArgsT = TypeVarTuple("PosArgsT")
+StrOrBytesPath: TypeAlias = str | bytes | PathLike[str] | PathLike[bytes]
+
+
+class AsyncBackend(metaclass=ABCMeta):
+ @classmethod
+ @abstractmethod
+ def run(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
+ args: tuple[Unpack[PosArgsT]],
+ kwargs: dict[str, Any],
+ options: dict[str, Any],
+ ) -> T_Retval:
+ """
+ Run the given coroutine function in an asynchronous event loop.
+
+ The current thread must not be already running an event loop.
+
+ :param func: a coroutine function
+ :param args: positional arguments to ``func``
+ :param kwargs: positional arguments to ``func``
+ :param options: keyword arguments to call the backend ``run()`` implementation
+ with
+ :return: the return value of the coroutine function
+ """
+
+ @classmethod
+ @abstractmethod
+ def current_token(cls) -> object:
+ """
+ Return an object that allows other threads to run code inside the event loop.
+
+ :return: a token object, specific to the event loop running in the current
+ thread
+ """
+
+ @classmethod
+ @abstractmethod
+ def current_time(cls) -> float:
+ """
+ Return the current value of the event loop's internal clock.
+
+ :return: the clock value (seconds)
+ """
+
+ @classmethod
+ @abstractmethod
+ def cancelled_exception_class(cls) -> type[BaseException]:
+ """Return the exception class that is raised in a task if it's cancelled."""
+
+ @classmethod
+ @abstractmethod
+ async def checkpoint(cls) -> None:
+ """
+ Check if the task has been cancelled, and allow rescheduling of other tasks.
+
+ This is effectively the same as running :meth:`checkpoint_if_cancelled` and then
+ :meth:`cancel_shielded_checkpoint`.
+ """
+
+ @classmethod
+ async def checkpoint_if_cancelled(cls) -> None:
+ """
+ Check if the current task group has been cancelled.
+
+ This will check if the task has been cancelled, but will not allow other tasks
+ to be scheduled if not.
+
+ """
+ if cls.current_effective_deadline() == -math.inf:
+ await cls.checkpoint()
+
+ @classmethod
+ async def cancel_shielded_checkpoint(cls) -> None:
+ """
+ Allow the rescheduling of other tasks.
+
+ This will give other tasks the opportunity to run, but without checking if the
+ current task group has been cancelled, unlike with :meth:`checkpoint`.
+
+ """
+ with cls.create_cancel_scope(shield=True):
+ await cls.sleep(0)
+
+ @classmethod
+ @abstractmethod
+ async def sleep(cls, delay: float) -> None:
+ """
+ Pause the current task for the specified duration.
+
+ :param delay: the duration, in seconds
+ """
+
+ @classmethod
+ @abstractmethod
+ def create_cancel_scope(
+ cls, *, deadline: float = math.inf, shield: bool = False
+ ) -> CancelScope:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def current_effective_deadline(cls) -> float:
+ """
+ Return the nearest deadline among all the cancel scopes effective for the
+ current task.
+
+ :return:
+ - a clock value from the event loop's internal clock
+ - ``inf`` if there is no deadline in effect
+ - ``-inf`` if the current scope has been cancelled
+ :rtype: float
+ """
+
+ @classmethod
+ @abstractmethod
+ def create_task_group(cls) -> TaskGroup:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def create_event(cls) -> Event:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def create_lock(cls, *, fast_acquire: bool) -> Lock:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def create_semaphore(
+ cls,
+ initial_value: int,
+ *,
+ max_value: int | None = None,
+ fast_acquire: bool = False,
+ ) -> Semaphore:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def create_capacity_limiter(cls, total_tokens: float) -> CapacityLimiter:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def run_sync_in_worker_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ abandon_on_cancel: bool = False,
+ limiter: CapacityLimiter | None = None,
+ ) -> T_Retval:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def check_cancelled(cls) -> None:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def run_async_from_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ args: tuple[Unpack[PosArgsT]],
+ token: object,
+ ) -> T_co:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def run_sync_from_thread(
+ cls,
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ token: object,
+ ) -> T_Retval:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def open_process(
+ cls,
+ command: StrOrBytesPath | Sequence[StrOrBytesPath],
+ *,
+ stdin: int | IO[Any] | None,
+ stdout: int | IO[Any] | None,
+ stderr: int | IO[Any] | None,
+ **kwargs: Any,
+ ) -> Process:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def setup_process_pool_exit_at_shutdown(cls, workers: set[Process]) -> None:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def connect_tcp(
+ cls, host: str, port: int, local_address: IPSockAddrType | None = None
+ ) -> SocketStream:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def connect_unix(cls, path: str | bytes) -> UNIXSocketStream:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def create_tcp_listener(cls, sock: socket) -> SocketListener:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def create_unix_listener(cls, sock: socket) -> SocketListener:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def create_udp_socket(
+ cls,
+ family: AddressFamily,
+ local_address: IPSockAddrType | None,
+ remote_address: IPSockAddrType | None,
+ reuse_port: bool,
+ ) -> UDPSocket | ConnectedUDPSocket:
+ pass
+
+ @classmethod
+ @overload
+ async def create_unix_datagram_socket(
+ cls, raw_socket: socket, remote_path: None
+ ) -> UNIXDatagramSocket: ...
+
+ @classmethod
+ @overload
+ async def create_unix_datagram_socket(
+ cls, raw_socket: socket, remote_path: str | bytes
+ ) -> ConnectedUNIXDatagramSocket: ...
+
+ @classmethod
+ @abstractmethod
+ async def create_unix_datagram_socket(
+ cls, raw_socket: socket, remote_path: str | bytes | None
+ ) -> UNIXDatagramSocket | ConnectedUNIXDatagramSocket:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def getaddrinfo(
+ cls,
+ host: bytes | str | None,
+ port: str | int | None,
+ *,
+ family: int | AddressFamily = 0,
+ type: int | SocketKind = 0,
+ proto: int = 0,
+ flags: int = 0,
+ ) -> Sequence[
+ tuple[
+ AddressFamily,
+ SocketKind,
+ int,
+ str,
+ tuple[str, int] | tuple[str, int, int, int] | tuple[int, bytes],
+ ]
+ ]:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def getnameinfo(
+ cls, sockaddr: IPSockAddrType, flags: int = 0
+ ) -> tuple[str, str]:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wait_readable(cls, obj: FileDescriptorLike) -> None:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wait_writable(cls, obj: FileDescriptorLike) -> None:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def notify_closing(cls, obj: FileDescriptorLike) -> None:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wrap_listener_socket(cls, sock: socket) -> SocketListener:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wrap_stream_socket(cls, sock: socket) -> SocketStream:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wrap_unix_stream_socket(cls, sock: socket) -> UNIXSocketStream:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wrap_udp_socket(cls, sock: socket) -> UDPSocket:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wrap_connected_udp_socket(cls, sock: socket) -> ConnectedUDPSocket:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wrap_unix_datagram_socket(cls, sock: socket) -> UNIXDatagramSocket:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wrap_connected_unix_datagram_socket(
+ cls, sock: socket
+ ) -> ConnectedUNIXDatagramSocket:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def current_default_thread_limiter(cls) -> CapacityLimiter:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def open_signal_receiver(
+ cls, *signals: Signals
+ ) -> AbstractContextManager[AsyncIterator[Signals]]:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def get_current_task(cls) -> TaskInfo:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def get_running_tasks(cls) -> Sequence[TaskInfo]:
+ pass
+
+ @classmethod
+ @abstractmethod
+ async def wait_all_tasks_blocked(cls) -> None:
+ pass
+
+ @classmethod
+ @abstractmethod
+ def create_test_runner(cls, options: dict[str, Any]) -> TestRunner:
+ pass
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/_resources.py b/venv/lib/python3.11/site-packages/anyio/abc/_resources.py
new file mode 100644
index 0000000000000000000000000000000000000000..10df115a7b9f975493476da763cc1e26dbd822e5
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/_resources.py
@@ -0,0 +1,33 @@
+from __future__ import annotations
+
+from abc import ABCMeta, abstractmethod
+from types import TracebackType
+from typing import TypeVar
+
+T = TypeVar("T")
+
+
+class AsyncResource(metaclass=ABCMeta):
+ """
+ Abstract base class for all closeable asynchronous resources.
+
+ Works as an asynchronous context manager which returns the instance itself on enter,
+ and calls :meth:`aclose` on exit.
+ """
+
+ __slots__ = ()
+
+ async def __aenter__(self: T) -> T:
+ return self
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ await self.aclose()
+
+ @abstractmethod
+ async def aclose(self) -> None:
+ """Close the resource."""
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/_sockets.py b/venv/lib/python3.11/site-packages/anyio/abc/_sockets.py
new file mode 100644
index 0000000000000000000000000000000000000000..feb26bd44a240acb20fd0f2498dff5631b8e2fb3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/_sockets.py
@@ -0,0 +1,399 @@
+from __future__ import annotations
+
+import errno
+import socket
+from abc import abstractmethod
+from collections.abc import Callable, Collection, Mapping
+from contextlib import AsyncExitStack
+from io import IOBase
+from ipaddress import IPv4Address, IPv6Address
+from socket import AddressFamily
+from typing import Any, TypeAlias, TypeVar
+
+from .._core._eventloop import get_async_backend
+from .._core._typedattr import (
+ TypedAttributeProvider,
+ TypedAttributeSet,
+ typed_attribute,
+)
+from ._streams import ByteStream, Listener, UnreliableObjectStream
+from ._tasks import TaskGroup
+
+IPAddressType: TypeAlias = str | IPv4Address | IPv6Address
+IPSockAddrType: TypeAlias = tuple[str, int]
+SockAddrType: TypeAlias = IPSockAddrType | str
+UDPPacketType: TypeAlias = tuple[bytes, IPSockAddrType]
+UNIXDatagramPacketType: TypeAlias = tuple[bytes, str]
+T_Retval = TypeVar("T_Retval")
+
+
+def _validate_socket(
+ sock_or_fd: socket.socket | int,
+ sock_type: socket.SocketKind,
+ addr_family: socket.AddressFamily = socket.AF_UNSPEC,
+ *,
+ require_connected: bool = False,
+ require_bound: bool = False,
+) -> socket.socket:
+ if isinstance(sock_or_fd, int):
+ try:
+ sock = socket.socket(fileno=sock_or_fd)
+ except OSError as exc:
+ if exc.errno == errno.ENOTSOCK:
+ raise ValueError(
+ "the file descriptor does not refer to a socket"
+ ) from exc
+ elif require_connected:
+ raise ValueError("the socket must be connected") from exc
+ elif require_bound:
+ raise ValueError("the socket must be bound to a local address") from exc
+ else:
+ raise
+ elif isinstance(sock_or_fd, socket.socket):
+ sock = sock_or_fd
+ else:
+ raise TypeError(
+ f"expected an int or socket, got {type(sock_or_fd).__qualname__} instead"
+ )
+
+ try:
+ if require_connected:
+ try:
+ sock.getpeername()
+ except OSError as exc:
+ raise ValueError("the socket must be connected") from exc
+
+ if require_bound:
+ try:
+ if sock.family in (socket.AF_INET, socket.AF_INET6):
+ bound_addr = sock.getsockname()[1]
+ else:
+ bound_addr = sock.getsockname()
+ except OSError:
+ bound_addr = None
+
+ if not bound_addr:
+ raise ValueError("the socket must be bound to a local address")
+
+ if addr_family != socket.AF_UNSPEC and sock.family != addr_family:
+ raise ValueError(
+ f"address family mismatch: expected {addr_family.name}, got "
+ f"{sock.family.name}"
+ )
+
+ if sock.type != sock_type:
+ raise ValueError(
+ f"socket type mismatch: expected {sock_type.name}, got {sock.type.name}"
+ )
+ except BaseException:
+ # Avoid ResourceWarning from the locally constructed socket object
+ if isinstance(sock_or_fd, int):
+ sock.detach()
+
+ raise
+
+ sock.setblocking(False)
+ return sock
+
+
+class SocketAttribute(TypedAttributeSet):
+ """
+ .. attribute:: family
+ :type: socket.AddressFamily
+
+ the address family of the underlying socket
+
+ .. attribute:: local_address
+ :type: tuple[str, int] | str
+
+ the local address the underlying socket is connected to
+
+ .. attribute:: local_port
+ :type: int
+
+ for IP based sockets, the local port the underlying socket is bound to
+
+ .. attribute:: raw_socket
+ :type: socket.socket
+
+ the underlying stdlib socket object
+
+ .. attribute:: remote_address
+ :type: tuple[str, int] | str
+
+ the remote address the underlying socket is connected to
+
+ .. attribute:: remote_port
+ :type: int
+
+ for IP based sockets, the remote port the underlying socket is connected to
+ """
+
+ family: AddressFamily = typed_attribute()
+ local_address: SockAddrType = typed_attribute()
+ local_port: int = typed_attribute()
+ raw_socket: socket.socket = typed_attribute()
+ remote_address: SockAddrType = typed_attribute()
+ remote_port: int = typed_attribute()
+
+
+class _SocketProvider(TypedAttributeProvider):
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ from .._core._sockets import convert_ipv6_sockaddr as convert
+
+ attributes: dict[Any, Callable[[], Any]] = {
+ SocketAttribute.family: lambda: self._raw_socket.family,
+ SocketAttribute.local_address: lambda: convert(
+ self._raw_socket.getsockname()
+ ),
+ SocketAttribute.raw_socket: lambda: self._raw_socket,
+ }
+ try:
+ peername: tuple[str, int] | None = convert(self._raw_socket.getpeername())
+ except OSError:
+ peername = None
+
+ # Provide the remote address for connected sockets
+ if peername is not None:
+ attributes[SocketAttribute.remote_address] = lambda: peername
+
+ # Provide local and remote ports for IP based sockets
+ if self._raw_socket.family in (AddressFamily.AF_INET, AddressFamily.AF_INET6):
+ attributes[SocketAttribute.local_port] = lambda: (
+ self._raw_socket.getsockname()[1]
+ )
+ if peername is not None:
+ remote_port = peername[1]
+ attributes[SocketAttribute.remote_port] = lambda: remote_port
+
+ return attributes
+
+ @property
+ @abstractmethod
+ def _raw_socket(self) -> socket.socket:
+ pass
+
+
+class SocketStream(ByteStream, _SocketProvider):
+ """
+ Transports bytes over a socket.
+
+ Supports all relevant extra attributes from :class:`~SocketAttribute`.
+ """
+
+ @classmethod
+ async def from_socket(cls, sock_or_fd: socket.socket | int) -> SocketStream:
+ """
+ Wrap an existing socket object or file descriptor as a socket stream.
+
+ The newly created socket wrapper takes ownership of the socket being passed in.
+ The existing socket must already be connected.
+
+ :param sock_or_fd: a socket object or file descriptor
+ :return: a socket stream
+
+ """
+ sock = _validate_socket(sock_or_fd, socket.SOCK_STREAM, require_connected=True)
+ return await get_async_backend().wrap_stream_socket(sock)
+
+
+class UNIXSocketStream(SocketStream):
+ @classmethod
+ async def from_socket(cls, sock_or_fd: socket.socket | int) -> UNIXSocketStream:
+ """
+ Wrap an existing socket object or file descriptor as a UNIX socket stream.
+
+ The newly created socket wrapper takes ownership of the socket being passed in.
+ The existing socket must already be connected.
+
+ :param sock_or_fd: a socket object or file descriptor
+ :return: a UNIX socket stream
+
+ """
+ sock = _validate_socket(
+ sock_or_fd, socket.SOCK_STREAM, socket.AF_UNIX, require_connected=True
+ )
+ return await get_async_backend().wrap_unix_stream_socket(sock)
+
+ @abstractmethod
+ async def send_fds(self, message: bytes, fds: Collection[int | IOBase]) -> None:
+ """
+ Send file descriptors along with a message to the peer.
+
+ :param message: a non-empty bytestring
+ :param fds: a collection of files (either numeric file descriptors or open file
+ or socket objects)
+ """
+
+ @abstractmethod
+ async def receive_fds(self, msglen: int, maxfds: int) -> tuple[bytes, list[int]]:
+ """
+ Receive file descriptors along with a message from the peer.
+
+ :param msglen: length of the message to expect from the peer
+ :param maxfds: maximum number of file descriptors to expect from the peer
+ :return: a tuple of (message, file descriptors)
+ """
+
+
+class SocketListener(Listener[SocketStream], _SocketProvider):
+ """
+ Listens to incoming socket connections.
+
+ Supports all relevant extra attributes from :class:`~SocketAttribute`.
+ """
+
+ @classmethod
+ async def from_socket(
+ cls,
+ sock_or_fd: socket.socket | int,
+ ) -> SocketListener:
+ """
+ Wrap an existing socket object or file descriptor as a socket listener.
+
+ The newly created listener takes ownership of the socket being passed in.
+
+ :param sock_or_fd: a socket object or file descriptor
+ :return: a socket listener
+
+ """
+ sock = _validate_socket(sock_or_fd, socket.SOCK_STREAM, require_bound=True)
+ return await get_async_backend().wrap_listener_socket(sock)
+
+ @abstractmethod
+ async def accept(self) -> SocketStream:
+ """Accept an incoming connection."""
+
+ async def serve(
+ self,
+ handler: Callable[[SocketStream], Any],
+ task_group: TaskGroup | None = None,
+ ) -> None:
+ from .. import create_task_group
+
+ async with AsyncExitStack() as stack:
+ if task_group is None:
+ task_group = await stack.enter_async_context(create_task_group())
+
+ while True:
+ stream = await self.accept()
+ task_group.start_soon(handler, stream)
+
+
+class UDPSocket(UnreliableObjectStream[UDPPacketType], _SocketProvider):
+ """
+ Represents an unconnected UDP socket.
+
+ Supports all relevant extra attributes from :class:`~SocketAttribute`.
+ """
+
+ @classmethod
+ async def from_socket(cls, sock_or_fd: socket.socket | int) -> UDPSocket:
+ """
+ Wrap an existing socket object or file descriptor as a UDP socket.
+
+ The newly created socket wrapper takes ownership of the socket being passed in.
+ The existing socket must be bound to a local address.
+
+ :param sock_or_fd: a socket object or file descriptor
+ :return: a UDP socket
+
+ """
+ sock = _validate_socket(sock_or_fd, socket.SOCK_DGRAM, require_bound=True)
+ return await get_async_backend().wrap_udp_socket(sock)
+
+ async def sendto(self, data: bytes, host: str, port: int) -> None:
+ """
+ Alias for :meth:`~.UnreliableObjectSendStream.send` ((data, (host, port))).
+
+ """
+ return await self.send((data, (host, port)))
+
+
+class ConnectedUDPSocket(UnreliableObjectStream[bytes], _SocketProvider):
+ """
+ Represents an connected UDP socket.
+
+ Supports all relevant extra attributes from :class:`~SocketAttribute`.
+ """
+
+ @classmethod
+ async def from_socket(cls, sock_or_fd: socket.socket | int) -> ConnectedUDPSocket:
+ """
+ Wrap an existing socket object or file descriptor as a connected UDP socket.
+
+ The newly created socket wrapper takes ownership of the socket being passed in.
+ The existing socket must already be connected.
+
+ :param sock_or_fd: a socket object or file descriptor
+ :return: a connected UDP socket
+
+ """
+ sock = _validate_socket(
+ sock_or_fd,
+ socket.SOCK_DGRAM,
+ require_connected=True,
+ )
+ return await get_async_backend().wrap_connected_udp_socket(sock)
+
+
+class UNIXDatagramSocket(
+ UnreliableObjectStream[UNIXDatagramPacketType], _SocketProvider
+):
+ """
+ Represents an unconnected Unix datagram socket.
+
+ Supports all relevant extra attributes from :class:`~SocketAttribute`.
+ """
+
+ @classmethod
+ async def from_socket(
+ cls,
+ sock_or_fd: socket.socket | int,
+ ) -> UNIXDatagramSocket:
+ """
+ Wrap an existing socket object or file descriptor as a UNIX datagram
+ socket.
+
+ The newly created socket wrapper takes ownership of the socket being passed in.
+
+ :param sock_or_fd: a socket object or file descriptor
+ :return: a UNIX datagram socket
+
+ """
+ sock = _validate_socket(sock_or_fd, socket.SOCK_DGRAM, socket.AF_UNIX)
+ return await get_async_backend().wrap_unix_datagram_socket(sock)
+
+ async def sendto(self, data: bytes, path: str) -> None:
+ """Alias for :meth:`~.UnreliableObjectSendStream.send` ((data, path))."""
+ return await self.send((data, path))
+
+
+class ConnectedUNIXDatagramSocket(UnreliableObjectStream[bytes], _SocketProvider):
+ """
+ Represents a connected Unix datagram socket.
+
+ Supports all relevant extra attributes from :class:`~SocketAttribute`.
+ """
+
+ @classmethod
+ async def from_socket(
+ cls,
+ sock_or_fd: socket.socket | int,
+ ) -> ConnectedUNIXDatagramSocket:
+ """
+ Wrap an existing socket object or file descriptor as a connected UNIX datagram
+ socket.
+
+ The newly created socket wrapper takes ownership of the socket being passed in.
+ The existing socket must already be connected.
+
+ :param sock_or_fd: a socket object or file descriptor
+ :return: a connected UNIX datagram socket
+
+ """
+ sock = _validate_socket(
+ sock_or_fd, socket.SOCK_DGRAM, socket.AF_UNIX, require_connected=True
+ )
+ return await get_async_backend().wrap_connected_unix_datagram_socket(sock)
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/_streams.py b/venv/lib/python3.11/site-packages/anyio/abc/_streams.py
new file mode 100644
index 0000000000000000000000000000000000000000..34ebfc1c0fa6f533e0621ac25659e13ffd5d7508
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/_streams.py
@@ -0,0 +1,235 @@
+from __future__ import annotations
+
+from abc import ABCMeta, abstractmethod
+from collections.abc import Callable
+from typing import Any, Generic, TypeAlias, TypeVar
+
+from .._core._exceptions import EndOfStream
+from .._core._typedattr import TypedAttributeProvider
+from ._resources import AsyncResource
+from ._tasks import TaskGroup
+
+T_Item = TypeVar("T_Item")
+T_co = TypeVar("T_co", covariant=True)
+T_contra = TypeVar("T_contra", contravariant=True)
+
+
+class UnreliableObjectReceiveStream(
+ Generic[T_co], AsyncResource, TypedAttributeProvider
+):
+ """
+ An interface for receiving objects.
+
+ This interface makes no guarantees that the received messages arrive in the order in
+ which they were sent, or that no messages are missed.
+
+ Asynchronously iterating over objects of this type will yield objects matching the
+ given type parameter.
+ """
+
+ def __aiter__(self) -> UnreliableObjectReceiveStream[T_co]:
+ return self
+
+ async def __anext__(self) -> T_co:
+ try:
+ return await self.receive()
+ except EndOfStream:
+ raise StopAsyncIteration from None
+
+ @abstractmethod
+ async def receive(self) -> T_co:
+ """
+ Receive the next item.
+
+ :raises ~anyio.ClosedResourceError: if the receive stream has been explicitly
+ closed
+ :raises ~anyio.EndOfStream: if this stream has been closed from the other end
+ :raises ~anyio.BrokenResourceError: if this stream has been rendered unusable
+ due to external causes
+ """
+
+
+class UnreliableObjectSendStream(
+ Generic[T_contra], AsyncResource, TypedAttributeProvider
+):
+ """
+ An interface for sending objects.
+
+ This interface makes no guarantees that the messages sent will reach the
+ recipient(s) in the same order in which they were sent, or at all.
+ """
+
+ @abstractmethod
+ async def send(self, item: T_contra) -> None:
+ """
+ Send an item to the peer(s).
+
+ :param item: the item to send
+ :raises ~anyio.ClosedResourceError: if the send stream has been explicitly
+ closed
+ :raises ~anyio.BrokenResourceError: if this stream has been rendered unusable
+ due to external causes
+ """
+
+
+class UnreliableObjectStream(
+ UnreliableObjectReceiveStream[T_Item], UnreliableObjectSendStream[T_Item]
+):
+ """
+ A bidirectional message stream which does not guarantee the order or reliability of
+ message delivery.
+ """
+
+
+class ObjectReceiveStream(UnreliableObjectReceiveStream[T_co]):
+ """
+ A receive message stream which guarantees that messages are received in the same
+ order in which they were sent, and that no messages are missed.
+ """
+
+
+class ObjectSendStream(UnreliableObjectSendStream[T_contra]):
+ """
+ A send message stream which guarantees that messages are delivered in the same order
+ in which they were sent, without missing any messages in the middle.
+ """
+
+
+class ObjectStream(
+ ObjectReceiveStream[T_Item],
+ ObjectSendStream[T_Item],
+ UnreliableObjectStream[T_Item],
+):
+ """
+ A bidirectional message stream which guarantees the order and reliability of message
+ delivery.
+ """
+
+ @abstractmethod
+ async def send_eof(self) -> None:
+ """
+ Send an end-of-file indication to the peer.
+
+ You should not try to send any further data to this stream after calling this
+ method. This method is idempotent (does nothing on successive calls).
+ """
+
+
+class ByteReceiveStream(AsyncResource, TypedAttributeProvider):
+ """
+ An interface for receiving bytes from a single peer.
+
+ Iterating this byte stream will yield a byte string of arbitrary length, but no more
+ than 65536 bytes.
+ """
+
+ def __aiter__(self) -> ByteReceiveStream:
+ return self
+
+ async def __anext__(self) -> bytes:
+ try:
+ return await self.receive()
+ except EndOfStream:
+ raise StopAsyncIteration from None
+
+ @abstractmethod
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ """
+ Receive at most ``max_bytes`` bytes from the peer.
+
+ .. note:: Implementers of this interface should not return an empty
+ :class:`bytes` object, and users should ignore them.
+
+ :param max_bytes: maximum number of bytes to receive (must be a positive
+ integer)
+ :return: the received bytes
+ :raises ValueError: if ``max_bytes`` is less than 1
+ :raises ~anyio.EndOfStream: if this stream has been closed from the other end
+ """
+
+
+class ByteSendStream(AsyncResource, TypedAttributeProvider):
+ """An interface for sending bytes to a single peer."""
+
+ @abstractmethod
+ async def send(self, item: bytes) -> None:
+ """
+ Send the given bytes to the peer.
+
+ :param item: the bytes to send
+ """
+
+
+class ByteStream(ByteReceiveStream, ByteSendStream):
+ """A bidirectional byte stream."""
+
+ @abstractmethod
+ async def send_eof(self) -> None:
+ """
+ Send an end-of-file indication to the peer.
+
+ You should not try to send any further data to this stream after calling this
+ method. This method is idempotent (does nothing on successive calls).
+ """
+
+
+#: Type alias for all unreliable bytes-oriented receive streams.
+AnyUnreliableByteReceiveStream: TypeAlias = (
+ UnreliableObjectReceiveStream[bytes] | ByteReceiveStream
+)
+#: Type alias for all unreliable bytes-oriented send streams.
+AnyUnreliableByteSendStream: TypeAlias = (
+ UnreliableObjectSendStream[bytes] | ByteSendStream
+)
+#: Type alias for all unreliable bytes-oriented streams.
+AnyUnreliableByteStream: TypeAlias = UnreliableObjectStream[bytes] | ByteStream
+#: Type alias for all bytes-oriented receive streams.
+AnyByteReceiveStream: TypeAlias = ObjectReceiveStream[bytes] | ByteReceiveStream
+#: Type alias for all bytes-oriented send streams.
+AnyByteSendStream: TypeAlias = ObjectSendStream[bytes] | ByteSendStream
+#: Type alias for all bytes-oriented streams.
+AnyByteStream: TypeAlias = ObjectStream[bytes] | ByteStream
+
+
+class Listener(Generic[T_co], AsyncResource, TypedAttributeProvider):
+ """An interface for objects that let you accept incoming connections."""
+
+ @abstractmethod
+ async def serve(
+ self, handler: Callable[[T_co], Any], task_group: TaskGroup | None = None
+ ) -> None:
+ """
+ Accept incoming connections as they come in and start tasks to handle them.
+
+ :param handler: a callable that will be used to handle each accepted connection
+ :param task_group: the task group that will be used to start tasks for handling
+ each accepted connection (if omitted, an ad-hoc task group will be created)
+ """
+
+
+class ObjectStreamConnectable(Generic[T_co], metaclass=ABCMeta):
+ @abstractmethod
+ async def connect(self) -> ObjectStream[T_co]:
+ """
+ Connect to the remote endpoint.
+
+ :return: an object stream connected to the remote end
+ :raises ConnectionFailed: if the connection fails
+ """
+
+
+class ByteStreamConnectable(metaclass=ABCMeta):
+ @abstractmethod
+ async def connect(self) -> ByteStream:
+ """
+ Connect to the remote endpoint.
+
+ :return: a bytestream connected to the remote end
+ :raises ConnectionFailed: if the connection fails
+ """
+
+
+#: Type alias for all connectables returning bytestreams or bytes-oriented object streams
+AnyByteStreamConnectable: TypeAlias = (
+ ObjectStreamConnectable[bytes] | ByteStreamConnectable
+)
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/_subprocesses.py b/venv/lib/python3.11/site-packages/anyio/abc/_subprocesses.py
new file mode 100644
index 0000000000000000000000000000000000000000..ce0564ceac8aac425675b5c8f7f7205d08061fd3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/_subprocesses.py
@@ -0,0 +1,79 @@
+from __future__ import annotations
+
+from abc import abstractmethod
+from signal import Signals
+
+from ._resources import AsyncResource
+from ._streams import ByteReceiveStream, ByteSendStream
+
+
+class Process(AsyncResource):
+ """An asynchronous version of :class:`subprocess.Popen`."""
+
+ @abstractmethod
+ async def wait(self) -> int:
+ """
+ Wait until the process exits.
+
+ :return: the exit code of the process
+ """
+
+ @abstractmethod
+ def terminate(self) -> None:
+ """
+ Terminates the process, gracefully if possible.
+
+ On Windows, this calls ``TerminateProcess()``.
+ On POSIX systems, this sends ``SIGTERM`` to the process.
+
+ .. seealso:: :meth:`subprocess.Popen.terminate`
+ """
+
+ @abstractmethod
+ def kill(self) -> None:
+ """
+ Kills the process.
+
+ On Windows, this calls ``TerminateProcess()``.
+ On POSIX systems, this sends ``SIGKILL`` to the process.
+
+ .. seealso:: :meth:`subprocess.Popen.kill`
+ """
+
+ @abstractmethod
+ def send_signal(self, signal: Signals) -> None:
+ """
+ Send a signal to the subprocess.
+
+ .. seealso:: :meth:`subprocess.Popen.send_signal`
+
+ :param signal: the signal number (e.g. :data:`signal.SIGHUP`)
+ """
+
+ @property
+ @abstractmethod
+ def pid(self) -> int:
+ """The process ID of the process."""
+
+ @property
+ @abstractmethod
+ def returncode(self) -> int | None:
+ """
+ The return code of the process. If the process has not yet terminated, this will
+ be ``None``.
+ """
+
+ @property
+ @abstractmethod
+ def stdin(self) -> ByteSendStream | None:
+ """The stream for the standard input of the process."""
+
+ @property
+ @abstractmethod
+ def stdout(self) -> ByteReceiveStream | None:
+ """The stream for the standard output of the process."""
+
+ @property
+ @abstractmethod
+ def stderr(self) -> ByteReceiveStream | None:
+ """The stream for the standard error output of the process."""
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/_tasks.py b/venv/lib/python3.11/site-packages/anyio/abc/_tasks.py
new file mode 100644
index 0000000000000000000000000000000000000000..44ee3a70028b609e0c264b10ef5cee2127437fee
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/_tasks.py
@@ -0,0 +1,209 @@
+from __future__ import annotations
+
+import sys
+from abc import ABCMeta, abstractmethod
+from collections.abc import Callable, Coroutine
+from contextvars import Context
+from types import TracebackType
+from typing import TYPE_CHECKING, Any, Literal, Protocol, final, overload
+
+if sys.version_info >= (3, 13):
+ from typing import TypeVar
+else:
+ from typing_extensions import TypeVar
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+if TYPE_CHECKING:
+ from .._core._tasks import CancelScope, TaskHandle
+
+T_co = TypeVar("T_co", covariant=True)
+T_contra = TypeVar("T_contra", contravariant=True, default=None)
+PosArgsT = TypeVarTuple("PosArgsT")
+
+
+def get_callable_name(func: Callable, override: object = None) -> str:
+ if override is not None:
+ return str(override)
+
+ module = getattr(func, "__module__", None)
+ qualname = getattr(func, "__qualname__", None)
+ return ".".join([x for x in (module, qualname) if x])
+
+
+def call_for_coroutine(
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ args: tuple[Unpack[PosArgsT]],
+ **kwargs: Any,
+) -> Coroutine[Any, Any, T_co]:
+ """
+ Call the given function with the given positional and keyword arguments.
+
+ :return: the resulting coroutine
+ :raises TypeError: if the return value was not a coroutine object
+
+ """
+ coro = func(*args, **kwargs)
+ if not isinstance(coro, Coroutine):
+ prefix = f"{func.__module__}." if hasattr(func, "__module__") else ""
+ raise TypeError(
+ f"Expected {prefix}{func.__qualname__}() to return a coroutine, but "
+ f"the return value ({coro!r}) is not a coroutine object"
+ )
+
+ return coro
+
+
+class TaskStatus(Protocol[T_contra]):
+ @overload
+ def started(self: TaskStatus[None]) -> None: ...
+
+ @overload
+ def started(self, value: T_contra) -> None: ...
+
+ def started(self, value: T_contra | None = None) -> None:
+ """
+ Signal that the task has started.
+
+ :param value: object passed back to the starter of the task
+ """
+
+
+class TaskGroup(metaclass=ABCMeta):
+ """
+ Groups several asynchronous tasks together.
+
+ :ivar cancel_scope: the cancel scope inherited by all child tasks
+ :vartype cancel_scope: CancelScope
+
+ .. note:: On asyncio, support for eager task factories is considered to be
+ **experimental**. In particular, they don't follow the usual semantics of new
+ tasks being scheduled on the next iteration of the event loop, and may thus
+ cause unexpected behavior in code that wasn't written with such semantics in
+ mind.
+ """
+
+ cancel_scope: CancelScope
+
+ def cancel(self, reason: str | None = None) -> None:
+ """
+ Cancel this task group's cancel scope immediately.
+
+ This is a shortcut for calling ``.cancel_scope.cancel()`` on the task group.
+
+ :param reason: a message describing the reason for the cancellation
+
+ .. versionadded:: 4.14.0
+
+ """
+ self.cancel_scope.cancel(reason)
+
+ @abstractmethod
+ def create_task(
+ self,
+ coro: Coroutine[Any, Any, T_co],
+ *,
+ name: object = None,
+ context: Context | None = None,
+ ) -> TaskHandle[T_co]:
+ """
+ Create a new task from a coroutine object and schedule it to run.
+
+ :param coro: a coroutine object
+ :param name: optional name to give the task
+ :param context: optional context to run the task in
+ :return: a task handle
+
+ .. versionadded:: 4.14.0
+ """
+
+ @final
+ def start_soon(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ *args: Unpack[PosArgsT],
+ name: object = None,
+ ) -> TaskHandle[T_co]:
+ """
+ Start a new task in this task group.
+
+ :param func: a coroutine function
+ :param args: positional arguments to call the function with
+ :param name: name of the task, for the purposes of introspection and debugging
+ :return: a task handle
+
+ .. versionadded:: 3.0
+ .. versionchanged:: 4.14.0
+ This method now returns a task handle.
+
+ """
+ final_name = get_callable_name(func, name)
+ return self.create_task(call_for_coroutine(func, args), name=final_name)
+
+ @overload
+ async def start(
+ self,
+ func: Callable[..., Coroutine[Any, Any, T_co]],
+ *args: object,
+ name: object = None,
+ return_handle: Literal[False] = ...,
+ ) -> Any: ...
+
+ @overload
+ async def start(
+ self,
+ func: Callable[..., Coroutine[Any, Any, T_co]],
+ *args: object,
+ name: object = None,
+ return_handle: Literal[True],
+ ) -> TaskHandle[T_co, Any]: ...
+
+ @abstractmethod
+ async def start(
+ self,
+ func: Callable[..., Coroutine[Any, Any, T_co]],
+ *args: object,
+ name: object = None,
+ return_handle: Literal[False] | Literal[True] = False,
+ ) -> Any:
+ """
+ Start a new task and wait until it signals for readiness.
+
+ The target callable must accept a keyword argument ``task_status`` (of type
+ :class:`TaskStatus`). Awaiting on this method will return whatever was passed to
+ ``task_status.started()`` (``None`` by default).
+
+ .. note:: The :class:`TaskStatus` class is generic, and the type argument should
+ indicate the type of the value that will be passed to
+ ``task_status.started()``.
+
+ :param func: a coroutine function that accepts the ``task_status`` keyword
+ argument
+ :param args: positional arguments to call the function with
+ :param name: an optional name for the task, for introspection and debugging
+ :param return_handle: if ``True``, return a :class:`TaskHandle` which also
+ contains the start value in ``start_value``
+ :return: the value passed to ``task_status.started()``
+ :raises RuntimeError: if the task finishes without calling
+ ``task_status.started()``
+
+ .. seealso:: :ref:`start_initialize`
+
+ .. versionadded:: 3.0
+ """
+
+ @abstractmethod
+ async def __aenter__(self) -> TaskGroup:
+ """Enter the task group context and allow starting new tasks."""
+
+ @abstractmethod
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool:
+ """Exit the task group context waiting for all tasks to finish."""
diff --git a/venv/lib/python3.11/site-packages/anyio/abc/_testing.py b/venv/lib/python3.11/site-packages/anyio/abc/_testing.py
new file mode 100644
index 0000000000000000000000000000000000000000..2a93fb7cc31533f08f5be52c0528e10147aaac57
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/abc/_testing.py
@@ -0,0 +1,73 @@
+from __future__ import annotations
+
+import types
+from abc import ABCMeta, abstractmethod
+from collections.abc import AsyncGenerator, Callable, Coroutine, Iterable
+from typing import Any, TypeVar
+
+_T = TypeVar("_T")
+
+
+class TestRunner(metaclass=ABCMeta):
+ """
+ Encapsulates a running event loop. Every call made through this object will use the
+ same event loop.
+ """
+
+ def __enter__(self) -> TestRunner:
+ return self
+
+ @abstractmethod
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: types.TracebackType | None,
+ ) -> bool | None: ...
+
+ @abstractmethod
+ def run_asyncgen_fixture(
+ self,
+ fixture_func: Callable[..., AsyncGenerator[_T, Any]],
+ kwargs: dict[str, Any],
+ ) -> Iterable[_T]:
+ """
+ Run an async generator fixture.
+
+ :param fixture_func: the fixture function
+ :param kwargs: keyword arguments to call the fixture function with
+ :return: an iterator yielding the value yielded from the async generator
+ """
+
+ @abstractmethod
+ def run_fixture(
+ self,
+ fixture_func: Callable[..., Coroutine[Any, Any, _T]],
+ kwargs: dict[str, Any],
+ ) -> _T:
+ """
+ Run an async fixture.
+
+ :param fixture_func: the fixture function
+ :param kwargs: keyword arguments to call the fixture function with
+ :return: the return value of the fixture function
+ """
+
+ @abstractmethod
+ def run_test(
+ self, test_func: Callable[..., Coroutine[Any, Any, Any]], kwargs: dict[str, Any]
+ ) -> None:
+ """
+ Run an async test function.
+
+ :param test_func: the test function
+ :param kwargs: keyword arguments to call the test function with
+ """
+
+ @abstractmethod
+ def is_running(self) -> bool:
+ """
+ Check if the test runner is running.
+
+ :return: ``True`` if the coroutine is currently being run, ``False`` otherwise.
+ """
diff --git a/venv/lib/python3.11/site-packages/anyio/from_thread.py b/venv/lib/python3.11/site-packages/anyio/from_thread.py
new file mode 100644
index 0000000000000000000000000000000000000000..8c7914c2ffff281fd5a0f0273e7d7f5d8a35e459
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/from_thread.py
@@ -0,0 +1,582 @@
+from __future__ import annotations
+
+__all__ = (
+ "BlockingPortal",
+ "BlockingPortalProvider",
+ "check_cancelled",
+ "run",
+ "run_sync",
+ "start_blocking_portal",
+)
+
+import sys
+from collections.abc import Awaitable, Callable, Coroutine, Generator
+from concurrent.futures import Future
+from contextlib import (
+ AbstractAsyncContextManager,
+ AbstractContextManager,
+ contextmanager,
+)
+from dataclasses import dataclass, field
+from functools import partial
+from inspect import isawaitable
+from threading import Lock, Thread, current_thread, get_ident
+from types import TracebackType
+from typing import (
+ Any,
+ Generic,
+ TypeVar,
+ cast,
+ overload,
+)
+
+from ._core._eventloop import (
+ get_cancelled_exc_class,
+ threadlocals,
+)
+from ._core._eventloop import run as run_eventloop
+from ._core._exceptions import NoEventLoopError
+from ._core._synchronization import Event
+from ._core._tasks import CancelScope, create_task_group
+from .abc._tasks import TaskStatus
+from .lowlevel import EventLoopToken, current_token
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+T_Retval = TypeVar("T_Retval")
+T_co = TypeVar("T_co", covariant=True)
+PosArgsT = TypeVarTuple("PosArgsT")
+
+
+def _token_or_error(token: EventLoopToken | None) -> EventLoopToken:
+ if token is not None:
+ return token
+
+ try:
+ return threadlocals.current_token
+ except AttributeError:
+ raise NoEventLoopError(
+ "Not running inside an AnyIO worker thread, and no event loop token was "
+ "provided"
+ ) from None
+
+
+def run(
+ func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
+ *args: Unpack[PosArgsT],
+ token: EventLoopToken | None = None,
+) -> T_co:
+ """
+ Call a coroutine function from a worker thread.
+
+ :param func: a coroutine function
+ :param args: positional arguments for the callable
+ :param token: an event loop token to use to get back to the event loop thread
+ (required if calling this function from outside an AnyIO worker thread)
+ :return: the return value of the coroutine function
+ :raises MissingTokenError: if no token was provided and called from outside an
+ AnyIO worker thread
+ :raises RunFinishedError: if the event loop tied to ``token`` is no longer running
+
+ .. versionchanged:: 4.11.0
+ Added the ``token`` parameter.
+
+ """
+ explicit_token = token is not None
+ token = _token_or_error(token)
+ return token.backend_class.run_async_from_thread(
+ func, args, token=token.native_token if explicit_token else None
+ )
+
+
+def run_sync(
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ *args: Unpack[PosArgsT],
+ token: EventLoopToken | None = None,
+) -> T_Retval:
+ """
+ Call a function in the event loop thread from a worker thread.
+
+ :param func: a callable
+ :param args: positional arguments for the callable
+ :param token: an event loop token to use to get back to the event loop thread
+ (required if calling this function from outside an AnyIO worker thread)
+ :return: the return value of the callable
+ :raises MissingTokenError: if no token was provided and called from outside an
+ AnyIO worker thread
+ :raises RunFinishedError: if the event loop tied to ``token`` is no longer running
+
+ .. versionchanged:: 4.11.0
+ Added the ``token`` parameter.
+
+ """
+ explicit_token = token is not None
+ token = _token_or_error(token)
+ return token.backend_class.run_sync_from_thread(
+ func, args, token=token.native_token if explicit_token else None
+ )
+
+
+class _BlockingAsyncContextManager(Generic[T_co], AbstractContextManager):
+ _enter_future: Future[T_co]
+ _exit_future: Future[bool | None]
+ _exit_event: Event
+ _exit_exc_info: tuple[
+ type[BaseException] | None, BaseException | None, TracebackType | None
+ ] = (None, None, None)
+
+ def __init__(
+ self, async_cm: AbstractAsyncContextManager[T_co], portal: BlockingPortal
+ ):
+ self._async_cm = async_cm
+ self._portal = portal
+
+ async def run_async_cm(self) -> bool | None:
+ try:
+ self._exit_event = Event()
+ value = await self._async_cm.__aenter__()
+ except BaseException as exc:
+ self._enter_future.set_exception(exc)
+ raise
+ else:
+ self._enter_future.set_result(value)
+
+ try:
+ # Wait for the sync context manager to exit.
+ # This next statement can raise `get_cancelled_exc_class()` if
+ # something went wrong in a task group in this async context
+ # manager.
+ await self._exit_event.wait()
+ finally:
+ # In case of cancellation, it could be that we end up here before
+ # `_BlockingAsyncContextManager.__exit__` is called, and an
+ # `_exit_exc_info` has been set.
+ result = await self._async_cm.__aexit__(*self._exit_exc_info)
+
+ return result
+
+ def __enter__(self) -> T_co:
+ self._enter_future = Future()
+ self._exit_future = self._portal.start_task_soon(self.run_async_cm)
+ return self._enter_future.result()
+
+ def __exit__(
+ self,
+ __exc_type: type[BaseException] | None,
+ __exc_value: BaseException | None,
+ __traceback: TracebackType | None,
+ ) -> bool | None:
+ self._exit_exc_info = __exc_type, __exc_value, __traceback
+ self._portal.call(self._exit_event.set)
+ return self._exit_future.result()
+
+
+class _BlockingPortalTaskStatus(TaskStatus):
+ def __init__(self, future: Future):
+ self._future = future
+
+ def started(self, value: object = None) -> None:
+ self._future.set_result(value)
+
+
+class BlockingPortal:
+ """
+ An object that lets external threads run code in an asynchronous event loop.
+
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+ """
+
+ def __init__(self) -> None:
+ self._token = current_token()
+ self._event_loop_thread_id: int | None = get_ident()
+ self._stop_event = Event()
+ self._task_group = create_task_group()
+
+ async def __aenter__(self) -> BlockingPortal:
+ await self._task_group.__aenter__()
+ return self
+
+ async def __aexit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool:
+ await self.stop()
+ return await self._task_group.__aexit__(exc_type, exc_val, exc_tb)
+
+ def _check_running(self) -> None:
+ if self._event_loop_thread_id is None:
+ raise RuntimeError("This portal is not running")
+ if self._event_loop_thread_id == get_ident():
+ raise RuntimeError(
+ "This method cannot be called from the event loop thread"
+ )
+
+ async def sleep_until_stopped(self) -> None:
+ """Sleep until :meth:`stop` is called."""
+ await self._stop_event.wait()
+
+ async def stop(self, cancel_remaining: bool = False) -> None:
+ """
+ Signal the portal to shut down.
+
+ This marks the portal as no longer accepting new calls and exits from
+ :meth:`sleep_until_stopped`.
+
+ :param cancel_remaining: ``True`` to cancel all the remaining tasks, ``False``
+ to let them finish before returning
+
+ """
+ self._event_loop_thread_id = None
+ self._stop_event.set()
+ if cancel_remaining:
+ self._task_group.cancel_scope.cancel("the blocking portal is shutting down")
+
+ async def _call_func(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ kwargs: dict[str, Any],
+ future: Future[T_Retval],
+ ) -> None:
+ event_loop_thread_id = self._event_loop_thread_id
+
+ def callback(f: Future[T_Retval]) -> None:
+ if f.cancelled():
+ if event_loop_thread_id == get_ident():
+ scope.cancel("the future was cancelled")
+ elif event_loop_thread_id is not None:
+ run_sync(
+ scope.cancel, "the future was cancelled", token=self._token
+ )
+
+ try:
+ retval_or_awaitable = func(*args, **kwargs)
+ if isawaitable(retval_or_awaitable):
+ with CancelScope() as scope:
+ future.add_done_callback(callback)
+ retval = await retval_or_awaitable
+ else:
+ retval = retval_or_awaitable
+ except get_cancelled_exc_class():
+ future.cancel()
+ future.set_running_or_notify_cancel()
+ except BaseException as exc:
+ if not future.cancelled():
+ future.set_exception(exc)
+
+ # Let base exceptions fall through
+ if not isinstance(exc, Exception):
+ raise
+ else:
+ if not future.cancelled():
+ future.set_result(retval)
+ finally:
+ scope = None # type: ignore[assignment]
+
+ def _spawn_task_from_thread(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval],
+ args: tuple[Unpack[PosArgsT]],
+ kwargs: dict[str, Any],
+ name: object,
+ future: Future[T_Retval],
+ ) -> None:
+ """
+ Spawn a new task using the given callable.
+
+ :param func: a callable
+ :param args: positional arguments to be passed to the callable
+ :param kwargs: keyword arguments to be passed to the callable
+ :param name: name of the task (will be coerced to a string if not ``None``)
+ :param future: a future that will resolve to the return value of the callable,
+ or the exception raised during its execution
+
+ """
+ run_sync(
+ partial(self._task_group.start_soon, name=name),
+ self._call_func,
+ func,
+ args,
+ kwargs,
+ future,
+ token=self._token,
+ )
+
+ @overload
+ def call(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
+ *args: Unpack[PosArgsT],
+ ) -> T_Retval: ...
+
+ @overload
+ def call(
+ self, func: Callable[[Unpack[PosArgsT]], T_Retval], *args: Unpack[PosArgsT]
+ ) -> T_Retval: ...
+
+ def call(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval],
+ *args: Unpack[PosArgsT],
+ ) -> T_Retval:
+ """
+ Call the given function in the event loop thread.
+
+ If the callable returns a coroutine object, it is awaited on.
+
+ :param func: any callable
+ :raises RuntimeError: if the portal is not running or if this method is called
+ from within the event loop thread
+
+ """
+ return cast(T_Retval, self.start_task_soon(func, *args).result())
+
+ @overload
+ def start_task_soon(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
+ *args: Unpack[PosArgsT],
+ name: object = None,
+ ) -> Future[T_Retval]: ...
+
+ @overload
+ def start_task_soon(
+ self,
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ *args: Unpack[PosArgsT],
+ name: object = None,
+ ) -> Future[T_Retval]: ...
+
+ def start_task_soon(
+ self,
+ func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval],
+ *args: Unpack[PosArgsT],
+ name: object = None,
+ ) -> Future[T_Retval]:
+ """
+ Start a task in the portal's task group.
+
+ The task will be run inside a cancel scope which can be cancelled by cancelling
+ the returned future.
+
+ :param func: the target function
+ :param args: positional arguments passed to ``func``
+ :param name: name of the task (will be coerced to a string if not ``None``)
+ :return: a future that resolves with the return value of the callable if the
+ task completes successfully, or with the exception raised in the task
+ :raises RuntimeError: if the portal is not running or if this method is called
+ from within the event loop thread
+ :rtype: concurrent.futures.Future[T_Retval]
+
+ .. versionadded:: 3.0
+
+ """
+ self._check_running()
+ f: Future[T_Retval] = Future()
+ self._spawn_task_from_thread(func, args, {}, name, f)
+ return f
+
+ def start_task(
+ self,
+ func: Callable[..., Awaitable[T_Retval]],
+ *args: object,
+ name: object = None,
+ ) -> tuple[Future[T_Retval], Any]:
+ """
+ Start a task in the portal's task group and wait until it signals for readiness.
+
+ This method works the same way as :meth:`.abc.TaskGroup.start`.
+
+ :param func: the target function
+ :param args: positional arguments passed to ``func``
+ :param name: name of the task (will be coerced to a string if not ``None``)
+ :return: a tuple of (future, task_status_value) where the ``task_status_value``
+ is the value passed to ``task_status.started()`` from within the target
+ function
+ :rtype: tuple[concurrent.futures.Future[T_Retval], Any]
+
+ .. versionadded:: 3.0
+
+ """
+
+ def task_done(future: Future[T_Retval]) -> None:
+ if not task_status_future.done():
+ if future.cancelled():
+ task_status_future.cancel()
+ elif future.exception():
+ task_status_future.set_exception(future.exception())
+ else:
+ exc = RuntimeError(
+ "Task exited without calling task_status.started()"
+ )
+ task_status_future.set_exception(exc)
+
+ self._check_running()
+ task_status_future: Future = Future()
+ task_status = _BlockingPortalTaskStatus(task_status_future)
+ f: Future = Future()
+ f.add_done_callback(task_done)
+ self._spawn_task_from_thread(func, args, {"task_status": task_status}, name, f)
+ return f, task_status_future.result()
+
+ def wrap_async_context_manager(
+ self, cm: AbstractAsyncContextManager[T_co]
+ ) -> AbstractContextManager[T_co]:
+ """
+ Wrap an async context manager as a synchronous context manager via this portal.
+
+ Spawns a task that will call both ``__aenter__()`` and ``__aexit__()``, stopping
+ in the middle until the synchronous context manager exits.
+
+ :param cm: an asynchronous context manager
+ :return: a synchronous context manager
+
+ .. versionadded:: 2.1
+
+ """
+ return _BlockingAsyncContextManager(cm, self)
+
+
+@dataclass
+class BlockingPortalProvider:
+ """
+ A manager for a blocking portal. Used as a context manager. The first thread to
+ enter this context manager causes a blocking portal to be started with the specific
+ parameters, and the last thread to exit causes the portal to be shut down. Thus,
+ there will be exactly one blocking portal running in this context as long as at
+ least one thread has entered this context manager.
+
+ The parameters are the same as for :func:`~anyio.run`.
+
+ :param backend: name of the backend
+ :param backend_options: backend options
+
+ .. versionadded:: 4.4
+ """
+
+ backend: str = "asyncio"
+ backend_options: dict[str, Any] | None = None
+ _lock: Lock = field(init=False, default_factory=Lock)
+ _leases: int = field(init=False, default=0)
+ _portal: BlockingPortal = field(init=False)
+ _portal_cm: AbstractContextManager[BlockingPortal] | None = field(
+ init=False, default=None
+ )
+
+ def __enter__(self) -> BlockingPortal:
+ with self._lock:
+ if self._portal_cm is None:
+ self._portal_cm = start_blocking_portal(
+ self.backend, self.backend_options
+ )
+ self._portal = self._portal_cm.__enter__()
+
+ self._leases += 1
+ return self._portal
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ portal_cm: AbstractContextManager[BlockingPortal] | None = None
+ with self._lock:
+ assert self._portal_cm
+ assert self._leases > 0
+ self._leases -= 1
+ if not self._leases:
+ portal_cm = self._portal_cm
+ self._portal_cm = None
+ del self._portal
+
+ if portal_cm:
+ portal_cm.__exit__(None, None, None)
+
+
+@contextmanager
+def start_blocking_portal(
+ backend: str = "asyncio",
+ backend_options: dict[str, Any] | None = None,
+ *,
+ name: str | None = None,
+) -> Generator[BlockingPortal, Any, None]:
+ """
+ Start a new event loop in a new thread and run a blocking portal in its main task.
+
+ The parameters are the same as for :func:`~anyio.run`.
+
+ :param backend: name of the backend
+ :param backend_options: backend options
+ :param name: name of the thread
+ :return: a context manager that yields a blocking portal
+
+ .. versionchanged:: 3.0
+ Usage as a context manager is now required.
+
+ """
+
+ async def run_portal() -> None:
+ async with BlockingPortal() as portal_:
+ if name is None:
+ current_thread().name = f"{backend}-portal-{id(portal_):x}"
+
+ future.set_result(portal_)
+ await portal_.sleep_until_stopped()
+
+ def run_blocking_portal() -> None:
+ if future.set_running_or_notify_cancel():
+ try:
+ run_eventloop(
+ run_portal, backend=backend, backend_options=backend_options
+ )
+ except BaseException as exc:
+ if not future.done():
+ future.set_exception(exc)
+
+ future: Future[BlockingPortal] = Future()
+ thread = Thread(target=run_blocking_portal, daemon=True, name=name)
+ thread.start()
+ try:
+ cancel_remaining_tasks = False
+ portal = future.result()
+ try:
+ yield portal
+ except BaseException:
+ cancel_remaining_tasks = True
+ raise
+ finally:
+ try:
+ portal.call(portal.stop, cancel_remaining_tasks)
+ except RuntimeError:
+ pass
+ finally:
+ thread.join()
+
+
+def check_cancelled() -> None:
+ """
+ Check if the cancel scope of the host task's running the current worker thread has
+ been cancelled.
+
+ If the host task's current cancel scope has indeed been cancelled, the
+ backend-specific cancellation exception will be raised.
+
+ :raises RuntimeError: if the current thread was not spawned by
+ :func:`.to_thread.run_sync`
+
+ """
+ try:
+ token: EventLoopToken = threadlocals.current_token
+ except AttributeError:
+ raise NoEventLoopError(
+ "This function can only be called inside an AnyIO worker thread"
+ ) from None
+
+ token.backend_class.check_cancelled()
diff --git a/venv/lib/python3.11/site-packages/anyio/functools.py b/venv/lib/python3.11/site-packages/anyio/functools.py
new file mode 100644
index 0000000000000000000000000000000000000000..b0bdfb4585efc4e4799388547668fba86fb5c687
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/functools.py
@@ -0,0 +1,400 @@
+from __future__ import annotations
+
+__all__ = (
+ "AsyncCacheInfo",
+ "AsyncCacheParameters",
+ "AsyncLRUCacheWrapper",
+ "cache",
+ "lru_cache",
+ "reduce",
+)
+
+import functools
+from collections import OrderedDict
+from collections.abc import (
+ AsyncIterable,
+ Awaitable,
+ Callable,
+ Coroutine,
+ Hashable,
+ Iterable,
+)
+from functools import update_wrapper
+from inspect import iscoroutinefunction
+from typing import (
+ Any,
+ Generic,
+ NamedTuple,
+ ParamSpec,
+ TypedDict,
+ TypeVar,
+ cast,
+ final,
+ overload,
+)
+from weakref import WeakKeyDictionary
+
+from ._core._eventloop import current_time
+from ._core._synchronization import Lock
+from .lowlevel import RunVar, checkpoint
+
+T = TypeVar("T")
+S = TypeVar("S")
+P = ParamSpec("P")
+lru_cache_items: RunVar[
+ WeakKeyDictionary[
+ AsyncLRUCacheWrapper[Any, Any],
+ OrderedDict[
+ Hashable,
+ tuple[_InitialMissingType, Lock, float | None]
+ | tuple[Any, None, float | None],
+ ],
+ ]
+] = RunVar("lru_cache_items")
+
+
+class _InitialMissingType:
+ pass
+
+
+initial_missing: _InitialMissingType = _InitialMissingType()
+
+
+class AsyncCacheInfo(NamedTuple):
+ hits: int
+ misses: int
+ maxsize: int | None
+ currsize: int
+ ttl: int | None
+
+
+class AsyncCacheParameters(TypedDict):
+ maxsize: int | None
+ typed: bool
+ always_checkpoint: bool
+ ttl: int | None
+
+
+class _LRUMethodWrapper(Generic[T]):
+ def __init__(self, wrapper: AsyncLRUCacheWrapper[..., T], instance: object):
+ self.__wrapper = wrapper
+ self.__instance = instance
+
+ def cache_info(self) -> AsyncCacheInfo:
+ return self.__wrapper.cache_info()
+
+ def cache_parameters(self) -> AsyncCacheParameters:
+ return self.__wrapper.cache_parameters()
+
+ def cache_clear(self) -> None:
+ self.__wrapper.cache_clear()
+
+ async def __call__(self, *args: Any, **kwargs: Any) -> T:
+ if self.__instance is None:
+ return await self.__wrapper(*args, **kwargs)
+
+ return await self.__wrapper(self.__instance, *args, **kwargs)
+
+
+@final
+class AsyncLRUCacheWrapper(Generic[P, T]):
+ def __init__(
+ self,
+ func: Callable[P, Awaitable[T]],
+ maxsize: int | None,
+ typed: bool,
+ always_checkpoint: bool,
+ ttl: int | None,
+ ):
+ self.__wrapped__ = func
+ self._hits: int = 0
+ self._misses: int = 0
+ self._maxsize = max(maxsize, 0) if maxsize is not None else None
+ self._currsize: int = 0
+ self._typed = typed
+ self._always_checkpoint = always_checkpoint
+ self._ttl = ttl
+ update_wrapper(self, func)
+
+ def cache_info(self) -> AsyncCacheInfo:
+ return AsyncCacheInfo(
+ self._hits, self._misses, self._maxsize, self._currsize, self._ttl
+ )
+
+ def cache_parameters(self) -> AsyncCacheParameters:
+ return {
+ "maxsize": self._maxsize,
+ "typed": self._typed,
+ "always_checkpoint": self._always_checkpoint,
+ "ttl": self._ttl,
+ }
+
+ def cache_clear(self) -> None:
+ if cache := lru_cache_items.get(None):
+ cache.pop(self, None)
+ self._hits = self._misses = self._currsize = 0
+
+ async def __call__(self, *args: P.args, **kwargs: P.kwargs) -> T:
+ # Easy case first: if maxsize == 0, no caching is done
+ if self._maxsize == 0:
+ value = await self.__wrapped__(*args, **kwargs)
+ self._misses += 1
+ return value
+
+ # The key is constructed as a flat tuple to avoid memory overhead
+ key: tuple[Any, ...] = args
+ if kwargs:
+ # initial_missing is used as a separator
+ key += (initial_missing,) + sum(kwargs.items(), ())
+
+ if self._typed:
+ key += tuple(type(arg) for arg in args)
+ if kwargs:
+ key += (initial_missing,) + tuple(type(val) for val in kwargs.values())
+
+ try:
+ cache = lru_cache_items.get()
+ except LookupError:
+ cache = WeakKeyDictionary()
+ lru_cache_items.set(cache)
+
+ try:
+ cache_entry = cache[self]
+ except KeyError:
+ cache_entry = cache[self] = OrderedDict()
+
+ cached_value: T | _InitialMissingType
+ try:
+ cached_value, lock, expires_at = cache_entry[key]
+ except KeyError:
+ # We're the first task to call this function
+ cached_value, lock, expires_at = (
+ initial_missing,
+ Lock(fast_acquire=not self._always_checkpoint),
+ None,
+ )
+ cache_entry[key] = cached_value, lock, expires_at
+
+ if lock is None:
+ if expires_at is not None and current_time() >= expires_at:
+ self._currsize -= 1
+ cached_value, lock, expires_at = (
+ initial_missing,
+ Lock(fast_acquire=not self._always_checkpoint),
+ None,
+ )
+ cache_entry[key] = cached_value, lock, expires_at
+ else:
+ # The value was already cached
+ self._hits += 1
+ cache_entry.move_to_end(key)
+ if self._always_checkpoint:
+ await checkpoint()
+
+ return cast(T, cached_value)
+
+ async with lock:
+ # Check if another task filled the cache while we acquired the lock
+ if (cached_value := cache_entry[key][0]) is initial_missing:
+ self._misses += 1
+ if self._maxsize is not None and self._currsize >= self._maxsize:
+ cache_entry.popitem(last=False)
+ else:
+ self._currsize += 1
+
+ value = await self.__wrapped__(*args, **kwargs)
+ expires_at = (
+ current_time() + self._ttl if self._ttl is not None else None
+ )
+ cache_entry[key] = value, None, expires_at
+ else:
+ # Another task filled the cache while we were waiting for the lock
+ self._hits += 1
+ cache_entry.move_to_end(key)
+ value = cast(T, cached_value)
+
+ return value
+
+ def __get__(
+ self, instance: object, owner: type | None = None
+ ) -> _LRUMethodWrapper[T]:
+ wrapper = _LRUMethodWrapper(self, instance)
+ update_wrapper(wrapper, self.__wrapped__)
+ return wrapper
+
+
+class _LRUCacheWrapper:
+ def __init__(
+ self, maxsize: int | None, typed: bool, always_checkpoint: bool, ttl: int | None
+ ):
+ self._maxsize = maxsize
+ self._typed = typed
+ self._always_checkpoint = always_checkpoint
+ self._ttl = ttl
+
+ @overload
+ def __call__( # type: ignore[overload-overlap]
+ self, func: Callable[P, Coroutine[Any, Any, T]], /
+ ) -> AsyncLRUCacheWrapper[P, T]: ...
+
+ @overload
+ def __call__(
+ self, func: Callable[..., T], /
+ ) -> functools._lru_cache_wrapper[T]: ...
+
+ def __call__(
+ self, f: Callable[P, Coroutine[Any, Any, T]] | Callable[..., T], /
+ ) -> AsyncLRUCacheWrapper[P, T] | functools._lru_cache_wrapper[T]:
+ if iscoroutinefunction(f):
+ return AsyncLRUCacheWrapper(
+ f, self._maxsize, self._typed, self._always_checkpoint, self._ttl
+ )
+
+ return functools.lru_cache(maxsize=self._maxsize, typed=self._typed)(f) # type: ignore[arg-type]
+
+
+@overload
+def cache( # type: ignore[overload-overlap]
+ func: Callable[P, Coroutine[Any, Any, T]], /
+) -> AsyncLRUCacheWrapper[P, T]: ...
+
+
+@overload
+def cache(func: Callable[..., T], /) -> functools._lru_cache_wrapper[T]: ...
+
+
+def cache(func: Callable[..., Any] | Callable[P, Coroutine[Any, Any, Any]], /) -> Any:
+ """
+ A convenient shortcut for :func:`lru_cache` with ``maxsize=None``.
+
+ This is the asynchronous equivalent to :func:`functools.cache`.
+
+ """
+ return lru_cache(maxsize=None)(func)
+
+
+@overload
+def lru_cache(
+ *,
+ maxsize: int | None = ...,
+ typed: bool = ...,
+ always_checkpoint: bool = ...,
+ ttl: int | None = ...,
+) -> _LRUCacheWrapper: ...
+
+
+@overload
+def lru_cache( # type: ignore[overload-overlap]
+ func: Callable[P, Coroutine[Any, Any, T]], /
+) -> AsyncLRUCacheWrapper[P, T]: ...
+
+
+@overload
+def lru_cache(func: Callable[..., T], /) -> functools._lru_cache_wrapper[T]: ...
+
+
+def lru_cache(
+ func: Callable[..., Coroutine[Any, Any, Any]] | Callable[..., Any] | None = None,
+ /,
+ *,
+ maxsize: int | None = 128,
+ typed: bool = False,
+ always_checkpoint: bool = False,
+ ttl: int | None = None,
+) -> Any:
+ """
+ An asynchronous version of :func:`functools.lru_cache`.
+
+ If a synchronous function is passed, the standard library
+ :func:`functools.lru_cache` is applied instead.
+
+ :param always_checkpoint: if ``True``, every call to the cached function will be
+ guaranteed to yield control to the event loop at least once
+ :param ttl: time in seconds after which to invalidate cache entries
+
+ .. note:: Caches and locks are managed on a per-event loop basis.
+
+ """
+ if func is None:
+ return _LRUCacheWrapper(maxsize, typed, always_checkpoint, ttl)
+
+ if not callable(func):
+ raise TypeError("the first argument must be callable")
+
+ return _LRUCacheWrapper(maxsize, typed, always_checkpoint, ttl)(func)
+
+
+@overload
+async def reduce(
+ function: Callable[[T, S], Awaitable[T]],
+ iterable: Iterable[S] | AsyncIterable[S],
+ /,
+ initial: T,
+) -> T: ...
+
+
+@overload
+async def reduce(
+ function: Callable[[T, T], Awaitable[T]],
+ iterable: Iterable[T] | AsyncIterable[T],
+ /,
+) -> T: ...
+
+
+async def reduce( # type: ignore[misc]
+ function: Callable[[T, T], Awaitable[T]] | Callable[[T, S], Awaitable[T]],
+ iterable: Iterable[T] | Iterable[S] | AsyncIterable[T] | AsyncIterable[S],
+ /,
+ initial: T | _InitialMissingType = initial_missing,
+) -> T:
+ """
+ Asynchronous version of :func:`functools.reduce`.
+
+ :param function: a coroutine function that takes two arguments: the accumulated
+ value and the next element from the iterable
+ :param iterable: an iterable or async iterable
+ :param initial: the initial value (if missing, the first element of the iterable is
+ used as the initial value)
+
+ """
+ element: Any
+ function_called = False
+ if isinstance(iterable, AsyncIterable):
+ async_it = iterable.__aiter__()
+ if initial is initial_missing:
+ try:
+ value = cast(T, await async_it.__anext__())
+ except StopAsyncIteration:
+ raise TypeError(
+ "reduce() of empty sequence with no initial value"
+ ) from None
+ else:
+ value = cast(T, initial)
+
+ async for element in async_it:
+ value = await function(value, element)
+ function_called = True
+ elif isinstance(iterable, Iterable):
+ it = iter(iterable)
+ if initial is initial_missing:
+ try:
+ value = cast(T, next(it))
+ except StopIteration:
+ raise TypeError(
+ "reduce() of empty sequence with no initial value"
+ ) from None
+ else:
+ value = cast(T, initial)
+
+ for element in it:
+ value = await function(value, element)
+ function_called = True
+ else:
+ raise TypeError("reduce() argument 2 must be an iterable or async iterable")
+
+ # Make sure there is at least one checkpoint, even if an empty iterable and an
+ # initial value were given
+ if not function_called:
+ await checkpoint()
+
+ return value
diff --git a/venv/lib/python3.11/site-packages/anyio/itertools.py b/venv/lib/python3.11/site-packages/anyio/itertools.py
new file mode 100644
index 0000000000000000000000000000000000000000..7e5248e4b8f99556cdbb98b024a188d65cdfce83
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/itertools.py
@@ -0,0 +1,626 @@
+from __future__ import annotations
+
+__all__ = (
+ "accumulate",
+ "batched",
+ "Chain",
+ "combinations",
+ "combinations_with_replacement",
+ "compress",
+ "count",
+ "cycle",
+ "dropwhile",
+ "filterfalse",
+ "groupby",
+ "islice",
+ "pairwise",
+ "permutations",
+ "product",
+ "repeat",
+ "starmap",
+ "tee",
+ "takewhile",
+ "zip_longest",
+)
+
+import itertools
+import operator
+import sys
+from collections.abc import (
+ AsyncGenerator,
+ AsyncIterable,
+ AsyncIterator,
+ Awaitable,
+ Callable,
+ Iterable,
+ Iterator,
+)
+from dataclasses import dataclass, field
+from typing import Any, Generic, TypeVar, cast, overload
+
+from ._core._synchronization import Lock
+from ._core._tasks import CancelScope
+from .lowlevel import cancel_shielded_checkpoint, checkpoint, checkpoint_if_cancelled
+
+T = TypeVar("T")
+R = TypeVar("R")
+_tee_end = object()
+
+
+@dataclass(eq=False)
+class _IterableAsyncIterator(AsyncIterator[T]):
+ iterator: Iterator[T]
+
+ async def __anext__(self) -> T:
+ await checkpoint_if_cancelled()
+ try:
+ result = next(self.iterator)
+ except StopIteration:
+ await cancel_shielded_checkpoint()
+ raise StopAsyncIteration from None
+
+ await cancel_shielded_checkpoint()
+ return result
+
+
+def _iterate(iterable: Iterable[T] | AsyncIterable[T]) -> AsyncIterator[T]:
+ if isinstance(iterable, AsyncIterator):
+ return iterable
+
+ if isinstance(iterable, AsyncIterable):
+ return iterable.__aiter__()
+
+ return _IterableAsyncIterator(iter(iterable))
+
+
+@dataclass(eq=False)
+class _TeeLink(Generic[T]):
+ value: object | None = None
+ next: _TeeLink[T] | None = None
+ filled: bool = False
+
+
+@dataclass(eq=False)
+class _TeeState(Generic[T]):
+ iterator: AsyncIterator[T]
+ lock: Lock = field(default_factory=Lock)
+
+ async def fill(self, link: _TeeLink[T]) -> bool:
+ if link.filled:
+ return False
+
+ async with self.lock:
+ if link.filled:
+ return True
+
+ link.value = await anext(self.iterator, _tee_end)
+ if link.value is not _tee_end:
+ link.next = _TeeLink()
+
+ link.filled = True
+ return True
+
+
+class _TeeAsyncIterator(AsyncIterator[T]):
+ _state: _TeeState[T]
+ _link: _TeeLink[T]
+ _element_yielded: bool
+
+ def __init__(
+ self, iterable: Iterable[T] | AsyncIterable[T] | _TeeAsyncIterator[T]
+ ) -> None:
+ if isinstance(iterable, _TeeAsyncIterator):
+ self._state = iterable._state
+ self._link = iterable._link
+ else:
+ self._state = _TeeState(_iterate(iterable))
+ self._link = _TeeLink()
+
+ self._element_yielded = False
+
+ async def __anext__(self) -> T:
+ had_yieldpoint = await self._state.fill(self._link)
+ if self._link.value is _tee_end:
+ if not self._element_yielded:
+ await checkpoint()
+
+ raise StopAsyncIteration
+
+ if not had_yieldpoint:
+ await checkpoint_if_cancelled()
+
+ self._element_yielded = True
+ value = cast(T, self._link.value)
+ next_link = self._link.next
+ assert next_link is not None
+ self._link = next_link
+ if not had_yieldpoint:
+ await cancel_shielded_checkpoint()
+
+ return value
+
+
+async def _operator_add(x: T, y: T) -> T:
+ return operator.add(x, y)
+
+
+async def accumulate(
+ iterable: Iterable[T] | AsyncIterable[T],
+ function: Callable[[T, T], Awaitable[T]] = _operator_add,
+ *,
+ initial: T | None = None,
+) -> AsyncGenerator[T, None]:
+ iterator = _iterate(iterable)
+ if initial is None:
+ try:
+ total = await anext(iterator)
+ except StopAsyncIteration:
+ await checkpoint()
+ return
+ else:
+ await checkpoint_if_cancelled()
+ total = initial
+ await cancel_shielded_checkpoint()
+
+ yield total
+
+ async for element in iterator:
+ total = await function(total, element)
+ yield total
+
+
+async def batched(
+ iterable: Iterable[T] | AsyncIterable[T], n: int, *, strict: bool = False
+) -> AsyncGenerator[tuple[T, ...], None]:
+ if n < 1:
+ raise ValueError("n must be at least one")
+
+ iterator = _iterate(iterable)
+
+ while True:
+ batch: list[T] = []
+ for _ in range(n):
+ try:
+ batch.append(await anext(iterator))
+ except StopAsyncIteration:
+ if not batch:
+ await checkpoint()
+ return
+ if strict:
+ raise ValueError("batched(): incomplete batch") from None
+
+ yield tuple(batch)
+ return
+
+ yield tuple(batch)
+
+
+class Chain:
+ def __call__(
+ self, *iterables: Iterable[T] | AsyncIterable[T]
+ ) -> AsyncGenerator[T, None]:
+ return self.from_iterable(iterables)
+
+ async def from_iterable(
+ self,
+ iterables: (
+ Iterable[Iterable[T] | AsyncIterable[T]]
+ | AsyncIterable[Iterable[T] | AsyncIterable[T]]
+ ),
+ ) -> AsyncGenerator[T, None]:
+ element_yielded = False
+ outer_iter = _iterate(iterables)
+
+ try:
+ async for iterable in outer_iter:
+ async for element in _iterate(iterable):
+ element_yielded = True
+ yield element
+ finally:
+ aclose = getattr(outer_iter, "aclose", None)
+ if aclose is not None:
+ with CancelScope(shield=True):
+ await aclose()
+
+ if not element_yielded:
+ await checkpoint()
+
+
+chain: Chain = Chain()
+
+
+async def combinations(
+ iterable: Iterable[T] | AsyncIterable[T], r: int
+) -> AsyncGenerator[tuple[T, ...], None]:
+ pool: list[T] = [element async for element in _iterate(iterable)]
+ async for combination in _iterate(itertools.combinations(pool, r)):
+ yield combination
+
+
+async def combinations_with_replacement(
+ iterable: Iterable[T] | AsyncIterable[T], r: int
+) -> AsyncGenerator[tuple[T, ...], None]:
+ pool: list[T] = [element async for element in _iterate(iterable)]
+ async for combination in _iterate(itertools.combinations_with_replacement(pool, r)):
+ yield combination
+
+
+async def compress(
+ data: Iterable[T] | AsyncIterable[T],
+ selectors: Iterable[object] | AsyncIterable[object],
+) -> AsyncGenerator[T, None]:
+ data_iterator = _iterate(data)
+ selector_iterator = _iterate(selectors)
+ element_yielded = False
+
+ while True:
+ try:
+ datum = await anext(data_iterator)
+ selector = await anext(selector_iterator)
+ except StopAsyncIteration:
+ if not element_yielded:
+ await checkpoint()
+
+ return
+
+ if selector:
+ element_yielded = True
+ yield datum
+
+
+async def count(start: int = 0, step: int = 1) -> AsyncGenerator[int, None]:
+ n = start
+ while True:
+ await checkpoint_if_cancelled()
+ value = n
+ n += step
+ await cancel_shielded_checkpoint()
+ yield value
+
+
+async def cycle(
+ iterable: Iterable[T] | AsyncIterable[T],
+) -> AsyncGenerator[T, None]:
+ saved: list[T] = []
+ async for element in _iterate(iterable):
+ saved.append(element)
+ yield element
+
+ if not saved:
+ await checkpoint()
+ return
+
+ while True:
+ for element in saved:
+ await checkpoint()
+ yield element
+
+
+async def dropwhile(
+ predicate: Callable[[T], Awaitable[object]],
+ iterable: Iterable[T] | AsyncIterable[T],
+) -> AsyncGenerator[T, None]:
+ element_yielded = False
+ dropping = True
+
+ async for element in _iterate(iterable):
+ if dropping and await predicate(element):
+ continue
+
+ dropping = False
+ element_yielded = True
+ yield element
+
+ if not element_yielded:
+ await checkpoint()
+
+
+async def filterfalse(
+ predicate: Callable[[T], Awaitable[object]],
+ iterable: Iterable[T] | AsyncIterable[T],
+) -> AsyncGenerator[T, None]:
+ element_yielded = False
+
+ async for element in _iterate(iterable):
+ if not await predicate(element):
+ element_yielded = True
+ yield element
+
+ if not element_yielded:
+ await checkpoint()
+
+
+@overload
+def groupby(
+ iterable: Iterable[T] | AsyncIterable[T],
+) -> AsyncGenerator[tuple[T, list[T]], None]: ...
+
+
+@overload
+def groupby(
+ iterable: Iterable[T] | AsyncIterable[T],
+ key: Callable[[T], Awaitable[R]],
+) -> AsyncGenerator[tuple[R, list[T]], None]: ...
+
+
+async def groupby(
+ iterable: Iterable[T] | AsyncIterable[T],
+ key: Callable[[T], Awaitable[object]] | None = None,
+) -> AsyncGenerator[tuple[object, list[T]], None]:
+ iterator = _iterate(iterable)
+ try:
+ element = await anext(iterator)
+ except StopAsyncIteration:
+ await checkpoint()
+ return
+
+ group_key = element if key is None else await key(element)
+ values = [element]
+
+ async for element in iterator:
+ next_key = element if key is None else await key(element)
+ if next_key != group_key:
+ completed_group = group_key, values
+ group_key = next_key
+ values = [element]
+ yield completed_group
+ else:
+ values.append(element)
+
+ yield group_key, values
+
+
+@overload
+def islice(
+ iterable: Iterable[T] | AsyncIterable[T],
+ stop: int | None,
+ /,
+) -> AsyncGenerator[T, None]: ...
+
+
+@overload
+def islice(
+ iterable: Iterable[T] | AsyncIterable[T],
+ start: int | None,
+ stop: int | None,
+ step: int | None = 1,
+ /,
+) -> AsyncGenerator[T, None]: ...
+
+
+async def islice(
+ iterable: Iterable[T] | AsyncIterable[T],
+ *args: int | None,
+) -> AsyncGenerator[T, None]:
+ if not args:
+ raise TypeError("islice expected at least 2 arguments, got 1")
+ if len(args) > 3:
+ raise TypeError(f"islice expected at most 4 arguments, got {len(args) + 1}")
+
+ slice_args = slice(*args)
+
+ start_message = (
+ "Indices for islice() must be None or an integer: 0 <= x <= sys.maxsize."
+ )
+ stop_message = (
+ "Stop argument for islice() must be None or an integer: 0 <= x <= sys.maxsize."
+ )
+ step_message = "Step for islice() must be a positive integer or None."
+
+ def normalize_index(value: object, message: str) -> int:
+ try:
+ index = operator.index(cast(Any, value))
+ except TypeError:
+ raise ValueError(message) from None
+
+ if index < 0 or index > sys.maxsize:
+ raise ValueError(message)
+
+ return index
+
+ start = (
+ 0
+ if slice_args.start is None
+ else normalize_index(slice_args.start, start_message)
+ )
+ stop = (
+ None
+ if slice_args.stop is None
+ else normalize_index(slice_args.stop, stop_message)
+ )
+ step = (
+ 1 if slice_args.step is None else normalize_index(slice_args.step, step_message)
+ )
+
+ if step <= 0:
+ raise ValueError(step_message)
+
+ if stop == 0 or start == stop:
+ await checkpoint()
+ return
+
+ iterator = _iterate(iterable)
+ index = 0
+ element_yielded = False
+
+ while stop is None or index < stop:
+ try:
+ element = await anext(iterator)
+ except StopAsyncIteration:
+ if not element_yielded:
+ await checkpoint()
+
+ return
+
+ if index >= start and (index - start) % step == 0:
+ index += 1
+ element_yielded = True
+ yield element
+ else:
+ index += 1
+
+ if not element_yielded:
+ await checkpoint()
+
+
+async def pairwise(
+ iterable: Iterable[T] | AsyncIterable[T],
+) -> AsyncGenerator[tuple[T, T], None]:
+ iterator = _iterate(iterable)
+ try:
+ previous = await anext(iterator)
+ except StopAsyncIteration:
+ await checkpoint()
+ return
+
+ element_yielded = False
+ async for element in iterator:
+ element_yielded = True
+ pair = (previous, element)
+ previous = element
+ yield pair
+
+ if not element_yielded:
+ await checkpoint()
+
+
+async def permutations(
+ iterable: Iterable[T] | AsyncIterable[T], r: int | None = None
+) -> AsyncGenerator[tuple[T, ...], None]:
+ pool: list[T] = [element async for element in _iterate(iterable)]
+ n = len(pool)
+ if r is None:
+ r = n
+ elif not isinstance(r, int):
+ raise TypeError("Expected int as r")
+ elif r < 0:
+ raise ValueError("r must be non-negative")
+
+ async for permutation in _iterate(itertools.permutations(pool, r)):
+ yield permutation
+
+
+async def product(
+ *iterables: Iterable[T] | AsyncIterable[T], repeat: int = 1
+) -> AsyncGenerator[tuple[T, ...], None]:
+ repeat = operator.index(repeat)
+ if repeat < 0:
+ raise ValueError("repeat argument cannot be negative")
+
+ pools: list[tuple[T, ...]] = []
+ for iterable in iterables:
+ pool: list[T] = [element async for element in _iterate(iterable)]
+ pools.append(tuple(pool))
+
+ async for value in _iterate(itertools.product(*pools, repeat=repeat)):
+ yield value
+
+
+async def repeat(element: T, times: int | None = None) -> AsyncGenerator[T, None]:
+ if times is None:
+ while True:
+ await checkpoint()
+ yield element
+
+ remaining = operator.index(cast(Any, times))
+ if remaining <= 0:
+ await checkpoint()
+ return
+
+ while remaining > 0:
+ await checkpoint_if_cancelled()
+ remaining -= 1
+ await cancel_shielded_checkpoint()
+ yield element
+
+
+async def starmap(
+ function: Callable[..., Awaitable[R]],
+ iterable: (
+ Iterable[Iterable[object] | AsyncIterable[object]]
+ | AsyncIterable[Iterable[object] | AsyncIterable[object]]
+ ),
+) -> AsyncGenerator[R, None]:
+ result_yielded = False
+
+ async for args_iterable in _iterate(iterable):
+ args = [element async for element in _iterate(args_iterable)]
+ result_yielded = True
+ yield await function(*args)
+
+ if not result_yielded:
+ await checkpoint()
+
+
+def tee(
+ iterable: Iterable[T] | AsyncIterable[T], n: int = 2
+) -> tuple[AsyncIterator[T], ...]:
+ n = operator.index(cast(Any, n))
+ if n < 0:
+ raise ValueError("n must be >= 0")
+ if n == 0:
+ return ()
+
+ iterator = _TeeAsyncIterator(iterable)
+ iterators: list[AsyncIterator[T]] = [iterator]
+ iterators.extend(_TeeAsyncIterator(iterator) for _ in range(n - 1))
+ return tuple(iterators)
+
+
+async def takewhile(
+ predicate: Callable[[T], Awaitable[object]],
+ iterable: Iterable[T] | AsyncIterable[T],
+) -> AsyncGenerator[T, None]:
+ element_yielded = False
+
+ async for element in _iterate(iterable):
+ if not await predicate(element):
+ if not element_yielded:
+ await checkpoint()
+
+ return
+
+ element_yielded = True
+ yield element
+
+ if not element_yielded:
+ await checkpoint()
+
+
+async def zip_longest(
+ *iterables: Iterable[object] | AsyncIterable[object],
+ fillvalue: object = None,
+) -> AsyncGenerator[tuple[object, ...], None]:
+ iterators = [_iterate(iterable) for iterable in iterables]
+ num_active = len(iterators)
+ if not num_active:
+ await checkpoint()
+ return
+
+ active = [True] * num_active
+ tuple_yielded = False
+
+ while True:
+ values: list[object] = []
+ for index, iterator in enumerate(iterators):
+ if not active[index]:
+ values.append(fillvalue)
+ continue
+
+ try:
+ value = await anext(iterator)
+ except StopAsyncIteration:
+ active[index] = False
+ num_active -= 1
+ if not num_active:
+ if not tuple_yielded:
+ await checkpoint()
+
+ return
+
+ value = fillvalue
+
+ values.append(value)
+
+ tuple_yielded = True
+ yield tuple(values)
diff --git a/venv/lib/python3.11/site-packages/anyio/lowlevel.py b/venv/lib/python3.11/site-packages/anyio/lowlevel.py
new file mode 100644
index 0000000000000000000000000000000000000000..ee111ecc8f22cf54e95bc7df18844e131730e572
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/lowlevel.py
@@ -0,0 +1,228 @@
+from __future__ import annotations
+
+__all__ = (
+ "EventLoopToken",
+ "RunvarToken",
+ "RunVar",
+ "checkpoint",
+ "checkpoint_if_cancelled",
+ "cancel_shielded_checkpoint",
+ "current_token",
+)
+
+import enum
+from dataclasses import dataclass
+from types import TracebackType
+from typing import TYPE_CHECKING, Any, Generic, Literal, TypeVar, final, overload
+from weakref import WeakKeyDictionary
+
+from ._core._eventloop import get_async_backend
+
+if TYPE_CHECKING:
+ from .abc import AsyncBackend
+
+T = TypeVar("T")
+D = TypeVar("D")
+
+
+async def checkpoint() -> None:
+ """
+ Check for cancellation and allow the scheduler to switch to another task.
+
+ Equivalent to (but more efficient than)::
+
+ await checkpoint_if_cancelled()
+ await cancel_shielded_checkpoint()
+
+ .. versionadded:: 3.0
+
+ """
+ await get_async_backend().checkpoint()
+
+
+async def checkpoint_if_cancelled() -> None:
+ """
+ Enter a checkpoint if the enclosing cancel scope has been cancelled.
+
+ This does not allow the scheduler to switch to a different task.
+
+ .. versionadded:: 3.0
+
+ """
+ await get_async_backend().checkpoint_if_cancelled()
+
+
+async def cancel_shielded_checkpoint() -> None:
+ """
+ Allow the scheduler to switch to another task but without checking for cancellation.
+
+ Equivalent to (but potentially more efficient than)::
+
+ with CancelScope(shield=True):
+ await checkpoint()
+
+ .. versionadded:: 3.0
+
+ """
+ await get_async_backend().cancel_shielded_checkpoint()
+
+
+@final
+@dataclass(frozen=True, repr=False)
+class EventLoopToken:
+ """
+ An opaque object that holds a reference to an event loop.
+
+ .. versionadded:: 4.11.0
+ """
+
+ backend_class: type[AsyncBackend]
+ native_token: object
+
+
+def current_token() -> EventLoopToken:
+ """
+ Return a token object that can be used to call code in the current event loop from
+ another thread.
+
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ .. versionadded:: 4.11.0
+
+ """
+ backend_class = get_async_backend()
+ raw_token = backend_class.current_token()
+ return EventLoopToken(backend_class, raw_token)
+
+
+_run_vars: WeakKeyDictionary[object, dict[RunVar[Any], Any]] = WeakKeyDictionary()
+
+
+class _NoValueSet(enum.Enum):
+ NO_VALUE_SET = enum.auto()
+
+
+class RunvarToken(Generic[T]):
+ """
+ A token that can be used to restore a :class:`RunVar` to its previous value.
+
+ Returned by :meth:`RunVar.set`. Can be used as a context manager to automatically
+ reset the variable on exit, or passed directly to :meth:`RunVar.reset`.
+ """
+
+ __slots__ = "_var", "_value", "_redeemed"
+
+ def __init__(self, var: RunVar[T], value: T | Literal[_NoValueSet.NO_VALUE_SET]):
+ self._var = var
+ self._value: T | Literal[_NoValueSet.NO_VALUE_SET] = value
+ self._redeemed = False
+
+ def __enter__(self) -> RunvarToken[T]:
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self._var.reset(self)
+
+
+class RunVar(Generic[T]):
+ """
+ Like a :class:`~contextvars.ContextVar`, except scoped to the running event loop.
+
+ Can be used as a context manager, Just like :class:`~contextvars.ContextVar`, that
+ will reset the variable to its previous value when the context block is exited.
+ """
+
+ __slots__ = "_name", "_default"
+
+ NO_VALUE_SET: Literal[_NoValueSet.NO_VALUE_SET] = _NoValueSet.NO_VALUE_SET
+
+ def __init__(
+ self, name: str, default: T | Literal[_NoValueSet.NO_VALUE_SET] = NO_VALUE_SET
+ ):
+ self._name = name
+ self._default = default
+
+ @property
+ def _current_vars(self) -> dict[RunVar[T], T]:
+ native_token = current_token().native_token
+ try:
+ return _run_vars[native_token]
+ except KeyError:
+ run_vars = _run_vars[native_token] = {}
+ return run_vars
+
+ @overload
+ def get(self, default: D) -> T | D: ...
+
+ @overload
+ def get(self) -> T: ...
+
+ def get(
+ self, default: D | Literal[_NoValueSet.NO_VALUE_SET] = NO_VALUE_SET
+ ) -> T | D:
+ """
+ Return the current value of this run variable.
+
+ :param default: a fallback value to return if no value has been set
+ :return: the current value, the provided default, or the variable's own default
+ :raises LookupError: if no value is set and no default is available
+
+ """
+ try:
+ return self._current_vars[self]
+ except KeyError:
+ if default is not RunVar.NO_VALUE_SET:
+ return default
+ elif self._default is not RunVar.NO_VALUE_SET:
+ return self._default
+
+ raise LookupError(
+ f'Run variable "{self._name}" has no value and no default set'
+ )
+
+ def set(self, value: T) -> RunvarToken[T]:
+ """
+ Set the value of this run variable for the current event loop.
+
+ :param value: the new value
+ :return: a token that can be used to restore the previous value
+
+ """
+ current_vars = self._current_vars
+ token = RunvarToken(self, current_vars.get(self, RunVar.NO_VALUE_SET))
+ current_vars[self] = value
+ return token
+
+ def reset(self, token: RunvarToken[T]) -> None:
+ """
+ Restore this run variable to the value it held before the matching :meth:`set`.
+
+ :param token: the token returned by :meth:`set`
+ :raises ValueError: if the token belongs to a different :class:`RunVar` or the token
+ has already been used
+
+ """
+ if token._var is not self:
+ raise ValueError("This token does not belong to this RunVar")
+
+ if token._redeemed:
+ raise ValueError("This token has already been used")
+
+ if token._value is _NoValueSet.NO_VALUE_SET:
+ try:
+ del self._current_vars[self]
+ except KeyError:
+ pass
+ else:
+ self._current_vars[self] = token._value
+
+ token._redeemed = True
+
+ def __repr__(self) -> str:
+ return f""
diff --git a/venv/lib/python3.11/site-packages/anyio/py.typed b/venv/lib/python3.11/site-packages/anyio/py.typed
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/anyio/pytest_plugin.py b/venv/lib/python3.11/site-packages/anyio/pytest_plugin.py
new file mode 100644
index 0000000000000000000000000000000000000000..5c667597d02f268b8f5427d8f80957116992c6ee
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/pytest_plugin.py
@@ -0,0 +1,375 @@
+from __future__ import annotations
+
+import dataclasses
+import socket
+import sys
+from collections.abc import Callable, Generator, Iterator
+from contextlib import ExitStack, contextmanager
+from inspect import isasyncgenfunction, iscoroutinefunction, ismethod
+from typing import Any, cast
+
+import pytest
+from _pytest.fixtures import FuncFixtureInfo, SubRequest
+from _pytest.outcomes import Exit
+from _pytest.python import CallSpec2
+from _pytest.scope import Scope
+
+from . import get_available_backends
+from ._core._eventloop import (
+ current_async_library,
+ get_async_backend,
+ reset_current_async_library,
+ set_current_async_library,
+)
+from ._core._exceptions import iterate_exceptions
+from .abc import TestRunner
+
+if sys.version_info < (3, 11):
+ from exceptiongroup import ExceptionGroup
+
+_current_runner: TestRunner | None = None
+_runner_stack: ExitStack | None = None
+_runner_leases = 0
+
+
+def extract_backend_and_options(backend: object) -> tuple[str, dict[str, Any]]:
+ if isinstance(backend, str):
+ return backend, {}
+ elif isinstance(backend, tuple) and len(backend) == 2:
+ if isinstance(backend[0], str) and isinstance(backend[1], dict):
+ return cast(tuple[str, dict[str, Any]], backend)
+
+ raise TypeError("anyio_backend must be either a string or tuple of (string, dict)")
+
+
+@contextmanager
+def get_runner(
+ backend_name: str, backend_options: dict[str, Any]
+) -> Iterator[TestRunner]:
+ global _current_runner, _runner_leases, _runner_stack
+ if _current_runner is None:
+ asynclib = get_async_backend(backend_name)
+ _runner_stack = ExitStack()
+ if current_async_library() is None:
+ # Since we're in control of the event loop, we can cache the name of the
+ # async library
+ token = set_current_async_library(backend_name)
+ _runner_stack.callback(reset_current_async_library, token)
+
+ backend_options = backend_options or {}
+ _current_runner = _runner_stack.enter_context(
+ asynclib.create_test_runner(backend_options)
+ )
+
+ _runner_leases += 1
+ try:
+ yield _current_runner
+ finally:
+ _runner_leases -= 1
+ if not _runner_leases:
+ assert _runner_stack is not None
+ _runner_stack.close()
+ _runner_stack = _current_runner = None
+
+
+def pytest_addoption(parser: pytest.Parser) -> None:
+ parser.addini(
+ "anyio_mode",
+ default="strict",
+ help='AnyIO plugin mode (either "strict" or "auto")',
+ )
+
+
+def pytest_configure(config: pytest.Config) -> None:
+ config.addinivalue_line(
+ "markers",
+ "anyio: mark the (coroutine function) test to be run asynchronously via anyio.",
+ )
+ if (
+ config.getini("anyio_mode") == "auto"
+ and config.pluginmanager.has_plugin("asyncio")
+ and config.getini("asyncio_mode") == "auto"
+ ):
+ config.issue_config_time_warning(
+ pytest.PytestConfigWarning(
+ "AnyIO auto mode has been enabled together with pytest-asyncio auto "
+ "mode. This may cause unexpected behavior."
+ ),
+ 1,
+ )
+
+
+@pytest.hookimpl(hookwrapper=True)
+def pytest_fixture_setup(fixturedef: Any, request: Any) -> Generator[Any]:
+ def wrapper(anyio_backend: Any, request: SubRequest, **kwargs: Any) -> Any:
+ # Rebind any fixture methods to the request instance
+ if (
+ request.instance
+ and ismethod(func)
+ and type(func.__self__) is type(request.instance)
+ ):
+ local_func = func.__func__.__get__(request.instance)
+ else:
+ local_func = func
+
+ backend_name, backend_options = extract_backend_and_options(anyio_backend)
+ if has_backend_arg:
+ kwargs["anyio_backend"] = anyio_backend
+
+ if has_request_arg:
+ kwargs["request"] = request
+
+ with get_runner(backend_name, backend_options) as runner:
+ # re-entrant call into the test runner detected. this happens when an async fixture
+ # is dynamically requested via request.getfixturevalue() from inside a running async
+ # test or fixture. on asyncio this raises RuntimeError: This event loop is already
+ # running, on trio the runner deadlocks - the host loop blocks waiting for the
+ # coroutine to return, but the coroutine is waiting for the host loop. raising here
+ # prevents the hang and gives a consistent error across backends.
+ if runner.is_running():
+ raise RuntimeError(
+ "Cannot schedule a coroutine in the test runner while another is already running; "
+ "likely caused by request.getfixturevalue() on an async fixture."
+ )
+
+ if isasyncgenfunction(local_func):
+ yield from runner.run_asyncgen_fixture(local_func, kwargs)
+ else:
+ yield runner.run_fixture(local_func, kwargs)
+
+ # Only apply this to coroutine functions and async generator functions in requests
+ # that involve the anyio_backend fixture
+ func = fixturedef.func
+ if isasyncgenfunction(func) or iscoroutinefunction(func):
+ if "anyio_backend" in request.fixturenames:
+ fixturedef.func = wrapper
+ original_argname = fixturedef.argnames
+
+ if not (has_backend_arg := "anyio_backend" in fixturedef.argnames):
+ fixturedef.argnames += ("anyio_backend",)
+
+ if not (has_request_arg := "request" in fixturedef.argnames):
+ fixturedef.argnames += ("request",)
+
+ try:
+ return (yield)
+ finally:
+ fixturedef.func = func
+ fixturedef.argnames = original_argname
+
+ return (yield)
+
+
+@pytest.hookimpl(tryfirst=True)
+def pytest_pycollect_makeitem(
+ collector: pytest.Module | pytest.Class, name: str, obj: object
+) -> None:
+ if collector.istestfunction(obj, name):
+ inner_func = obj.hypothesis.inner_test if hasattr(obj, "hypothesis") else obj
+ if iscoroutinefunction(inner_func):
+ anyio_auto_mode = collector.config.getini("anyio_mode") == "auto"
+ marker = collector.get_closest_marker("anyio")
+ own_markers = getattr(obj, "pytestmark", ())
+ if (
+ anyio_auto_mode
+ or marker
+ or any(marker.name == "anyio" for marker in own_markers)
+ ):
+ pytest.mark.usefixtures("anyio_backend")(obj)
+
+
+def pytest_collection_finish(session: pytest.Session) -> None:
+ for i, item in reversed(list(enumerate(session.items))):
+ if (
+ isinstance(item, pytest.Function)
+ and iscoroutinefunction(item.function)
+ and item.get_closest_marker("anyio") is not None
+ and "anyio_backend" not in item.fixturenames
+ ):
+ new_items = []
+ try:
+ cs_fields = {f.name for f in dataclasses.fields(CallSpec2)}
+ except TypeError:
+ cs_fields = set()
+
+ for param_index, backend in enumerate(get_available_backends()):
+ if "_arg2scope" in cs_fields: # pytest >= 8
+ callspec = CallSpec2(
+ params={"anyio_backend": backend},
+ indices={"anyio_backend": param_index},
+ _arg2scope={"anyio_backend": Scope.Module},
+ _idlist=[backend],
+ marks=[],
+ )
+ else: # pytest 7.x
+ callspec = CallSpec2( # type: ignore[call-arg]
+ funcargs={},
+ params={"anyio_backend": backend},
+ indices={"anyio_backend": param_index},
+ arg2scope={"anyio_backend": Scope.Module},
+ idlist=[backend],
+ marks=[],
+ )
+
+ fi = item._fixtureinfo
+ new_names_closure = list(fi.names_closure)
+ if "anyio_backend" not in new_names_closure:
+ new_names_closure.append("anyio_backend")
+
+ new_fixtureinfo = FuncFixtureInfo(
+ argnames=fi.argnames,
+ initialnames=fi.initialnames,
+ names_closure=new_names_closure,
+ name2fixturedefs=fi.name2fixturedefs,
+ )
+ new_item = pytest.Function.from_parent(
+ item.parent,
+ name=f"{item.originalname}[{backend}]",
+ callspec=callspec,
+ callobj=item.obj,
+ fixtureinfo=new_fixtureinfo,
+ keywords=item.keywords,
+ originalname=item.originalname,
+ )
+ new_items.append(new_item)
+
+ session.items[i : i + 1] = new_items
+
+
+@pytest.hookimpl(tryfirst=True)
+def pytest_pyfunc_call(pyfuncitem: Any) -> bool | None:
+ def run_with_hypothesis(**kwargs: Any) -> None:
+ with get_runner(backend_name, backend_options) as runner:
+ runner.run_test(original_func, kwargs)
+
+ backend = pyfuncitem.funcargs.get("anyio_backend")
+ if backend:
+ backend_name, backend_options = extract_backend_and_options(backend)
+
+ if hasattr(pyfuncitem.obj, "hypothesis"):
+ # Wrap the inner test function unless it's already wrapped
+ original_func = pyfuncitem.obj.hypothesis.inner_test
+ if original_func.__qualname__ != run_with_hypothesis.__qualname__:
+ if iscoroutinefunction(original_func):
+ pyfuncitem.obj.hypothesis.inner_test = run_with_hypothesis
+
+ return None
+
+ if iscoroutinefunction(pyfuncitem.obj):
+ funcargs = pyfuncitem.funcargs
+ testargs = {arg: funcargs[arg] for arg in pyfuncitem._fixtureinfo.argnames}
+ with get_runner(backend_name, backend_options) as runner:
+ try:
+ runner.run_test(pyfuncitem.obj, testargs)
+ except ExceptionGroup as excgrp:
+ for exc in iterate_exceptions(excgrp):
+ if isinstance(exc, (Exit, KeyboardInterrupt, SystemExit)):
+ raise exc from excgrp
+
+ raise
+
+ return True
+
+ return None
+
+
+@pytest.fixture(scope="module", params=get_available_backends())
+def anyio_backend(request: Any) -> Any:
+ return request.param
+
+
+@pytest.fixture
+def anyio_backend_name(anyio_backend: Any) -> str:
+ if isinstance(anyio_backend, str):
+ return anyio_backend
+ else:
+ return anyio_backend[0]
+
+
+@pytest.fixture
+def anyio_backend_options(anyio_backend: Any) -> dict[str, Any]:
+ if isinstance(anyio_backend, str):
+ return {}
+ else:
+ return anyio_backend[1]
+
+
+class FreePortFactory:
+ """
+ Manages port generation based on specified socket kind, ensuring no duplicate
+ ports are generated.
+
+ This class provides functionality for generating available free ports on the
+ system. It is initialized with a specific socket kind and can generate ports
+ for given address families while avoiding reuse of previously generated ports.
+
+ Users should not instantiate this class directly, but use the
+ ``free_tcp_port_factory`` and ``free_udp_port_factory`` fixtures instead. For simple
+ uses cases, ``free_tcp_port`` and ``free_udp_port`` can be used instead.
+ """
+
+ def __init__(self, kind: socket.SocketKind) -> None:
+ self._kind = kind
+ self._generated = set[int]()
+
+ @property
+ def kind(self) -> socket.SocketKind:
+ """
+ The type of socket connection (e.g., :data:`~socket.SOCK_STREAM` or
+ :data:`~socket.SOCK_DGRAM`) used to bind for checking port availability
+
+ """
+ return self._kind
+
+ def __call__(self, family: socket.AddressFamily | None = None) -> int:
+ """
+ Return an unbound port for the given address family.
+
+ :param family: if omitted, both IPv4 and IPv6 addresses will be tried
+ :return: a port number
+
+ """
+ if family is not None:
+ families = [family]
+ else:
+ families = [socket.AF_INET]
+ if socket.has_ipv6:
+ families.append(socket.AF_INET6)
+
+ while True:
+ port = 0
+ with ExitStack() as stack:
+ for family in families:
+ sock = stack.enter_context(socket.socket(family, self._kind))
+ addr = "::1" if family == socket.AF_INET6 else "127.0.0.1"
+ try:
+ sock.bind((addr, port))
+ except OSError:
+ break
+
+ if not port:
+ port = sock.getsockname()[1]
+ else:
+ if port not in self._generated:
+ self._generated.add(port)
+ return port
+
+
+@pytest.fixture(scope="session")
+def free_tcp_port_factory() -> FreePortFactory:
+ return FreePortFactory(socket.SOCK_STREAM)
+
+
+@pytest.fixture(scope="session")
+def free_udp_port_factory() -> FreePortFactory:
+ return FreePortFactory(socket.SOCK_DGRAM)
+
+
+@pytest.fixture
+def free_tcp_port(free_tcp_port_factory: Callable[[], int]) -> int:
+ return free_tcp_port_factory()
+
+
+@pytest.fixture
+def free_udp_port(free_udp_port_factory: Callable[[], int]) -> int:
+ return free_udp_port_factory()
diff --git a/venv/lib/python3.11/site-packages/anyio/streams/__init__.py b/venv/lib/python3.11/site-packages/anyio/streams/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/anyio/streams/buffered.py b/venv/lib/python3.11/site-packages/anyio/streams/buffered.py
new file mode 100644
index 0000000000000000000000000000000000000000..a3b07f73ac44d5876c239fd0317846564ef06f91
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/streams/buffered.py
@@ -0,0 +1,201 @@
+from __future__ import annotations
+
+__all__ = (
+ "BufferedByteReceiveStream",
+ "BufferedByteStream",
+ "BufferedConnectable",
+)
+
+import sys
+from collections.abc import Callable, Iterable, Mapping
+from dataclasses import dataclass, field
+from typing import Any, SupportsIndex
+
+from .. import ClosedResourceError, DelimiterNotFound, EndOfStream, IncompleteRead
+from ..abc import (
+ AnyByteReceiveStream,
+ AnyByteStream,
+ AnyByteStreamConnectable,
+ ByteReceiveStream,
+ ByteStream,
+ ByteStreamConnectable,
+)
+
+if sys.version_info >= (3, 12):
+ from typing import override
+else:
+ from typing_extensions import override
+
+
+@dataclass(eq=False)
+class BufferedByteReceiveStream(ByteReceiveStream):
+ """
+ Wraps any bytes-based receive stream and uses a buffer to provide sophisticated
+ receiving capabilities in the form of a byte stream.
+ """
+
+ receive_stream: AnyByteReceiveStream
+ _buffer: bytearray = field(init=False, default_factory=bytearray)
+ _closed: bool = field(init=False, default=False)
+
+ async def aclose(self) -> None:
+ await self.receive_stream.aclose()
+ self._closed = True
+
+ @property
+ def buffer(self) -> bytes:
+ """The bytes currently in the buffer."""
+ return bytes(self._buffer)
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return self.receive_stream.extra_attributes
+
+ def feed_data(self, data: Iterable[SupportsIndex], /) -> None:
+ """
+ Append data directly into the buffer.
+
+ Any data in the buffer will be consumed by receive operations before receiving
+ anything from the wrapped stream.
+
+ :param data: the data to append to the buffer (can be bytes or anything else
+ that supports ``__index__()``)
+
+ """
+ self._buffer.extend(data)
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ if self._closed:
+ raise ClosedResourceError
+
+ if self._buffer:
+ chunk = bytes(self._buffer[:max_bytes])
+ del self._buffer[:max_bytes]
+ return chunk
+ elif isinstance(self.receive_stream, ByteReceiveStream):
+ return await self.receive_stream.receive(max_bytes)
+ else:
+ # With a bytes-oriented object stream, we need to handle any surplus bytes
+ # we get from the receive() call
+ chunk = await self.receive_stream.receive()
+ if len(chunk) > max_bytes:
+ # Save the surplus bytes in the buffer
+ self._buffer.extend(chunk[max_bytes:])
+ return chunk[:max_bytes]
+ else:
+ return chunk
+
+ async def receive_exactly(self, nbytes: int) -> bytes:
+ """
+ Read exactly the given amount of bytes from the stream.
+
+ :param nbytes: the number of bytes to read
+ :return: the bytes read
+ :raises ~anyio.IncompleteRead: if the stream was closed before the requested
+ amount of bytes could be read from the stream
+
+ """
+ while True:
+ remaining = nbytes - len(self._buffer)
+ if remaining <= 0:
+ retval = self._buffer[:nbytes]
+ del self._buffer[:nbytes]
+ return bytes(retval)
+
+ try:
+ if isinstance(self.receive_stream, ByteReceiveStream):
+ chunk = await self.receive_stream.receive(remaining)
+ else:
+ chunk = await self.receive_stream.receive()
+ except EndOfStream as exc:
+ raise IncompleteRead from exc
+
+ self._buffer.extend(chunk)
+
+ async def receive_until(self, delimiter: bytes, max_bytes: int) -> bytes:
+ """
+ Read from the stream until the delimiter is found or max_bytes have been read.
+
+ :param delimiter: the marker to look for in the stream
+ :param max_bytes: maximum number of bytes that will be read before raising
+ :exc:`~anyio.DelimiterNotFound`
+ :return: the bytes read (not including the delimiter)
+ :raises ~anyio.IncompleteRead: if the stream was closed before the delimiter
+ was found
+ :raises ~anyio.DelimiterNotFound: if the delimiter is not found within the
+ bytes read up to the maximum allowed
+
+ """
+ delimiter_size = len(delimiter)
+ offset = 0
+ while True:
+ # Check if the delimiter can be found in the current buffer
+ index = self._buffer.find(delimiter, offset)
+ if index >= 0:
+ found = self._buffer[:index]
+ del self._buffer[: index + len(delimiter) :]
+ return bytes(found)
+
+ # Check if the buffer is already at or over the limit
+ if len(self._buffer) >= max_bytes:
+ raise DelimiterNotFound(max_bytes)
+
+ # Read more data into the buffer from the socket
+ try:
+ data = await self.receive_stream.receive()
+ except EndOfStream as exc:
+ raise IncompleteRead from exc
+
+ # Move the offset forward and add the new data to the buffer
+ offset = max(len(self._buffer) - delimiter_size + 1, 0)
+ self._buffer.extend(data)
+
+
+class BufferedByteStream(BufferedByteReceiveStream, ByteStream):
+ """
+ A full-duplex variant of :class:`BufferedByteReceiveStream`. All writes are passed
+ through to the wrapped stream as-is.
+ """
+
+ def __init__(self, stream: AnyByteStream):
+ """
+ :param stream: the stream to be wrapped
+
+ """
+ super().__init__(stream)
+ self._stream = stream
+
+ @override
+ async def send_eof(self) -> None:
+ await self._stream.send_eof()
+
+ @override
+ async def send(self, item: bytes) -> None:
+ await self._stream.send(item)
+
+
+class BufferedConnectable(ByteStreamConnectable):
+ """
+ Wraps a byte stream connectable to produce :class:`BufferedByteStream` connections.
+
+ Use this when you want the streams returned by :meth:`connect` to have the buffered
+ receive API (e.g. :meth:`~BufferedByteReceiveStream.receive_exactly` and
+ :meth:`~BufferedByteReceiveStream.receive_until`).
+
+ :param connectable: the byte stream connectable to wrap
+ """
+
+ def __init__(self, connectable: AnyByteStreamConnectable):
+ """
+ :param connectable: the connectable to wrap
+
+ """
+ self.connectable = connectable
+
+ @override
+ async def connect(self) -> BufferedByteStream:
+ stream = await self.connectable.connect()
+ return BufferedByteStream(stream)
diff --git a/venv/lib/python3.11/site-packages/anyio/streams/file.py b/venv/lib/python3.11/site-packages/anyio/streams/file.py
new file mode 100644
index 0000000000000000000000000000000000000000..b9d7ca4f3341a7ad7e44553e4e15134f223cb750
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/streams/file.py
@@ -0,0 +1,157 @@
+from __future__ import annotations
+
+__all__ = (
+ "FileReadStream",
+ "FileStreamAttribute",
+ "FileWriteStream",
+)
+
+from collections.abc import Callable, Mapping
+from io import SEEK_SET, UnsupportedOperation
+from os import PathLike
+from pathlib import Path
+from typing import IO, Any
+
+from .. import (
+ BrokenResourceError,
+ ClosedResourceError,
+ EndOfStream,
+ TypedAttributeSet,
+ to_thread,
+ typed_attribute,
+)
+from ..abc import ByteReceiveStream, ByteSendStream
+
+
+class FileStreamAttribute(TypedAttributeSet):
+ #: the open file descriptor
+ file: IO[bytes] = typed_attribute()
+ #: the path of the file on the file system, if available (file must be a real file)
+ path: Path = typed_attribute()
+ #: the file number, if available (file must be a real file or a TTY)
+ fileno: int = typed_attribute()
+
+
+class _BaseFileStream:
+ def __init__(self, file: IO[bytes]):
+ self._file = file
+
+ async def aclose(self) -> None:
+ await to_thread.run_sync(self._file.close)
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ attributes: dict[Any, Callable[[], Any]] = {
+ FileStreamAttribute.file: lambda: self._file,
+ }
+
+ if hasattr(self._file, "name"):
+ attributes[FileStreamAttribute.path] = lambda: Path(self._file.name)
+
+ try:
+ self._file.fileno()
+ except UnsupportedOperation:
+ pass
+ else:
+ attributes[FileStreamAttribute.fileno] = lambda: self._file.fileno()
+
+ return attributes
+
+
+class FileReadStream(_BaseFileStream, ByteReceiveStream):
+ """
+ A byte stream that reads from a file in the file system.
+
+ :param file: a file that has been opened for reading in binary mode
+
+ .. versionadded:: 3.0
+ """
+
+ @classmethod
+ async def from_path(cls, path: str | PathLike[str]) -> FileReadStream:
+ """
+ Create a file read stream by opening the given file.
+
+ :param path: path of the file to read from
+
+ """
+ file = await to_thread.run_sync(Path(path).open, "rb")
+ return cls(file)
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ try:
+ data = await to_thread.run_sync(self._file.read, max_bytes)
+ except ValueError:
+ raise ClosedResourceError from None
+ except OSError as exc:
+ raise BrokenResourceError from exc
+
+ if data:
+ return data
+ else:
+ raise EndOfStream
+
+ async def seek(self, position: int, whence: int = SEEK_SET) -> int:
+ """
+ Seek the file to the given position.
+
+ .. seealso:: :meth:`io.IOBase.seek`
+
+ .. note:: Not all file descriptors are seekable.
+
+ :param position: position to seek the file to
+ :param whence: controls how ``position`` is interpreted
+ :return: the new absolute position
+ :raises OSError: if the file is not seekable
+
+ """
+ return await to_thread.run_sync(self._file.seek, position, whence)
+
+ async def tell(self) -> int:
+ """
+ Return the current stream position.
+
+ .. note:: Not all file descriptors are seekable.
+
+ :return: the current absolute position
+ :raises OSError: if the file is not seekable
+
+ """
+ return await to_thread.run_sync(self._file.tell)
+
+
+class FileWriteStream(_BaseFileStream, ByteSendStream):
+ """
+ A byte stream that writes to a file in the file system.
+
+ :param file: a file that has been opened for writing in binary mode
+
+ .. versionadded:: 3.0
+ """
+
+ @classmethod
+ async def from_path(
+ cls, path: str | PathLike[str], append: bool = False
+ ) -> FileWriteStream:
+ """
+ Create a file write stream by opening the given file for writing.
+
+ :param path: path of the file to write to
+ :param append: if ``True``, open the file for appending; if ``False``, any
+ existing file at the given path will be truncated
+
+ """
+ mode = "ab" if append else "wb"
+ file = await to_thread.run_sync(Path(path).open, mode)
+ return cls(file)
+
+ async def send(self, item: bytes) -> None:
+ try:
+ await to_thread.run_sync(self._file.write, item)
+ except ValueError:
+ raise ClosedResourceError from None
+ except OSError as exc:
+ raise BrokenResourceError from exc
diff --git a/venv/lib/python3.11/site-packages/anyio/streams/memory.py b/venv/lib/python3.11/site-packages/anyio/streams/memory.py
new file mode 100644
index 0000000000000000000000000000000000000000..d8c205b82636d40064445e68748dc54682ff3e97
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/streams/memory.py
@@ -0,0 +1,326 @@
+from __future__ import annotations
+
+__all__ = (
+ "MemoryObjectReceiveStream",
+ "MemoryObjectSendStream",
+ "MemoryObjectStreamStatistics",
+)
+
+import warnings
+from collections import OrderedDict, deque
+from dataclasses import dataclass, field
+from types import TracebackType
+from typing import Generic, NamedTuple, TypeVar
+
+from .. import (
+ BrokenResourceError,
+ ClosedResourceError,
+ EndOfStream,
+ WouldBlock,
+)
+from .._core._synchronization import Event
+from .._core._testing import TaskInfo, get_current_task
+from ..abc import ObjectReceiveStream, ObjectSendStream
+from ..lowlevel import checkpoint
+
+T_Item = TypeVar("T_Item")
+T_co = TypeVar("T_co", covariant=True)
+T_contra = TypeVar("T_contra", contravariant=True)
+
+
+class MemoryObjectStreamStatistics(NamedTuple):
+ current_buffer_used: int #: number of items stored in the buffer
+ #: maximum number of items that can be stored on this stream (or :data:`math.inf`)
+ max_buffer_size: float
+ open_send_streams: int #: number of unclosed clones of the send stream
+ open_receive_streams: int #: number of unclosed clones of the receive stream
+ #: number of tasks blocked on :meth:`MemoryObjectSendStream.send`
+ tasks_waiting_send: int
+ #: number of tasks blocked on :meth:`MemoryObjectReceiveStream.receive`
+ tasks_waiting_receive: int
+
+
+@dataclass(eq=False)
+class _MemoryObjectItemReceiver(Generic[T_Item]):
+ task_info: TaskInfo = field(init=False, default_factory=get_current_task)
+ item: T_Item = field(init=False)
+
+ def __repr__(self) -> str:
+ # When item is not defined, we get following error with default __repr__:
+ # AttributeError: 'MemoryObjectItemReceiver' object has no attribute 'item'
+ item = getattr(self, "item", None)
+ return f"{self.__class__.__name__}(task_info={self.task_info}, item={item!r})"
+
+
+@dataclass(eq=False)
+class _MemoryObjectStreamState(Generic[T_Item]):
+ max_buffer_size: float = field()
+ buffer: deque[T_Item] = field(init=False, default_factory=deque)
+ open_send_channels: int = field(init=False, default=0)
+ open_receive_channels: int = field(init=False, default=0)
+ waiting_receivers: OrderedDict[Event, _MemoryObjectItemReceiver[T_Item]] = field(
+ init=False, default_factory=OrderedDict
+ )
+ waiting_senders: OrderedDict[Event, T_Item] = field(
+ init=False, default_factory=OrderedDict
+ )
+
+ def statistics(self) -> MemoryObjectStreamStatistics:
+ return MemoryObjectStreamStatistics(
+ len(self.buffer),
+ self.max_buffer_size,
+ self.open_send_channels,
+ self.open_receive_channels,
+ len(self.waiting_senders),
+ len(self.waiting_receivers),
+ )
+
+
+@dataclass(eq=False)
+class MemoryObjectReceiveStream(Generic[T_co], ObjectReceiveStream[T_co]):
+ _state: _MemoryObjectStreamState[T_co]
+ _closed: bool = field(init=False, default=False)
+
+ def __post_init__(self) -> None:
+ self._state.open_receive_channels += 1
+
+ def receive_nowait(self) -> T_co:
+ """
+ Receive the next item if it can be done without waiting.
+
+ :return: the received item
+ :raises ~anyio.ClosedResourceError: if this send stream has been closed
+ :raises ~anyio.EndOfStream: if the buffer is empty and this stream has been
+ closed from the sending end
+ :raises ~anyio.WouldBlock: if there are no items in the buffer and no tasks
+ waiting to send
+
+ """
+ if self._closed:
+ raise ClosedResourceError
+
+ if self._state.waiting_senders:
+ # Get the item from the next sender
+ send_event, item = self._state.waiting_senders.popitem(last=False)
+ self._state.buffer.append(item)
+ send_event.set()
+
+ if self._state.buffer:
+ return self._state.buffer.popleft()
+ elif not self._state.open_send_channels:
+ raise EndOfStream
+
+ raise WouldBlock
+
+ async def receive(self) -> T_co:
+ await checkpoint()
+ try:
+ return self.receive_nowait()
+ except WouldBlock:
+ # Add ourselves in the queue
+ receive_event = Event()
+ receiver = _MemoryObjectItemReceiver[T_co]()
+ self._state.waiting_receivers[receive_event] = receiver
+
+ try:
+ await receive_event.wait()
+ finally:
+ self._state.waiting_receivers.pop(receive_event, None)
+
+ try:
+ return receiver.item
+ except AttributeError:
+ raise EndOfStream from None
+
+ def clone(self) -> MemoryObjectReceiveStream[T_co]:
+ """
+ Create a clone of this receive stream.
+
+ Each clone can be closed separately. Only when all clones have been closed will
+ the receiving end of the memory stream be considered closed by the sending ends.
+
+ :return: the cloned stream
+
+ """
+ if self._closed:
+ raise ClosedResourceError
+
+ return MemoryObjectReceiveStream(_state=self._state)
+
+ def close(self) -> None:
+ """
+ Close the stream.
+
+ This works the exact same way as :meth:`aclose`, but is provided as a special
+ case for the benefit of synchronous callbacks.
+
+ """
+ if not self._closed:
+ self._closed = True
+ self._state.open_receive_channels -= 1
+ if self._state.open_receive_channels == 0:
+ send_events = list(self._state.waiting_senders.keys())
+ for event in send_events:
+ event.set()
+
+ async def aclose(self) -> None:
+ self.close()
+
+ def statistics(self) -> MemoryObjectStreamStatistics:
+ """
+ Return statistics about the current state of this stream.
+
+ .. versionadded:: 3.0
+ """
+ return self._state.statistics()
+
+ def __enter__(self) -> MemoryObjectReceiveStream[T_co]:
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self.close()
+
+ def __del__(self) -> None:
+ if not self._closed:
+ warnings.warn(
+ f"Unclosed <{self.__class__.__name__} at {id(self):x}>",
+ ResourceWarning,
+ stacklevel=1,
+ source=self,
+ )
+
+
+@dataclass(eq=False)
+class MemoryObjectSendStream(Generic[T_contra], ObjectSendStream[T_contra]):
+ _state: _MemoryObjectStreamState[T_contra]
+ _closed: bool = field(init=False, default=False)
+
+ def __post_init__(self) -> None:
+ self._state.open_send_channels += 1
+
+ def send_nowait(self, item: T_contra) -> None:
+ """
+ Send an item immediately if it can be done without waiting.
+
+ :param item: the item to send
+ :raises ~anyio.ClosedResourceError: if this send stream has been closed
+ :raises ~anyio.BrokenResourceError: if the stream has been closed from the
+ receiving end
+ :raises ~anyio.WouldBlock: if the buffer is full and there are no tasks waiting
+ to receive
+
+ """
+ if self._closed:
+ raise ClosedResourceError
+ if not self._state.open_receive_channels:
+ raise BrokenResourceError
+
+ while self._state.waiting_receivers:
+ receive_event, receiver = self._state.waiting_receivers.popitem(last=False)
+ if not receiver.task_info.has_pending_cancellation():
+ receiver.item = item
+ receive_event.set()
+ return
+
+ if len(self._state.buffer) < self._state.max_buffer_size:
+ self._state.buffer.append(item)
+ else:
+ raise WouldBlock
+
+ async def send(self, item: T_contra) -> None:
+ """
+ Send an item to the stream.
+
+ If the buffer is full, this method blocks until there is again room in the
+ buffer or the item can be sent directly to a receiver.
+
+ :param item: the item to send
+ :raises ~anyio.ClosedResourceError: if this send stream has been closed
+ :raises ~anyio.BrokenResourceError: if the stream has been closed from the
+ receiving end
+
+ """
+ await checkpoint()
+ try:
+ self.send_nowait(item)
+ except WouldBlock:
+ # Wait until there's someone on the receiving end
+ send_event = Event()
+ self._state.waiting_senders[send_event] = item
+ try:
+ await send_event.wait()
+ except BaseException:
+ self._state.waiting_senders.pop(send_event, None)
+ raise
+
+ if send_event in self._state.waiting_senders:
+ del self._state.waiting_senders[send_event]
+ raise BrokenResourceError from None
+
+ def clone(self) -> MemoryObjectSendStream[T_contra]:
+ """
+ Create a clone of this send stream.
+
+ Each clone can be closed separately. Only when all clones have been closed will
+ the sending end of the memory stream be considered closed by the receiving ends.
+
+ :return: the cloned stream
+
+ """
+ if self._closed:
+ raise ClosedResourceError
+
+ return MemoryObjectSendStream(_state=self._state)
+
+ def close(self) -> None:
+ """
+ Close the stream.
+
+ This works the exact same way as :meth:`aclose`, but is provided as a special
+ case for the benefit of synchronous callbacks.
+
+ """
+ if not self._closed:
+ self._closed = True
+ self._state.open_send_channels -= 1
+ if self._state.open_send_channels == 0:
+ receive_events = list(self._state.waiting_receivers.keys())
+ self._state.waiting_receivers.clear()
+ for event in receive_events:
+ event.set()
+
+ async def aclose(self) -> None:
+ self.close()
+
+ def statistics(self) -> MemoryObjectStreamStatistics:
+ """
+ Return statistics about the current state of this stream.
+
+ .. versionadded:: 3.0
+ """
+ return self._state.statistics()
+
+ def __enter__(self) -> MemoryObjectSendStream[T_contra]:
+ return self
+
+ def __exit__(
+ self,
+ exc_type: type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> None:
+ self.close()
+
+ def __del__(self) -> None:
+ if not self._closed:
+ warnings.warn(
+ f"Unclosed <{self.__class__.__name__} at {id(self):x}>",
+ ResourceWarning,
+ stacklevel=1,
+ source=self,
+ )
diff --git a/venv/lib/python3.11/site-packages/anyio/streams/stapled.py b/venv/lib/python3.11/site-packages/anyio/streams/stapled.py
new file mode 100644
index 0000000000000000000000000000000000000000..0a3c53da25ab33e502ac047b619b3f5ebb3c93d3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/streams/stapled.py
@@ -0,0 +1,150 @@
+from __future__ import annotations
+
+__all__ = (
+ "MultiListener",
+ "StapledByteStream",
+ "StapledObjectStream",
+)
+
+from collections.abc import Callable, Mapping, Sequence
+from dataclasses import dataclass
+from typing import Any, Generic, TypeVar
+
+from ..abc import (
+ ByteReceiveStream,
+ ByteSendStream,
+ ByteStream,
+ Listener,
+ ObjectReceiveStream,
+ ObjectSendStream,
+ ObjectStream,
+ TaskGroup,
+)
+
+T_Item = TypeVar("T_Item")
+T_Stream = TypeVar("T_Stream")
+
+
+@dataclass(eq=False)
+class StapledByteStream(ByteStream):
+ """
+ Combines two byte streams into a single, bidirectional byte stream.
+
+ Extra attributes will be provided from both streams, with the receive stream
+ providing the values in case of a conflict.
+
+ :param ByteSendStream send_stream: the sending byte stream
+ :param ByteReceiveStream receive_stream: the receiving byte stream
+ """
+
+ send_stream: ByteSendStream
+ receive_stream: ByteReceiveStream
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ return await self.receive_stream.receive(max_bytes)
+
+ async def send(self, item: bytes) -> None:
+ await self.send_stream.send(item)
+
+ async def send_eof(self) -> None:
+ await self.send_stream.aclose()
+
+ async def aclose(self) -> None:
+ await self.send_stream.aclose()
+ await self.receive_stream.aclose()
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return {
+ **self.send_stream.extra_attributes,
+ **self.receive_stream.extra_attributes,
+ }
+
+
+@dataclass(eq=False)
+class StapledObjectStream(Generic[T_Item], ObjectStream[T_Item]):
+ """
+ Combines two object streams into a single, bidirectional object stream.
+
+ Extra attributes will be provided from both streams, with the receive stream
+ providing the values in case of a conflict.
+
+ :param ObjectSendStream send_stream: the sending object stream
+ :param ObjectReceiveStream receive_stream: the receiving object stream
+ """
+
+ send_stream: ObjectSendStream[T_Item]
+ receive_stream: ObjectReceiveStream[T_Item]
+
+ async def receive(self) -> T_Item:
+ return await self.receive_stream.receive()
+
+ async def send(self, item: T_Item) -> None:
+ await self.send_stream.send(item)
+
+ async def send_eof(self) -> None:
+ await self.send_stream.aclose()
+
+ async def aclose(self) -> None:
+ await self.send_stream.aclose()
+ await self.receive_stream.aclose()
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return {
+ **self.send_stream.extra_attributes,
+ **self.receive_stream.extra_attributes,
+ }
+
+
+@dataclass(eq=False)
+class MultiListener(Generic[T_Stream], Listener[T_Stream]):
+ """
+ Combines multiple listeners into one, serving connections from all of them at once.
+
+ Any MultiListeners in the given collection of listeners will have their listeners
+ moved into this one.
+
+ Extra attributes are provided from each listener, with each successive listener
+ overriding any conflicting attributes from the previous one.
+
+ :param listeners: listeners to serve
+ :type listeners: Sequence[Listener[T_Stream]]
+ """
+
+ listeners: Sequence[Listener[T_Stream]]
+
+ def __post_init__(self) -> None:
+ listeners: list[Listener[T_Stream]] = []
+ for listener in self.listeners:
+ if isinstance(listener, MultiListener):
+ listeners.extend(listener.listeners)
+ del listener.listeners[:] # type: ignore[attr-defined]
+ else:
+ listeners.append(listener)
+
+ self.listeners = listeners
+
+ async def serve(
+ self, handler: Callable[[T_Stream], Any], task_group: TaskGroup | None = None
+ ) -> None:
+ from .. import create_task_group
+
+ async with create_task_group() as tg:
+ for listener in self.listeners:
+ tg.start_soon(listener.serve, handler, task_group)
+
+ async def aclose(self) -> None:
+ for listener in self.listeners:
+ await listener.aclose()
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ attributes: dict = {}
+ for listener in self.listeners:
+ attributes.update(listener.extra_attributes)
+
+ return attributes
diff --git a/venv/lib/python3.11/site-packages/anyio/streams/text.py b/venv/lib/python3.11/site-packages/anyio/streams/text.py
new file mode 100644
index 0000000000000000000000000000000000000000..296cd250459f3848bb333301fff1ac32973f219a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/streams/text.py
@@ -0,0 +1,176 @@
+from __future__ import annotations
+
+__all__ = (
+ "TextConnectable",
+ "TextReceiveStream",
+ "TextSendStream",
+ "TextStream",
+)
+
+import codecs
+import sys
+from collections.abc import Callable, Mapping
+from dataclasses import InitVar, dataclass, field
+from typing import Any
+
+from ..abc import (
+ AnyByteReceiveStream,
+ AnyByteSendStream,
+ AnyByteStream,
+ AnyByteStreamConnectable,
+ ObjectReceiveStream,
+ ObjectSendStream,
+ ObjectStream,
+ ObjectStreamConnectable,
+)
+
+if sys.version_info >= (3, 12):
+ from typing import override
+else:
+ from typing_extensions import override
+
+
+@dataclass(eq=False)
+class TextReceiveStream(ObjectReceiveStream[str]):
+ """
+ Stream wrapper that decodes bytes to strings using the given encoding.
+
+ Decoding is done using :class:`~codecs.IncrementalDecoder` which returns any
+ completely received unicode characters as soon as they come in.
+
+ :param transport_stream: any bytes-based receive stream
+ :param encoding: character encoding to use for decoding bytes to strings (defaults
+ to ``utf-8``)
+ :param errors: handling scheme for decoding errors (defaults to ``strict``; see the
+ `codecs module documentation`_ for a comprehensive list of options)
+
+ .. _codecs module documentation:
+ https://docs.python.org/3/library/codecs.html#codec-objects
+ """
+
+ transport_stream: AnyByteReceiveStream
+ encoding: InitVar[str] = "utf-8"
+ errors: InitVar[str] = "strict"
+ _decoder: codecs.IncrementalDecoder = field(init=False)
+
+ def __post_init__(self, encoding: str, errors: str) -> None:
+ decoder_class = codecs.getincrementaldecoder(encoding)
+ self._decoder = decoder_class(errors=errors)
+
+ async def receive(self) -> str:
+ while True:
+ chunk = await self.transport_stream.receive()
+ decoded = self._decoder.decode(chunk)
+ if decoded:
+ return decoded
+
+ async def aclose(self) -> None:
+ await self.transport_stream.aclose()
+ self._decoder.reset()
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return self.transport_stream.extra_attributes
+
+
+@dataclass(eq=False)
+class TextSendStream(ObjectSendStream[str]):
+ """
+ Sends strings to the wrapped stream as bytes using the given encoding.
+
+ :param AnyByteSendStream transport_stream: any bytes-based send stream
+ :param str encoding: character encoding to use for encoding strings to bytes
+ (defaults to ``utf-8``)
+ :param str errors: handling scheme for encoding errors (defaults to ``strict``; see
+ the `codecs module documentation`_ for a comprehensive list of options)
+
+ .. _codecs module documentation:
+ https://docs.python.org/3/library/codecs.html#codec-objects
+ """
+
+ transport_stream: AnyByteSendStream
+ encoding: InitVar[str] = "utf-8"
+ errors: str = "strict"
+ _encoder: Callable[..., tuple[bytes, int]] = field(init=False)
+
+ def __post_init__(self, encoding: str) -> None:
+ self._encoder = codecs.getencoder(encoding)
+
+ async def send(self, item: str) -> None:
+ encoded = self._encoder(item, self.errors)[0]
+ await self.transport_stream.send(encoded)
+
+ async def aclose(self) -> None:
+ await self.transport_stream.aclose()
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return self.transport_stream.extra_attributes
+
+
+@dataclass(eq=False)
+class TextStream(ObjectStream[str]):
+ """
+ A bidirectional stream that decodes bytes to strings on receive and encodes strings
+ to bytes on send.
+
+ Extra attributes will be provided from both streams, with the receive stream
+ providing the values in case of a conflict.
+
+ :param AnyByteStream transport_stream: any bytes-based stream
+ :param str encoding: character encoding to use for encoding/decoding strings to/from
+ bytes (defaults to ``utf-8``)
+ :param str errors: handling scheme for encoding errors (defaults to ``strict``; see
+ the `codecs module documentation`_ for a comprehensive list of options)
+
+ .. _codecs module documentation:
+ https://docs.python.org/3/library/codecs.html#codec-objects
+ """
+
+ transport_stream: AnyByteStream
+ encoding: InitVar[str] = "utf-8"
+ errors: InitVar[str] = "strict"
+ _receive_stream: TextReceiveStream = field(init=False)
+ _send_stream: TextSendStream = field(init=False)
+
+ def __post_init__(self, encoding: str, errors: str) -> None:
+ self._receive_stream = TextReceiveStream(
+ self.transport_stream, encoding=encoding, errors=errors
+ )
+ self._send_stream = TextSendStream(
+ self.transport_stream, encoding=encoding, errors=errors
+ )
+
+ async def receive(self) -> str:
+ return await self._receive_stream.receive()
+
+ async def send(self, item: str) -> None:
+ await self._send_stream.send(item)
+
+ async def send_eof(self) -> None:
+ await self.transport_stream.send_eof()
+
+ async def aclose(self) -> None:
+ await self._send_stream.aclose()
+ await self._receive_stream.aclose()
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return {
+ **self._send_stream.extra_attributes,
+ **self._receive_stream.extra_attributes,
+ }
+
+
+class TextConnectable(ObjectStreamConnectable[str]):
+ def __init__(self, connectable: AnyByteStreamConnectable):
+ """
+ :param connectable: the bytestream endpoint to wrap
+
+ """
+ self.connectable = connectable
+
+ @override
+ async def connect(self) -> TextStream:
+ stream = await self.connectable.connect()
+ return TextStream(stream)
diff --git a/venv/lib/python3.11/site-packages/anyio/streams/tls.py b/venv/lib/python3.11/site-packages/anyio/streams/tls.py
new file mode 100644
index 0000000000000000000000000000000000000000..282174c71d6d6672eaff99b919a0bb5cf7d418b2
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/streams/tls.py
@@ -0,0 +1,436 @@
+from __future__ import annotations
+
+__all__ = (
+ "TLSAttribute",
+ "TLSConnectable",
+ "TLSListener",
+ "TLSStream",
+)
+
+import logging
+import re
+import ssl
+import sys
+from collections.abc import Callable, Mapping
+from dataclasses import dataclass
+from functools import wraps
+from ssl import SSLContext
+from typing import Any, TypeAlias, TypeVar
+
+from .. import (
+ BrokenResourceError,
+ EndOfStream,
+ aclose_forcefully,
+ get_cancelled_exc_class,
+ to_thread,
+)
+from .._core._typedattr import TypedAttributeSet, typed_attribute
+from ..abc import (
+ AnyByteStream,
+ AnyByteStreamConnectable,
+ ByteStream,
+ ByteStreamConnectable,
+ Listener,
+ TaskGroup,
+)
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+if sys.version_info >= (3, 12):
+ from typing import override
+else:
+ from typing_extensions import override
+
+T_Retval = TypeVar("T_Retval")
+PosArgsT = TypeVarTuple("PosArgsT")
+_PCTRTT: TypeAlias = tuple[tuple[str, str], ...]
+_PCTRTTT: TypeAlias = tuple[_PCTRTT, ...]
+
+
+class TLSAttribute(TypedAttributeSet):
+ """Contains Transport Layer Security related attributes."""
+
+ #: the selected ALPN protocol
+ alpn_protocol: str | None = typed_attribute()
+ #: the channel binding for type ``tls-unique``
+ channel_binding_tls_unique: bytes = typed_attribute()
+ #: the selected cipher
+ cipher: tuple[str, str, int] = typed_attribute()
+ #: the peer certificate in dictionary form (see :meth:`ssl.SSLSocket.getpeercert`
+ # for more information)
+ peer_certificate: None | (dict[str, str | _PCTRTTT | _PCTRTT]) = typed_attribute()
+ #: the peer certificate in binary form
+ peer_certificate_binary: bytes | None = typed_attribute()
+ #: ``True`` if this is the server side of the connection
+ server_side: bool = typed_attribute()
+ #: ciphers shared by the client during the TLS handshake (``None`` if this is the
+ #: client side)
+ shared_ciphers: list[tuple[str, str, int]] | None = typed_attribute()
+ #: the :class:`~ssl.SSLObject` used for encryption
+ ssl_object: ssl.SSLObject = typed_attribute()
+ #: ``True`` if this stream does (and expects) a closing TLS handshake when the
+ #: stream is being closed
+ standard_compatible: bool = typed_attribute()
+ #: the TLS protocol version (e.g. ``TLSv1.2``)
+ tls_version: str = typed_attribute()
+
+
+@dataclass(eq=False)
+class TLSStream(ByteStream):
+ """
+ A stream wrapper that encrypts all sent data and decrypts received data.
+
+ This class has no public initializer; use :meth:`wrap` instead.
+ All extra attributes from :class:`~TLSAttribute` are supported.
+
+ :var AnyByteStream transport_stream: the wrapped stream
+
+ """
+
+ transport_stream: AnyByteStream
+ standard_compatible: bool
+ _ssl_object: ssl.SSLObject
+ _read_bio: ssl.MemoryBIO
+ _write_bio: ssl.MemoryBIO
+
+ @classmethod
+ async def wrap(
+ cls,
+ transport_stream: AnyByteStream,
+ *,
+ server_side: bool | None = None,
+ hostname: str | None = None,
+ ssl_context: ssl.SSLContext | None = None,
+ standard_compatible: bool = True,
+ ) -> TLSStream:
+ """
+ Wrap an existing stream with Transport Layer Security.
+
+ This performs a TLS handshake with the peer.
+
+ :param transport_stream: a bytes-transporting stream to wrap
+ :param server_side: ``True`` if this is the server side of the connection,
+ ``False`` if this is the client side (if omitted, will be set to ``False``
+ if ``hostname`` has been provided, ``False`` otherwise). Used only to create
+ a default context when an explicit context has not been provided.
+ :param hostname: host name of the peer (if host name checking is desired)
+ :param ssl_context: the SSLContext object to use (if not provided, a secure
+ default will be created)
+ :param standard_compatible: if ``False``, skip the closing handshake when
+ closing the connection, and don't raise an exception if the peer does the
+ same
+ :raises ~ssl.SSLError: if the TLS handshake fails
+
+ """
+ if server_side is None:
+ server_side = not hostname
+
+ if not ssl_context:
+ purpose = (
+ ssl.Purpose.CLIENT_AUTH if server_side else ssl.Purpose.SERVER_AUTH
+ )
+ ssl_context = ssl.create_default_context(purpose)
+
+ # Re-enable detection of unexpected EOFs if it was disabled by Python
+ if hasattr(ssl, "OP_IGNORE_UNEXPECTED_EOF"):
+ ssl_context.options &= ~ssl.OP_IGNORE_UNEXPECTED_EOF
+
+ bio_in = ssl.MemoryBIO()
+ bio_out = ssl.MemoryBIO()
+
+ # Resolve international host names using IDNA 2008.
+ # Otherwise wrap_bio() would resolve them with IDNA 2003.
+ if hostname is not None:
+ from .._core._sockets import idna2008_resolve
+
+ server_hostname: bytes | None = idna2008_resolve(hostname)
+ else:
+ server_hostname = None
+
+ # External SSLContext implementations may do blocking I/O in wrap_bio(),
+ # but the standard library implementation won't
+ if type(ssl_context) is ssl.SSLContext:
+ ssl_object = ssl_context.wrap_bio(
+ bio_in,
+ bio_out,
+ server_side=server_side,
+ server_hostname=server_hostname,
+ )
+ else:
+ ssl_object = await to_thread.run_sync(
+ ssl_context.wrap_bio,
+ bio_in,
+ bio_out,
+ server_side,
+ server_hostname,
+ None,
+ )
+
+ wrapper = cls(
+ transport_stream=transport_stream,
+ standard_compatible=standard_compatible,
+ _ssl_object=ssl_object,
+ _read_bio=bio_in,
+ _write_bio=bio_out,
+ )
+ await wrapper._call_sslobject_method(ssl_object.do_handshake)
+ return wrapper
+
+ async def _call_sslobject_method(
+ self, func: Callable[[Unpack[PosArgsT]], T_Retval], *args: Unpack[PosArgsT]
+ ) -> T_Retval:
+ while True:
+ try:
+ result = func(*args)
+ except ssl.SSLWantReadError:
+ try:
+ # Flush any pending writes first
+ if self._write_bio.pending:
+ await self.transport_stream.send(self._write_bio.read())
+
+ data = await self.transport_stream.receive()
+ except EndOfStream:
+ self._read_bio.write_eof()
+ except OSError as exc:
+ self._read_bio.write_eof()
+ self._write_bio.write_eof()
+ raise BrokenResourceError from exc
+ else:
+ self._read_bio.write(data)
+ except ssl.SSLWantWriteError:
+ await self.transport_stream.send(self._write_bio.read())
+ except ssl.SSLSyscallError as exc:
+ self._read_bio.write_eof()
+ self._write_bio.write_eof()
+ raise BrokenResourceError from exc
+ except ssl.SSLError as exc:
+ self._read_bio.write_eof()
+ self._write_bio.write_eof()
+ if isinstance(exc, ssl.SSLEOFError) or (
+ exc.strerror and "UNEXPECTED_EOF_WHILE_READING" in exc.strerror
+ ):
+ if self.standard_compatible:
+ raise BrokenResourceError from exc
+ else:
+ raise EndOfStream from None
+
+ raise
+ else:
+ # Flush any pending writes first
+ if self._write_bio.pending:
+ await self.transport_stream.send(self._write_bio.read())
+
+ return result
+
+ async def unwrap(self) -> tuple[AnyByteStream, bytes]:
+ """
+ Does the TLS closing handshake.
+
+ :return: a tuple of (wrapped byte stream, bytes left in the read buffer)
+
+ """
+ await self._call_sslobject_method(self._ssl_object.unwrap)
+ self._read_bio.write_eof()
+ self._write_bio.write_eof()
+ return self.transport_stream, self._read_bio.read()
+
+ async def aclose(self) -> None:
+ if self.standard_compatible:
+ try:
+ await self.unwrap()
+ except BaseException:
+ await aclose_forcefully(self.transport_stream)
+ raise
+
+ await self.transport_stream.aclose()
+
+ async def receive(self, max_bytes: int = 65536) -> bytes:
+ if max_bytes < 1:
+ raise ValueError("max_bytes must be a positive integer")
+
+ data = await self._call_sslobject_method(self._ssl_object.read, max_bytes)
+ if not data:
+ raise EndOfStream
+
+ return data
+
+ async def send(self, item: bytes) -> None:
+ await self._call_sslobject_method(self._ssl_object.write, item)
+
+ async def send_eof(self) -> None:
+ tls_version = self.extra(TLSAttribute.tls_version)
+ match = re.match(r"TLSv(\d+)(?:\.(\d+))?", tls_version)
+ if match:
+ major, minor = int(match.group(1)), int(match.group(2) or 0)
+ if (major, minor) < (1, 3):
+ raise NotImplementedError(
+ f"send_eof() requires at least TLSv1.3; current "
+ f"session uses {tls_version}"
+ )
+
+ raise NotImplementedError(
+ "send_eof() has not yet been implemented for TLS streams"
+ )
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return {
+ **self.transport_stream.extra_attributes,
+ TLSAttribute.alpn_protocol: self._ssl_object.selected_alpn_protocol,
+ TLSAttribute.channel_binding_tls_unique: (
+ self._ssl_object.get_channel_binding
+ ),
+ TLSAttribute.cipher: self._ssl_object.cipher,
+ TLSAttribute.peer_certificate: lambda: self._ssl_object.getpeercert(False),
+ TLSAttribute.peer_certificate_binary: lambda: self._ssl_object.getpeercert(
+ True
+ ),
+ TLSAttribute.server_side: lambda: self._ssl_object.server_side,
+ TLSAttribute.shared_ciphers: lambda: (
+ self._ssl_object.shared_ciphers()
+ if self._ssl_object.server_side
+ else None
+ ),
+ TLSAttribute.standard_compatible: lambda: self.standard_compatible,
+ TLSAttribute.ssl_object: lambda: self._ssl_object,
+ TLSAttribute.tls_version: self._ssl_object.version,
+ }
+
+
+@dataclass(eq=False)
+class TLSListener(Listener[TLSStream]):
+ """
+ A convenience listener that wraps another listener and auto-negotiates a TLS session
+ on every accepted connection.
+
+ If the TLS handshake times out or raises an exception,
+ :meth:`handle_handshake_error` is called to do whatever post-mortem processing is
+ deemed necessary.
+
+ Supports only the :attr:`~TLSAttribute.standard_compatible` extra attribute.
+
+ :param Listener listener: the listener to wrap
+ :param ssl_context: the SSL context object
+ :param standard_compatible: a flag passed through to :meth:`TLSStream.wrap`
+ :param handshake_timeout: time limit for the TLS handshake
+ (passed to :func:`~anyio.fail_after`)
+ """
+
+ listener: Listener[Any]
+ ssl_context: ssl.SSLContext
+ standard_compatible: bool = True
+ handshake_timeout: float = 30
+
+ @staticmethod
+ async def handle_handshake_error(exc: BaseException, stream: AnyByteStream) -> None:
+ """
+ Handle an exception raised during the TLS handshake.
+
+ This method does 3 things:
+
+ #. Forcefully closes the original stream
+ #. Logs the exception (unless it was a cancellation exception) using the
+ ``anyio.streams.tls`` logger
+ #. Reraises the exception if it was a base exception or a cancellation exception
+
+ :param exc: the exception
+ :param stream: the original stream
+
+ """
+ await aclose_forcefully(stream)
+
+ # Log all except cancellation exceptions
+ if not isinstance(exc, get_cancelled_exc_class()):
+ # CPython (as of 3.11.5) returns incorrect `sys.exc_info()` here when using
+ # any asyncio implementation, so we explicitly pass the exception to log
+ # (https://github.com/python/cpython/issues/108668). Trio does not have this
+ # issue because it works around the CPython bug.
+ logging.getLogger(__name__).exception(
+ "Error during TLS handshake", exc_info=exc
+ )
+
+ # Only reraise base exceptions and cancellation exceptions
+ if not isinstance(exc, Exception) or isinstance(exc, get_cancelled_exc_class()):
+ raise
+
+ async def serve(
+ self,
+ handler: Callable[[TLSStream], Any],
+ task_group: TaskGroup | None = None,
+ ) -> None:
+ @wraps(handler)
+ async def handler_wrapper(stream: AnyByteStream) -> None:
+ from .. import fail_after
+
+ try:
+ with fail_after(self.handshake_timeout):
+ wrapped_stream = await TLSStream.wrap(
+ stream,
+ ssl_context=self.ssl_context,
+ standard_compatible=self.standard_compatible,
+ )
+ except BaseException as exc:
+ await self.handle_handshake_error(exc, stream)
+ else:
+ await handler(wrapped_stream)
+
+ await self.listener.serve(handler_wrapper, task_group)
+
+ async def aclose(self) -> None:
+ await self.listener.aclose()
+
+ @property
+ def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
+ return {
+ TLSAttribute.standard_compatible: lambda: self.standard_compatible,
+ }
+
+
+class TLSConnectable(ByteStreamConnectable):
+ """
+ Wraps another connectable and does TLS negotiation after a successful connection.
+
+ :param connectable: the connectable to wrap
+ :param hostname: host name of the server (if host name checking is desired)
+ :param ssl_context: the SSLContext object to use (if not provided, a secure default
+ will be created)
+ :param standard_compatible: if ``False``, skip the closing handshake when closing
+ the connection, and don't raise an exception if the server does the same
+ """
+
+ def __init__(
+ self,
+ connectable: AnyByteStreamConnectable,
+ *,
+ hostname: str | None = None,
+ ssl_context: ssl.SSLContext | None = None,
+ standard_compatible: bool = True,
+ ) -> None:
+ self.connectable = connectable
+ self.ssl_context: SSLContext = ssl_context or ssl.create_default_context(
+ ssl.Purpose.SERVER_AUTH
+ )
+ if not isinstance(self.ssl_context, ssl.SSLContext):
+ raise TypeError(
+ "ssl_context must be an instance of ssl.SSLContext, not "
+ f"{type(self.ssl_context).__name__}"
+ )
+ self.hostname = hostname
+ self.standard_compatible = standard_compatible
+
+ @override
+ async def connect(self) -> TLSStream:
+ stream = await self.connectable.connect()
+ try:
+ return await TLSStream.wrap(
+ stream,
+ hostname=self.hostname,
+ ssl_context=self.ssl_context,
+ standard_compatible=self.standard_compatible,
+ )
+ except BaseException:
+ await aclose_forcefully(stream)
+ raise
diff --git a/venv/lib/python3.11/site-packages/anyio/to_interpreter.py b/venv/lib/python3.11/site-packages/anyio/to_interpreter.py
new file mode 100644
index 0000000000000000000000000000000000000000..694dbe77bc8581032ee72316afe4e0590311ba00
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/to_interpreter.py
@@ -0,0 +1,246 @@
+from __future__ import annotations
+
+__all__ = (
+ "run_sync",
+ "current_default_interpreter_limiter",
+)
+
+import atexit
+import os
+import sys
+from collections import deque
+from collections.abc import Callable
+from typing import Any, Final, TypeVar
+
+from . import current_time, to_thread
+from ._core._exceptions import BrokenWorkerInterpreter
+from ._core._synchronization import CapacityLimiter
+from .lowlevel import RunVar
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+if sys.version_info >= (3, 14):
+ from concurrent.interpreters import ExecutionFailed, create
+
+ def _interp_call(
+ func: Callable[..., Any], args: tuple[Any, ...]
+ ) -> tuple[Any, bool]:
+ try:
+ retval = func(*args)
+ except BaseException as exc:
+ return exc, True
+ else:
+ return retval, False
+
+ class _Worker:
+ last_used: float = 0
+
+ def __init__(self) -> None:
+ self._interpreter = create()
+
+ def destroy(self) -> None:
+ self._interpreter.close()
+
+ def call(
+ self,
+ func: Callable[..., T_Retval],
+ args: tuple[Any, ...],
+ ) -> T_Retval:
+ try:
+ res, is_exception = self._interpreter.call(_interp_call, func, args)
+ except ExecutionFailed as exc:
+ raise BrokenWorkerInterpreter(exc.excinfo) from exc
+
+ if is_exception:
+ raise res
+
+ return res
+elif sys.version_info >= (3, 13):
+ import _interpqueues
+ import _interpreters
+
+ UNBOUND: Final = 2 # I have no clue how this works, but it was used in the stdlib
+ FMT_UNPICKLED: Final = 0
+ FMT_PICKLED: Final = 1
+ QUEUE_PICKLE_ARGS: Final = (FMT_PICKLED, UNBOUND)
+ QUEUE_UNPICKLE_ARGS: Final = (FMT_UNPICKLED, UNBOUND)
+
+ _run_func = compile(
+ """
+import _interpqueues
+from _interpreters import NotShareableError
+from pickle import loads, dumps, HIGHEST_PROTOCOL
+
+QUEUE_PICKLE_ARGS = (1, 2)
+QUEUE_UNPICKLE_ARGS = (0, 2)
+
+item = _interpqueues.get(queue_id)[0]
+try:
+ func, args = loads(item)
+ retval = func(*args)
+except BaseException as exc:
+ is_exception = True
+ retval = exc
+else:
+ is_exception = False
+
+try:
+ _interpqueues.put(queue_id, (retval, is_exception), *QUEUE_UNPICKLE_ARGS)
+except NotShareableError:
+ retval = dumps(retval, HIGHEST_PROTOCOL)
+ _interpqueues.put(queue_id, (retval, is_exception), *QUEUE_PICKLE_ARGS)
+ """,
+ "",
+ "exec",
+ )
+
+ class _Worker:
+ last_used: float = 0
+
+ def __init__(self) -> None:
+ self._interpreter_id = _interpreters.create()
+ self._queue_id = _interpqueues.create(1, *QUEUE_UNPICKLE_ARGS)
+ _interpreters.set___main___attrs(
+ self._interpreter_id, {"queue_id": self._queue_id}
+ )
+
+ def destroy(self) -> None:
+ _interpqueues.destroy(self._queue_id)
+ _interpreters.destroy(self._interpreter_id)
+
+ def call(
+ self,
+ func: Callable[..., T_Retval],
+ args: tuple[Any, ...],
+ ) -> T_Retval:
+ import pickle
+
+ item = pickle.dumps((func, args), pickle.HIGHEST_PROTOCOL)
+ _interpqueues.put(self._queue_id, item, *QUEUE_PICKLE_ARGS)
+ exc_info = _interpreters.exec(self._interpreter_id, _run_func)
+ if exc_info:
+ raise BrokenWorkerInterpreter(exc_info)
+
+ res = _interpqueues.get(self._queue_id)
+ (res, is_exception), fmt = res[:2]
+ if fmt == FMT_PICKLED:
+ res = pickle.loads(res)
+
+ if is_exception:
+ raise res
+
+ return res
+else:
+
+ class _Worker:
+ last_used: float = 0
+
+ def __init__(self) -> None:
+ raise RuntimeError("subinterpreters require at least Python 3.13")
+
+ def call(
+ self,
+ func: Callable[..., T_Retval],
+ args: tuple[Any, ...],
+ ) -> T_Retval:
+ raise NotImplementedError
+
+ def destroy(self) -> None:
+ pass
+
+
+DEFAULT_CPU_COUNT: Final = 8 # this is just an arbitrarily selected value
+MAX_WORKER_IDLE_TIME = (
+ 30 # seconds a subinterpreter can be idle before becoming eligible for pruning
+)
+
+T_Retval = TypeVar("T_Retval")
+PosArgsT = TypeVarTuple("PosArgsT")
+
+_idle_workers = RunVar[deque[_Worker]]("_available_workers")
+_default_interpreter_limiter = RunVar[CapacityLimiter]("_default_interpreter_limiter")
+
+
+def _stop_workers(workers: deque[_Worker]) -> None:
+ for worker in workers:
+ worker.destroy()
+
+ workers.clear()
+
+
+async def run_sync(
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ *args: Unpack[PosArgsT],
+ limiter: CapacityLimiter | None = None,
+) -> T_Retval:
+ """
+ Call the given function with the given arguments in a subinterpreter.
+
+ .. warning:: On Python 3.13, the :mod:`concurrent.interpreters` module was not yet
+ available, so the code path for that Python version relies on an undocumented,
+ private API. As such, it is recommended to not rely on this function for anything
+ mission-critical on Python 3.13.
+
+ :param func: a callable
+ :param args: the positional arguments for the callable
+ :param limiter: capacity limiter to use to limit the total number of subinterpreters
+ running (if omitted, the default limiter is used)
+ :return: the result of the call
+ :raises BrokenWorkerInterpreter: if there's an internal error in a subinterpreter
+
+ """
+ if limiter is None:
+ limiter = current_default_interpreter_limiter()
+
+ try:
+ idle_workers = _idle_workers.get()
+ except LookupError:
+ idle_workers = deque()
+ _idle_workers.set(idle_workers)
+ atexit.register(_stop_workers, idle_workers)
+
+ async with limiter:
+ try:
+ worker = idle_workers.pop()
+ except IndexError:
+ worker = _Worker()
+
+ try:
+ return await to_thread.run_sync(
+ worker.call,
+ func,
+ args,
+ limiter=limiter,
+ )
+ finally:
+ # Prune workers that have been idle for too long
+ now = current_time()
+ while idle_workers:
+ if now - idle_workers[0].last_used <= MAX_WORKER_IDLE_TIME:
+ break
+
+ await to_thread.run_sync(idle_workers.popleft().destroy, limiter=limiter)
+
+ worker.last_used = current_time()
+ idle_workers.append(worker)
+
+
+def current_default_interpreter_limiter() -> CapacityLimiter:
+ """
+ Return the capacity limiter used by default to limit the number of concurrently
+ running subinterpreters.
+
+ Defaults to the number of CPU cores.
+
+ :return: a capacity limiter object
+
+ """
+ try:
+ return _default_interpreter_limiter.get()
+ except LookupError:
+ limiter = CapacityLimiter(os.cpu_count() or DEFAULT_CPU_COUNT)
+ _default_interpreter_limiter.set(limiter)
+ return limiter
diff --git a/venv/lib/python3.11/site-packages/anyio/to_process.py b/venv/lib/python3.11/site-packages/anyio/to_process.py
new file mode 100644
index 0000000000000000000000000000000000000000..8d356fbd2a3df39a9b44eacc3c6c924420c1a7b8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/to_process.py
@@ -0,0 +1,269 @@
+from __future__ import annotations
+
+__all__ = (
+ "current_default_process_limiter",
+ "process_worker",
+ "run_sync",
+)
+
+import os
+import pickle
+import runpy
+import subprocess
+import sys
+from collections import deque
+from collections.abc import Callable
+from types import ModuleType
+from typing import TypeVar, cast
+
+from ._core._eventloop import current_time, get_async_backend, get_cancelled_exc_class
+from ._core._exceptions import BrokenWorkerProcess
+from ._core._subprocesses import open_process
+from ._core._synchronization import CapacityLimiter
+from ._core._tasks import CancelScope, fail_after
+from .abc import ByteReceiveStream, ByteSendStream, Process
+from .lowlevel import RunVar, checkpoint_if_cancelled
+from .streams.buffered import BufferedByteReceiveStream
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+WORKER_MAX_IDLE_TIME = 300 # 5 minutes
+
+T_Retval = TypeVar("T_Retval")
+PosArgsT = TypeVarTuple("PosArgsT")
+
+_process_pool_workers: RunVar[set[Process]] = RunVar("_process_pool_workers")
+_process_pool_idle_workers: RunVar[deque[tuple[Process, float]]] = RunVar(
+ "_process_pool_idle_workers"
+)
+_default_process_limiter: RunVar[CapacityLimiter] = RunVar("_default_process_limiter")
+
+
+async def run_sync( # type: ignore[return]
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ *args: Unpack[PosArgsT],
+ cancellable: bool = False,
+ limiter: CapacityLimiter | None = None,
+) -> T_Retval:
+ """
+ Call the given function with the given arguments in a worker process.
+
+ If the ``cancellable`` option is enabled and the task waiting for its completion is
+ cancelled, the worker process running it will be abruptly terminated using SIGKILL
+ (or ``terminateProcess()`` on Windows).
+
+ :param func: a callable
+ :param args: positional arguments for the callable
+ :param cancellable: ``True`` to allow cancellation of the operation while it's
+ running
+ :param limiter: capacity limiter to use to limit the total amount of processes
+ running (if omitted, the default limiter is used)
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+ :return: an awaitable that yields the return value of the function.
+
+ """
+
+ async def send_raw_command(pickled_cmd: bytes) -> object:
+ try:
+ await stdin.send(pickled_cmd)
+ response = await buffered.receive_until(b"\n", 50)
+ status, length = response.split(b" ")
+ if status not in (b"RETURN", b"EXCEPTION"):
+ raise RuntimeError(
+ f"Worker process returned unexpected response: {response!r}"
+ )
+
+ pickled_response = await buffered.receive_exactly(int(length))
+ except BaseException as exc:
+ workers.discard(process)
+ try:
+ process.kill()
+ with CancelScope(shield=True):
+ await process.aclose()
+ except ProcessLookupError:
+ pass
+
+ if isinstance(exc, get_cancelled_exc_class()):
+ raise
+ else:
+ raise BrokenWorkerProcess from exc
+
+ retval = pickle.loads(pickled_response)
+ if status == b"EXCEPTION":
+ assert isinstance(retval, BaseException)
+ raise retval
+ else:
+ return retval
+
+ # First pickle the request before trying to reserve a worker process
+ await checkpoint_if_cancelled()
+ request = pickle.dumps(("run", func, args), protocol=pickle.HIGHEST_PROTOCOL)
+
+ # If this is the first run in this event loop thread, set up the necessary variables
+ try:
+ workers = _process_pool_workers.get()
+ idle_workers = _process_pool_idle_workers.get()
+ except LookupError:
+ workers = set()
+ idle_workers = deque()
+ _process_pool_workers.set(workers)
+ _process_pool_idle_workers.set(idle_workers)
+ get_async_backend().setup_process_pool_exit_at_shutdown(workers)
+
+ async with limiter or current_default_process_limiter():
+ # Pop processes from the pool (starting from the most recently used) until we
+ # find one that hasn't exited yet
+ process: Process
+ while idle_workers:
+ process, idle_since = idle_workers.pop()
+ if process.returncode is None:
+ stdin = cast(ByteSendStream, process.stdin)
+ buffered = BufferedByteReceiveStream(
+ cast(ByteReceiveStream, process.stdout)
+ )
+
+ # Prune any other workers that have been idle for WORKER_MAX_IDLE_TIME
+ # seconds or longer
+ now = current_time()
+ killed_processes: list[Process] = []
+ while idle_workers:
+ if now - idle_workers[0][1] < WORKER_MAX_IDLE_TIME:
+ break
+
+ process_to_kill, idle_since = idle_workers.popleft()
+ process_to_kill.kill()
+ workers.remove(process_to_kill)
+ killed_processes.append(process_to_kill)
+
+ with CancelScope(shield=True):
+ for killed_process in killed_processes:
+ await killed_process.aclose()
+
+ break
+
+ workers.remove(process)
+ else:
+ command = [sys.executable, "-u", "-m", __name__]
+ process = await open_process(
+ command, stdin=subprocess.PIPE, stdout=subprocess.PIPE
+ )
+ try:
+ stdin = cast(ByteSendStream, process.stdin)
+ buffered = BufferedByteReceiveStream(
+ cast(ByteReceiveStream, process.stdout)
+ )
+ with fail_after(20):
+ message = await buffered.receive(6)
+
+ if message != b"READY\n":
+ raise BrokenWorkerProcess(
+ f"Worker process returned unexpected response: {message!r}"
+ )
+
+ main_module_path = getattr(sys.modules["__main__"], "__file__", None)
+ pickled = pickle.dumps(
+ ("init", sys.path, main_module_path),
+ protocol=pickle.HIGHEST_PROTOCOL,
+ )
+ await send_raw_command(pickled)
+ except (BrokenWorkerProcess, get_cancelled_exc_class()):
+ raise
+ except BaseException as exc:
+ process.kill()
+ raise BrokenWorkerProcess(
+ "Error during worker process initialization"
+ ) from exc
+
+ workers.add(process)
+
+ with CancelScope(shield=not cancellable):
+ try:
+ return cast(T_Retval, await send_raw_command(request))
+ finally:
+ if process in workers:
+ idle_workers.append((process, current_time()))
+
+
+def current_default_process_limiter() -> CapacityLimiter:
+ """
+ Return the capacity limiter that is used by default to limit the number of worker
+ processes.
+
+ :return: a capacity limiter object
+
+ """
+ try:
+ return _default_process_limiter.get()
+ except LookupError:
+ limiter = CapacityLimiter(os.cpu_count() or 2)
+ _default_process_limiter.set(limiter)
+ return limiter
+
+
+def process_worker() -> None:
+ # Redirect standard streams to os.devnull so that user code won't interfere with the
+ # parent-worker communication
+ stdin = sys.stdin
+ stdout = sys.stdout
+ sys.stdin = open(os.devnull)
+ sys.stdout = open(os.devnull, "w")
+ sys.stderr = open(os.devnull, "w")
+
+ stdout.buffer.write(b"READY\n")
+ while True:
+ retval = exception = None
+ try:
+ command, *args = pickle.load(stdin.buffer)
+ except EOFError:
+ return
+ except BaseException as exc:
+ exception = exc
+ else:
+ if command == "run":
+ func, args = args
+ try:
+ retval = func(*args)
+ except BaseException as exc:
+ exception = exc
+ elif command == "init":
+ main_module_path: str | None
+ sys.path, main_module_path = args
+ del sys.modules["__main__"]
+ if main_module_path and os.path.isfile(main_module_path):
+ # Load the parent's main module but as __mp_main__ instead of
+ # __main__ (like multiprocessing does) to avoid infinite recursion
+ try:
+ main = ModuleType("__mp_main__")
+ main_content = runpy.run_path(
+ main_module_path, run_name="__mp_main__"
+ )
+ main.__dict__.update(main_content)
+ sys.modules["__main__"] = sys.modules["__mp_main__"] = main
+ except BaseException as exc:
+ exception = exc
+ try:
+ if exception is not None:
+ status = b"EXCEPTION"
+ pickled = pickle.dumps(exception, pickle.HIGHEST_PROTOCOL)
+ else:
+ status = b"RETURN"
+ pickled = pickle.dumps(retval, pickle.HIGHEST_PROTOCOL)
+ except BaseException as exc:
+ exception = exc
+ status = b"EXCEPTION"
+ pickled = pickle.dumps(exc, pickle.HIGHEST_PROTOCOL)
+
+ stdout.buffer.write(b"%s %d\n" % (status, len(pickled)))
+ stdout.buffer.write(pickled)
+
+ # Respect SIGTERM
+ if isinstance(exception, SystemExit):
+ raise exception
+
+
+if __name__ == "__main__":
+ process_worker()
diff --git a/venv/lib/python3.11/site-packages/anyio/to_thread.py b/venv/lib/python3.11/site-packages/anyio/to_thread.py
new file mode 100644
index 0000000000000000000000000000000000000000..a01f24f6e776ef30836fb5d82639776820bf194f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/anyio/to_thread.py
@@ -0,0 +1,80 @@
+from __future__ import annotations
+
+__all__ = (
+ "run_sync",
+ "current_default_thread_limiter",
+)
+
+import sys
+from collections.abc import Callable
+from typing import TYPE_CHECKING, TypeVar
+from warnings import warn
+
+from ._core._eventloop import get_async_backend
+
+if TYPE_CHECKING:
+ from ._core._synchronization import CapacityLimiter
+
+if sys.version_info >= (3, 11):
+ from typing import TypeVarTuple, Unpack
+else:
+ from typing_extensions import TypeVarTuple, Unpack
+
+T_Retval = TypeVar("T_Retval")
+PosArgsT = TypeVarTuple("PosArgsT")
+
+
+async def run_sync(
+ func: Callable[[Unpack[PosArgsT]], T_Retval],
+ *args: Unpack[PosArgsT],
+ abandon_on_cancel: bool = False,
+ cancellable: bool | None = None,
+ limiter: CapacityLimiter | None = None,
+) -> T_Retval:
+ """
+ Call the given function with the given arguments in a worker thread.
+
+ If the ``abandon_on_cancel`` option is enabled and the task waiting for its
+ completion is cancelled, the thread will still run its course but its
+ return value (or any raised exception) will be ignored.
+
+ :param func: a callable
+ :param args: positional arguments for the callable
+ :param abandon_on_cancel: ``True`` to abandon the thread (leaving it to run
+ unchecked on own) if the host task is cancelled, ``False`` to ignore
+ cancellations in the host task until the operation has completed in the worker
+ thread
+ :param cancellable: deprecated alias of ``abandon_on_cancel``; will override
+ ``abandon_on_cancel`` if both parameters are passed
+ :param limiter: capacity limiter to use to limit the total amount of threads running
+ (if omitted, the default limiter is used)
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+ :return: an awaitable that yields the return value of the function.
+
+ """
+ if cancellable is not None:
+ abandon_on_cancel = cancellable
+ warn(
+ "The `cancellable=` keyword argument to `anyio.to_thread.run_sync` is "
+ "deprecated since AnyIO 4.1.0; use `abandon_on_cancel=` instead",
+ DeprecationWarning,
+ stacklevel=2,
+ )
+
+ return await get_async_backend().run_sync_in_worker_thread(
+ func, args, abandon_on_cancel=abandon_on_cancel, limiter=limiter
+ )
+
+
+def current_default_thread_limiter() -> CapacityLimiter:
+ """
+ Return the capacity limiter that is used by default to limit the number of
+ concurrent threads.
+
+ :return: a capacity limiter object
+ :raises NoEventLoopError: if no supported asynchronous event loop is running in the
+ current thread
+
+ """
+ return get_async_backend().current_default_thread_limiter()
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/AUTHORS.py b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/AUTHORS.py
new file mode 100644
index 0000000000000000000000000000000000000000..f17b31aeb7fd8e9f5b80727e0ee0d7ec65a9debf
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/AUTHORS.py
@@ -0,0 +1,112 @@
+import math
+import subprocess
+
+print(
+ """Contributors
+============
+
+All contributors (by number of commits):
+"""
+)
+
+
+email_map = {
+ # Maintainers.
+ "git@mikeboers.com": "github@mikeboers.com",
+ "mboers@keypics.com": "github@mikeboers.com",
+ "mikeb@loftysky.com": "github@mikeboers.com",
+ "mikeb@markmedia.co": "github@mikeboers.com",
+ "westernx@mikeboers.com": "github@mikeboers.com",
+ # Junk.
+ "mark@mark-VirtualBox.(none)": None,
+ # Aliases.
+ "a.davoudi@aut.ac.ir": "davoudialireza@gmail.com",
+ "tcaswell@bnl.gov": "tcaswell@gmail.com",
+ "xxr3376@gmail.com": "xxr@megvii.com",
+ "dallan@pha.jhu.edu": "daniel.b.allan@gmail.com",
+ "61652821+laggykiller@users.noreply.github.com": "chaudominic2@gmail.com",
+}
+
+name_map = {
+ "caspervdw@gmail.com": "Casper van der Wel",
+ "daniel.b.allan@gmail.com": "Dan Allan",
+ "mgoacolou@cls.fr": "Manuel Goacolou",
+ "mindmark@gmail.com": "Mark Reid",
+ "moritzkassner@gmail.com": "Moritz Kassner",
+ "vidartf@gmail.com": "Vidar Tonaas Fauske",
+ "xxr@megvii.com": "Xinran Xu",
+}
+
+github_map = {
+ "billy.shambrook@gmail.com": "billyshambrook",
+ "daniel.b.allan@gmail.com": "danielballan",
+ "davoudialireza@gmail.com": "adavoudi",
+ "github@mikeboers.com": "mikeboers",
+ "jeremy.laine@m4x.org": "jlaine",
+ "kalle.litterfeldt@gmail.com": "litterfeldt",
+ "mindmark@gmail.com": "markreidvfx",
+ "moritzkassner@gmail.com": "mkassner",
+ "rush@logic.cz": "radek-senfeld",
+ "self@brendanlong.com": "brendanlong",
+ "tcaswell@gmail.com": "tacaswell",
+ "ulrik.mikaelsson@magine.com": "rawler",
+ "vidartf@gmail.com": "vidartf",
+ "willpatera@gmail.com": "willpatera",
+ "xxr@megvii.com": "xxr3376",
+ "chaudominic2@gmail.com": "laggykiller",
+ "wyattblue@auto-editor.com": "WyattBlue",
+}
+
+
+email_count = {}
+for line in (
+ subprocess.check_output(["git", "log", "--format=%aN,%aE"]).decode().splitlines()
+):
+ name, email = line.strip().rsplit(",", 1)
+
+ email = email_map.get(email, email)
+ if not email:
+ continue
+
+ names = name_map.setdefault(email, set())
+ if isinstance(names, set):
+ names.add(name)
+
+ email_count[email] = email_count.get(email, 0) + 1
+
+
+last = None
+block_i = 0
+for email, count in sorted(email_count.items(), key=lambda x: (-x[1], x[0])):
+
+ # This is the natural log, because of course it should be. ;)
+ order = int(math.log(count))
+ if last and last != order:
+ block_i += 1
+ print()
+ last = order
+
+ names = name_map[email]
+ if isinstance(names, set):
+ name = ", ".join(sorted(names))
+ else:
+ name = names
+
+ github = github_map.get(email)
+
+ # The '-' vs '*' is so that Sphinx treats them as different lists, and
+ # introduces a gap bettween them.
+ if github:
+ print(
+ "%s %s <%s>; `@%s `_"
+ % ("-*"[block_i % 2], name, email, github, github)
+ )
+ else:
+ print(
+ "%s %s <%s>"
+ % (
+ "-*"[block_i % 2],
+ name,
+ email,
+ )
+ )
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/AUTHORS.rst b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/AUTHORS.rst
new file mode 100644
index 0000000000000000000000000000000000000000..d5b4cb7eecb61b2f0f666ecbf1aa05658fb1b0c3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/AUTHORS.rst
@@ -0,0 +1,101 @@
+Contributors
+============
+
+All contributors (by number of commits):
+
+- Mike Boers ; `@mikeboers `_
+
+* Jeremy Lainé ; `@jlaine `_
+* WyattBlue ; `@WyattBlue `_
+
+- Mark Reid ; `@markreidvfx `_
+
+* Vidar Tonaas Fauske ; `@vidartf `_
+* laggykiller ; `@laggykiller `_
+* Billy Shambrook ; `@billyshambrook `_
+* Casper van der Wel
+* Philip de Nier
+* Tadas Dailyda
+* JoeUgly <41972063+JoeUgly@users.noreply.github.com>
+* Justin Wong <46082645+uvjustin@users.noreply.github.com>
+* Mark Harfouche
+
+- Alba Mendez
+- Dave Johansen
+- Xinran Xu ; `@xxr3376 `_
+- Dan Allan ; `@danielballan `_
+- Moonsik Park
+- Santtu Keskinen
+- Marc Mueller <30130371+cdce8p@users.noreply.github.com>
+- Christoph Rackwitz
+- David Plowman
+- Alireza Davoudi ; `@adavoudi `_
+- Jonathan Drolet
+- Moritz Kassner ; `@mkassner `_
+- Thomas A Caswell ; `@tacaswell `_
+- Ulrik Mikaelsson ; `@rawler `_
+- Wel C. van der
+- Will Patera ; `@willpatera `_
+
+* Dexer <73297572+DexerBR@users.noreply.github.com>
+* rutsh
+* Felix Vollmer
+* Santiago Castro
+* Christian Clauss
+* Ihor Liubymov
+* Johannes Erdfelt
+* Karl Litterfeldt ; `@litterfeldt `_
+* Martin Larralde
+* Simon-Martin Schröder
+* Matteo Destro
+* mephi42
+* Miles Kaufmann
+* Pablo Prietz
+* Andrew Wason
+* Radek Senfeld ; `@radek-senfeld `_
+* Benjamin Chrétien <2742231+bchretien@users.noreply.github.com>
+* zzjjbb <31069326+zzjjbb@users.noreply.github.com>
+* davidplowman <38045873+davidplowman@users.noreply.github.com>
+* Hanz <40712686+HanzCEO@users.noreply.github.com>
+* Joe Schiff <41972063+JoeSchiff@users.noreply.github.com>
+* Artturin
+* Ian Lee
+* Ryan Huang
+* Arthur Barros
+* Carlos Ruiz
+* Carlos Ruiz
+* Maxime Desroches
+* egao1980
+* Eric Kalosa-Kenyon
+* elxy
+* Gemfield
+* henri-gasc
+* Jonathan Martin
+* Johan Jeppsson Karlin
+* Philipp Klaus
+* Lukas Geiger
+* Mattias Wadman
+* Manuel Goacolou
+* Julian Schweizer
+* Ömer Sezgin Uğurlu
+* Orivej Desh
+* Philipp Krähenbühl
+* ramoncaldeira
+* Roland van Laar
+* Santiago Castro
+* Kengo Sawatsu
+* FirefoxMetzger
+* hyenal
+* Brendan Long ; `@brendanlong `_
+* Семён Марьясин
+* Stephen.Y
+* Tom Flanagan
+* Tim O'Shea
+* Tim Ahpee
+* Jonas Tingeborn
+* Pino Toscano
+* Ulrik Mikaelsson
+* Vasiliy Kotov
+* Koichi Akabe
+* David Joy
+* Sviatoslav Sydorenko (Святослав Сидоренко)
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/INSTALLER b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/INSTALLER
new file mode 100644
index 0000000000000000000000000000000000000000..5c69047b2eb8235994febeeae1da4a82365a240a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/INSTALLER
@@ -0,0 +1 @@
+uv
\ No newline at end of file
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/LICENSE.txt b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/LICENSE.txt
new file mode 100644
index 0000000000000000000000000000000000000000..9db1661e7d701c7fd262ed77efc2acc3831e41a2
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/LICENSE.txt
@@ -0,0 +1,23 @@
+Copyright retained by original committers. All rights reserved.
+
+Redistribution and use in source and binary forms, with or without
+modification, are permitted provided that the following conditions are met:
+ * Redistributions of source code must retain the above copyright
+ notice, this list of conditions and the following disclaimer.
+ * Redistributions in binary form must reproduce the above copyright
+ notice, this list of conditions and the following disclaimer in the
+ documentation and/or other materials provided with the distribution.
+ * Neither the name of the project nor the names of its contributors may be
+ used to endorse or promote products derived from this software without
+ specific prior written permission.
+
+THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
+AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
+IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
+DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDERS BE LIABLE FOR ANY DIRECT,
+INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
+BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
+DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY
+OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING
+NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE,
+EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/METADATA b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/METADATA
new file mode 100644
index 0000000000000000000000000000000000000000..3e3f98b5ebd2cc0e005f4bc25678674a4cdecf04
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/METADATA
@@ -0,0 +1,123 @@
+Metadata-Version: 2.1
+Name: av
+Version: 14.0.1
+Summary: Pythonic bindings for FFmpeg's libraries.
+Home-page: https://github.com/PyAV-Org/PyAV
+Author: Mike Boers
+Author-email: pyav@mikeboers.com
+License: BSD
+Project-URL: Bug Reports, https://github.com/PyAV-Org/PyAV/discussions/new?category=4-bugs
+Project-URL: Documentation, https://pyav.basswood-io.com
+Project-URL: Download, https://pypi.org/project/av
+Classifier: Development Status :: 5 - Production/Stable
+Classifier: Intended Audience :: Developers
+Classifier: License :: OSI Approved :: BSD License
+Classifier: Natural Language :: English
+Classifier: Operating System :: MacOS :: MacOS X
+Classifier: Operating System :: POSIX
+Classifier: Operating System :: Unix
+Classifier: Operating System :: Microsoft :: Windows
+Classifier: Programming Language :: Cython
+Classifier: Programming Language :: Python :: 3.9
+Classifier: Programming Language :: Python :: 3.10
+Classifier: Programming Language :: Python :: 3.11
+Classifier: Programming Language :: Python :: 3.12
+Classifier: Programming Language :: Python :: 3.13
+Classifier: Topic :: Software Development :: Libraries :: Python Modules
+Classifier: Topic :: Multimedia :: Sound/Audio
+Classifier: Topic :: Multimedia :: Sound/Audio :: Conversion
+Classifier: Topic :: Multimedia :: Video
+Classifier: Topic :: Multimedia :: Video :: Conversion
+Requires-Python: >=3.9
+Description-Content-Type: text/markdown
+License-File: LICENSE.txt
+License-File: AUTHORS.py
+License-File: AUTHORS.rst
+
+PyAV
+====
+
+PyAV is a Pythonic binding for the [FFmpeg][ffmpeg] libraries. We aim to provide all of the power and control of the underlying library, but manage the gritty details as much as possible.
+
+---
+
+[![GitHub Test Status][github-tests-badge]][github-tests] [![Documentation][docs-badge]][docs] [![Python Package Index][pypi-badge]][pypi] [![Conda Forge][conda-badge]][conda]
+
+PyAV is for direct and precise access to your media via containers, streams, packets, codecs, and frames. It exposes a few transformations of that data, and helps you get your data to/from other packages (e.g. Numpy and Pillow).
+
+This power does come with some responsibility as working with media is horrendously complicated and PyAV can't abstract it away or make all the best decisions for you. If the `ffmpeg` command does the job without you bending over backwards, PyAV is likely going to be more of a hindrance than a help.
+
+But where you can't work without it, PyAV is a critical tool.
+
+
+Installation
+------------
+
+Due to the complexity of the dependencies, PyAV is not always the easiest Python package to install from source. Since release 8.0.0 binary wheels are provided on [PyPI][pypi] for Linux, Mac and Windows linked against a modern FFmpeg. You can install these wheels by running:
+
+```bash
+pip install av
+```
+
+If you want to use your existing FFmpeg, the source version of PyAV is on [PyPI][pypi] too:
+
+```bash
+pip install av --no-binary av
+```
+
+Installing from source is not supported on Windows.
+
+Alternative installation methods
+--------------------------------
+
+Another way of installing PyAV is via [conda-forge][conda-forge]:
+
+```bash
+conda install av -c conda-forge
+```
+
+See the [Conda install][conda-install] docs to get started with (mini)Conda.
+
+And if you want to build from the absolute source (POSIX only):
+
+```bash
+git clone https://github.com/PyAV-Org/PyAV.git
+cd PyAV
+source scripts/activate.sh
+
+# Build ffmpeg from source. You can skip this step
+# if ffmpeg is already installed.
+./scripts/build-deps
+
+# Build PyAV
+make
+
+# Testing
+make test
+
+# Install globally
+deactivate
+pip install .
+```
+
+---
+
+Have fun, [read the docs][docs], [come chat with us][discuss], and good luck!
+
+
+
+[conda-badge]: https://img.shields.io/conda/vn/conda-forge/av.svg?colorB=CCB39A
+[conda]: https://anaconda.org/conda-forge/av
+[docs-badge]: https://img.shields.io/badge/docs-on%20pyav.basswood--io.com-blue.svg
+[docs]: https://pyav.basswood-io.com
+[pypi-badge]: https://img.shields.io/pypi/v/av.svg?colorB=CCB39A
+[pypi]: https://pypi.org/project/av
+[discuss]: https://github.com/PyAV-Org/PyAV/discussions
+
+[github-tests-badge]: https://github.com/PyAV-Org/PyAV/workflows/tests/badge.svg
+[github-tests]: https://github.com/PyAV-Org/PyAV/actions?workflow=tests
+[github]: https://github.com/PyAV-Org/PyAV
+
+[ffmpeg]: https://ffmpeg.org/
+[conda-forge]: https://conda-forge.github.io/
+[conda-install]: https://docs.conda.io/projects/conda/en/latest/user-guide/install/index.html
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/RECORD b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/RECORD
new file mode 100644
index 0000000000000000000000000000000000000000..cb5cb9007d78b05f9f24ff439c935e926677fb7c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/RECORD
@@ -0,0 +1,240 @@
+../../../bin/pyav,sha256=vZYciBlqPAYiTyqSbdBjoq4YYvGIu2foBykDi2UTB9U,304
+av-14.0.1.dist-info/AUTHORS.py,sha256=JaO6FdSYqANsqEwO4yBrTO-_xaLI-aNnZqvuVsy8XXU,3121
+av-14.0.1.dist-info/AUTHORS.rst,sha256=lCQZ6h5NmZyk0DEAu4PDuGPepLN7SnsBrNI-W3jrq_k,4790
+av-14.0.1.dist-info/INSTALLER,sha256=5hhM4Q4mYTT9z6QB6PGpUAW81PGNFrYrdXMj4oM_6ak,2
+av-14.0.1.dist-info/LICENSE.txt,sha256=dq8EYf-5LhnxwURJ6VVX2Dot-qG68gLUnl8dh0bA2hk,1505
+av-14.0.1.dist-info/METADATA,sha256=vJVTH6xCqAsrvDSMCa7vJxX9IWzYB4nh-1MpjjzRm0Q,4454
+av-14.0.1.dist-info/RECORD,,
+av-14.0.1.dist-info/REQUESTED,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av-14.0.1.dist-info/WHEEL,sha256=9BFfIe-Zq441iQ0ehutX65O5faGDpmB1Uw3WaQGk4f0,151
+av-14.0.1.dist-info/entry_points.txt,sha256=3XMdM30ih673nLSRVzDsHLBmGYNlt7wQ1xyW8xoHAzg,42
+av-14.0.1.dist-info/top_level.txt,sha256=TuQF-stvFHN8ilfr36ctqc7_MR5IOhUqrR0i6i5gNR8,3
+av.libs/libXau-00ec42fe.so.6.0.0,sha256=JjysEtjYterX3CORw1X-n8k5lA4eoi7ZjuVqjLYc5oQ,17049
+av.libs/libaom-e9efed4a.so.3.2.0,sha256=EW7wZ1XwM3WixaiE-WjsuNz0SzFB0Jg6LsItS8DwaVE,7285769
+av.libs/libavcodec-fb6c662d.so.61.19.100,sha256=wXuZmgZP83BjyQGV06R-qPScDkHqvQeBd4WZYqeWq90,14588393
+av.libs/libavdevice-84390fdd.so.61.3.100,sha256=jyq1bYBz_lV2cSPqHEzC4PHGJWHWVf2xhY6hNaZRqSQ,90793
+av.libs/libavfilter-3235a7c8.so.10.4.100,sha256=2tbod_zF1S_wm0vRO4tcH5uAkCu8mk3Xz-2xVQcBiNA,4526257
+av.libs/libavformat-071c54bd.so.61.7.100,sha256=7WmyXJVOPlif4mSErNEe9CyyzmFrZOXi05ZKlMS0ZyE,2782073
+av.libs/libavutil-2749f1ba.so.59.39.100,sha256=m-y3DokQq8oc-NI3ClFIloS_UbiPirIFBSshuqMD8H8,1028897
+av.libs/libdav1d-1b53ef2f.so.7.0.0,sha256=4SP0_2wJb_EwziPGVeM_RlGWJuyDUJnbEpbWiA5Zo2c,1922417
+av.libs/libdrm-827b956f.so.2.4.0,sha256=NjKUhxMlaqzB81yWUVWn3yvIOT_2wwijGiGzZ2JF5zI,79233
+av.libs/libgmp-a4b719d5.so.10.5.0,sha256=LJ5FsfdgN3V1b-tkjYRSvwj-HpSRmeLUTFOXT7dB0As,507625
+av.libs/libgnutls-b9b94016.so.30.36.0,sha256=a5ui8nn704ifbzOrsawf1rMNXxmoFTD9P1vjFCb-aF8,2218121
+av.libs/libhogweed-9544c08c.so.6.8,sha256=G4hk_Eb1r0bQV1ya870jfgDLDuOmv9KsjdvAPfdiEO8,316505
+av.libs/liblzma-af70179d.so.5.4.4,sha256=avYlwaSXKgpM0PhiClHvFFSoGY7s9XfpsLv2kXi0AQw,196537
+av.libs/libmp3lame-3ecc6556.so.0.0.0,sha256=hwkyn-fZ3hrQjb8jGMyha0PJq_PHJGw-6AHYwvi93bM,417001
+av.libs/libnettle-ecd2e589.so.8.8,sha256=4bBU9Ip79XDtCqyzbxv_zoV4vgx0xJ3gHqnuKkxeI_A,351169
+av.libs/libogg-bbd52b06.so.0.8.5,sha256=xHNvLO2BMzkRUQQKoEJRZI_wZexIHOlN1c9EksrzM00,43049
+av.libs/libopencore-amrnb-393dbae2.so.0.0.3,sha256=mMZxoKhEV6Ub0RJ8g4wOo_R4xHwUOjFwCKgn6uQ_PKI,172889
+av.libs/libopencore-amrwb-9db94aa9.so.0.0.3,sha256=BQb_y6-G8d5a7tOf037vF0t-thXBsyTFQ2SZVBv53LE,82609
+av.libs/libopus-59fcdf85.so.0.9.0,sha256=NxQeslQKOFrBRthCxmBj6aLkZ-VpYxN4s9MTjgc6Q1w,375521
+av.libs/libpostproc-dde0c2af.so.58.3.100,sha256=BI-ky890WWYsb03WkRcxAZjvgg75j5FEZNENd8UzLSA,75001
+av.libs/libsharpyuv-6e5d6e5b.so.0.1.0,sha256=t5DM4iiLGfVhiNSUu4aE1QKaKTEMol0ql5nsg1hzXZc,41969
+av.libs/libspeex-2370356a.so.1.5.2,sha256=HYJwjYCFxTMjf00YHwP-WOW0v0OC58N2YjBB8VWmDg4,108521
+av.libs/libswresample-da7d062e.so.5.3.100,sha256=gTmj_DAVVicH7EaoC2Ch8Bw1w6U63iiaDASjz1WvXsI,128265
+av.libs/libswscale-9212cf18.so.8.3.100,sha256=7kZzWyp1_YsjeN7Y1o65weKSWk-3TjTByb3pEk_3VA0,619929
+av.libs/libtwolame-72d74ef7.so.0.0.0,sha256=2Pco2Y3IRX0gBHso2uUH_Nb2sGNpBpO1pl00Apeia1Y,143513
+av.libs/libunistring-214e3d6e.so.5.1.0,sha256=gqnJUsuF1--retPHbHVwYN6WQQdots1OzqTkLd4aTq4,1815865
+av.libs/libvorbis-f4a9a6fd.so.0.4.9,sha256=GlMQnIUKSHX3454cyXfk-FEtAngMjI2y6gi7H5Sq380,244425
+av.libs/libvorbisenc-0d9d5bdf.so.2.0.12,sha256=790Y-5nbxz09d07ZD9n2v0Vp9Ae2hs1KWn0_qRJsfk8,713145
+av.libs/libvpx-832f6f52.so.9.0.0,sha256=LzlXQW2hRMb8pxcSWpsVIaNxQvq5L6FtxmQ8hQ-5zDg,2209833
+av.libs/libwebp-e16038c7.so.7.1.9,sha256=Yw3fF7-wR0Rqo73o3doCaem34Q7UOV68qSYBJxrKRcw,637017
+av.libs/libwebpmux-03713f16.so.3.1.0,sha256=2Z5WGCJuu6ElYe1eJEYy9QAUgmyL0kfglry3SzDz9Gg,58705
+av.libs/libx264-2a4c6f6d.so.164,sha256=hN-S_XNGR-eCIRwe016sppSbtsTI7YiJHtvExkQAKyw,2267921
+av.libs/libx265-d8690e8d.so.199,sha256=UOeGAXaS309xhdZdgpe_STOcvifGGXrXAGYlAgWs7go,19369233
+av.libs/libxcb-65da195c.so.1.1.0,sha256=zcPTuH8Ot2mTvEdVf9My_9J12Qj3nV5XwTgpiXVCe74,210465
+av.libs/libxcb-shape-25c2b258.so.0.0.0,sha256=Uw-Fq_rtzVc2cp-157WYQEZerrkf0y5k5jh5w1IQAgM,21801
+av.libs/libxcb-shm-7a199f70.so.0.0.0,sha256=NmOIxYVi6C80QPRpx7Kvs3KB2sCPJnIfpwL0mRXjri4,21401
+av.libs/libxcb-xfixes-9be3ba6f.so.0.0.0,sha256=Of6rBxt3xRw-7RrccGJajNQMo6DDcF1Tee-oFF_6WSw,45361
+av.libs/libxml2-cb941fce.so.2.9.13,sha256=U5fq_CgNSHGFFSXq6ZYKXmy7JncJEfCWR1ycyHHAHjk,1609769
+av/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/__init__.py,sha256=5_Ef3RbOBFc8NKfOopPG714KkyqU9YYld_VTNMOYOkA,2084
+av/__main__.py,sha256=hkZUood1yfn5SZfa_Sh-DGdPTKtUB0BL92ISHD1zJMY,1158
+av/_core.cpython-311-x86_64-linux-gnu.so,sha256=N9tZqs_ZW_KTxOvQaYJLSQ8iECm7pBybKKFPpnKFvUY,183497
+av/_core.pyi,sha256=FAlEwvbG4HgiXmupnWpi1jebTWRqVvE8QFOTxZhqbp0,251
+av/about.py,sha256=4sFnGHq13KvjziF6BIgyOblKiuj6VfT-h1DH3Sokz-o,23
+av/attachments/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/attachments/stream.cpython-311-x86_64-linux-gnu.so,sha256=LipN0GwJ_SkCOWKzHZTPxH8NHECde-1Zf9-sUcmJ19o,355185
+av/attachments/stream.pxd,sha256=UnV7BiZNp3gWpEV0zsY9tRPpzcxMZWm2glJ3KABk4GE,78
+av/attachments/stream.pyi,sha256=6q-XxZ5PNJM8jnjPnXXbgau9zGNfpeP02L0gFRqsMyU,178
+av/audio/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/audio/__init__.py,sha256=iqa8GHOGWz_CvVGp1WahggT77Wez0vXxBRSYSBE2rzM,62
+av/audio/__init__.pyi,sha256=TnYNY6cjWTYe9M7TFNXtoMPNq6nleziXWSGaO93rmls,103
+av/audio/codeccontext.cpython-311-x86_64-linux-gnu.so,sha256=2JofT07nW8UbE06xIYGfgO_MTdg-g2jfYkWUIXPL2RQ,474297
+av/audio/codeccontext.pxd,sha256=Yf6KxyJRo33v_ceKrqlk404otGtTvxZY-q3W_TISF6M,334
+av/audio/codeccontext.pyi,sha256=k-_y8gFQGb8PfDRQQrQ3ygXpLKNRWv1VlcL8Keua-58,1014
+av/audio/fifo.cpython-311-x86_64-linux-gnu.so,sha256=wrLug1wrNXfW_iRP_8EEx7r93zOQJ78P7iPLtkSwfho,720329
+av/audio/fifo.pxd,sha256=ZTO40EMP15fO1Adm34Qh4TzgQDXeVLLAn08DQlF5Q9c,461
+av/audio/fifo.pyi,sha256=I_nzPMFexnhzLcPZVjFQvIMevxrEHsEQxSjKmBdvh4E,712
+av/audio/format.cpython-311-x86_64-linux-gnu.so,sha256=IWDVomgjIFROKgmH1ib-PlnH1ZvNV-gcewiFMghV-2s,392577
+av/audio/format.pxd,sha256=Tl74C2mVPotu_X7veiBGvlNSpGwEMC26jvB6_DubKXs,203
+av/audio/format.pyi,sha256=SrRDgZkJ8kzKtywq2htlbJKK3d2ae9PIL-Qcnyoui8s,236
+av/audio/frame.cpython-311-x86_64-linux-gnu.so,sha256=yZSMYTiTZdP9EKyECpwQ-0BwabYANV4skLFp1DohzK8,978673
+av/audio/frame.pxd,sha256=yRuyThzIm-1HI5qCGAQkhawW77vKaDTADBfb793C5Zc,739
+av/audio/frame.pyi,sha256=Q8gi_l02RFpq5wuL14nOlLujpQCVpK2N1y6EnmOuEwI,1346
+av/audio/layout.cpython-311-x86_64-linux-gnu.so,sha256=A64TwmBdIFnBvI8qfpN3S2xWsxLQ79FUBXyq3b6H5bs,474721
+av/audio/layout.pxd,sha256=KUqev70Dwc-bH3IBWbhO_U3k_QQk_BNnUE5hy2XAasw,197
+av/audio/layout.pyi,sha256=cc0Y9RCsei824HWwZUt00tYAus92FKgIHuUk07lpEZw,250
+av/audio/plane.cpython-311-x86_64-linux-gnu.so,sha256=y4qsd9ZEy34qLRfZPnjVYozzS99dHkPZ27KxpjpYPSU,388025
+av/audio/plane.pxd,sha256=48WrkcprpCoOZsNpFT6C5R8ZIYeZwML_ndMQrAXLTnc,134
+av/audio/plane.pyi,sha256=84IuSlDcyAhLh1wR-QPer4s0mfFVvF-yc88v0qfPntQ,74
+av/audio/resampler.cpython-311-x86_64-linux-gnu.so,sha256=-P8bPr-dEHK7RBYadQIG8dsMZ4al2rFcUt2Lj3Khztc,777601
+av/audio/resampler.pxd,sha256=z7minznjFiwFsiriZL_A6RQW8APkA4fOlkjxpbnf6a8,488
+av/audio/resampler.pyi,sha256=dVyEvLOJ3IUN3AXsSYfg77H2Yaz2sIW38sJASH5fHRg,542
+av/audio/stream.cpython-311-x86_64-linux-gnu.so,sha256=gl6QVw4BDjAfqbo6a0MKiGxUatl90_XjwI5L3U0U_Cs,498817
+av/audio/stream.pxd,sha256=E38lNdZrVkiNbVLUJo8sgqFLQYb_UVTjOLmhM-cjbbg,209
+av/audio/stream.pyi,sha256=9OcxV5qNCwxGBN1vrommvbMfDwHtuumEPsujVw2ysbI,989
+av/bitstream.cpython-311-x86_64-linux-gnu.so,sha256=YWCmD-ww9X696ALHHSfrS-sQiv_hlKXEX0U1lvGnGCc,454017
+av/bitstream.pxd,sha256=hi30kwIdgdCQEVjtUqXo2L4gioHU_1_yWgKu3SqCgTc,184
+av/bitstream.pyi,sha256=CuumhPlDJV91iKtUpJDcuvPxighE9q2FzYmc0qtIN7A,389
+av/buffer.cpython-311-x86_64-linux-gnu.so,sha256=9iS-1LILo7qilgOEK7FoNbFg5RoEQ9KZPADnjjpRMWU,523633
+av/buffer.pxd,sha256=lUJVzgIKsCMeTie-IXPdOQGBpjI03U23fgbJ7506KIY,126
+av/buffer.pyi,sha256=0JCfkCBhhP3a3-_1AIhMf4sbNYJ1Uyy8GOrQ5-U4EVs,316
+av/bytesource.cpython-311-x86_64-linux-gnu.so,sha256=Sd_ayFTsTKHzg2PvwwRpXV_0JomRO2Lz6Q4fqfOnGTc,318561
+av/bytesource.pxd,sha256=dXfcjCPLXEQ4dCE12Cdo9t0kE-A_m5S6VW_fFxRRfuc,241
+av/codec/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/codec/__init__.py,sha256=e7SWLDOAb_79kT3sp20aZsXshp4s1UfJt2mocwfM7HM,255
+av/codec/codec.cpython-311-x86_64-linux-gnu.so,sha256=k5-MQBs_gVR_PigcA1HxqAAQY-lNbZeID5DdbCpe_Vg,856025
+av/codec/codec.pxd,sha256=pyzKMSmgS-3_ECiHKXkQTl6ifUJtaYKrZd8U4POjEgA,229
+av/codec/codec.pyi,sha256=cx-kawvhflHGEfhzpgK9JKl8aXoHR7kejfrI-Adq5vQ,3045
+av/codec/context.cpython-311-x86_64-linux-gnu.so,sha256=I1hRosk9XhzaarGwCQqPKH4U3Vu-tDuIHLmfJwFAqrM,1372465
+av/codec/context.pxd,sha256=r-sN1eF3_TdqLbITVaMSC5WN0NeVshnvpQL2L1ebjXc,2204
+av/codec/context.pyi,sha256=ho_m4Om06T55FBubHVXfaUqr4CtpWUW3D386P1iAdh0,2355
+av/container/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/container/__init__.py,sha256=71UExHmjXzXjVvUlD2ivW72Vb7FUW98ohbB4OLtMQ9I,111
+av/container/__init__.pyi,sha256=OvEGnWBxXGuuGVG9n8raY8mfhsO5JKM7PiG9NF1xHzQ,63
+av/container/core.cpython-311-x86_64-linux-gnu.so,sha256=hGBylRAgx5ksh3_7WWcHfbHcnPG9EDseVitZhijM8k4,1188113
+av/container/core.pxd,sha256=bQkHTPayqXJygVva3Y9KYbGm9OJvSFoxh1lkntIzMTg,1235
+av/container/core.pyi,sha256=5Y0weckc1Sz9KPSWrwVldVoybIFmX40zvqssYBwTHcM,3638
+av/container/input.cpython-311-x86_64-linux-gnu.so,sha256=fv63R9gbSWyvEOp1BYRob2CteniF3lJfvlfNEQ1fkMw,917521
+av/container/input.pxd,sha256=2K4YyQsAw_4l8RH3itgAW4Chguiw2k9QKwMmNnm3Jqk,163
+av/container/input.pyi,sha256=EK5c2mvUIsbH9o9oU9TtqHyD6oCw1HgF93hWLktmT2c,1610
+av/container/output.cpython-311-x86_64-linux-gnu.so,sha256=lh4hrHjAAmXhknpEAlMpFjOEtrWm0RzXd-jXnI88sMQ,1089689
+av/container/output.pxd,sha256=bznvdoePdkKchnHth3vUQJ3Lo-ynPpld7HU5NWiEtKs,244
+av/container/output.pyi,sha256=B-MZd4lWQwQs6R49nAKMRRQfTo9PssUFS28AzE1AB5s,1908
+av/container/pyio.cpython-311-x86_64-linux-gnu.so,sha256=YrFcAcG3LLKPOb2pH-fnbwqDEabnhwIyxKdeQi0D4Tg,712129
+av/container/pyio.pxd,sha256=Q1wOsyR09k5h4dS-wjk_CjkASabD3hDQ71W8FXvGjTQ,729
+av/container/streams.cpython-311-x86_64-linux-gnu.so,sha256=hwUE_wpVAfu71J_Y6w0E91wSVP5twAFUDnZsZ0TCxrI,929537
+av/container/streams.pxd,sha256=3sO8KT2Gp3PZi66YOAOuPbWtf9yoxJnV7aNuX41aeCA,514
+av/container/streams.pyi,sha256=AbYLuljXZTjXmpONkMfvPm0IJI2XN19un4JN4yr2l1Q,1229
+av/data/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/data/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/data/stream.cpython-311-x86_64-linux-gnu.so,sha256=YWwTdwxSY2UNzmemfh9D0XFN9bDVSuoDeS00LJmZgDo,388193
+av/data/stream.pxd,sha256=6zBdSTmLf77bjLDJqJu8IynSmRQkTJHEmQDnVhYgKpE,72
+av/data/stream.pyi,sha256=tTCCtKATxUAPiR0SZiM0zlf8ZOHLrE6VDNEfDJxENLA,133
+av/datasets.py,sha256=crxN4-UwGSOG6eHNfli7kSnOgwc3SK8fnY-p8Hje4Cc,3094
+av/descriptor.cpython-311-x86_64-linux-gnu.so,sha256=pRNaievJgetH3Myt_YI2s4MbreFVtHdi5-9xMDiJIbc,359561
+av/descriptor.pxd,sha256=nav7Vs2kmV1G4HpMBTSBR2TPtPsGZ5uYUMh1tfTmSHA,519
+av/descriptor.pyi,sha256=MuNobj8b0buKrhZ3cjdnZe3ggOLsDuUxGN-JD6iDbZo,121
+av/dictionary.cpython-311-x86_64-linux-gnu.so,sha256=mrL7BcxhXG6oy8XxKuOtViBMU0P9GBGhE4vpH7hpHO8,663409
+av/dictionary.pxd,sha256=iE0ZE3ZT5NDsmiT9HLQ8ACQxUvv6RhzSB5u3NeFRI3o,174
+av/dictionary.pyi,sha256=a51nyivuIlouh2LqbxHjnv9cQneros71w4a9Xzhxh-o,388
+av/error.cpython-311-x86_64-linux-gnu.so,sha256=44ief-G-rGI-IfsYxigiRcmCu2MUrZPICCxrr2caXf4,1802113
+av/error.pxd,sha256=07gJZT8560oSMfM24pm8jMTtTp4O3X_GPxLmOyypmd0,89
+av/error.pyi,sha256=sPH3uVvmVB7nxXkni0_LfImb1UZoQLrjmCsCYsVGPDk,3186
+av/filter/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/filter/__init__.py,sha256=wkkcm2Mo7dr9ZwMRKSz1ut3Mp8b3slR35i5C2JCsydk,131
+av/filter/__init__.pyi,sha256=Jkfg79hkGdMuvHM5_rIK_YyP1flJJqlpSrV6w-znzig,90
+av/filter/context.cpython-311-x86_64-linux-gnu.so,sha256=XZI27aKwkA2PPjYZfbpFuhPAtDarZJzqPYJRqyHM8QM,835217
+av/filter/context.pxd,sha256=ytFXp6WCjHWm89b2hs4xQpVEv9H9MqiOQFtYhrCyxho,388
+av/filter/context.pyi,sha256=iSXEfdSM3pytvaiD9AGKTdt2evllYupku5NeGuyTASc,536
+av/filter/filter.cpython-311-x86_64-linux-gnu.so,sha256=Ztav_b37DV3hExeq8cnjfp2E7KUN2rfVa6hV-99AyE8,982713
+av/filter/filter.pxd,sha256=ArARpXHdaAd2fz5BdZRV-wIf5AhkjynBi_eXah_G5CI,248
+av/filter/filter.pyi,sha256=lYbyrvJXs2OpEpa6-uGCb1nM0Pu6YmJKu61AoP9XgIE,500
+av/filter/graph.cpython-311-x86_64-linux-gnu.so,sha256=6-tX_GAmOus8qRZ8Mi2OG_3tF_beV_3GoNFAzGbvbsk,1077057
+av/filter/graph.pxd,sha256=keZlIspoOPiIqSR9_Fk2xXo3PFgwLON1oKbIVWHfxUk,518
+av/filter/graph.pyi,sha256=jnxZF-9DLv6uaHY9u33a26GMUpLPBxrTUom6_IxIviI,1660
+av/filter/link.cpython-311-x86_64-linux-gnu.so,sha256=nxnr1w8Z16vsPG5cM4whI4pNGC3yXwwQ96nuc1CdJv4,404601
+av/filter/link.pxd,sha256=DEG229dm5_xf3VApBc-UyBe26q-IWKhjbPLrqOv8-dI,327
+av/filter/link.pyi,sha256=4psi3xhfutTxLD82vWVDdTgB_QsuqDT3gFsaV1IiMkg,110
+av/filter/loudnorm.cpython-311-x86_64-linux-gnu.so,sha256=ghQhCKshJVEKzujL5VNXK59u8E8cIKAl3k5D56FZm4I,446337
+av/filter/loudnorm.pxd,sha256=NEHEd-3ot9wa6ANJ9Tw7a2V6BTPVw3GKt1gGS3dpLQ0,100
+av/filter/loudnorm.pyi,sha256=N_HuXpQlTaZLpnKjfqMk68ac3djymdy3x_9en33sA74,106
+av/filter/pad.cpython-311-x86_64-linux-gnu.so,sha256=-mskz9RPIAhyknVY3z8-Io4Pd5o98TUnAKbNiG0_Ygw,490817
+av/filter/pad.pxd,sha256=HCbx0vtexIbyTsBv_RmADZRM_v2y83CXnvMeJOnkPVM,518
+av/filter/pad.pyi,sha256=dtJ-oUfUjIuh-1b-eyPAL1lmFccBaHlDDXBDlyK4KGI,195
+av/format.cpython-311-x86_64-linux-gnu.so,sha256=p3_6qbg1-PLkq9iN5tX0uHk-3BgKVoSaiia148FUeDg,527977
+av/format.pxd,sha256=7g7FhC8pBeFxDNErwJVQUnGtNaK1QnPbqwYMH-R9FmE,234
+av/format.pyi,sha256=4-R8otDwbsvNGUzLoWoZC4Pz1xhqO6WrU42ZboSNiSw,1193
+av/frame.cpython-311-x86_64-linux-gnu.so,sha256=lbpDhBDSQ53MwCUdGFmnDlIq1dhAP8CJtFWD7u8kSKY,503273
+av/frame.pxd,sha256=MdtAffABVLDFi_c5oQFm3lT-n6xVbya8IBEsppbf6tU,411
+av/frame.pyi,sha256=11NZOtXaPeAyqhNwbN1LxMIkitJvPTwAhxgZGmJK9PQ,507
+av/include/libav.pxd,sha256=dtszNF-Rn1WNtwjU6uYrm8Gn47z-hbtDaWpVuvH9NAA,799
+av/include/libavcodec/avcodec.pxd,sha256=keoIHXlP5qbo8Vqv6ENZlxgRWOUsIGITgW1WL9gnHSM,15055
+av/include/libavcodec/bsf.pxd,sha256=xyTavBNhAE4oTkovaZumeU1JDSxzkIge4Q2xNemBtM4,894
+av/include/libavdevice/avdevice.pxd,sha256=PY8CO4hKCGPwVGg6Act8qQOnX6AaKHUHp40Cs9aoDa0,465
+av/include/libavfilter/avfilter.pxd,sha256=z2xgL2BuHJd-NTwwT7nnM3Ocz0fWsVMA0AqpvgewWZ8,2775
+av/include/libavfilter/avfiltergraph.pxd,sha256=7S_DjuOXiQzdTHgTcybQjInjmeFeHOnRQdnr4ZgvYFE,1278
+av/include/libavfilter/buffersink.pxd,sha256=6LM7cAayzCA0_DAFUUIduZGA62pn9tdHRSvQuZv-GP4,144
+av/include/libavfilter/buffersrc.pxd,sha256=kT-0WFhgFQHT-71N_fV3GucZj3ATJlUKif8vcYKGsAE,150
+av/include/libavformat/avformat.pxd,sha256=scYWbY5u2G0v42NVFFsfDW5_E1YLKc3rQ1ptIvGbmB0,8323
+av/include/libavutil/avutil.pxd,sha256=3zgEyxs2rxXmCYReHaZMcLCTFUH2CJpIYht8eQGoh-U,10041
+av/include/libavutil/buffer.pxd,sha256=xBBm37NMvQOp7DKI3GfDXGaY0ulMexA-efadAuwehu8,306
+av/include/libavutil/channel_layout.pxd,sha256=f0FM5O9zclV0okfNep_7QSLd2Jou45imhR4f5CkM7wg,412
+av/include/libavutil/dict.pxd,sha256=vdDof0JwBUSdaBIUa394LpWYvw_D01uQ0IjdnT1ADN8,832
+av/include/libavutil/error.pxd,sha256=mvsofkeCxD0tbA846JDc5Sl9o0TZ9iCOKP6l8zgU7Vw,1277
+av/include/libavutil/frame.pxd,sha256=qW5BCgqDxVVMg5YzeSlrJbeee2ERctBwuCTpnF4uCEU,718
+av/include/libavutil/motion_vector.pxd,sha256=VZ7Yl0Tp_p0l0rUNeHTzT4YP_4KlGNfelWdV1G4iMIU,450
+av/include/libavutil/samplefmt.pxd,sha256=rO7gCeqaMUr955eYZGttT6xP20W9Jk6uBrGSaaktUz4,1647
+av/include/libswresample/swresample.pxd,sha256=9ncJdDi6rwjhFju7eg8QCI1LIdymlltUk6Ev6wfD8d0,1078
+av/include/libswscale/swscale.pxd,sha256=6T2-QvuOpL5nErO9jhzG2g1yCv9ZdqlFLY4n9o1RRRE,2307
+av/logging.cpython-311-x86_64-linux-gnu.so,sha256=YKbiwei5yRP4Cf589QNBE6msurMB4kKyqefb4owj6_8,1044313
+av/logging.pxd,sha256=9WCF9ygPhR9OO5RxYNHBa6Mz_KqGWNFbIr1D54r2B8o,24
+av/logging.pyi,sha256=Afqq5ud46eh617UM6b_D3xpIisdz9uyxvKK4rK1n6RA,885
+av/opaque.cpython-311-x86_64-linux-gnu.so,sha256=6CJusNXukGnhSsR9bpfYzalg3mgguUUZf-4Kil_Bleo,347297
+av/opaque.pxd,sha256=fFkinAiXKDuHp8E0f-ZQ0sskpFPvj8hRCLf_KK6PgFM,237
+av/option.cpython-311-x86_64-linux-gnu.so,sha256=yQSRlGIsUZ7wyy5eqOWi_RFLSr6qQj-CphClU2Y1RJ0,642377
+av/option.pxd,sha256=J3fikrPq554_N2x4tw0b0n9LDY-TqFd-wXXrhAkSoEA,366
+av/option.pyi,sha256=i7x22Ntv8MiL0gAibXPRuE4kU4_2XpZH1tSpHj6KmN0,977
+av/packet.cpython-311-x86_64-linux-gnu.so,sha256=3T_iwso8kZgJi2Mao35YGuhxOyGIVV7sgQmAXKN2m6Q,597449
+av/packet.pxd,sha256=z2UQIRbz-wHHOuGploQg7K2g1i0afym32ig-ZcYZAgo,447
+av/packet.pyi,sha256=BUBVujWi7ZkNlI4A8cfyJUjhvSwbOTWkAM2YuvStFbQ,559
+av/plane.cpython-311-x86_64-linux-gnu.so,sha256=xtvWmLYs6dt1a5XlooFQ2C6CO3U_Jq6V8nrwk4XGlXU,425121
+av/plane.pxd,sha256=qtt0qmUOXmL5Ak8TlFYCLi4oiesC1kU4Ilesqk7hNlY,196
+av/plane.pyi,sha256=qJLw7M73ekThXf05LqTllf1BLIfF3gckxEy6rEjKx4U,169
+av/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/sidedata/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/sidedata/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/sidedata/motionvectors.cpython-311-x86_64-linux-gnu.so,sha256=Uh-_1ctDKIJYwXCZ78gJbR6WXAdcSsPqMzRGtSbZ6oY,703953
+av/sidedata/motionvectors.pxd,sha256=Ly2xMmZ0Bj1An5-eYViZgNhJtSgSkSI7E78EiS4PKkE,266
+av/sidedata/motionvectors.pyi,sha256=InIFdsNMbPU0apf_ufDtKoD0KAr0VCXaPkwGD9fLHCg,597
+av/sidedata/sidedata.cpython-311-x86_64-linux-gnu.so,sha256=tCt-bvhD4P3HGGkClXLvAMbsclVGA33zF5NlhUN4nBU,831089
+av/sidedata/sidedata.pxd,sha256=DqadUtnTm0yN1pq5NII_W2COFAAAOk6YooiUzxOR4hk,409
+av/sidedata/sidedata.pyi,sha256=GaKV71PuL24eY5xczIMQGoUnW7u55U2svoqqTZNwF8o,1627
+av/stream.cpython-311-x86_64-linux-gnu.so,sha256=ndZqpXqFfyX7TensqSxZIqUiB2wyjdaGCwJUH2kC8i4,609681
+av/stream.pxd,sha256=hDl0Xr5bO3xgsiVRiB43BZDv397GbbyZTPeDt92cdXI,635
+av/stream.pyi,sha256=87jtzAlp3m9ysSachGPW4zXEx2Qy8leroVJmLjSqcLo,638
+av/subtitles/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/subtitles/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/subtitles/codeccontext.cpython-311-x86_64-linux-gnu.so,sha256=OjbVCkHCz_6zpkkvrqbsfI0ZQ4TOXm5U6RtleH3pxpY,371697
+av/subtitles/codeccontext.pxd,sha256=vnN6tA3y9vQxD4ZAk51sfDZl5mbTwOU4VamAUeuBiW0,101
+av/subtitles/codeccontext.pyi,sha256=DtCjVm3TqN6jOUc2IZPBB-sFIVxUNEcFa2EwNL8_4Cs,143
+av/subtitles/stream.cpython-311-x86_64-linux-gnu.so,sha256=-ivDHYn44y0KoIRVvL6yzXc5SUMD1A_bWL6HoUxPGbM,400337
+av/subtitles/stream.pxd,sha256=Lo4RBTTOWErzKgNqc07AoNDaAgkUsF_Q25DJaAIo52Y,137
+av/subtitles/stream.pyi,sha256=SeFOe9EPAw3SWfwOpP7JPHOSsoovxC95reqeIfurngQ,212
+av/subtitles/subtitle.cpython-311-x86_64-linux-gnu.so,sha256=LWnc_YJmMCMwipL5elqlR6dZS0zuhu3_8n36dVhbkz0,1011585
+av/subtitles/subtitle.pxd,sha256=SE6M-1ylGQB6F0xnUwLTs7xODznkIs81MYKsVrzGBd0,592
+av/subtitles/subtitle.pyi,sha256=XHFT-7hgtc4ISnTrt6eslSiFCcc0xt1312LUnVEIjNk,804
+av/utils.cpython-311-x86_64-linux-gnu.so,sha256=DXNZ1ZwWSOYjjF8wMDYKrSCDFTT0pzFcQcPXVrShxlQ,289681
+av/utils.pxd,sha256=mHNd9sR9mNP12Yfm4Derd4LULoyzlJeU7Gdf86ZvjpQ,452
+av/video/__init__.pxd,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
+av/video/__init__.py,sha256=BRH_qu6VrLKnGq1hpmXmYQVO_o-r32fnsR-V6NmbLis,62
+av/video/__init__.pyi,sha256=0IQg23uvxHelDoV-DJDoT-QuVGAnktL6WmWo4cmvT6U,103
+av/video/codeccontext.cpython-311-x86_64-linux-gnu.so,sha256=Sii-8O-SvrYs3BFLU4gNaX1o_s5O3J9CYPi6l8cxFFY,761129
+av/video/codeccontext.pxd,sha256=ORf03B8jXhyRmZpPjD02Xzy1LR6xBC9zagKD_Zco8o8,483
+av/video/codeccontext.pyi,sha256=V3D3vRTGUJlZM7UGqv5L4XEdYwX9BIbtYNUmKCMJfEo,968
+av/video/format.cpython-311-x86_64-linux-gnu.so,sha256=2puA96W0pzpmD4yD0SwuC5XHgfh0FE3F0bprCCrRABc,892793
+av/video/format.pxd,sha256=jyu_VoT833m7msVQ3PGRfZ6r_-2PLbDSMJ5O9ZCQyPQ,727
+av/video/format.pyi,sha256=cnkyQNxgCbpcoDK2hc-ixSJ3l9ospSd5XBAdjE4FPSQ,694
+av/video/frame.cpython-311-x86_64-linux-gnu.so,sha256=5eN_N5EwaVGAV_j_GMFbrJuzAg82-npPWJTZ5AgvI4c,3289633
+av/video/frame.pxd,sha256=5jBtLdUmvtxJ7ERSOX7hhzd9yp-rDLHPxeLhxNTnzbo,621
+av/video/frame.pyi,sha256=najv_Wjq0MeNE8iwn9JeItDXbmAYezyEN58gjwi4eiE,2082
+av/video/plane.cpython-311-x86_64-linux-gnu.so,sha256=94AmWzwDuoNSGuy4Gpq73w172EUWRkZbCWuRRCbfe8k,470233
+av/video/plane.pxd,sha256=K7jqsI_cz-zPl6-UMS-6zNLSj3gD9BuLC_deNnegPC0,193
+av/video/plane.pyi,sha256=x-RqS_bUUBPv2ymfroBYEz4ROsEw1bpxlnKe5GEvbZ8,223
+av/video/reformatter.cpython-311-x86_64-linux-gnu.so,sha256=b8_St_CncCpZAwvPTYtfFgPr6w2C-_sHMOtCUVtyqbc,810697
+av/video/reformatter.pxd,sha256=ugLZ44hKS8-OYLJCsvfLrL1OhtQzMPI-oqCvrYhpRuc,373
+av/video/reformatter.pyi,sha256=HGkfXHcON_JN-ploOKOYOvGBuZaZbJ24q25IO1OgzsY,1062
+av/video/stream.cpython-311-x86_64-linux-gnu.so,sha256=SC8NEYD71kV3d0Tv-a0cg6RdOCAuUnkDTYutZF4qHPQ,601377
+av/video/stream.pxd,sha256=-VhPgrv_87HGxdBkfwR5qIO3jwwyQqMwRvbTH7cunuA,209
+av/video/stream.pyi,sha256=JZKIjhHzj90bs-_iN9dP6pmw1M3TSPbp_XlD73Cv7H4,1187
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/REQUESTED b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/REQUESTED
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/WHEEL b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/WHEEL
new file mode 100644
index 0000000000000000000000000000000000000000..8b0363604372a86351ec451bb3c5afa26234e6f0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/WHEEL
@@ -0,0 +1,6 @@
+Wheel-Version: 1.0
+Generator: setuptools (75.6.0)
+Root-Is-Purelib: false
+Tag: cp311-cp311-manylinux_2_17_x86_64
+Tag: cp311-cp311-manylinux2014_x86_64
+
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/entry_points.txt b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/entry_points.txt
new file mode 100644
index 0000000000000000000000000000000000000000..60fa2439e030a3e73d660335bb0f93033d182c92
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/entry_points.txt
@@ -0,0 +1,2 @@
+[console_scripts]
+pyav = av.__main__:main
diff --git a/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/top_level.txt b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/top_level.txt
new file mode 100644
index 0000000000000000000000000000000000000000..dc1ce6e8a1399e86a202114c4837e05d5ee491c7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av-14.0.1.dist-info/top_level.txt
@@ -0,0 +1 @@
+av
diff --git a/venv/lib/python3.11/site-packages/av.libs/libXau-00ec42fe.so.6.0.0 b/venv/lib/python3.11/site-packages/av.libs/libXau-00ec42fe.so.6.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..d05343446689f8639c6ba8ab99d1573827f74fef
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libXau-00ec42fe.so.6.0.0 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libaom-e9efed4a.so.3.2.0 b/venv/lib/python3.11/site-packages/av.libs/libaom-e9efed4a.so.3.2.0
new file mode 100644
index 0000000000000000000000000000000000000000..eb7a34884d23373a5c7d7fa2f60b26070898708e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libaom-e9efed4a.so.3.2.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:116ef06755f03375a2c5a884f968ecb8dcf44b3141d0983a2ec22d4bc0f06951
+size 7285769
diff --git a/venv/lib/python3.11/site-packages/av.libs/libavcodec-fb6c662d.so.61.19.100 b/venv/lib/python3.11/site-packages/av.libs/libavcodec-fb6c662d.so.61.19.100
new file mode 100644
index 0000000000000000000000000000000000000000..5e943172bb11f042782de7ed5c38798d7d3720f0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libavcodec-fb6c662d.so.61.19.100
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:c17b999a064ff37063c90195d3a47ea8f49c0e41eabd078177859962a796abdd
+size 14588393
diff --git a/venv/lib/python3.11/site-packages/av.libs/libavdevice-84390fdd.so.61.3.100 b/venv/lib/python3.11/site-packages/av.libs/libavdevice-84390fdd.so.61.3.100
new file mode 100644
index 0000000000000000000000000000000000000000..b17db2476bf6833dadbd21ccda771b4be75f85d0
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libavdevice-84390fdd.so.61.3.100 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libavfilter-3235a7c8.so.10.4.100 b/venv/lib/python3.11/site-packages/av.libs/libavfilter-3235a7c8.so.10.4.100
new file mode 100644
index 0000000000000000000000000000000000000000..96c41fa7f3a79d9003d65f51d689110fa1470c72
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libavfilter-3235a7c8.so.10.4.100
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:dad6e877fcc5d52ff09b4bd13b8b5c1f9b80902bbc9a4dd7cfedb155070188d0
+size 4526257
diff --git a/venv/lib/python3.11/site-packages/av.libs/libavformat-071c54bd.so.61.7.100 b/venv/lib/python3.11/site-packages/av.libs/libavformat-071c54bd.so.61.7.100
new file mode 100644
index 0000000000000000000000000000000000000000..db75e0176678b57f71cd3d08cddeb01848924b87
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libavformat-071c54bd.so.61.7.100
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:ed69b25c954e3e589fe26484acd11ef42cb2ce616b64e5e2d3964a94c4b46721
+size 2782073
diff --git a/venv/lib/python3.11/site-packages/av.libs/libavutil-2749f1ba.so.59.39.100 b/venv/lib/python3.11/site-packages/av.libs/libavutil-2749f1ba.so.59.39.100
new file mode 100644
index 0000000000000000000000000000000000000000..fc1822955b8ebb9961f2d0b84385e8e7df2464d3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libavutil-2749f1ba.so.59.39.100
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:9becb70e8910abca1cf8d2370a51489684bf51b88f8ab205052b21baa303f07f
+size 1028897
diff --git a/venv/lib/python3.11/site-packages/av.libs/libdav1d-1b53ef2f.so.7.0.0 b/venv/lib/python3.11/site-packages/av.libs/libdav1d-1b53ef2f.so.7.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..c7e8ea75d04cf61c536fed0f518bc31a856a309f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libdav1d-1b53ef2f.so.7.0.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:e123f4ff6c096ff130ce23c655e33f46519626ec835099db1296d6880e59a367
+size 1922417
diff --git a/venv/lib/python3.11/site-packages/av.libs/libdrm-827b956f.so.2.4.0 b/venv/lib/python3.11/site-packages/av.libs/libdrm-827b956f.so.2.4.0
new file mode 100644
index 0000000000000000000000000000000000000000..109074994a77380b2d3b593671b504e3d6bcac61
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libdrm-827b956f.so.2.4.0 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libgmp-a4b719d5.so.10.5.0 b/venv/lib/python3.11/site-packages/av.libs/libgmp-a4b719d5.so.10.5.0
new file mode 100644
index 0000000000000000000000000000000000000000..4064aa24e0db5d3171c710394f0300e9397b33d1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libgmp-a4b719d5.so.10.5.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:2c9e45b1f7603775756feb648d8452bf08fe1e949199e2d44c53974fb741d00b
+size 507625
diff --git a/venv/lib/python3.11/site-packages/av.libs/libgnutls-b9b94016.so.30.36.0 b/venv/lib/python3.11/site-packages/av.libs/libgnutls-b9b94016.so.30.36.0
new file mode 100644
index 0000000000000000000000000000000000000000..fbb93bf58bf2cd687426982bbabf8c7924dc5f0d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libgnutls-b9b94016.so.30.36.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:6b9ba2f279fbd3889f6f33abb1ac1fd6b30d5f19a81530fd3f5be31426fe685f
+size 2218121
diff --git a/venv/lib/python3.11/site-packages/av.libs/libhogweed-9544c08c.so.6.8 b/venv/lib/python3.11/site-packages/av.libs/libhogweed-9544c08c.so.6.8
new file mode 100644
index 0000000000000000000000000000000000000000..4f415a8cdbdad6c2c64e108a9f50650b1f0060f3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libhogweed-9544c08c.so.6.8
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:1b8864fc46f5af46d0575c9af3bd237e00cb0ee3a6bfd2ac8ddbc03df76210ef
+size 316505
diff --git a/venv/lib/python3.11/site-packages/av.libs/liblzma-af70179d.so.5.4.4 b/venv/lib/python3.11/site-packages/av.libs/liblzma-af70179d.so.5.4.4
new file mode 100644
index 0000000000000000000000000000000000000000..82896ccfef8a1ccd47706d80d137d5649c0b51b8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/liblzma-af70179d.so.5.4.4
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:6af625c1a4972a0a4cd0f8620a51ef1454a8198eecf577e9b0bbf69178b4010c
+size 196537
diff --git a/venv/lib/python3.11/site-packages/av.libs/libmp3lame-3ecc6556.so.0.0.0 b/venv/lib/python3.11/site-packages/av.libs/libmp3lame-3ecc6556.so.0.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..8af610051833c5810f65388bdd2bb17d4eaa0cce
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libmp3lame-3ecc6556.so.0.0.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:8709329fe7d9de1ad08dbf2318cca16b43c9abf3c7246c3ee801d8c2f8bdddb3
+size 417001
diff --git a/venv/lib/python3.11/site-packages/av.libs/libnettle-ecd2e589.so.8.8 b/venv/lib/python3.11/site-packages/av.libs/libnettle-ecd2e589.so.8.8
new file mode 100644
index 0000000000000000000000000000000000000000..2070193b0415b025d0472e0f02a435a87225764d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libnettle-ecd2e589.so.8.8
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:e1b054f48a7bf570ed0aacb36f1bffce8578be0c74c49de01ea9ee2a4c5e23f0
+size 351169
diff --git a/venv/lib/python3.11/site-packages/av.libs/libogg-bbd52b06.so.0.8.5 b/venv/lib/python3.11/site-packages/av.libs/libogg-bbd52b06.so.0.8.5
new file mode 100644
index 0000000000000000000000000000000000000000..b85bd1c734b7ce4c4587a627a5bbe76ab98542b0
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libogg-bbd52b06.so.0.8.5 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libopencore-amrnb-393dbae2.so.0.0.3 b/venv/lib/python3.11/site-packages/av.libs/libopencore-amrnb-393dbae2.so.0.0.3
new file mode 100644
index 0000000000000000000000000000000000000000..f5a64de9351e23f08b4d8b8f0a5c39d1d34c4cc6
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libopencore-amrnb-393dbae2.so.0.0.3
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:98c671a0a84457a51bd1127c838c0ea3f478c47c143a317008a827eae43f3ca2
+size 172889
diff --git a/venv/lib/python3.11/site-packages/av.libs/libopencore-amrwb-9db94aa9.so.0.0.3 b/venv/lib/python3.11/site-packages/av.libs/libopencore-amrwb-9db94aa9.so.0.0.3
new file mode 100644
index 0000000000000000000000000000000000000000..4edb967ea1775c4c6dedfb646dc0e26a0e803779
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libopencore-amrwb-9db94aa9.so.0.0.3 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libopus-59fcdf85.so.0.9.0 b/venv/lib/python3.11/site-packages/av.libs/libopus-59fcdf85.so.0.9.0
new file mode 100644
index 0000000000000000000000000000000000000000..b564c75b5afc0f7648e7e609c9feb52e75ca80a0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libopus-59fcdf85.so.0.9.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:37141eb2540a385ac146d842c66063e9a2e467e569631378b3d3138e073a435c
+size 375521
diff --git a/venv/lib/python3.11/site-packages/av.libs/libpostproc-dde0c2af.so.58.3.100 b/venv/lib/python3.11/site-packages/av.libs/libpostproc-dde0c2af.so.58.3.100
new file mode 100644
index 0000000000000000000000000000000000000000..326f3ab114a271d8eed7ee728dbb8f38e4bce8a2
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libpostproc-dde0c2af.so.58.3.100 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libsharpyuv-6e5d6e5b.so.0.1.0 b/venv/lib/python3.11/site-packages/av.libs/libsharpyuv-6e5d6e5b.so.0.1.0
new file mode 100644
index 0000000000000000000000000000000000000000..834cb7a515ecfd808a12f85eaa034354041c549c
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libsharpyuv-6e5d6e5b.so.0.1.0 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libspeex-2370356a.so.1.5.2 b/venv/lib/python3.11/site-packages/av.libs/libspeex-2370356a.so.1.5.2
new file mode 100644
index 0000000000000000000000000000000000000000..0d60947536fb224c19289bb9567c5b2110fd5830
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libspeex-2370356a.so.1.5.2
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:1d82708d8085c533237f4d181f03fe58e5b4bf4382e7c376623041f155a60e0e
+size 108521
diff --git a/venv/lib/python3.11/site-packages/av.libs/libswresample-da7d062e.so.5.3.100 b/venv/lib/python3.11/site-packages/av.libs/libswresample-da7d062e.so.5.3.100
new file mode 100644
index 0000000000000000000000000000000000000000..a9bf9e2f2584bf15b45c6328a1ae9c3ec755b23e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libswresample-da7d062e.so.5.3.100
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:8139a3fc3015562707ec46a80b60a1f01c35c3a53ade289a0c04a3cf55af5ec2
+size 128265
diff --git a/venv/lib/python3.11/site-packages/av.libs/libswscale-9212cf18.so.8.3.100 b/venv/lib/python3.11/site-packages/av.libs/libswscale-9212cf18.so.8.3.100
new file mode 100644
index 0000000000000000000000000000000000000000..43f5e126ba816ad4a1503502efae9d09000cc348
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libswscale-9212cf18.so.8.3.100
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:ee46735b2a75fd8b2378ded8d68eb9c1e2925a4fb74e34c1c9bde9124ff7540d
+size 619929
diff --git a/venv/lib/python3.11/site-packages/av.libs/libtwolame-72d74ef7.so.0.0.0 b/venv/lib/python3.11/site-packages/av.libs/libtwolame-72d74ef7.so.0.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..9e254b2b595c2a2b1390678c5b330f2e3e7472fe
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libtwolame-72d74ef7.so.0.0.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:d8f728d98dc8457d20047b28dae507fcd6f6b063690693b5a65d340297a26b56
+size 143513
diff --git a/venv/lib/python3.11/site-packages/av.libs/libunistring-214e3d6e.so.5.1.0 b/venv/lib/python3.11/site-packages/av.libs/libunistring-214e3d6e.so.5.1.0
new file mode 100644
index 0000000000000000000000000000000000000000..ed9dfbf1598e6731afc1439dd0dbe7e789a3d573
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libunistring-214e3d6e.so.5.1.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:82a9c952cb85d7efab7ad3c76c757060de96410768b6cd4ecea4e42dde1a4eae
+size 1815865
diff --git a/venv/lib/python3.11/site-packages/av.libs/libvorbis-f4a9a6fd.so.0.4.9 b/venv/lib/python3.11/site-packages/av.libs/libvorbis-f4a9a6fd.so.0.4.9
new file mode 100644
index 0000000000000000000000000000000000000000..857ecd918faaaff950b8ec964957a3632fea4afd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libvorbis-f4a9a6fd.so.0.4.9
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:1a53109c850a4875f7e39e1cc977e4f8512d02780c8c8db2ea08bb1f94aadfcd
+size 244425
diff --git a/venv/lib/python3.11/site-packages/av.libs/libvorbisenc-0d9d5bdf.so.2.0.12 b/venv/lib/python3.11/site-packages/av.libs/libvorbisenc-0d9d5bdf.so.2.0.12
new file mode 100644
index 0000000000000000000000000000000000000000..4824e0f3b181daf0e3fd92756ff3056194a55452
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libvorbisenc-0d9d5bdf.so.2.0.12
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:efdd18fb99dbc73d3d774ed90fd9f6bf4569f407b686cd4a5a7d3fa9126c7e4f
+size 713145
diff --git a/venv/lib/python3.11/site-packages/av.libs/libvpx-832f6f52.so.9.0.0 b/venv/lib/python3.11/site-packages/av.libs/libvpx-832f6f52.so.9.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..4ba90f4362f555e216a3a7edcb9febe30ede0923
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libvpx-832f6f52.so.9.0.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:2f3957416da144c6fca717125a9b1521a37142fab92fa16dc6643c850fb9cc38
+size 2209833
diff --git a/venv/lib/python3.11/site-packages/av.libs/libwebp-e16038c7.so.7.1.9 b/venv/lib/python3.11/site-packages/av.libs/libwebp-e16038c7.so.7.1.9
new file mode 100644
index 0000000000000000000000000000000000000000..9d4fe3c13a63ea1ac8bdb9734c3dd995b789ce50
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libwebp-e16038c7.so.7.1.9
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:630ddf17bfb047446aa3bde8ddda0269e9b7e10ed4395ebca92601271aca45cc
+size 637017
diff --git a/venv/lib/python3.11/site-packages/av.libs/libwebpmux-03713f16.so.3.1.0 b/venv/lib/python3.11/site-packages/av.libs/libwebpmux-03713f16.so.3.1.0
new file mode 100644
index 0000000000000000000000000000000000000000..b585420cde22fd8496beccbbb6e9d5a807ebfef2
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libwebpmux-03713f16.so.3.1.0 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libx264-2a4c6f6d.so.164 b/venv/lib/python3.11/site-packages/av.libs/libx264-2a4c6f6d.so.164
new file mode 100644
index 0000000000000000000000000000000000000000..6376abd58d607a1e4a3d2e339241672629252161
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libx264-2a4c6f6d.so.164
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:84df92fd734647e782211c1ed35eaca6949bb6c4c8ed88891edbc4c644002b2c
+size 2267921
diff --git a/venv/lib/python3.11/site-packages/av.libs/libx265-d8690e8d.so.199 b/venv/lib/python3.11/site-packages/av.libs/libx265-d8690e8d.so.199
new file mode 100644
index 0000000000000000000000000000000000000000..2e75ef5f6bdcd7d2d7810b2fdaace499c97df3a0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libx265-d8690e8d.so.199
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:50e786017692df4f7185d65d8297bf49339cbe27c6197ad70066250205acee0a
+size 19369233
diff --git a/venv/lib/python3.11/site-packages/av.libs/libxcb-65da195c.so.1.1.0 b/venv/lib/python3.11/site-packages/av.libs/libxcb-65da195c.so.1.1.0
new file mode 100644
index 0000000000000000000000000000000000000000..2e361f0dda26b610340f764c2c62e802551156d4
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libxcb-65da195c.so.1.1.0
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:cdc3d3b87f0eb76993bc47557fd332ffd275d908f79d5e57c138298975427bbe
+size 210465
diff --git a/venv/lib/python3.11/site-packages/av.libs/libxcb-shape-25c2b258.so.0.0.0 b/venv/lib/python3.11/site-packages/av.libs/libxcb-shape-25c2b258.so.0.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..eabcf6681072d53fed0aac80bf8b8187f9abab41
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libxcb-shape-25c2b258.so.0.0.0 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libxcb-shm-7a199f70.so.0.0.0 b/venv/lib/python3.11/site-packages/av.libs/libxcb-shm-7a199f70.so.0.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..ef12fcd5a6100030095e01bdd0366003530e7254
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libxcb-shm-7a199f70.so.0.0.0 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libxcb-xfixes-9be3ba6f.so.0.0.0 b/venv/lib/python3.11/site-packages/av.libs/libxcb-xfixes-9be3ba6f.so.0.0.0
new file mode 100644
index 0000000000000000000000000000000000000000..18fe7c8898a2493c6f5b4f36b0c71c6f8fb80f4d
Binary files /dev/null and b/venv/lib/python3.11/site-packages/av.libs/libxcb-xfixes-9be3ba6f.so.0.0.0 differ
diff --git a/venv/lib/python3.11/site-packages/av.libs/libxml2-cb941fce.so.2.9.13 b/venv/lib/python3.11/site-packages/av.libs/libxml2-cb941fce.so.2.9.13
new file mode 100644
index 0000000000000000000000000000000000000000..530a5f42c08a7c3fc4685e37a58c5eaa946516b9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av.libs/libxml2-cb941fce.so.2.9.13
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:5397eafc280d4871851525eae9960a5e6cbb26770911f096475c9cc871c01e39
+size 1609769
diff --git a/venv/lib/python3.11/site-packages/av/__init__.pxd b/venv/lib/python3.11/site-packages/av/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/__init__.py b/venv/lib/python3.11/site-packages/av/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..cbc3c8a2ff1b80c8207bc493d1ad337e95f126b9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/__init__.py
@@ -0,0 +1,68 @@
+# MUST import the core before anything else in order to initialize the underlying
+# library that is being wrapped.
+from av._core import time_base, library_versions, ffmpeg_version_info
+
+# Capture logging (by importing it).
+from av import logging
+
+# For convenience, import all common attributes.
+from av.about import __version__
+from av.audio.codeccontext import AudioCodecContext
+from av.audio.fifo import AudioFifo
+from av.audio.format import AudioFormat
+from av.audio.frame import AudioFrame
+from av.audio.layout import AudioLayout
+from av.audio.resampler import AudioResampler
+from av.audio.stream import AudioStream
+from av.bitstream import BitStreamFilterContext, bitstream_filters_available
+from av.codec.codec import Codec, codecs_available
+from av.codec.context import CodecContext
+from av.container import open
+from av.format import ContainerFormat, formats_available
+from av.packet import Packet
+from av.error import * # noqa: F403; This is limited to exception types.
+from av.video.codeccontext import VideoCodecContext
+from av.video.format import VideoFormat
+from av.video.frame import VideoFrame
+from av.video.stream import VideoStream
+
+__all__ = (
+ "__version__",
+ "time_base",
+ "ffmpeg_version_info",
+ "library_versions",
+ "AudioCodecContext",
+ "AudioFifo",
+ "AudioFormat",
+ "AudioFrame",
+ "AudioLayout",
+ "AudioResampler",
+ "AudioStream",
+ "BitStreamFilterContext",
+ "bitstream_filters_available",
+ "Codec",
+ "codecs_available",
+ "CodecContext",
+ "open",
+ "ContainerFormat",
+ "formats_available",
+ "Packet",
+ "VideoCodecContext",
+ "VideoFormat",
+ "VideoFrame",
+ "VideoStream",
+)
+
+
+def get_include() -> str:
+ """
+ Returns the path to the `include` folder to be used when building extensions to av.
+ """
+ import os
+
+ # Installed package
+ include_path = os.path.join(os.path.dirname(__file__), "include")
+ if os.path.exists(include_path):
+ return include_path
+ # Running from source directory
+ return os.path.join(os.path.dirname(__file__), os.pardir, "include")
diff --git a/venv/lib/python3.11/site-packages/av/__main__.py b/venv/lib/python3.11/site-packages/av/__main__.py
new file mode 100644
index 0000000000000000000000000000000000000000..bc353d1472c578ef5901705afcaecb28efd901f0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/__main__.py
@@ -0,0 +1,40 @@
+from __future__ import annotations
+
+import argparse
+
+
+def main() -> None:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--codecs", action="store_true")
+ parser.add_argument("--version", action="store_true")
+ args = parser.parse_args()
+
+ if args.version:
+ import av
+ import av._core
+
+ print(f"PyAV v{av.__version__}")
+
+ by_config: dict = {}
+ for libname, config in sorted(av._core.library_meta.items()):
+ version = config["version"]
+ if version[0] >= 0:
+ by_config.setdefault(
+ (config["configuration"], config["license"]), []
+ ).append((libname, config))
+
+ for (config, license), libs in sorted(by_config.items()):
+ print("library configuration:", config)
+ print("library license:", license)
+ for libname, config in libs:
+ version = config["version"]
+ print(f"{libname:<13} {version[0]:3d}.{version[1]:3d}.{version[2]:3d}")
+
+ if args.codecs:
+ from av.codec.codec import dump_codecs
+
+ dump_codecs()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/venv/lib/python3.11/site-packages/av/_core.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/_core.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..ac913908029bc851bf05296c58a4790513f1d6fd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/_core.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:37db59aacfd95bf293c4ebd069824b490f221029bba41c9b28a14fa67285bd46
+size 183497
diff --git a/venv/lib/python3.11/site-packages/av/_core.pyi b/venv/lib/python3.11/site-packages/av/_core.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..0ee0a562677bdb1c4a98c52722af95e0a4f08d00
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/_core.pyi
@@ -0,0 +1,12 @@
+from typing import TypedDict
+
+class _Meta(TypedDict):
+ version: tuple[int, int, int]
+ configuration: str
+ license: str
+
+library_meta: dict[str, _Meta]
+library_versions: dict[str, tuple[int, int, int]]
+ffmpeg_version_info: str
+
+time_base: int
diff --git a/venv/lib/python3.11/site-packages/av/about.py b/venv/lib/python3.11/site-packages/av/about.py
new file mode 100644
index 0000000000000000000000000000000000000000..4fcf9b8bb23d8e8c934236789c0b4a0758662409
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/about.py
@@ -0,0 +1 @@
+__version__ = "14.0.1"
diff --git a/venv/lib/python3.11/site-packages/av/attachments/__init__.py b/venv/lib/python3.11/site-packages/av/attachments/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/attachments/stream.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/attachments/stream.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..c2f0f8b04f9289aa813514ee4950f0c3b1414f8b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/attachments/stream.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:2e2a4dd06c09fd29023962b31d94cfc47f0d1c409d7bed597fdfac51c989d7da
+size 355185
diff --git a/venv/lib/python3.11/site-packages/av/attachments/stream.pxd b/venv/lib/python3.11/site-packages/av/attachments/stream.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..81f788b7741fc3d148f2d55e3d33aead5bcc9f62
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/attachments/stream.pxd
@@ -0,0 +1,5 @@
+from av.stream cimport Stream
+
+
+cdef class AttachmentStream(Stream):
+ pass
diff --git a/venv/lib/python3.11/site-packages/av/attachments/stream.pyi b/venv/lib/python3.11/site-packages/av/attachments/stream.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..3d660e4a0eb6afbe9be9d912082893d9c8973b16
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/attachments/stream.pyi
@@ -0,0 +1,8 @@
+from typing import Literal
+
+from av.stream import Stream
+
+class AttachmentStream(Stream):
+ type: Literal["attachment"]
+ @property
+ def mimetype(self) -> str | None: ...
diff --git a/venv/lib/python3.11/site-packages/av/audio/__init__.pxd b/venv/lib/python3.11/site-packages/av/audio/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/audio/__init__.py b/venv/lib/python3.11/site-packages/av/audio/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..74ddf696497af9293024c39df458c9b03dd0cf0a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/__init__.py
@@ -0,0 +1,2 @@
+from .frame import AudioFrame
+from .stream import AudioStream
diff --git a/venv/lib/python3.11/site-packages/av/audio/__init__.pyi b/venv/lib/python3.11/site-packages/av/audio/__init__.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..73f2eebddf52ed45cdb7cd819e75b3ea7d735f4a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/__init__.pyi
@@ -0,0 +1,4 @@
+from .frame import AudioFrame
+from .stream import AudioStream
+
+__all__ = ("AudioFrame", "AudioStream")
diff --git a/venv/lib/python3.11/site-packages/av/audio/codeccontext.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/codeccontext.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..ca44da64ac7610979235259932821c8b9b7021d0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/codeccontext.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:d89a1f4f4ee75bc51b134eb121819f80efcc4dd83e8368df6245942173cbd914
+size 474297
diff --git a/venv/lib/python3.11/site-packages/av/audio/codeccontext.pxd b/venv/lib/python3.11/site-packages/av/audio/codeccontext.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..55ad15e9f0404d2b77810682774fd25ab434fadb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/codeccontext.pxd
@@ -0,0 +1,11 @@
+
+from av.audio.frame cimport AudioFrame
+from av.audio.resampler cimport AudioResampler
+from av.codec.context cimport CodecContext
+
+
+cdef class AudioCodecContext(CodecContext):
+ # Hold onto the frames that we will decode until we have a full one.
+ cdef AudioFrame next_frame
+ # For encoding.
+ cdef AudioResampler resampler
diff --git a/venv/lib/python3.11/site-packages/av/audio/codeccontext.pyi b/venv/lib/python3.11/site-packages/av/audio/codeccontext.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..b3ec3ce6e877a88a9fcf7661c9aded27201fc8fb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/codeccontext.pyi
@@ -0,0 +1,29 @@
+from typing import Iterator, Literal
+
+from av.codec.context import CodecContext
+from av.packet import Packet
+
+from .format import AudioFormat
+from .frame import AudioFrame
+from .layout import AudioLayout
+
+class _Format:
+ def __get__(self, i: object | None, owner: type | None = None) -> AudioFormat: ...
+ def __set__(self, instance: object, value: AudioFormat | str) -> None: ...
+
+class _Layout:
+ def __get__(self, i: object | None, owner: type | None = None) -> AudioLayout: ...
+ def __set__(self, instance: object, value: AudioLayout | str) -> None: ...
+
+class AudioCodecContext(CodecContext):
+ frame_size: int
+ sample_rate: int
+ rate: int
+ type: Literal["audio"]
+ format: _Format
+ layout: _Layout
+ @property
+ def channels(self) -> int: ...
+ def encode(self, frame: AudioFrame | None = None) -> list[Packet]: ...
+ def encode_lazy(self, frame: AudioFrame | None = None) -> Iterator[Packet]: ...
+ def decode(self, packet: Packet | None = None) -> list[AudioFrame]: ...
diff --git a/venv/lib/python3.11/site-packages/av/audio/fifo.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/fifo.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..f7500e7a277cdfe28ebd36054deabd94100aa6e1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/fifo.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:c2b2ee835c2b3577d6fe244fffc104c7bafddf339027bf0fee23cbb644b07e1a
+size 720329
diff --git a/venv/lib/python3.11/site-packages/av/audio/fifo.pxd b/venv/lib/python3.11/site-packages/av/audio/fifo.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..0ace5e4b1ec84daaeb09f27f6f9689a56d026912
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/fifo.pxd
@@ -0,0 +1,19 @@
+cimport libav as lib
+from libc.stdint cimport int64_t, uint64_t
+
+from av.audio.frame cimport AudioFrame
+
+
+cdef class AudioFifo:
+
+ cdef lib.AVAudioFifo *ptr
+
+ cdef AudioFrame template
+
+ cdef readonly uint64_t samples_written
+ cdef readonly uint64_t samples_read
+ cdef readonly double pts_per_sample
+
+ cpdef write(self, AudioFrame frame)
+ cpdef read(self, int samples=*, bint partial=*)
+ cpdef read_many(self, int samples, bint partial=*)
diff --git a/venv/lib/python3.11/site-packages/av/audio/fifo.pyi b/venv/lib/python3.11/site-packages/av/audio/fifo.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..085ed4bba1039b77675da29e26ef9ccbd539c8e1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/fifo.pyi
@@ -0,0 +1,22 @@
+from .format import AudioFormat
+from .frame import AudioFrame
+from .layout import AudioLayout
+
+class AudioFifo:
+ def write(self, frame: AudioFrame) -> None: ...
+ def read(self, samples: int = 0, partial: bool = False) -> AudioFrame | None: ...
+ def read_many(self, samples: int, partial: bool = False) -> list[AudioFrame]: ...
+ @property
+ def format(self) -> AudioFormat: ...
+ @property
+ def layout(self) -> AudioLayout: ...
+ @property
+ def sample_rate(self) -> int: ...
+ @property
+ def samples(self) -> int: ...
+ @property
+ def samples_written(self) -> int: ...
+ @property
+ def samples_read(self) -> int: ...
+ @property
+ def pts_per_sample(self) -> float: ...
diff --git a/venv/lib/python3.11/site-packages/av/audio/format.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/format.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..cc6357a3a537e02bb1da9370d50d5750c77766b0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/format.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:2160d5a2682320544e2a0987d626fe3e59c7d59bcd57e81c7b0885320855fb6b
+size 392577
diff --git a/venv/lib/python3.11/site-packages/av/audio/format.pxd b/venv/lib/python3.11/site-packages/av/audio/format.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..4160aa85b000714890d8a1d341becc0689d4091a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/format.pxd
@@ -0,0 +1,11 @@
+cimport libav as lib
+
+
+cdef class AudioFormat:
+
+ cdef lib.AVSampleFormat sample_fmt
+
+ cdef _init(self, lib.AVSampleFormat sample_fmt)
+
+
+cdef AudioFormat get_audio_format(lib.AVSampleFormat format)
diff --git a/venv/lib/python3.11/site-packages/av/audio/format.pyi b/venv/lib/python3.11/site-packages/av/audio/format.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..5f7e322edc85d831b5b6cc9979fd4dfca4b994c1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/format.pyi
@@ -0,0 +1,11 @@
+class AudioFormat:
+ name: str
+ bytes: int
+ bits: int
+ is_planar: bool
+ is_packed: bool
+ planar: AudioFormat
+ packed: AudioFormat
+ container_name: str
+
+ def __init__(self, name: str | AudioFormat) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/audio/frame.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/frame.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..beaba4a032de4be97cae8dd2b659bb2ba74257bb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/frame.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:c9948c61389365d3fd10ac840a9c10fb407069b600355e2c90b169d43a21ccaf
+size 978673
diff --git a/venv/lib/python3.11/site-packages/av/audio/frame.pxd b/venv/lib/python3.11/site-packages/av/audio/frame.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..398d76d33ea1173e31a5e13809ba3910385572b3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/frame.pxd
@@ -0,0 +1,31 @@
+cimport libav as lib
+from libc.stdint cimport uint8_t, uint64_t
+
+from av.audio.format cimport AudioFormat
+from av.audio.layout cimport AudioLayout
+from av.frame cimport Frame
+
+
+cdef class AudioFrame(Frame):
+ # For raw storage of the frame's data; don't ever touch this.
+ cdef uint8_t *_buffer
+ cdef size_t _buffer_size
+
+ cdef readonly AudioLayout layout
+ """
+ The audio channel layout.
+
+ :type: AudioLayout
+ """
+
+ cdef readonly AudioFormat format
+ """
+ The audio sample format.
+
+ :type: AudioFormat
+ """
+
+ cdef _init(self, lib.AVSampleFormat format, lib.AVChannelLayout layout, unsigned int nb_samples, unsigned int align)
+ cdef _init_user_attributes(self)
+
+cdef AudioFrame alloc_audio_frame()
diff --git a/venv/lib/python3.11/site-packages/av/audio/frame.pyi b/venv/lib/python3.11/site-packages/av/audio/frame.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..7f61e4e6dcccb5ae54aa27269c6adddf172a52a8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/frame.pyi
@@ -0,0 +1,47 @@
+from typing import Any, Union
+
+import numpy as np
+
+from av.frame import Frame
+
+from .format import AudioFormat
+from .layout import AudioLayout
+from .plane import AudioPlane
+
+format_dtypes: dict[str, str]
+_SupportedNDarray = Union[
+ np.ndarray[Any, np.dtype[np.float64]], # f8
+ np.ndarray[Any, np.dtype[np.float32]], # f4
+ np.ndarray[Any, np.dtype[np.int32]], # i4
+ np.ndarray[Any, np.dtype[np.int16]], # i2
+ np.ndarray[Any, np.dtype[np.uint8]], # u1
+]
+
+class _Format:
+ def __get__(self, i: object | None, owner: type | None = None) -> AudioFormat: ...
+ def __set__(self, instance: object, value: AudioFormat | str) -> None: ...
+
+class _Layout:
+ def __get__(self, i: object | None, owner: type | None = None) -> AudioLayout: ...
+ def __set__(self, instance: object, value: AudioLayout | str) -> None: ...
+
+class AudioFrame(Frame):
+ planes: tuple[AudioPlane, ...]
+ samples: int
+ sample_rate: int
+ rate: int
+ format: _Format
+ layout: _Layout
+
+ def __init__(
+ self,
+ format: str = "s16",
+ layout: str = "stereo",
+ samples: int = 0,
+ align: int = 1,
+ ) -> None: ...
+ @staticmethod
+ def from_ndarray(
+ array: _SupportedNDarray, format: str = "s16", layout: str = "stereo"
+ ) -> AudioFrame: ...
+ def to_ndarray(self) -> _SupportedNDarray: ...
diff --git a/venv/lib/python3.11/site-packages/av/audio/layout.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/layout.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..dcd95d9f7a9f2a00d7fe9434e52674f129a857b3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/layout.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:03ae13c2605d2059c1bc8f2a7e93774b6c56b312d0efd154057caaddbe87e5bb
+size 474721
diff --git a/venv/lib/python3.11/site-packages/av/audio/layout.pxd b/venv/lib/python3.11/site-packages/av/audio/layout.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..c7a2368f15840dd241940868eda157cb6fb1527e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/layout.pxd
@@ -0,0 +1,8 @@
+cimport libav as lib
+
+
+cdef class AudioLayout:
+ cdef lib.AVChannelLayout layout
+ cdef _init(self, lib.AVChannelLayout layout)
+
+cdef AudioLayout get_audio_layout(lib.AVChannelLayout c_layout)
diff --git a/venv/lib/python3.11/site-packages/av/audio/layout.pyi b/venv/lib/python3.11/site-packages/av/audio/layout.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..9fdf0ac1531ff4d74ac60ee191ab899755443cf1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/layout.pyi
@@ -0,0 +1,12 @@
+from dataclasses import dataclass
+
+class AudioLayout:
+ name: str
+ nb_channels: int
+ channels: tuple[AudioChannel, ...]
+ def __init__(self, layout: str | AudioLayout): ...
+
+@dataclass
+class AudioChannel:
+ name: str
+ description: str
diff --git a/venv/lib/python3.11/site-packages/av/audio/plane.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/plane.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..925e146dc6f3ae4a00797222996beb00ac109883
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/plane.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:cb8aac77d644cb7e2a2d17d93e78d5628cf34bdf5d1e43d9dbb2b1a63a583d25
+size 388025
diff --git a/venv/lib/python3.11/site-packages/av/audio/plane.pxd b/venv/lib/python3.11/site-packages/av/audio/plane.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..316c84031c3da5b35135132227133b9ee5694bab
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/plane.pxd
@@ -0,0 +1,8 @@
+from av.plane cimport Plane
+
+
+cdef class AudioPlane(Plane):
+
+ cdef readonly size_t buffer_size
+
+ cdef size_t _buffer_size(self)
diff --git a/venv/lib/python3.11/site-packages/av/audio/plane.pyi b/venv/lib/python3.11/site-packages/av/audio/plane.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..64524dcdb5170ebfbee5aaafff86f37f9a476444
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/plane.pyi
@@ -0,0 +1,4 @@
+from av.plane import Plane
+
+class AudioPlane(Plane):
+ buffer_size: int
diff --git a/venv/lib/python3.11/site-packages/av/audio/resampler.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/resampler.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..642ba2da629a682be937cb102fa63f771d6090cf
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/resampler.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:f8ff1b3ebf9d1072bb44161a750206f1db0c6786a5dab15c52dd8b8f72a1ced7
+size 777601
diff --git a/venv/lib/python3.11/site-packages/av/audio/resampler.pxd b/venv/lib/python3.11/site-packages/av/audio/resampler.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..d3601403d96bcc8774e993979109b3ff94495889
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/resampler.pxd
@@ -0,0 +1,21 @@
+from av.audio.format cimport AudioFormat
+from av.audio.frame cimport AudioFrame
+from av.audio.layout cimport AudioLayout
+from av.filter.graph cimport Graph
+
+
+cdef class AudioResampler:
+
+ cdef readonly bint is_passthrough
+
+ cdef AudioFrame template
+
+ # Destination descriptors
+ cdef readonly AudioFormat format
+ cdef readonly AudioLayout layout
+ cdef readonly int rate
+ cdef readonly unsigned int frame_size
+
+ cdef Graph graph
+
+ cpdef resample(self, AudioFrame)
diff --git a/venv/lib/python3.11/site-packages/av/audio/resampler.pyi b/venv/lib/python3.11/site-packages/av/audio/resampler.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..cbf2134aa13911ebf13c7e2e676fadb96230f530
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/resampler.pyi
@@ -0,0 +1,20 @@
+from av.filter.graph import Graph
+
+from .format import AudioFormat
+from .frame import AudioFrame
+from .layout import AudioLayout
+
+class AudioResampler:
+ rate: int
+ frame_size: int
+ format: AudioFormat
+ graph: Graph | None
+
+ def __init__(
+ self,
+ format: str | int | AudioFormat | None = None,
+ layout: str | int | AudioLayout | None = None,
+ rate: int | None = None,
+ frame_size: int | None = None,
+ ) -> None: ...
+ def resample(self, frame: AudioFrame | None) -> list[AudioFrame]: ...
diff --git a/venv/lib/python3.11/site-packages/av/audio/stream.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/audio/stream.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..57e7b80b78b4f3d535c529fc62aee2ada79877a9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/stream.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:825e90570e010e301fa9ba3a6b430a886c546ad97dd3f5e3c08e4bdd4d14fc2b
+size 498817
diff --git a/venv/lib/python3.11/site-packages/av/audio/stream.pxd b/venv/lib/python3.11/site-packages/av/audio/stream.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..8462061f8e0da6cd418527a73ac80d8ce59d7f0e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/stream.pxd
@@ -0,0 +1,9 @@
+from av.packet cimport Packet
+from av.stream cimport Stream
+
+from .frame cimport AudioFrame
+
+
+cdef class AudioStream(Stream):
+ cpdef encode(self, AudioFrame frame=?)
+ cpdef decode(self, Packet packet=?)
diff --git a/venv/lib/python3.11/site-packages/av/audio/stream.pyi b/venv/lib/python3.11/site-packages/av/audio/stream.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..f92fb52ba2a941f1f539b690bb3b26fc434a8ec3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/audio/stream.pyi
@@ -0,0 +1,32 @@
+from typing import Literal
+
+from av.packet import Packet
+from av.stream import Stream
+
+from .codeccontext import AudioCodecContext
+from .format import AudioFormat
+from .frame import AudioFrame
+from .layout import AudioLayout
+
+class _Format:
+ def __get__(self, i: object | None, owner: type | None = None) -> AudioFormat: ...
+ def __set__(self, instance: object, value: AudioFormat | str) -> None: ...
+
+class _Layout:
+ def __get__(self, i: object | None, owner: type | None = None) -> AudioLayout: ...
+ def __set__(self, instance: object, value: AudioLayout | str) -> None: ...
+
+class AudioStream(Stream):
+ codec_context: AudioCodecContext
+ def encode(self, frame: AudioFrame | None = None) -> list[Packet]: ...
+ def decode(self, packet: Packet | None = None) -> list[AudioFrame]: ...
+
+ # From codec context
+ frame_size: int
+ sample_rate: int
+ bit_rate: int
+ rate: int
+ channels: int
+ type: Literal["audio"]
+ format: _Format
+ layout: _Layout
diff --git a/venv/lib/python3.11/site-packages/av/bitstream.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/bitstream.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..9420fcede9785372fcd89a795ea85baa0015df28
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/bitstream.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:6160a60fec30f57ebde802c71d27eb4beb108affe194a5c45f453596f1a71827
+size 454017
diff --git a/venv/lib/python3.11/site-packages/av/bitstream.pxd b/venv/lib/python3.11/site-packages/av/bitstream.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..dbb89c9844551a1b17f839e03604b7821294f1a0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/bitstream.pxd
@@ -0,0 +1,11 @@
+cimport libav as lib
+
+from av.packet cimport Packet
+
+
+cdef class BitStreamFilterContext:
+
+ cdef lib.AVBSFContext *ptr
+
+ cpdef filter(self, Packet packet=?)
+ cpdef flush(self)
diff --git a/venv/lib/python3.11/site-packages/av/bitstream.pyi b/venv/lib/python3.11/site-packages/av/bitstream.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..477c65f2d2356b2f9fd160d24af5b10cfae39c7f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/bitstream.pyi
@@ -0,0 +1,14 @@
+from .packet import Packet
+from .stream import Stream
+
+class BitStreamFilterContext:
+ def __init__(
+ self,
+ filter_description: str | bytes,
+ in_stream: Stream | None = None,
+ out_stream: Stream | None = None,
+ ): ...
+ def filter(self, packet: Packet | None) -> list[Packet]: ...
+ def flush(self) -> None: ...
+
+bitstream_filters_available: set[str]
diff --git a/venv/lib/python3.11/site-packages/av/buffer.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/buffer.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..b17bf2ef2bc62a19876982c204788e8e20f9031c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/buffer.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:f624bed4b20ba3baa29603842bb16835b160e51a0443d2993c00e78e3a513165
+size 523633
diff --git a/venv/lib/python3.11/site-packages/av/buffer.pxd b/venv/lib/python3.11/site-packages/av/buffer.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..cfab07ca065b699eb1ce4f9f4b4b7a068f04d766
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/buffer.pxd
@@ -0,0 +1,6 @@
+
+cdef class Buffer:
+
+ cdef size_t _buffer_size(self)
+ cdef void* _buffer_ptr(self)
+ cdef bint _buffer_writable(self)
diff --git a/venv/lib/python3.11/site-packages/av/buffer.pyi b/venv/lib/python3.11/site-packages/av/buffer.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..bc1090d1deac54dcfc364fea49ee31bca526c6f8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/buffer.pyi
@@ -0,0 +1,9 @@
+# When Python 3.12 becomes our lowest supported version, we could make this
+# class inherit `collections.abc.Buffer`.
+
+class Buffer:
+ buffer_size: int
+ buffer_ptr: int
+ def update(self, input: bytes) -> None: ...
+ def __buffer__(self, flags: int) -> memoryview: ...
+ def __bytes__(self) -> bytes: ...
diff --git a/venv/lib/python3.11/site-packages/av/bytesource.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/bytesource.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..63488fcb6df904ce28edc8aa06d127e32720185a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/bytesource.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:49dfdac854ec4ca1f38363efc304695d5ff42689913b62f3e90e1fa9f3a71937
+size 318561
diff --git a/venv/lib/python3.11/site-packages/av/bytesource.pxd b/venv/lib/python3.11/site-packages/av/bytesource.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..050baab352f74227f7cc253a80b0b125e45894ce
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/bytesource.pxd
@@ -0,0 +1,14 @@
+from cpython.buffer cimport Py_buffer
+
+
+cdef class ByteSource:
+
+ cdef object owner
+
+ cdef bint has_view
+ cdef Py_buffer view
+
+ cdef unsigned char *ptr
+ cdef size_t length
+
+cdef ByteSource bytesource(object, bint allow_none=*)
diff --git a/venv/lib/python3.11/site-packages/av/codec/__init__.pxd b/venv/lib/python3.11/site-packages/av/codec/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/codec/__init__.py b/venv/lib/python3.11/site-packages/av/codec/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..f35f9b7d4e5b36ea62d1d6268aea5f8dc0007b5a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/codec/__init__.py
@@ -0,0 +1,11 @@
+from .codec import Capabilities, Codec, Properties, codec_descriptor, codecs_available
+from .context import CodecContext
+
+__all__ = (
+ "Capabilities",
+ "Codec",
+ "Properties",
+ "codec_descriptor",
+ "codecs_available",
+ "CodecContext",
+)
diff --git a/venv/lib/python3.11/site-packages/av/codec/codec.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/codec/codec.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..3279e632915c44753db2cc09dc45c2280c4fd7a5
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/codec/codec.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:939f8c401b3f81547f3e281c0351f1a8001063e94d6d97880f90dd6c2a5efd58
+size 856025
diff --git a/venv/lib/python3.11/site-packages/av/codec/codec.pxd b/venv/lib/python3.11/site-packages/av/codec/codec.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..b9925df132a6a3d416016a57d973eb46f8e4169c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/codec/codec.pxd
@@ -0,0 +1,13 @@
+cimport libav as lib
+
+
+cdef class Codec:
+
+ cdef const lib.AVCodec *ptr
+ cdef const lib.AVCodecDescriptor *desc
+ cdef readonly bint is_encoder
+
+ cdef _init(self, name=?)
+
+
+cdef Codec wrap_codec(const lib.AVCodec *ptr)
diff --git a/venv/lib/python3.11/site-packages/av/codec/codec.pyi b/venv/lib/python3.11/site-packages/av/codec/codec.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..32736c080a65b4a079b96471bc42454ed52b018b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/codec/codec.pyi
@@ -0,0 +1,110 @@
+from enum import Flag, IntEnum
+from fractions import Fraction
+from typing import ClassVar, Literal, overload
+
+from av.audio.codeccontext import AudioCodecContext
+from av.audio.format import AudioFormat
+from av.descriptor import Descriptor
+from av.subtitles.codeccontext import SubtitleCodecContext
+from av.video.codeccontext import VideoCodecContext
+from av.video.format import VideoFormat
+
+from .context import CodecContext
+
+class Properties(Flag):
+ NONE: ClassVar[Properties]
+ INTRA_ONLY: ClassVar[Properties]
+ LOSSY: ClassVar[Properties]
+ LOSSLESS: ClassVar[Properties]
+ REORDER: ClassVar[Properties]
+ BITMAP_SUB: ClassVar[Properties]
+ TEXT_SUB: ClassVar[Properties]
+
+class Capabilities(IntEnum):
+ none: int
+ draw_horiz_band: int
+ dr1: int
+ hwaccel: int
+ delay: int
+ small_last_frame: int
+ hwaccel_vdpau: int
+ subframes: int
+ experimental: int
+ channel_conf: int
+ neg_linesizes: int
+ frame_threads: int
+ slice_threads: int
+ param_change: int
+ auto_threads: int
+ variable_frame_size: int
+ avoid_probing: int
+ hardware: int
+ hybrid: int
+ encoder_reordered_opaque: int
+ encoder_flush: int
+ encoder_recon_frame: int
+
+class UnknownCodecError(ValueError): ...
+
+class Codec:
+ @property
+ def is_encoder(self) -> bool: ...
+ @property
+ def is_decoder(self) -> bool: ...
+ descriptor: Descriptor
+ @property
+ def name(self) -> str: ...
+ @property
+ def long_name(self) -> str: ...
+ @property
+ def type(self) -> Literal["video", "audio", "data", "subtitle", "attachment"]: ...
+ @property
+ def id(self) -> int: ...
+ frame_rates: list[Fraction] | None
+ audio_rates: list[int] | None
+ video_formats: list[VideoFormat] | None
+ audio_formats: list[AudioFormat] | None
+
+ @property
+ def properties(self) -> int: ...
+ @property
+ def intra_only(self) -> bool: ...
+ @property
+ def lossy(self) -> bool: ...
+ @property
+ def lossless(self) -> bool: ...
+ @property
+ def reorder(self) -> bool: ...
+ @property
+ def bitmap_sub(self) -> bool: ...
+ @property
+ def text_sub(self) -> bool: ...
+ @property
+ def capabilities(self) -> int: ...
+ @property
+ def experimental(self) -> bool: ...
+ @property
+ def delay(self) -> bool: ...
+ def __init__(self, name: str, mode: Literal["r", "w"] = "r") -> None: ...
+ @overload
+ def create(self, kind: Literal["video"]) -> VideoCodecContext: ...
+ @overload
+ def create(self, kind: Literal["audio"]) -> AudioCodecContext: ...
+ @overload
+ def create(self, kind: Literal["subtitle"]) -> SubtitleCodecContext: ...
+ @overload
+ def create(self, kind: None = None) -> CodecContext: ...
+ @overload
+ def create(
+ self, kind: Literal["video", "audio", "subtitle"] | None = None
+ ) -> (
+ VideoCodecContext | AudioCodecContext | SubtitleCodecContext | CodecContext
+ ): ...
+
+class codec_descriptor:
+ name: str
+ options: tuple[int, ...]
+
+codecs_available: set[str]
+
+def dump_codecs() -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/codec/context.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/codec/context.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..ff677b91c58b01a2af97b4358f7ca427209cc55f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/codec/context.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:235851a2c93d5e1cda6ab1b0090a8f287e14dd5bbeb43b881cb99f270140aab3
+size 1372465
diff --git a/venv/lib/python3.11/site-packages/av/codec/context.pxd b/venv/lib/python3.11/site-packages/av/codec/context.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..42b2d63e7e203b35877a3c59bc1d15cfb9dc8535
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/codec/context.pxd
@@ -0,0 +1,58 @@
+cimport libav as lib
+from libc.stdint cimport int64_t
+
+from av.bytesource cimport ByteSource
+from av.codec.codec cimport Codec
+from av.frame cimport Frame
+from av.packet cimport Packet
+
+
+cdef class CodecContext:
+ cdef lib.AVCodecContext *ptr
+
+ # Whether AVCodecContext.extradata should be de-allocated upon destruction.
+ cdef bint extradata_set
+
+ # Used as a signal that this is within a stream, and also for us to access that
+ # stream. This is set "manually" by the stream after constructing this object.
+ cdef int stream_index
+
+ cdef lib.AVCodecParserContext *parser
+ cdef _init(self, lib.AVCodecContext *ptr, const lib.AVCodec *codec)
+
+ # Public API.
+ cdef readonly bint is_open
+ cdef readonly Codec codec
+ cdef public dict options
+ cpdef open(self, bint strict=?)
+
+ # Wraps both versions of the transcode API, returning lists.
+ cpdef encode(self, Frame frame=?)
+ cpdef decode(self, Packet packet=?)
+ cpdef flush_buffers(self)
+
+ # Used by both transcode APIs to setup user-land objects.
+ # TODO: Remove the `Packet` from `_setup_decoded_frame` (because flushing packets
+ # are bogus). It should take all info it needs from the context and/or stream.
+ cdef _prepare_and_time_rebase_frames_for_encode(self, Frame frame)
+ cdef _prepare_frames_for_encode(self, Frame frame)
+ cdef _setup_encoded_packet(self, Packet)
+ cdef _setup_decoded_frame(self, Frame, Packet)
+
+ # Implemented by base for the generic send/recv API.
+ # Note that the user cannot send without receiving. This is because
+ # `_prepare_frames_for_encode` may expand a frame into multiple (e.g. when
+ # resampling audio to a higher rate but with fixed size frames), and the
+ # send/recv buffer may be limited to a single frame. Ergo, we need to flush
+ # the buffer as often as possible.
+ cdef _recv_packet(self)
+ cdef _send_packet_and_recv(self, Packet packet)
+ cdef _recv_frame(self)
+
+ # Implemented by children for the generic send/recv API, so we have the
+ # correct subclass of Frame.
+ cdef Frame _next_frame
+ cdef Frame _alloc_next_frame(self)
+
+
+cdef CodecContext wrap_codec_context(lib.AVCodecContext*, const lib.AVCodec*)
diff --git a/venv/lib/python3.11/site-packages/av/codec/context.pyi b/venv/lib/python3.11/site-packages/av/codec/context.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..a6ca9647ebc6ed658aaf2a9421a0992cbac8ed4e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/codec/context.pyi
@@ -0,0 +1,94 @@
+from enum import Flag, IntEnum
+from fractions import Fraction
+from typing import ClassVar, Literal
+
+from av.packet import Packet
+
+from .codec import Codec
+
+class ThreadType(Flag):
+ NONE: ClassVar[ThreadType]
+ FRAME: ClassVar[ThreadType]
+ SLICE: ClassVar[ThreadType]
+ AUTO: ClassVar[ThreadType]
+ def __get__(self, i: object | None, owner: type | None = None) -> ThreadType: ...
+ def __set__(self, instance: object, value: int | str | ThreadType) -> None: ...
+
+class Flags(IntEnum):
+ unaligned: int
+ qscale: int
+ four_mv: int
+ output_corrupt: int
+ qpel: int
+ drop_changed: int
+ recon_frame: int
+ copy_opaque: int
+ frame_duration: int
+ pass1: int
+ pass2: int
+ loop_filter: int
+ gray: int
+ psnr: int
+ interlaced_dct: int
+ low_delay: int
+ global_header: int
+ bitexact: int
+ ac_pred: int
+ interlaced_me: int
+ closed_gop: int
+
+class Flags2(IntEnum):
+ fast: int
+ no_output: int
+ local_header: int
+ chunks: int
+ ignore_crop: int
+ show_all: int
+ export_mvs: int
+ skip_manual: int
+ ro_flush_noop: int
+
+class CodecContext:
+ name: str
+ type: Literal["video", "audio", "data", "subtitle", "attachment"]
+ options: dict[str, str]
+ profile: str | None
+ @property
+ def profiles(self) -> list[str]: ...
+ extradata: bytes | None
+ time_base: Fraction
+ codec_tag: str
+ bit_rate: int | None
+ bit_rate_tolerance: int
+ thread_count: int
+ thread_type: ThreadType
+ skip_frame: Literal[
+ "NONE", "DEFAULT", "NONREF", "BIDIR", "NONINTRA", "NONKEY", "ALL"
+ ]
+ flags: int
+ qscale: bool
+ copy_opaque: bool
+ flags2: int
+ @property
+ def is_open(self) -> bool: ...
+ @property
+ def is_encoder(self) -> bool: ...
+ @property
+ def is_decoder(self) -> bool: ...
+ @property
+ def codec(self) -> Codec: ...
+ @property
+ def max_bit_rate(self) -> int | None: ...
+ @property
+ def delay(self) -> bool: ...
+ @property
+ def extradata_size(self) -> int: ...
+ def open(self, strict: bool = True) -> None: ...
+ @staticmethod
+ def create(
+ codec: str | Codec, mode: Literal["r", "w"] | None = None
+ ) -> CodecContext: ...
+ def parse(
+ self, raw_input: bytes | bytearray | memoryview | None = None
+ ) -> list[Packet]: ...
+ def flush_buffers(self) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/container/__init__.pxd b/venv/lib/python3.11/site-packages/av/container/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/container/__init__.py b/venv/lib/python3.11/site-packages/av/container/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..98d49dd4e89c66b2ab3a92cd675ca86b4b719c90
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/__init__.py
@@ -0,0 +1,3 @@
+from .core import Container, Flags, open
+from .input import InputContainer
+from .output import OutputContainer
diff --git a/venv/lib/python3.11/site-packages/av/container/__init__.pyi b/venv/lib/python3.11/site-packages/av/container/__init__.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..7160777cc06b6ebf854ad7988da98d9740b7a36b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/__init__.pyi
@@ -0,0 +1,3 @@
+from .core import *
+from .input import *
+from .output import *
diff --git a/venv/lib/python3.11/site-packages/av/container/core.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/container/core.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..a32133df13bdb76b52501e1f9f5704b877a5e9f3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/core.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:846072951020c7992c877ffb5967077db1dc9cf1bd103b1e562b598628ccf24e
+size 1188113
diff --git a/venv/lib/python3.11/site-packages/av/container/core.pxd b/venv/lib/python3.11/site-packages/av/container/core.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..1aed54b906d91269a320f25482f6d353c5b825c9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/core.pxd
@@ -0,0 +1,48 @@
+cimport libav as lib
+
+from av.container.pyio cimport PyIOFile
+from av.container.streams cimport StreamContainer
+from av.dictionary cimport _Dictionary
+from av.format cimport ContainerFormat
+from av.stream cimport Stream
+
+# Interrupt callback information, times are in seconds.
+ctypedef struct timeout_info:
+ double start_time
+ double timeout
+
+
+cdef class Container:
+
+ cdef readonly bint writeable
+ cdef lib.AVFormatContext *ptr
+
+ cdef readonly object name
+ cdef readonly str metadata_encoding
+ cdef readonly str metadata_errors
+
+ cdef readonly PyIOFile file
+ cdef int buffer_size
+ cdef bint input_was_opened
+ cdef readonly object io_open
+ cdef readonly object open_files
+
+ cdef readonly ContainerFormat format
+
+ cdef readonly dict options
+ cdef readonly dict container_options
+ cdef readonly list stream_options
+
+ cdef readonly StreamContainer streams
+ cdef readonly dict metadata
+
+ # Private API.
+ cdef _assert_open(self)
+ cdef int err_check(self, int value) except -1
+
+ # Timeouts
+ cdef readonly object open_timeout
+ cdef readonly object read_timeout
+ cdef timeout_info interrupt_callback_info
+ cdef set_timeout(self, object)
+ cdef start_timeout(self)
diff --git a/venv/lib/python3.11/site-packages/av/container/core.pyi b/venv/lib/python3.11/site-packages/av/container/core.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..227a7d32a8af57308c5af9aee57bf3140549f818
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/core.pyi
@@ -0,0 +1,118 @@
+from enum import Flag
+from fractions import Fraction
+from pathlib import Path
+from types import TracebackType
+from typing import Any, Callable, ClassVar, Literal, Type, overload
+
+from av.format import ContainerFormat
+
+from .input import InputContainer
+from .output import OutputContainer
+from .streams import StreamContainer
+
+Real = int | float | Fraction
+
+class Flags(Flag):
+ gen_pts: ClassVar[Flags]
+ ign_idx: ClassVar[Flags]
+ non_block: ClassVar[Flags]
+ ign_dts: ClassVar[Flags]
+ no_fillin: ClassVar[Flags]
+ no_parse: ClassVar[Flags]
+ no_buffer: ClassVar[Flags]
+ custom_io: ClassVar[Flags]
+ discard_corrupt: ClassVar[Flags]
+ flush_packets: ClassVar[Flags]
+ bitexact: ClassVar[Flags]
+ sort_dts: ClassVar[Flags]
+ fast_seek: ClassVar[Flags]
+ shortest: ClassVar[Flags]
+ auto_bsf: ClassVar[Flags]
+
+class Container:
+ writeable: bool
+ name: str
+ metadata_encoding: str
+ metadata_errors: str
+ file: Any
+ buffer_size: int
+ input_was_opened: bool
+ io_open: Any
+ open_files: Any
+ format: ContainerFormat
+ options: dict[str, str]
+ container_options: dict[str, str]
+ stream_options: list[dict[str, str]]
+ streams: StreamContainer
+ metadata: dict[str, str]
+ open_timeout: Real | None
+ read_timeout: Real | None
+ flags: int
+
+ def __enter__(self) -> Container: ...
+ def __exit__(
+ self,
+ exc_type: Type[BaseException] | None,
+ exc_val: BaseException | None,
+ exc_tb: TracebackType | None,
+ ) -> bool: ...
+ def err_check(self, value: int) -> int: ...
+ def set_timeout(self, timeout: Real | None) -> None: ...
+ def start_timeout(self) -> None: ...
+
+@overload
+def open(
+ file: Any,
+ mode: Literal["r"],
+ format: str | None = None,
+ options: dict[str, str] | None = None,
+ container_options: dict[str, str] | None = None,
+ stream_options: list[str] | None = None,
+ metadata_encoding: str = "utf-8",
+ metadata_errors: str = "strict",
+ buffer_size: int = 32768,
+ timeout: Real | None | tuple[Real | None, Real | None] = None,
+ io_open: Callable[..., Any] | None = None,
+) -> InputContainer: ...
+@overload
+def open(
+ file: str | Path,
+ mode: Literal["r"] | None = None,
+ format: str | None = None,
+ options: dict[str, str] | None = None,
+ container_options: dict[str, str] | None = None,
+ stream_options: list[str] | None = None,
+ metadata_encoding: str = "utf-8",
+ metadata_errors: str = "strict",
+ buffer_size: int = 32768,
+ timeout: Real | None | tuple[Real | None, Real | None] = None,
+ io_open: Callable[..., Any] | None = None,
+) -> InputContainer: ...
+@overload
+def open(
+ file: Any,
+ mode: Literal["w"],
+ format: str | None = None,
+ options: dict[str, str] | None = None,
+ container_options: dict[str, str] | None = None,
+ stream_options: list[str] | None = None,
+ metadata_encoding: str = "utf-8",
+ metadata_errors: str = "strict",
+ buffer_size: int = 32768,
+ timeout: Real | None | tuple[Real | None, Real | None] = None,
+ io_open: Callable[..., Any] | None = None,
+) -> OutputContainer: ...
+@overload
+def open(
+ file: Any,
+ mode: Literal["r", "w"] | None = None,
+ format: str | None = None,
+ options: dict[str, str] | None = None,
+ container_options: dict[str, str] | None = None,
+ stream_options: list[str] | None = None,
+ metadata_encoding: str = "utf-8",
+ metadata_errors: str = "strict",
+ buffer_size: int = 32768,
+ timeout: Real | None | tuple[Real | None, Real | None] = None,
+ io_open: Callable[..., Any] | None = None,
+) -> InputContainer | OutputContainer: ...
diff --git a/venv/lib/python3.11/site-packages/av/container/input.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/container/input.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..3ef083c4b4143f1f7e0a8ef06342499162660f8c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/input.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:7efeb747d81b496caf10ea750584686f60ad7a7885de525fbe57cd110d5f90cc
+size 917521
diff --git a/venv/lib/python3.11/site-packages/av/container/input.pxd b/venv/lib/python3.11/site-packages/av/container/input.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..8c369d8ad24b97981e6f6c34ceea21986ec94681
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/input.pxd
@@ -0,0 +1,9 @@
+cimport libav as lib
+
+from av.container.core cimport Container
+from av.stream cimport Stream
+
+
+cdef class InputContainer(Container):
+
+ cdef flush_buffers(self)
diff --git a/venv/lib/python3.11/site-packages/av/container/input.pyi b/venv/lib/python3.11/site-packages/av/container/input.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..90154c331746478d7cd8d22100c21164b37331b7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/input.pyi
@@ -0,0 +1,49 @@
+from typing import Any, Iterator, overload
+
+from av.audio.frame import AudioFrame
+from av.audio.stream import AudioStream
+from av.packet import Packet
+from av.stream import Stream
+from av.subtitles.stream import SubtitleStream
+from av.subtitles.subtitle import SubtitleSet
+from av.video.frame import VideoFrame
+from av.video.stream import VideoStream
+
+from .core import Container
+
+class InputContainer(Container):
+ start_time: int
+ duration: int | None
+ bit_rate: int
+ size: int
+
+ def __enter__(self) -> InputContainer: ...
+ def close(self) -> None: ...
+ def demux(self, *args: Any, **kwargs: Any) -> Iterator[Packet]: ...
+ @overload
+ def decode(self, video: int) -> Iterator[VideoFrame]: ...
+ @overload
+ def decode(self, audio: int) -> Iterator[AudioFrame]: ...
+ @overload
+ def decode(self, subtitles: int) -> Iterator[SubtitleSet]: ...
+ @overload
+ def decode(self, *args: VideoStream) -> Iterator[VideoFrame]: ...
+ @overload
+ def decode(self, *args: AudioStream) -> Iterator[AudioFrame]: ...
+ @overload
+ def decode(self, *args: SubtitleStream) -> Iterator[SubtitleSet]: ...
+ @overload
+ def decode(
+ self, *args: Any, **kwargs: Any
+ ) -> Iterator[VideoFrame | AudioFrame | SubtitleSet]: ...
+ def seek(
+ self,
+ offset: int,
+ *,
+ backward: bool = True,
+ any_frame: bool = False,
+ stream: Stream | VideoStream | AudioStream | None = None,
+ unsupported_frame_offset: bool = False,
+ unsupported_byte_offset: bool = False,
+ ) -> None: ...
+ def flush_buffers(self) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/container/output.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/container/output.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..c9c705e844c2413bd3c6d69cb77da45d013440c7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/output.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:961e21ac78c00265e1927a44025329163384b6b5a6d11cd777e8d79c8f3cb0c4
+size 1089689
diff --git a/venv/lib/python3.11/site-packages/av/container/output.pxd b/venv/lib/python3.11/site-packages/av/container/output.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..a4299891c4a680f3592cfab013c084914be25af4
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/output.pxd
@@ -0,0 +1,13 @@
+cimport libav as lib
+
+from av.container.core cimport Container
+from av.stream cimport Stream
+
+
+cdef class OutputContainer(Container):
+
+ cdef bint _started
+ cdef bint _done
+ cdef lib.AVPacket *packet_ptr
+
+ cpdef start_encoding(self)
diff --git a/venv/lib/python3.11/site-packages/av/container/output.pyi b/venv/lib/python3.11/site-packages/av/container/output.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..a7c89452f4d539c8f722363412125a38bde3e58c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/output.pyi
@@ -0,0 +1,56 @@
+from fractions import Fraction
+from typing import Literal, Sequence, TypeVar, Union, overload
+
+from av.audio.stream import AudioStream
+from av.data.stream import DataStream
+from av.packet import Packet
+from av.stream import Stream
+from av.subtitles.stream import SubtitleStream
+from av.video.stream import VideoStream
+
+from .core import Container
+
+_StreamT = TypeVar("_StreamT", bound=Union[VideoStream, AudioStream, SubtitleStream])
+
+class OutputContainer(Container):
+ def __enter__(self) -> OutputContainer: ...
+ @overload
+ def add_stream(
+ self,
+ codec_name: Literal["pcm_s16le", "aac", "mp3", "mp2"],
+ rate: int | None = None,
+ options: dict[str, str] | None = None,
+ **kwargs,
+ ) -> AudioStream: ...
+ @overload
+ def add_stream(
+ self,
+ codec_name: Literal["h264", "hevc", "mpeg4", "png", "gif", "qtrle"],
+ rate: Fraction | int | None = None,
+ options: dict[str, str] | None = None,
+ **kwargs,
+ ) -> VideoStream: ...
+ @overload
+ def add_stream(
+ self,
+ codec_name: str,
+ rate: Fraction | int | None = None,
+ options: dict[str, str] | None = None,
+ **kwargs,
+ ) -> VideoStream | AudioStream | SubtitleStream: ...
+ def add_stream_from_template(self, template: _StreamT, **kwargs) -> _StreamT: ...
+ def add_data_stream(
+ self, codec_name: str | None = None, options: dict[str, str] | None = None
+ ) -> DataStream: ...
+ def start_encoding(self) -> None: ...
+ def close(self) -> None: ...
+ def mux(self, packets: Packet | Sequence[Packet]) -> None: ...
+ def mux_one(self, packet: Packet) -> None: ...
+ @property
+ def default_video_codec(self) -> str: ...
+ @property
+ def default_audio_codec(self) -> str: ...
+ @property
+ def default_subtitle_codec(self) -> str: ...
+ @property
+ def supported_codecs(self) -> set[str]: ...
diff --git a/venv/lib/python3.11/site-packages/av/container/pyio.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/container/pyio.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..74e06bf46819417319d05c57efa7332d9e36412a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/pyio.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:62b15c01c1b72cb28f39bda91fe7e76f0a8311a6e7870232c4a75e422d03e138
+size 712129
diff --git a/venv/lib/python3.11/site-packages/av/container/pyio.pxd b/venv/lib/python3.11/site-packages/av/container/pyio.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..80edc8a6b7a95f6b3568fd24d93a0a4036f135b2
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/pyio.pxd
@@ -0,0 +1,24 @@
+cimport libav as lib
+from libc.stdint cimport int64_t, uint8_t
+
+
+cdef int pyio_read(void *opaque, uint8_t *buf, int buf_size) noexcept nogil
+cdef int pyio_write(void *opaque, const uint8_t *buf, int buf_size) noexcept nogil
+cdef int64_t pyio_seek(void *opaque, int64_t offset, int whence) noexcept nogil
+cdef int pyio_close_gil(lib.AVIOContext *pb)
+cdef int pyio_close_custom_gil(lib.AVIOContext *pb)
+
+cdef class PyIOFile:
+ # File-like source.
+ cdef readonly object file
+ cdef object fread
+ cdef object fwrite
+ cdef object fseek
+ cdef object ftell
+ cdef object fclose
+
+ # Custom IO for above.
+ cdef lib.AVIOContext *iocontext
+ cdef unsigned char *buffer
+ cdef long pos
+ cdef bint pos_is_valid
diff --git a/venv/lib/python3.11/site-packages/av/container/streams.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/container/streams.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..9a0c73437261008ee7f2660d2797beb300b75599
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/streams.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:870504ff0a5501fbbbd49fd8eb0d04f75c1254fe6dc001540e766c6744c2c6b2
+size 929537
diff --git a/venv/lib/python3.11/site-packages/av/container/streams.pxd b/venv/lib/python3.11/site-packages/av/container/streams.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..097176e10b3e2c0f9b48c6b852f3829b03977b2b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/streams.pxd
@@ -0,0 +1,21 @@
+cimport libav as lib
+
+from av.stream cimport Stream
+
+from .core cimport Container
+
+
+cdef class StreamContainer:
+ cdef list _streams
+
+ # For the different types.
+ cdef readonly tuple video
+ cdef readonly tuple audio
+ cdef readonly tuple subtitles
+ cdef readonly tuple attachments
+ cdef readonly tuple data
+ cdef readonly tuple other
+
+ cdef add_stream(self, Stream stream)
+ cdef int _get_best_stream_index(self, Container container, lib.AVMediaType type_enum, Stream related) noexcept
+
diff --git a/venv/lib/python3.11/site-packages/av/container/streams.pyi b/venv/lib/python3.11/site-packages/av/container/streams.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..fbaf1b67f2a6c3fb105246eee64522cd5080be40
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/container/streams.pyi
@@ -0,0 +1,37 @@
+from typing import Iterator, Literal, overload
+
+from av.attachments.stream import AttachmentStream
+from av.audio.stream import AudioStream
+from av.data.stream import DataStream
+from av.stream import Stream
+from av.subtitles.stream import SubtitleStream
+from av.video.stream import VideoStream
+
+class StreamContainer:
+ video: tuple[VideoStream, ...]
+ audio: tuple[AudioStream, ...]
+ subtitles: tuple[SubtitleStream, ...]
+ attachments: tuple[AttachmentStream, ...]
+ data: tuple[DataStream, ...]
+ other: tuple[Stream, ...]
+
+ def __init__(self) -> None: ...
+ def __len__(self) -> int: ...
+ def __iter__(self) -> Iterator[Stream]: ...
+ @overload
+ def __getitem__(self, index: int) -> Stream: ...
+ @overload
+ def __getitem__(self, index: slice) -> list[Stream]: ...
+ @overload
+ def __getitem__(self, index: int | slice) -> Stream | list[Stream]: ...
+ def get(
+ self,
+ *args: int | Stream | dict[str, int | tuple[int, ...]],
+ **kwargs: int | tuple[int, ...],
+ ) -> list[Stream]: ...
+ def best(
+ self,
+ type: Literal["video", "audio", "subtitle", "data", "attachment"],
+ /,
+ related: Stream | None = None,
+ ) -> Stream | None: ...
diff --git a/venv/lib/python3.11/site-packages/av/data/__init__.pxd b/venv/lib/python3.11/site-packages/av/data/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/data/__init__.py b/venv/lib/python3.11/site-packages/av/data/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/data/stream.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/data/stream.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..870ef383b93375df604fa24bda1c18391f0681a1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/data/stream.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:616c13770c5263650dce67a67e1f43d1714df5b0d54aea03792d342c9999803a
+size 388193
diff --git a/venv/lib/python3.11/site-packages/av/data/stream.pxd b/venv/lib/python3.11/site-packages/av/data/stream.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..012792a4aacb8e3c4a91dfb2114a92a6c0a380e8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/data/stream.pxd
@@ -0,0 +1,5 @@
+from av.stream cimport Stream
+
+
+cdef class DataStream(Stream):
+ pass
diff --git a/venv/lib/python3.11/site-packages/av/data/stream.pyi b/venv/lib/python3.11/site-packages/av/data/stream.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..45a669d4fdc703e078287fd9525b14d7d6be6ed2
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/data/stream.pyi
@@ -0,0 +1,6 @@
+from av.frame import Frame
+from av.packet import Packet
+from av.stream import Stream
+
+class DataStream(Stream):
+ name: str | None
diff --git a/venv/lib/python3.11/site-packages/av/datasets.py b/venv/lib/python3.11/site-packages/av/datasets.py
new file mode 100644
index 0000000000000000000000000000000000000000..5954a9c98370d8424c9becce2cd6010b03700688
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/datasets.py
@@ -0,0 +1,123 @@
+import errno
+import logging
+import os
+import sys
+from typing import Iterator
+from urllib.request import urlopen
+
+log = logging.getLogger(__name__)
+
+
+def iter_data_dirs(check_writable: bool = False) -> Iterator[str]:
+ try:
+ yield os.environ["PYAV_TESTDATA_DIR"]
+ except KeyError:
+ pass
+
+ if os.name == "nt":
+ yield os.path.join(sys.prefix, "pyav", "datasets")
+ return
+
+ bases = [
+ "/usr/local/share",
+ "/usr/local/lib",
+ "/usr/share",
+ "/usr/lib",
+ ]
+
+ # Prefer the local virtualenv.
+ if hasattr(sys, "real_prefix"):
+ bases.insert(0, sys.prefix)
+
+ for base in bases:
+ dir_ = os.path.join(base, "pyav", "datasets")
+ if check_writable:
+ if os.path.exists(dir_):
+ if not os.access(dir_, os.W_OK):
+ continue
+ else:
+ if not os.access(base, os.W_OK):
+ continue
+ yield dir_
+
+ yield os.path.join(os.path.expanduser("~"), ".pyav", "datasets")
+
+
+def cached_download(url: str, name: str) -> str:
+ """Download the data at a URL, and cache it under the given name.
+
+ The file is stored under `pyav/test` with the given name in the directory
+ :envvar:`PYAV_TESTDATA_DIR`, or the first that is writeable of:
+
+ - the current virtualenv
+ - ``/usr/local/share``
+ - ``/usr/local/lib``
+ - ``/usr/share``
+ - ``/usr/lib``
+ - the user's home
+
+ """
+
+ clean_name = os.path.normpath(name)
+ if clean_name != name:
+ raise ValueError(f"{name} is not normalized.")
+
+ for dir_ in iter_data_dirs():
+ path = os.path.join(dir_, name)
+ if os.path.exists(path):
+ return path
+
+ dir_ = next(iter_data_dirs(True))
+ path = os.path.join(dir_, name)
+
+ log.info(f"Downloading {url} to {path}")
+
+ response = urlopen(url)
+ if response.getcode() != 200:
+ raise ValueError(f"HTTP {response.getcode()}")
+
+ dir_ = os.path.dirname(path)
+ try:
+ os.makedirs(dir_)
+ except OSError as e:
+ if e.errno != errno.EEXIST:
+ raise
+
+ tmp_path = path + ".tmp"
+ with open(tmp_path, "wb") as fh:
+ while True:
+ chunk = response.read(8196)
+ if chunk:
+ fh.write(chunk)
+ else:
+ break
+
+ os.rename(tmp_path, path)
+
+ return path
+
+
+def fate(name: str) -> str:
+ """Download and return a path to a sample from the FFmpeg test suite.
+
+ Data is handled by :func:`cached_download`.
+
+ See the `FFmpeg Automated Test Environment `_
+
+ """
+ return cached_download(
+ "http://fate.ffmpeg.org/fate-suite/" + name,
+ os.path.join("fate-suite", name.replace("/", os.path.sep)),
+ )
+
+
+def curated(name: str) -> str:
+ """Download and return a path to a sample that is curated by the PyAV developers.
+
+ Data is handled by :func:`cached_download`.
+
+ """
+ return cached_download(
+ "https://pyav.org/datasets/" + name,
+ os.path.join("pyav-curated", name.replace("/", os.path.sep)),
+ )
diff --git a/venv/lib/python3.11/site-packages/av/descriptor.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/descriptor.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..630b8593b0c91e9e51cfba415acd548ce6649943
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/descriptor.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:a5135a89ebc981eb47dcccadfd8236b3831bade155b47762e7ef7130388921b7
+size 359561
diff --git a/venv/lib/python3.11/site-packages/av/descriptor.pxd b/venv/lib/python3.11/site-packages/av/descriptor.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..404f646afce60e3703814c4ae6c493dc8757b5c9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/descriptor.pxd
@@ -0,0 +1,20 @@
+cimport libav as lib
+
+
+cdef class Descriptor:
+
+ # These are present as:
+ # - AVCodecContext.av_class (same as avcodec_get_class())
+ # - AVFormatContext.av_class (same as avformat_get_class())
+ # - AVFilterContext.av_class (same as avfilter_get_class())
+ # - AVCodec.priv_class
+ # - AVOutputFormat.priv_class
+ # - AVInputFormat.priv_class
+ # - AVFilter.priv_class
+
+ cdef const lib.AVClass *ptr
+
+ cdef object _options # Option list cache.
+
+
+cdef Descriptor wrap_avclass(const lib.AVClass*)
diff --git a/venv/lib/python3.11/site-packages/av/descriptor.pyi b/venv/lib/python3.11/site-packages/av/descriptor.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..ae1998391263a4d332a3f8c01c87495d4ef5c1cd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/descriptor.pyi
@@ -0,0 +1,7 @@
+from typing import NoReturn
+
+from .option import Option
+
+class Descriptor:
+ name: str
+ options: tuple[Option, ...]
diff --git a/venv/lib/python3.11/site-packages/av/dictionary.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/dictionary.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..4305cef795b5e6e71d4746e64a44530d48c70689
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/dictionary.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:9ab2fb05cc615c6ea8cbc5f12ae3ad56204c5343fd1811a1138be91fb8691cef
+size 663409
diff --git a/venv/lib/python3.11/site-packages/av/dictionary.pxd b/venv/lib/python3.11/site-packages/av/dictionary.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..1c59df448a20d3f8929b50c7b386a587d8afb175
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/dictionary.pxd
@@ -0,0 +1,11 @@
+cimport libav as lib
+
+
+cdef class _Dictionary:
+
+ cdef lib.AVDictionary *ptr
+
+ cpdef _Dictionary copy(self)
+
+
+cdef _Dictionary wrap_dictionary(lib.AVDictionary *input_)
diff --git a/venv/lib/python3.11/site-packages/av/dictionary.pyi b/venv/lib/python3.11/site-packages/av/dictionary.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..a6868bea2c9e122963a58410707f9cbd865b49fe
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/dictionary.pyi
@@ -0,0 +1,10 @@
+from collections.abc import MutableMapping
+from typing import Iterator
+
+class Dictionary(MutableMapping[str, str]):
+ def __getitem__(self, key: str) -> str: ...
+ def __setitem__(self, key: str, value: str) -> None: ...
+ def __delitem__(self, key: str) -> None: ...
+ def __len__(self) -> int: ...
+ def __iter__(self) -> Iterator[str]: ...
+ def __repr__(self) -> str: ...
diff --git a/venv/lib/python3.11/site-packages/av/error.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/error.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..e427e10417b619d75711dd88344b07957788784d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/error.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:e3889e7fe1beac623e21fb18c6282245c982bb6314ad93c8082c6baf671a5dfe
+size 1802113
diff --git a/venv/lib/python3.11/site-packages/av/error.pxd b/venv/lib/python3.11/site-packages/av/error.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..d9a542a36ab6db47d7226da67b07933f3b569ceb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/error.pxd
@@ -0,0 +1,3 @@
+
+cdef int stash_exception(exc_info=*)
+cpdef int err_check(int res, filename=*) except -1
diff --git a/venv/lib/python3.11/site-packages/av/error.pyi b/venv/lib/python3.11/site-packages/av/error.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..abbe2188ca8245346cc9e51801c1b384c4f56ba4
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/error.pyi
@@ -0,0 +1,72 @@
+import builtins
+from enum import Enum
+
+classes: dict[int, Exception]
+
+def code_to_tag(code: int) -> bytes: ...
+def tag_to_code(tag: bytes) -> int: ...
+def err_check(res: int, filename: str | None = None) -> int: ...
+
+class FFmpegError(Exception):
+ errno: int
+ strerror: str
+ filename: str
+ log: tuple[int, tuple[int, str, str] | None]
+
+ def __init__(
+ self,
+ code: int,
+ message: str,
+ filename: str | None = None,
+ log: tuple[int, tuple[int, str, str] | None] | None = None,
+ ) -> None: ...
+
+class LookupError(FFmpegError): ...
+class HTTPError(FFmpegError): ...
+class HTTPClientError(FFmpegError): ...
+class UndefinedError(FFmpegError): ...
+class InvalidDataError(FFmpegError, builtins.ValueError): ...
+class BugError(FFmpegError, builtins.RuntimeError): ...
+class BufferTooSmallError(FFmpegError, builtins.ValueError): ...
+class BSFNotFoundError(LookupError): ...
+class DecoderNotFoundError(LookupError): ...
+class DemuxerNotFoundError(LookupError): ...
+class EncoderNotFoundError(LookupError): ...
+class ExitError(FFmpegError): ...
+class ExternalError(FFmpegError): ...
+class FilterNotFoundError(LookupError): ...
+class MuxerNotFoundError(LookupError): ...
+class OptionNotFoundError(LookupError): ...
+class PatchWelcomeError(FFmpegError): ...
+class ProtocolNotFoundError(LookupError): ...
+class UnknownError(FFmpegError): ...
+class ExperimentalError(FFmpegError): ...
+class InputChangedError(FFmpegError): ...
+class OutputChangedError(FFmpegError): ...
+class HTTPBadRequestError(HTTPClientError): ...
+class HTTPUnauthorizedError(HTTPClientError): ...
+class HTTPForbiddenError(HTTPClientError): ...
+class HTTPNotFoundError(HTTPClientError): ...
+class HTTPOtherClientError(HTTPClientError): ...
+class HTTPServerError(HTTPError): ...
+class PyAVCallbackError(FFmpegError, builtins.RuntimeError): ...
+class BrokenPipeError(FFmpegError, builtins.BrokenPipeError): ...
+class ChildProcessError(FFmpegError, builtins.ChildProcessError): ...
+class ConnectionAbortedError(FFmpegError, builtins.ConnectionAbortedError): ...
+class ConnectionRefusedError(FFmpegError, builtins.ConnectionRefusedError): ...
+class ConnectionResetError(FFmpegError, builtins.ConnectionResetError): ...
+class BlockingIOError(FFmpegError, builtins.BlockingIOError): ...
+class EOFError(FFmpegError, builtins.EOFError): ...
+class FileExistsError(FFmpegError, builtins.FileExistsError): ...
+class FileNotFoundError(FFmpegError, builtins.FileNotFoundError): ...
+class InterruptedError(FFmpegError, builtins.InterruptedError): ...
+class IsADirectoryError(FFmpegError, builtins.IsADirectoryError): ...
+class MemoryError(FFmpegError, builtins.MemoryError): ...
+class NotADirectoryError(FFmpegError, builtins.NotADirectoryError): ...
+class NotImplementedError(FFmpegError, builtins.NotImplementedError): ...
+class OverflowError(FFmpegError, builtins.OverflowError): ...
+class OSError(FFmpegError, builtins.OSError): ...
+class PermissionError(FFmpegError, builtins.PermissionError): ...
+class ProcessLookupError(FFmpegError, builtins.ProcessLookupError): ...
+class TimeoutError(FFmpegError, builtins.TimeoutError): ...
+class ValueError(FFmpegError, builtins.ValueError): ...
diff --git a/venv/lib/python3.11/site-packages/av/filter/__init__.pxd b/venv/lib/python3.11/site-packages/av/filter/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/filter/__init__.py b/venv/lib/python3.11/site-packages/av/filter/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..5dd4430d474e27d1714b54a0f218709993def47c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/__init__.py
@@ -0,0 +1,3 @@
+from .filter import Filter, FilterFlags, filter_descriptor, filters_available
+from .graph import Graph
+from .loudnorm import stats
diff --git a/venv/lib/python3.11/site-packages/av/filter/__init__.pyi b/venv/lib/python3.11/site-packages/av/filter/__init__.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..5be1326c99efab97108a8dba6887cda669b6996a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/__init__.pyi
@@ -0,0 +1,4 @@
+from .context import *
+from .filter import *
+from .graph import *
+from .loudnorm import *
diff --git a/venv/lib/python3.11/site-packages/av/filter/context.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/filter/context.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..d5d71c9ac9dd97c01a67445f382b27ee6b472f35
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/context.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:5d9236eda2b0900d8f3e36197dba45ba13c0b436ab649cea3d8251ab21ccf103
+size 835217
diff --git a/venv/lib/python3.11/site-packages/av/filter/context.pxd b/venv/lib/python3.11/site-packages/av/filter/context.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..ae9f27c99f85e2de9ee526095b391464049f9efc
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/context.pxd
@@ -0,0 +1,19 @@
+cimport libav as lib
+
+from av.filter.filter cimport Filter
+from av.filter.graph cimport Graph
+
+
+cdef class FilterContext:
+
+ cdef lib.AVFilterContext *ptr
+ cdef readonly object _graph
+ cdef readonly Filter filter
+
+ cdef object _inputs
+ cdef object _outputs
+
+ cdef bint inited
+
+
+cdef FilterContext wrap_filter_context(Graph graph, Filter filter, lib.AVFilterContext *ptr)
diff --git a/venv/lib/python3.11/site-packages/av/filter/context.pyi b/venv/lib/python3.11/site-packages/av/filter/context.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..7c00087a928d37033f1e20a3692c0ede9276dfaa
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/context.pyi
@@ -0,0 +1,18 @@
+from av.filter import Graph
+from av.frame import Frame
+
+from .pad import FilterContextPad
+
+class FilterContext:
+ name: str | None
+ inputs: tuple[FilterContextPad, ...]
+ outputs: tuple[FilterContextPad, ...]
+
+ def init(self, args: str | None = None, **kwargs: str | None) -> None: ...
+ def link_to(
+ self, input_: FilterContext, output_idx: int = 0, input_idx: int = 0
+ ) -> None: ...
+ @property
+ def graph(self) -> Graph: ...
+ def push(self, frame: Frame) -> None: ...
+ def pull(self) -> Frame: ...
diff --git a/venv/lib/python3.11/site-packages/av/filter/filter.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/filter/filter.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..5261deb3c28edcfd669d106e426d65aeb0058435
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/filter.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:66d6affdbdfb0d5de11317aaf1c9e37e9d84eca50ddab7d56ba855fbdf40c84f
+size 982713
diff --git a/venv/lib/python3.11/site-packages/av/filter/filter.pxd b/venv/lib/python3.11/site-packages/av/filter/filter.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..27501ae575f7ed16741cee1ab497eb8a7c4a5bb7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/filter.pxd
@@ -0,0 +1,15 @@
+cimport libav as lib
+
+from av.descriptor cimport Descriptor
+
+
+cdef class Filter:
+
+ cdef const lib.AVFilter *ptr
+
+ cdef object _inputs
+ cdef object _outputs
+ cdef Descriptor _descriptor
+
+
+cdef Filter wrap_filter(const lib.AVFilter *ptr)
diff --git a/venv/lib/python3.11/site-packages/av/filter/filter.pyi b/venv/lib/python3.11/site-packages/av/filter/filter.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..2751e973cf5a8350a5f266604b65270ab6045324
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/filter.pyi
@@ -0,0 +1,23 @@
+from av.descriptor import Descriptor
+from av.option import Option
+
+from .pad import FilterPad
+
+class Filter:
+ name: str
+ description: str
+
+ descriptor: Descriptor
+ options: tuple[Option, ...] | None
+ flags: int
+ dynamic_inputs: bool
+ dynamic_outputs: bool
+ timeline_support: bool
+ slice_threads: bool
+ command_support: bool
+ inputs: tuple[FilterPad, ...]
+ outputs: tuple[FilterPad, ...]
+
+ def __init__(self, name: str) -> None: ...
+
+filters_available: set[str]
diff --git a/venv/lib/python3.11/site-packages/av/filter/graph.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/filter/graph.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..7f5f0c7c3da36774bd5f910f5c7c735deefbac55
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/graph.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:ebeb57fc60263aeb3ca9167c322d8e1bfded17f6de57fdc6a0d140cc66ef6ec9
+size 1077057
diff --git a/venv/lib/python3.11/site-packages/av/filter/graph.pxd b/venv/lib/python3.11/site-packages/av/filter/graph.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..2e52bd6ec3a67992dc0f2a40bbd5da8e5ed41fda
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/graph.pxd
@@ -0,0 +1,22 @@
+cimport libav as lib
+
+from av.filter.context cimport FilterContext
+
+
+cdef class Graph:
+ cdef object __weakref__
+
+ cdef lib.AVFilterGraph *ptr
+
+ cdef readonly bint configured
+ cpdef configure(self, bint auto_buffer=*, bint force=*)
+
+ cdef dict _name_counts
+ cdef str _get_unique_name(self, str name)
+
+ cdef _register_context(self, FilterContext)
+ cdef _auto_register(self)
+ cdef int _nb_filters_seen
+ cdef dict _context_by_ptr
+ cdef dict _context_by_name
+ cdef dict _context_by_type
diff --git a/venv/lib/python3.11/site-packages/av/filter/graph.pyi b/venv/lib/python3.11/site-packages/av/filter/graph.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..e170c2ce77da8ba8fa81561c24285886200d3d0d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/graph.pyi
@@ -0,0 +1,47 @@
+from fractions import Fraction
+from typing import Any
+
+from av.audio.format import AudioFormat
+from av.audio.frame import AudioFrame
+from av.audio.layout import AudioLayout
+from av.audio.stream import AudioStream
+from av.video.format import VideoFormat
+from av.video.frame import VideoFrame
+from av.video.stream import VideoStream
+
+from .context import FilterContext
+from .filter import Filter
+
+class Graph:
+ configured: bool
+
+ def __init__(self) -> None: ...
+ def configure(self, auto_buffer: bool = True, force: bool = False) -> None: ...
+ def link_nodes(self, *nodes: FilterContext) -> Graph: ...
+ def add(
+ self, filter: str | Filter, args: Any = None, **kwargs: str
+ ) -> FilterContext: ...
+ def add_buffer(
+ self,
+ template: VideoStream | None = None,
+ width: int | None = None,
+ height: int | None = None,
+ format: VideoFormat | None = None,
+ name: str | None = None,
+ time_base: Fraction | None = None,
+ ) -> FilterContext: ...
+ def add_abuffer(
+ self,
+ template: AudioStream | None = None,
+ sample_rate: int | None = None,
+ format: AudioFormat | str | None = None,
+ layout: AudioLayout | str | None = None,
+ channels: int | None = None,
+ name: str | None = None,
+ time_base: Fraction | None = None,
+ ) -> FilterContext: ...
+ def set_audio_frame_size(self, frame_size: int) -> None: ...
+ def push(self, frame: None | AudioFrame | VideoFrame) -> None: ...
+ def pull(self) -> VideoFrame | AudioFrame: ...
+ def vpush(self, frame: VideoFrame | None) -> None: ...
+ def vpull(self) -> VideoFrame: ...
diff --git a/venv/lib/python3.11/site-packages/av/filter/link.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/filter/link.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..3433eccaabb62824bc3dca75547953241d079245
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/link.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:9f19ebd70f19d7abec3c6e5c338c21238a4d182df25f0c10f7a9ee73509d26fe
+size 404601
diff --git a/venv/lib/python3.11/site-packages/av/filter/link.pxd b/venv/lib/python3.11/site-packages/av/filter/link.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..a6a4b1c092163e4739735026b6c26502003daa24
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/link.pxd
@@ -0,0 +1,16 @@
+cimport libav as lib
+
+from av.filter.graph cimport Graph
+from av.filter.pad cimport FilterContextPad
+
+
+cdef class FilterLink:
+
+ cdef readonly Graph graph
+ cdef lib.AVFilterLink *ptr
+
+ cdef FilterContextPad _input
+ cdef FilterContextPad _output
+
+
+cdef FilterLink wrap_filter_link(Graph graph, lib.AVFilterLink *ptr)
diff --git a/venv/lib/python3.11/site-packages/av/filter/link.pyi b/venv/lib/python3.11/site-packages/av/filter/link.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..dd420ad91e94856492769f1efea588d2ee9e30ff
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/link.pyi
@@ -0,0 +1,5 @@
+from .pad import FilterContextPad
+
+class FilterLink:
+ input: FilterContextPad
+ output: FilterContextPad
diff --git a/venv/lib/python3.11/site-packages/av/filter/loudnorm.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/filter/loudnorm.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..59ebc18242fa427c2f53104db8ca01526662f619
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/loudnorm.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:82142108ab2125510acee8cbe553572b9f6ef04f1c20a025de4e43e7a1599b82
+size 446337
diff --git a/venv/lib/python3.11/site-packages/av/filter/loudnorm.pxd b/venv/lib/python3.11/site-packages/av/filter/loudnorm.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..b08d3502fc9a3e39fd3273fb474fd8e974d65568
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/loudnorm.pxd
@@ -0,0 +1,4 @@
+from av.audio.stream cimport AudioStream
+
+
+cpdef bytes stats(str loudnorm_args, AudioStream stream)
diff --git a/venv/lib/python3.11/site-packages/av/filter/loudnorm.pyi b/venv/lib/python3.11/site-packages/av/filter/loudnorm.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..c680f638d43a8908438d81c56c07387524488412
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/loudnorm.pyi
@@ -0,0 +1,3 @@
+from av.audio.stream import AudioStream
+
+def stats(loudnorm_args: str, stream: AudioStream) -> bytes: ...
diff --git a/venv/lib/python3.11/site-packages/av/filter/pad.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/filter/pad.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..9e9b1a2a4bfc932e4e10d2706b4bb32cb123921a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/pad.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:fa6b24cfd44f200872927558df3f3e228e0f779a3df1352700a6cd886d3f620c
+size 490817
diff --git a/venv/lib/python3.11/site-packages/av/filter/pad.pxd b/venv/lib/python3.11/site-packages/av/filter/pad.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..15ac950fca1a3bcec2602b177775a456d9c20280
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/pad.pxd
@@ -0,0 +1,23 @@
+cimport libav as lib
+
+from av.filter.context cimport FilterContext
+from av.filter.filter cimport Filter
+from av.filter.link cimport FilterLink
+
+
+cdef class FilterPad:
+
+ cdef readonly Filter filter
+ cdef readonly FilterContext context
+ cdef readonly bint is_input
+ cdef readonly int index
+
+ cdef const lib.AVFilterPad *base_ptr
+
+
+cdef class FilterContextPad(FilterPad):
+
+ cdef FilterLink _link
+
+
+cdef tuple alloc_filter_pads(Filter, const lib.AVFilterPad *ptr, bint is_input, FilterContext context=?)
diff --git a/venv/lib/python3.11/site-packages/av/filter/pad.pyi b/venv/lib/python3.11/site-packages/av/filter/pad.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..1a6c9bda66dbc730f34b86104e8a4001d90e9dfd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/filter/pad.pyi
@@ -0,0 +1,10 @@
+from .link import FilterLink
+
+class FilterPad:
+ is_output: bool
+ name: str
+ type: str
+
+class FilterContextPad(FilterPad):
+ link: FilterLink | None
+ linked: FilterContextPad | None
diff --git a/venv/lib/python3.11/site-packages/av/format.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/format.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..4c096a1dbdbf204710858235d068f966513de7d0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/format.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:a77ffaa9b835f8f2e4abd88de6d5f4b8793edc180a56849a8a26b5e3c1547838
+size 527977
diff --git a/venv/lib/python3.11/site-packages/av/format.pxd b/venv/lib/python3.11/site-packages/av/format.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..31cac50aab569887a8c50d8c2883815f89ba22e0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/format.pxd
@@ -0,0 +1,12 @@
+cimport libav as lib
+
+
+cdef class ContainerFormat:
+
+ cdef readonly str name
+
+ cdef lib.AVInputFormat *iptr
+ cdef lib.AVOutputFormat *optr
+
+
+cdef ContainerFormat build_container_format(lib.AVInputFormat*, lib.AVOutputFormat*)
diff --git a/venv/lib/python3.11/site-packages/av/format.pyi b/venv/lib/python3.11/site-packages/av/format.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..b30a84bf67f59142986bc652255ece8d00fe604d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/format.pyi
@@ -0,0 +1,42 @@
+__all__ = ("Flags", "ContainerFormat", "formats_available")
+
+from enum import Flag
+from typing import ClassVar, Literal
+
+class Flags(Flag):
+ no_file: ClassVar[Flags]
+ need_number: ClassVar[Flags]
+ show_ids: ClassVar[Flags]
+ global_header: ClassVar[Flags]
+ no_timestamps: ClassVar[Flags]
+ generic_index: ClassVar[Flags]
+ ts_discont: ClassVar[Flags]
+ variable_fps: ClassVar[Flags]
+ no_dimensions: ClassVar[Flags]
+ no_streams: ClassVar[Flags]
+ no_bin_search: ClassVar[Flags]
+ no_gen_search: ClassVar[Flags]
+ no_byte_seek: ClassVar[Flags]
+ allow_flush: ClassVar[Flags]
+ ts_nonstrict: ClassVar[Flags]
+ ts_negative: ClassVar[Flags]
+ seek_to_pts: ClassVar[Flags]
+
+class ContainerFormat:
+ def __init__(self, name: str, mode: Literal["r", "w"] | None = None) -> None: ...
+ @property
+ def name(self) -> str: ...
+ @property
+ def long_name(self) -> str: ...
+ @property
+ def is_input(self) -> bool: ...
+ @property
+ def is_output(self) -> bool: ...
+ @property
+ def extensions(self) -> set[str]: ...
+ @property
+ def flags(self) -> int: ...
+ @property
+ def no_file(self) -> bool: ...
+
+formats_available: set[str]
diff --git a/venv/lib/python3.11/site-packages/av/frame.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/frame.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..9eca3ef43c22799aebec75acd38efaac24309aa9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/frame.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:95ba438410d2439dccc0251d1859a70e522ad5d8403fc089b45583eeef2448a6
+size 503273
diff --git a/venv/lib/python3.11/site-packages/av/frame.pxd b/venv/lib/python3.11/site-packages/av/frame.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..6d7214b7caf4edae8f59075a32c805574ccdc514
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/frame.pxd
@@ -0,0 +1,14 @@
+cimport libav as lib
+
+from av.packet cimport Packet
+from av.sidedata.sidedata cimport _SideDataContainer
+
+
+cdef class Frame:
+ cdef lib.AVFrame *ptr
+ # We define our own time.
+ cdef lib.AVRational _time_base
+ cdef _rebase_time(self, lib.AVRational)
+ cdef _SideDataContainer _side_data
+ cdef _copy_internal_attributes(self, Frame source, bint data_layout=?)
+ cdef _init_user_attributes(self)
diff --git a/venv/lib/python3.11/site-packages/av/frame.pyi b/venv/lib/python3.11/site-packages/av/frame.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..9af81dcfe41c0d2ccba741604cef7c22eb2ac529
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/frame.pyi
@@ -0,0 +1,21 @@
+from fractions import Fraction
+from typing import TypedDict
+
+from av.sidedata.motionvectors import MotionVectors
+
+class SideData(TypedDict, total=False):
+ MOTION_VECTORS: MotionVectors
+
+class Frame:
+ dts: int | None
+ pts: int | None
+ time_base: Fraction
+ side_data: SideData
+ opaque: object
+ @property
+ def time(self) -> float | None: ...
+ @property
+ def is_corrupt(self) -> bool: ...
+ @property
+ def key_frame(self) -> bool: ...
+ def make_writable(self) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/include/libav.pxd b/venv/lib/python3.11/site-packages/av/include/libav.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..c793b9988168f40fdf2f46274bc276fe270d23f9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libav.pxd
@@ -0,0 +1,26 @@
+include "libavutil/avutil.pxd"
+include "libavutil/buffer.pxd"
+include "libavutil/channel_layout.pxd"
+include "libavutil/dict.pxd"
+include "libavutil/error.pxd"
+include "libavutil/frame.pxd"
+include "libavutil/samplefmt.pxd"
+include "libavutil/motion_vector.pxd"
+
+include "libavcodec/avcodec.pxd"
+include "libavcodec/bsf.pxd"
+include "libavdevice/avdevice.pxd"
+include "libavformat/avformat.pxd"
+include "libswresample/swresample.pxd"
+include "libswscale/swscale.pxd"
+
+include "libavfilter/avfilter.pxd"
+include "libavfilter/avfiltergraph.pxd"
+include "libavfilter/buffersink.pxd"
+include "libavfilter/buffersrc.pxd"
+
+
+cdef extern from "stdio.h" nogil:
+
+ cdef int snprintf(char *output, int n, const char *format, ...)
+ cdef int vsnprintf(char *output, int n, const char *format, va_list args)
diff --git a/venv/lib/python3.11/site-packages/av/include/libavcodec/avcodec.pxd b/venv/lib/python3.11/site-packages/av/include/libavcodec/avcodec.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..172c9cc65652693362534d10214764939c0eabb9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavcodec/avcodec.pxd
@@ -0,0 +1,549 @@
+from libc.stdint cimport int8_t, int64_t, uint16_t, uint32_t
+
+cdef extern from "libavcodec/codec.h":
+ struct AVCodecTag:
+ pass
+
+cdef extern from "libavcodec/codec_id.h":
+ AVCodecID av_codec_get_id(const AVCodecTag *const *tags, uint32_t tag)
+
+
+cdef extern from "libavutil/channel_layout.h":
+ ctypedef enum AVChannelOrder:
+ AV_CHANNEL_ORDER_UNSPEC
+ AV_CHANNEL_ORDER_NATIVE
+ AV_CHANNEL_ORDER_CUSTOM
+ AV_CHANNEL_ORDER_AMBISONIC
+
+ ctypedef enum AVChannel:
+ AV_CHAN_NONE = -1
+ AV_CHAN_FRONT_LEFT
+ AV_CHAN_FRONT_RIGHT
+ AV_CHAN_FRONT_CENTER
+ # ... other channel enum values ...
+
+ ctypedef struct AVChannelCustom:
+ AVChannel id
+ char name[16]
+ void *opaque
+
+ ctypedef struct AVChannelLayout:
+ AVChannelOrder order
+ int nb_channels
+ uint64_t mask
+ # union:
+ # uint64_t mask
+ # AVChannelCustom *map
+ void *opaque
+
+ int av_channel_layout_default(AVChannelLayout *ch_layout, int nb_channels)
+ int av_channel_layout_from_mask(AVChannelLayout *channel_layout, uint64_t mask)
+ int av_channel_layout_from_string(AVChannelLayout *channel_layout, const char *str)
+ void av_channel_layout_uninit(AVChannelLayout *channel_layout)
+ int av_channel_layout_copy(AVChannelLayout *dst, const AVChannelLayout *src)
+ int av_channel_layout_describe(const AVChannelLayout *channel_layout, char *buf, size_t buf_size)
+ int av_channel_name(char *buf, size_t buf_size, AVChannel channel_id)
+ int av_channel_description(char *buf, size_t buf_size, AVChannel channel_id)
+ AVChannel av_channel_layout_channel_from_index(AVChannelLayout *channel_layout, unsigned int idx)
+
+
+cdef extern from "libavcodec/avcodec.h" nogil:
+ cdef set pyav_get_available_codecs()
+
+ cdef int avcodec_version()
+ cdef char* avcodec_configuration()
+ cdef char* avcodec_license()
+
+ cdef size_t AV_INPUT_BUFFER_PADDING_SIZE
+ cdef int64_t AV_NOPTS_VALUE
+
+ # AVCodecDescriptor.props
+ cdef enum:
+ AV_CODEC_PROP_INTRA_ONLY
+ AV_CODEC_PROP_LOSSY
+ AV_CODEC_PROP_LOSSLESS
+ AV_CODEC_PROP_REORDER
+ AV_CODEC_PROP_BITMAP_SUB
+ AV_CODEC_PROP_TEXT_SUB
+
+ # AVCodec.capabilities
+ cdef enum:
+ AV_CODEC_CAP_DRAW_HORIZ_BAND
+ AV_CODEC_CAP_DR1
+ # AV_CODEC_CAP_HWACCEL
+ AV_CODEC_CAP_DELAY
+ AV_CODEC_CAP_SMALL_LAST_FRAME
+ # AV_CODEC_CAP_HWACCEL_VDPAU
+ AV_CODEC_CAP_SUBFRAMES
+ AV_CODEC_CAP_EXPERIMENTAL
+ AV_CODEC_CAP_CHANNEL_CONF
+ # AV_CODEC_CAP_NEG_LINESIZES
+ AV_CODEC_CAP_FRAME_THREADS
+ AV_CODEC_CAP_SLICE_THREADS
+ AV_CODEC_CAP_PARAM_CHANGE
+ AV_CODEC_CAP_OTHER_THREADS
+ AV_CODEC_CAP_VARIABLE_FRAME_SIZE
+ AV_CODEC_CAP_AVOID_PROBING
+ AV_CODEC_CAP_HARDWARE
+ AV_CODEC_CAP_HYBRID
+ AV_CODEC_CAP_ENCODER_REORDERED_OPAQUE
+
+ cdef enum:
+ FF_THREAD_FRAME
+ FF_THREAD_SLICE
+
+ cdef enum:
+ AV_CODEC_FLAG_UNALIGNED
+ AV_CODEC_FLAG_QSCALE
+ AV_CODEC_FLAG_4MV
+ AV_CODEC_FLAG_OUTPUT_CORRUPT
+ AV_CODEC_FLAG_QPEL
+ AV_CODEC_FLAG_DROPCHANGED
+ AV_CODEC_FLAG_RECON_FRAME
+ AV_CODEC_FLAG_COPY_OPAQUE
+ AV_CODEC_FLAG_FRAME_DURATION
+ AV_CODEC_FLAG_PASS1
+ AV_CODEC_FLAG_PASS2
+ AV_CODEC_FLAG_LOOP_FILTER
+ AV_CODEC_FLAG_GRAY
+ AV_CODEC_FLAG_PSNR
+ AV_CODEC_FLAG_INTERLACED_DCT
+ AV_CODEC_FLAG_LOW_DELAY
+ AV_CODEC_FLAG_GLOBAL_HEADER
+ AV_CODEC_FLAG_BITEXACT
+ AV_CODEC_FLAG_AC_PRED
+ AV_CODEC_FLAG_INTERLACED_ME
+ AV_CODEC_FLAG_CLOSED_GOP
+
+ cdef enum:
+ AV_CODEC_FLAG2_FAST
+ AV_CODEC_FLAG2_NO_OUTPUT
+ AV_CODEC_FLAG2_LOCAL_HEADER
+ AV_CODEC_FLAG2_CHUNKS
+ AV_CODEC_FLAG2_IGNORE_CROP
+ AV_CODEC_FLAG2_SHOW_ALL
+ AV_CODEC_FLAG2_EXPORT_MVS
+ AV_CODEC_FLAG2_SKIP_MANUAL
+ AV_CODEC_FLAG2_RO_FLUSH_NOOP
+
+ cdef enum:
+ AV_PKT_FLAG_KEY
+ AV_PKT_FLAG_CORRUPT
+ AV_PKT_FLAG_DISCARD
+ AV_PKT_FLAG_TRUSTED
+ AV_PKT_FLAG_DISPOSABLE
+
+ cdef enum:
+ AV_FRAME_FLAG_CORRUPT
+ AV_FRAME_FLAG_KEY
+ AV_FRAME_FLAG_DISCARD
+ AV_FRAME_FLAG_INTERLACED
+
+ cdef enum:
+ FF_COMPLIANCE_VERY_STRICT
+ FF_COMPLIANCE_STRICT
+ FF_COMPLIANCE_NORMAL
+ FF_COMPLIANCE_UNOFFICIAL
+ FF_COMPLIANCE_EXPERIMENTAL
+
+ cdef enum:
+ FF_PROFILE_UNKNOWN = -99
+
+ cdef enum AVCodecID:
+ AV_CODEC_ID_NONE
+ AV_CODEC_ID_MPEG2VIDEO
+ AV_CODEC_ID_MPEG1VIDEO
+
+ cdef enum AVDiscard:
+ AVDISCARD_NONE
+ AVDISCARD_DEFAULT
+ AVDISCARD_NONREF
+ AVDISCARD_BIDIR
+ AVDISCARD_NONINTRA
+ AVDISCARD_NONKEY
+ AVDISCARD_ALL
+
+ cdef struct AVCodec:
+ char *name
+ char *long_name
+ AVMediaType type
+ AVCodecID id
+
+ int capabilities
+
+ AVRational* supported_framerates
+ AVSampleFormat* sample_fmts
+ AVPixelFormat* pix_fmts
+ int* supported_samplerates
+
+ AVClass *priv_class
+
+ cdef int av_codec_is_encoder(AVCodec*)
+ cdef int av_codec_is_decoder(AVCodec*)
+
+ cdef struct AVProfile:
+ int profile
+ char *name
+
+ cdef struct AVCodecDescriptor:
+ AVCodecID id
+ char *name
+ char *long_name
+ int props
+ char **mime_types
+ AVProfile *profiles
+
+ AVCodecDescriptor* avcodec_descriptor_get(AVCodecID)
+
+ cdef struct AVCodecContext:
+ AVClass *av_class
+
+ AVMediaType codec_type
+ char codec_name[32]
+ unsigned int codec_tag
+ AVCodecID codec_id
+
+ int flags
+ int flags2
+ int thread_count
+ int thread_type
+
+ int profile
+ AVDiscard skip_frame
+
+ AVFrame* coded_frame
+
+ int bit_rate
+ int bit_rate_tolerance
+ int mb_decision
+
+ int bits_per_coded_sample
+ int global_quality
+ int compression_level
+
+ int qmin
+ int qmax
+ int rc_max_rate
+ int rc_min_rate
+ int rc_buffer_size
+ float rc_max_available_vbv_use
+ float rc_min_vbv_overflow_use
+
+ AVRational framerate
+ AVRational pkt_timebase
+ AVRational time_base
+
+ int extradata_size
+ uint8_t *extradata
+
+ int delay
+
+ AVCodec *codec
+
+ # Video.
+ int width
+ int height
+ int coded_width
+ int coded_height
+
+ AVPixelFormat pix_fmt
+ AVRational sample_aspect_ratio
+ int gop_size # The number of pictures in a group of pictures, or 0 for intra_only.
+ int max_b_frames
+ int has_b_frames
+ AVColorRange color_range
+ AVColorPrimaries color_primaries
+ AVColorTransferCharacteristic color_trc
+ AVColorSpace colorspace
+
+ # Audio.
+ AVSampleFormat sample_fmt
+ int sample_rate
+ AVChannelLayout ch_layout
+ int frame_size
+
+ #: .. todo:: ``get_buffer`` is deprecated for get_buffer2 in newer versions of FFmpeg.
+ int get_buffer(AVCodecContext *ctx, AVFrame *frame)
+ void release_buffer(AVCodecContext *ctx, AVFrame *frame)
+
+ # User Data
+ void *opaque
+
+ cdef AVCodecContext* avcodec_alloc_context3(AVCodec *codec)
+ cdef void avcodec_free_context(AVCodecContext **ctx)
+
+ cdef AVClass* avcodec_get_class()
+
+ cdef AVCodec* avcodec_find_decoder(AVCodecID id)
+ cdef AVCodec* avcodec_find_encoder(AVCodecID id)
+
+ cdef AVCodec* avcodec_find_decoder_by_name(char *name)
+ cdef AVCodec* avcodec_find_encoder_by_name(char *name)
+
+ cdef const AVCodec* av_codec_iterate(void **opaque)
+
+ cdef AVCodecDescriptor* avcodec_descriptor_get (AVCodecID id)
+ cdef AVCodecDescriptor* avcodec_descriptor_get_by_name (char *name)
+
+ cdef char* avcodec_get_name(AVCodecID id)
+
+ cdef char* av_get_profile_name(AVCodec *codec, int profile)
+
+ cdef int avcodec_open2(
+ AVCodecContext *ctx,
+ AVCodec *codec,
+ AVDictionary **options,
+ )
+
+ cdef int AV_NUM_DATA_POINTERS
+
+ cdef enum AVPacketSideDataType:
+ AV_PKT_DATA_PALETTE
+ AV_PKT_DATA_NEW_EXTRADATA
+ AV_PKT_DATA_PARAM_CHANGE
+ AV_PKT_DATA_H263_MB_INFO
+ AV_PKT_DATA_REPLAYGAIN
+ AV_PKT_DATA_DISPLAYMATRIX
+ AV_PKT_DATA_STEREO3D
+ AV_PKT_DATA_AUDIO_SERVICE_TYPE
+ AV_PKT_DATA_QUALITY_STATS
+ AV_PKT_DATA_FALLBACK_TRACK
+ AV_PKT_DATA_CPB_PROPERTIES
+ AV_PKT_DATA_SKIP_SAMPLES
+ AV_PKT_DATA_JP_DUALMONO
+ AV_PKT_DATA_STRINGS_METADATA
+ AV_PKT_DATA_SUBTITLE_POSITION
+ AV_PKT_DATA_MATROSKA_BLOCKADDITIONAL
+ AV_PKT_DATA_WEBVTT_IDENTIFIER
+ AV_PKT_DATA_WEBVTT_SETTINGS
+ AV_PKT_DATA_METADATA_UPDATE
+ AV_PKT_DATA_MPEGTS_STREAM_ID
+ AV_PKT_DATA_MASTERING_DISPLAY_METADATA
+ AV_PKT_DATA_SPHERICAL
+ AV_PKT_DATA_CONTENT_LIGHT_LEVEL
+ AV_PKT_DATA_A53_CC
+ AV_PKT_DATA_ENCRYPTION_INIT_INFO
+ AV_PKT_DATA_ENCRYPTION_INFO
+ AV_PKT_DATA_AFD
+ AV_PKT_DATA_PRFT
+ AV_PKT_DATA_ICC_PROFILE
+ AV_PKT_DATA_DOVI_CONF
+ AV_PKT_DATA_S12M_TIMECODE
+ AV_PKT_DATA_DYNAMIC_HDR10_PLUS
+ AV_PKT_DATA_NB
+
+ cdef struct AVPacketSideData:
+ uint8_t *data;
+ size_t size;
+ AVPacketSideDataType type;
+
+ cdef enum AVFrameSideDataType:
+ AV_FRAME_DATA_PANSCAN
+ AV_FRAME_DATA_A53_CC
+ AV_FRAME_DATA_STEREO3D
+ AV_FRAME_DATA_MATRIXENCODING
+ AV_FRAME_DATA_DOWNMIX_INFO
+ AV_FRAME_DATA_REPLAYGAIN
+ AV_FRAME_DATA_DISPLAYMATRIX
+ AV_FRAME_DATA_AFD
+ AV_FRAME_DATA_MOTION_VECTORS
+ AV_FRAME_DATA_SKIP_SAMPLES
+ AV_FRAME_DATA_AUDIO_SERVICE_TYPE
+ AV_FRAME_DATA_MASTERING_DISPLAY_METADATA
+ AV_FRAME_DATA_GOP_TIMECODE
+ AV_FRAME_DATA_SPHERICAL
+ AV_FRAME_DATA_CONTENT_LIGHT_LEVEL
+ AV_FRAME_DATA_ICC_PROFILE
+ AV_FRAME_DATA_S12M_TIMECODE
+ AV_FRAME_DATA_DYNAMIC_HDR_PLUS
+ AV_FRAME_DATA_REGIONS_OF_INTEREST
+ AV_FRAME_DATA_VIDEO_ENC_PARAMS
+ AV_FRAME_DATA_SEI_UNREGISTERED
+ AV_FRAME_DATA_FILM_GRAIN_PARAMS
+ AV_FRAME_DATA_DETECTION_BBOXES
+ AV_FRAME_DATA_DOVI_RPU_BUFFER
+ AV_FRAME_DATA_DOVI_METADATA
+ AV_FRAME_DATA_DYNAMIC_HDR_VIVID
+ AV_FRAME_DATA_AMBIENT_VIEWING_ENVIRONMENT
+ AV_FRAME_DATA_VIDEO_HINT
+
+ cdef struct AVFrameSideData:
+ AVFrameSideDataType type
+ uint8_t *data
+ int size
+ AVDictionary *metadata
+
+ # See: http://ffmpeg.org/doxygen/trunk/structAVFrame.html
+ cdef struct AVFrame:
+ uint8_t *data[4]
+ int linesize[4]
+ uint8_t **extended_data
+
+ int format # Should be AVPixelFormat or AVSampleFormat
+ AVPictureType pict_type
+
+ int width
+ int height
+
+ int nb_side_data
+ AVFrameSideData **side_data
+
+ int nb_samples
+ int sample_rate
+
+ AVChannelLayout ch_layout
+
+ int64_t pts
+ int64_t pkt_dts
+
+ int pkt_size
+
+ uint8_t **base
+ void *opaque
+ AVBufferRef *opaque_ref
+ AVDictionary *metadata
+ int flags
+ int decode_error_flags
+ AVColorRange color_range
+ AVColorPrimaries color_primaries
+ AVColorTransferCharacteristic color_trc
+ AVColorSpace colorspace
+
+ cdef AVFrame* avcodec_alloc_frame()
+
+ cdef struct AVPacket:
+
+ int64_t pts
+ int64_t dts
+ uint8_t *data
+
+ int size
+ int stream_index
+ int flags
+
+ int duration
+
+ int64_t pos
+
+ void *opaque
+ AVBufferRef *opaque_ref
+
+
+ cdef int avcodec_fill_audio_frame(
+ AVFrame *frame,
+ int nb_channels,
+ AVSampleFormat sample_fmt,
+ uint8_t *buf,
+ int buf_size,
+ int align
+ )
+
+ cdef void avcodec_free_frame(AVFrame **frame)
+
+ cdef AVPacket* av_packet_alloc()
+ cdef void av_packet_free(AVPacket **)
+ cdef int av_new_packet(AVPacket*, int)
+ cdef int av_packet_ref(AVPacket *dst, const AVPacket *src)
+ cdef void av_packet_rescale_ts(AVPacket *pkt, AVRational src_tb, AVRational dst_tb)
+
+ cdef enum AVSubtitleType:
+ SUBTITLE_NONE
+ SUBTITLE_BITMAP
+ SUBTITLE_TEXT
+ SUBTITLE_ASS
+
+ cdef struct AVSubtitleRect:
+ int x
+ int y
+ int w
+ int h
+ int nb_colors
+ uint8_t *data[4]
+ int linesize[4]
+ AVSubtitleType type
+ char *text
+ char *ass
+ int flags
+
+ cdef struct AVSubtitle:
+ uint16_t format
+ uint32_t start_display_time
+ uint32_t end_display_time
+ unsigned int num_rects
+ AVSubtitleRect **rects
+ int64_t pts
+
+ cdef int avcodec_decode_subtitle2(
+ AVCodecContext *ctx,
+ AVSubtitle *sub,
+ int *done,
+ AVPacket *pkt,
+ )
+
+ cdef int avcodec_encode_subtitle(
+ AVCodecContext *avctx,
+ uint8_t *buf,
+ int buf_size,
+ AVSubtitle *sub
+ )
+
+ cdef void avsubtitle_free(AVSubtitle*)
+
+ cdef void avcodec_get_frame_defaults(AVFrame* frame)
+
+ cdef void avcodec_flush_buffers(AVCodecContext *ctx)
+
+ # TODO: avcodec_default_get_buffer is deprecated for avcodec_default_get_buffer2 in newer versions of FFmpeg
+ cdef int avcodec_default_get_buffer(AVCodecContext *ctx, AVFrame *frame)
+ cdef void avcodec_default_release_buffer(AVCodecContext *ctx, AVFrame *frame)
+
+ # === New-style Transcoding
+ cdef int avcodec_send_packet(AVCodecContext *avctx, AVPacket *packet)
+ cdef int avcodec_receive_frame(AVCodecContext *avctx, AVFrame *frame)
+ cdef int avcodec_send_frame(AVCodecContext *avctx, AVFrame *frame)
+ cdef int avcodec_receive_packet(AVCodecContext *avctx, AVPacket *avpkt)
+
+ # === Parsers
+
+ cdef struct AVCodecParser:
+ int codec_ids[5]
+
+ cdef AVCodecParser* av_parser_next(AVCodecParser *c)
+
+ cdef struct AVCodecParserContext:
+ pass
+
+ cdef AVCodecParserContext *av_parser_init(int codec_id)
+ cdef int av_parser_parse2(
+ AVCodecParserContext *s,
+ AVCodecContext *avctx,
+ uint8_t **poutbuf, int *poutbuf_size,
+ const uint8_t *buf, int buf_size,
+ int64_t pts, int64_t dts,
+ int64_t pos
+ )
+ cdef int av_parser_change(
+ AVCodecParserContext *s,
+ AVCodecContext *avctx,
+ uint8_t **poutbuf, int *poutbuf_size,
+ const uint8_t *buf, int buf_size,
+ int keyframe
+ )
+ cdef void av_parser_close(AVCodecParserContext *s)
+
+ cdef struct AVCodecParameters:
+ AVMediaType codec_type
+ AVCodecID codec_id
+
+ cdef int avcodec_parameters_copy(
+ AVCodecParameters *dst,
+ const AVCodecParameters *src
+ )
+ cdef int avcodec_parameters_from_context(
+ AVCodecParameters *par,
+ const AVCodecContext *codec,
+ )
+ cdef int avcodec_parameters_to_context(
+ AVCodecContext *codec,
+ const AVCodecParameters *par
+ )
diff --git a/venv/lib/python3.11/site-packages/av/include/libavcodec/bsf.pxd b/venv/lib/python3.11/site-packages/av/include/libavcodec/bsf.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..4a558b4789da6a80abc8a30022f9b0d6fcb0f087
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavcodec/bsf.pxd
@@ -0,0 +1,40 @@
+
+cdef extern from "libavcodec/bsf.h" nogil:
+
+ cdef struct AVBitStreamFilter:
+ const char *name
+ AVCodecID *codec_ids
+
+ cdef struct AVCodecParameters:
+ pass
+
+ cdef struct AVBSFContext:
+ const AVBitStreamFilter *filter
+ const AVCodecParameters *par_in
+ const AVCodecParameters *par_out
+
+ cdef const AVBitStreamFilter* av_bsf_get_by_name(const char *name)
+
+ cdef int av_bsf_list_parse_str(
+ const char *str,
+ AVBSFContext **bsf
+ )
+
+ cdef int av_bsf_init(AVBSFContext *ctx)
+ cdef void av_bsf_free(AVBSFContext **ctx)
+
+ cdef AVBitStreamFilter* av_bsf_iterate(void **opaque)
+
+ cdef int av_bsf_send_packet(
+ AVBSFContext *ctx,
+ AVPacket *pkt
+ )
+
+ cdef int av_bsf_receive_packet(
+ AVBSFContext *ctx,
+ AVPacket *pkt
+ )
+
+ cdef void av_bsf_flush(
+ AVBSFContext *ctx
+ )
diff --git a/venv/lib/python3.11/site-packages/av/include/libavdevice/avdevice.pxd b/venv/lib/python3.11/site-packages/av/include/libavdevice/avdevice.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..bc9b4adea1ae9e3aaad90a21177a227bff69d3ac
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavdevice/avdevice.pxd
@@ -0,0 +1,12 @@
+
+cdef extern from "libavdevice/avdevice.h" nogil:
+
+ cdef int avdevice_version()
+ cdef char* avdevice_configuration()
+ cdef char* avdevice_license()
+ void avdevice_register_all()
+
+ AVInputFormat * av_input_audio_device_next(AVInputFormat *d)
+ AVInputFormat * av_input_video_device_next(AVInputFormat *d)
+ AVOutputFormat * av_output_audio_device_next(AVOutputFormat *d)
+ AVOutputFormat * av_output_video_device_next(AVOutputFormat *d)
diff --git a/venv/lib/python3.11/site-packages/av/include/libavfilter/avfilter.pxd b/venv/lib/python3.11/site-packages/av/include/libavfilter/avfilter.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..dd3e91ddf3ff97645f2b633859668fc5f0d9ea7e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavfilter/avfilter.pxd
@@ -0,0 +1,91 @@
+
+cdef extern from "libavfilter/avfilter.h" nogil:
+ """
+ #if (LIBAVFILTER_VERSION_INT >= 525156)
+ // avfilter_filter_pad_count is available since version 8.3.100 of libavfilter (FFmpeg 5.0)
+ #define _avfilter_get_num_pads(filter, is_output, pads) (avfilter_filter_pad_count(filter, is_output))
+ #else
+ // avfilter_filter_pad_count has been deprecated as of version 8.3.100 of libavfilter (FFmpeg 5.0)
+ #define _avfilter_get_num_pads(filter, is_output, pads) (avfilter_pad_count(pads))
+ #endif
+ """
+ cdef int avfilter_version()
+ cdef char* avfilter_configuration()
+ cdef char* avfilter_license()
+
+ cdef struct AVFilterPad:
+ # This struct is opaque.
+ pass
+
+ const char* avfilter_pad_get_name(const AVFilterPad *pads, int index)
+ AVMediaType avfilter_pad_get_type(const AVFilterPad *pads, int index)
+
+ int pyav_get_num_pads "_avfilter_get_num_pads" (const AVFilter *filter, int is_output, const AVFilterPad *pads)
+
+ cdef struct AVFilter:
+
+ AVClass *priv_class
+
+ const char *name
+ const char *description
+
+ const int flags
+
+ const AVFilterPad *inputs
+ const AVFilterPad *outputs
+ int (*process_command)(AVFilterContext *, const char *cmd, const char *arg, char *res, int res_len, int flags)
+
+ cdef enum:
+ AVFILTER_FLAG_DYNAMIC_INPUTS
+ AVFILTER_FLAG_DYNAMIC_OUTPUTS
+ AVFILTER_FLAG_SLICE_THREADS
+ AVFILTER_FLAG_SUPPORT_TIMELINE_GENERIC
+ AVFILTER_FLAG_SUPPORT_TIMELINE_INTERNAL
+
+ cdef AVFilter* avfilter_get_by_name(const char *name)
+ cdef const AVFilter* av_filter_iterate(void **opaque)
+
+ cdef struct AVFilterLink # Defined later.
+
+ cdef struct AVFilterContext:
+
+ AVClass *av_class
+ AVFilter *filter
+
+ char *name
+
+ unsigned int nb_inputs
+ AVFilterPad *input_pads
+ AVFilterLink **inputs
+
+ unsigned int nb_outputs
+ AVFilterPad *output_pads
+ AVFilterLink **outputs
+
+ cdef int avfilter_init_str(AVFilterContext *ctx, const char *args)
+ cdef int avfilter_init_dict(AVFilterContext *ctx, AVDictionary **options)
+ cdef void avfilter_free(AVFilterContext*)
+ cdef AVClass* avfilter_get_class()
+
+ cdef struct AVFilterLink:
+
+ AVFilterContext *src
+ AVFilterPad *srcpad
+ AVFilterContext *dst
+ AVFilterPad *dstpad
+
+ AVMediaType Type
+ int w
+ int h
+ AVRational sample_aspect_ratio
+ uint64_t channel_layout
+ int sample_rate
+ int format
+ AVRational time_base
+
+ # custom
+ cdef set pyav_get_available_filters()
+
+
+cdef extern from "libavfilter/buffersink.h" nogil:
+ cdef void av_buffersink_set_frame_size(AVFilterContext *ctx, unsigned frame_size)
diff --git a/venv/lib/python3.11/site-packages/av/include/libavfilter/avfiltergraph.pxd b/venv/lib/python3.11/site-packages/av/include/libavfilter/avfiltergraph.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..b773063f90cae11865555c1ee846ee1ddc0e07d8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavfilter/avfiltergraph.pxd
@@ -0,0 +1,50 @@
+
+cdef extern from "libavfilter/avfilter.h" nogil:
+
+ cdef struct AVFilterGraph:
+ int nb_filters
+ AVFilterContext **filters
+
+ cdef struct AVFilterInOut:
+ char *name
+ AVFilterContext *filter_ctx
+ int pad_idx
+ AVFilterInOut *next
+
+ cdef AVFilterGraph* avfilter_graph_alloc()
+ cdef void avfilter_graph_free(AVFilterGraph **ptr)
+
+ cdef int avfilter_graph_parse2(
+ AVFilterGraph *graph,
+ const char *filter_str,
+ AVFilterInOut **inputs,
+ AVFilterInOut **outputs
+ )
+
+ cdef AVFilterContext* avfilter_graph_alloc_filter(
+ AVFilterGraph *graph,
+ const AVFilter *filter,
+ const char *name
+ )
+
+ cdef int avfilter_graph_create_filter(
+ AVFilterContext **filt_ctx,
+ AVFilter *filt,
+ const char *name,
+ const char *args,
+ void *opaque,
+ AVFilterGraph *graph_ctx
+ )
+
+ cdef int avfilter_link(
+ AVFilterContext *src,
+ unsigned int srcpad,
+ AVFilterContext *dst,
+ unsigned int dstpad
+ )
+
+ cdef int avfilter_graph_config(AVFilterGraph *graph, void *logctx)
+
+ cdef char* avfilter_graph_dump(AVFilterGraph *graph, const char *options)
+
+ cdef void avfilter_inout_free(AVFilterInOut **inout_list)
diff --git a/venv/lib/python3.11/site-packages/av/include/libavfilter/buffersink.pxd b/venv/lib/python3.11/site-packages/av/include/libavfilter/buffersink.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..84ea56c687d69155fd754edacf42e7f34e6b3639
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavfilter/buffersink.pxd
@@ -0,0 +1,6 @@
+cdef extern from "libavfilter/buffersink.h" nogil:
+
+ int av_buffersink_get_frame(
+ AVFilterContext *ctx,
+ AVFrame *frame
+ )
diff --git a/venv/lib/python3.11/site-packages/av/include/libavfilter/buffersrc.pxd b/venv/lib/python3.11/site-packages/av/include/libavfilter/buffersrc.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..db057662ad3eaa6c7b8269d1874a2618c6b8ee8e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavfilter/buffersrc.pxd
@@ -0,0 +1,6 @@
+cdef extern from "libavfilter/buffersrc.h" nogil:
+
+ int av_buffersrc_write_frame(
+ AVFilterContext *ctx,
+ const AVFrame *frame
+ )
diff --git a/venv/lib/python3.11/site-packages/av/include/libavformat/avformat.pxd b/venv/lib/python3.11/site-packages/av/include/libavformat/avformat.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..5fa25043aa67ba8e2c91745ccadf80ba1e5f14eb
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavformat/avformat.pxd
@@ -0,0 +1,345 @@
+from libc.stdint cimport int64_t, uint64_t
+
+
+cdef extern from "libavformat/avformat.h" nogil:
+
+ cdef int avformat_version()
+ cdef char* avformat_configuration()
+ cdef char* avformat_license()
+ cdef void avformat_network_init()
+
+ cdef int64_t INT64_MIN
+
+ cdef int AV_TIME_BASE
+ cdef int AVSEEK_FLAG_BACKWARD
+ cdef int AVSEEK_FLAG_BYTE
+ cdef int AVSEEK_FLAG_ANY
+ cdef int AVSEEK_FLAG_FRAME
+
+ cdef int AVIO_FLAG_WRITE
+
+ cdef enum AVMediaType:
+ AVMEDIA_TYPE_UNKNOWN
+ AVMEDIA_TYPE_VIDEO
+ AVMEDIA_TYPE_AUDIO
+ AVMEDIA_TYPE_DATA
+ AVMEDIA_TYPE_SUBTITLE
+ AVMEDIA_TYPE_ATTACHMENT
+ AVMEDIA_TYPE_NB
+
+ cdef struct AVStream:
+ int index
+ int id
+
+ AVCodecParameters *codecpar
+
+ AVRational time_base
+
+ int64_t start_time
+ int64_t duration
+ int64_t nb_frames
+ int64_t cur_dts
+
+ AVDictionary *metadata
+
+ AVRational avg_frame_rate
+ AVRational r_frame_rate
+ AVRational sample_aspect_ratio
+
+ # http://ffmpeg.org/doxygen/trunk/structAVIOContext.html
+ cdef struct AVIOContext:
+ unsigned char* buffer
+ int buffer_size
+ int write_flag
+ int direct
+ int seekable
+ int max_packet_size
+ void *opaque
+
+ # http://ffmpeg.org/doxygen/trunk/structAVIOInterruptCB.html
+ cdef struct AVIOInterruptCB:
+ int (*callback)(void*)
+ void *opaque
+
+ cdef int AVIO_FLAG_DIRECT
+ cdef int AVIO_SEEKABLE_NORMAL
+
+ cdef int SEEK_SET
+ cdef int SEEK_CUR
+ cdef int SEEK_END
+ cdef int AVSEEK_SIZE
+
+ cdef AVIOContext* avio_alloc_context(
+ unsigned char *buffer,
+ int buffer_size,
+ int write_flag,
+ void *opaque,
+ int(*read_packet)(void *opaque, uint8_t *buf, int buf_size),
+ int(*write_packet)(void *opaque, const uint8_t *buf, int buf_size),
+ int64_t(*seek)(void *opaque, int64_t offset, int whence)
+ )
+
+ # http://ffmpeg.org/doxygen/trunk/structAVInputFormat.html
+ cdef struct AVInputFormat:
+ const char *name
+ const char *long_name
+ const char *extensions
+ int flags
+ # const AVCodecTag* const *codec_tag
+ const AVClass *priv_class
+
+ cdef struct AVProbeData:
+ unsigned char *buf
+ int buf_size
+ const char *filename
+
+ cdef AVInputFormat* av_probe_input_format(
+ AVProbeData *pd,
+ int is_opened
+ )
+
+ # http://ffmpeg.org/doxygen/trunk/structAVOutputFormat.html
+ cdef struct AVOutputFormat:
+ const char *name
+ const char *long_name
+ const char *extensions
+ AVCodecID video_codec
+ AVCodecID audio_codec
+ AVCodecID subtitle_codec
+ int flags
+ # const AVCodecTag* const *codec_tag
+ const AVClass *priv_class
+
+ int avformat_query_codec(const AVOutputFormat *oformat, AVCodecID codec_id, int std_compliance)
+
+ # AVInputFormat.flags and AVOutputFormat.flags
+ cdef enum:
+ AVFMT_NOFILE
+ AVFMT_NEEDNUMBER
+ AVFMT_SHOW_IDS
+ AVFMT_GLOBALHEADER
+ AVFMT_NOTIMESTAMPS
+ AVFMT_GENERIC_INDEX
+ AVFMT_TS_DISCONT
+ AVFMT_VARIABLE_FPS
+ AVFMT_NODIMENSIONS
+ AVFMT_NOSTREAMS
+ AVFMT_NOBINSEARCH
+ AVFMT_NOGENSEARCH
+ AVFMT_NO_BYTE_SEEK
+ AVFMT_ALLOW_FLUSH
+ AVFMT_TS_NONSTRICT
+ AVFMT_TS_NEGATIVE
+ AVFMT_SEEK_TO_PTS
+
+ # AVFormatContext.flags
+ cdef enum:
+ AVFMT_FLAG_GENPTS
+ AVFMT_FLAG_IGNIDX
+ AVFMT_FLAG_NONBLOCK
+ AVFMT_FLAG_IGNDTS
+ AVFMT_FLAG_NOFILLIN
+ AVFMT_FLAG_NOPARSE
+ AVFMT_FLAG_NOBUFFER
+ AVFMT_FLAG_CUSTOM_IO
+ AVFMT_FLAG_DISCARD_CORRUPT
+ AVFMT_FLAG_FLUSH_PACKETS
+ AVFMT_FLAG_BITEXACT
+ AVFMT_FLAG_SORT_DTS
+ AVFMT_FLAG_FAST_SEEK
+ AVFMT_FLAG_SHORTEST
+ AVFMT_FLAG_AUTO_BSF
+
+ cdef int av_probe_input_buffer(
+ AVIOContext *pb,
+ AVInputFormat **fmt,
+ const char *filename,
+ void *logctx,
+ unsigned int offset,
+ unsigned int max_probe_size
+ )
+
+ cdef int av_find_best_stream(
+ AVFormatContext *ic,
+ AVMediaType type,
+ int wanted_stream_nb,
+ int related_stream,
+ AVCodec **decoder_ret,
+ int flags
+ )
+
+ cdef AVInputFormat* av_find_input_format(const char *name)
+
+ # http://ffmpeg.org/doxygen/trunk/structAVFormatContext.html
+ cdef struct AVFormatContext:
+
+ # Streams.
+ unsigned int nb_streams
+ AVStream **streams
+
+ AVInputFormat *iformat
+ AVOutputFormat *oformat
+
+ AVIOContext *pb
+ AVIOInterruptCB interrupt_callback
+
+ AVDictionary *metadata
+
+ char filename
+ int64_t start_time
+ int64_t duration
+ int bit_rate
+
+ int flags
+ int64_t max_analyze_duration
+
+ void *opaque
+
+ int (*io_open)(
+ AVFormatContext *s,
+ AVIOContext **pb,
+ const char *url,
+ int flags,
+ AVDictionary **options
+ )
+ int (*io_close2)(
+ AVFormatContext *s,
+ AVIOContext *pb
+ )
+
+ cdef AVFormatContext* avformat_alloc_context()
+
+ # .. c:function:: avformat_open_input(...)
+ #
+ # Options are passed via :func:`av.open`.
+ #
+ # .. seealso:: FFmpeg's docs: :ffmpeg:`avformat_open_input`
+ #
+ cdef int avformat_open_input(
+ AVFormatContext **ctx, # NULL will allocate for you.
+ char *filename,
+ AVInputFormat *format, # Can be NULL.
+ AVDictionary **options # Can be NULL.
+ )
+
+ cdef int avformat_close_input(AVFormatContext **ctx)
+
+ # .. c:function:: avformat_write_header(...)
+ #
+ # Options are passed via :func:`av.open`; called in
+ # :meth:`av.container.OutputContainer.start_encoding`.
+ #
+ # .. seealso:: FFmpeg's docs: :ffmpeg:`avformat_write_header`
+ #
+ cdef int avformat_write_header(
+ AVFormatContext *ctx,
+ AVDictionary **options # Can be NULL
+ )
+
+ cdef int av_write_trailer(AVFormatContext *ctx)
+
+ cdef int av_interleaved_write_frame(
+ AVFormatContext *ctx,
+ AVPacket *pkt
+ )
+
+ cdef int av_write_frame(
+ AVFormatContext *ctx,
+ AVPacket *pkt
+ )
+
+ cdef int avio_open(
+ AVIOContext **s,
+ char *url,
+ int flags
+ )
+
+ cdef int64_t avio_size(
+ AVIOContext *s
+ )
+
+ cdef AVOutputFormat* av_guess_format(
+ char *short_name,
+ char *filename,
+ char *mime_type
+ )
+
+ cdef int avformat_query_codec(
+ AVOutputFormat *ofmt,
+ AVCodecID codec_id,
+ int std_compliance
+ )
+
+ cdef void avio_flush(AVIOContext *s)
+
+ cdef int avio_close(AVIOContext *s)
+
+ cdef int avio_closep(AVIOContext **s)
+
+ cdef int avformat_find_stream_info(
+ AVFormatContext *ctx,
+ AVDictionary **options, # Can be NULL.
+ )
+
+ cdef AVStream* avformat_new_stream(
+ AVFormatContext *ctx,
+ AVCodec *c
+ )
+
+ cdef int avformat_alloc_output_context2(
+ AVFormatContext **ctx,
+ AVOutputFormat *oformat,
+ char *format_name,
+ char *filename
+ )
+
+ cdef int avformat_free_context(AVFormatContext *ctx)
+
+ cdef AVClass* avformat_get_class()
+
+ cdef void av_dump_format(
+ AVFormatContext *ctx,
+ int index,
+ char *url,
+ int is_output,
+ )
+
+ cdef int av_read_frame(
+ AVFormatContext *ctx,
+ AVPacket *packet,
+ )
+
+ cdef int av_seek_frame(
+ AVFormatContext *ctx,
+ int stream_index,
+ int64_t timestamp,
+ int flags
+ )
+
+ cdef int avformat_seek_file(
+ AVFormatContext *ctx,
+ int stream_index,
+ int64_t min_ts,
+ int64_t ts,
+ int64_t max_ts,
+ int flags
+ )
+
+ cdef AVRational av_guess_frame_rate(
+ AVFormatContext *ctx,
+ AVStream *stream,
+ AVFrame *frame
+ )
+
+ cdef AVRational av_guess_sample_aspect_ratio(
+ AVFormatContext *ctx,
+ AVStream *stream,
+ AVFrame *frame
+ )
+
+ cdef const AVInputFormat* av_demuxer_iterate(void **opaque)
+ cdef const AVOutputFormat* av_muxer_iterate(void **opaque)
+
+ # custom
+
+ cdef set pyav_get_available_formats()
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/avutil.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/avutil.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..ed281aeafd55a69ae7d6e0753ea6f1f8e3eb6121
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/avutil.pxd
@@ -0,0 +1,390 @@
+from libc.stdint cimport int64_t, uint8_t, uint64_t, int32_t
+
+
+cdef extern from "libavutil/mathematics.h" nogil:
+ pass
+
+cdef extern from "libavutil/rational.h" nogil:
+ cdef int av_reduce(int *dst_num, int *dst_den, int64_t num, int64_t den, int64_t max)
+
+cdef extern from "libavutil/avutil.h" nogil:
+
+ cdef const char* av_version_info()
+ cdef int avutil_version()
+ cdef char* avutil_configuration()
+ cdef char* avutil_license()
+
+ cdef enum AVPictureType:
+ AV_PICTURE_TYPE_NONE
+ AV_PICTURE_TYPE_I
+ AV_PICTURE_TYPE_P
+ AV_PICTURE_TYPE_B
+ AV_PICTURE_TYPE_S
+ AV_PICTURE_TYPE_SI
+ AV_PICTURE_TYPE_SP
+ AV_PICTURE_TYPE_BI
+
+ cdef enum AVPixelFormat:
+ AV_PIX_FMT_NONE
+ AV_PIX_FMT_YUV420P
+ AV_PIX_FMT_RGB24
+ PIX_FMT_RGB24
+ PIX_FMT_RGBA
+
+ cdef enum AVRounding:
+ AV_ROUND_ZERO
+ AV_ROUND_INF
+ AV_ROUND_DOWN
+ AV_ROUND_UP
+ AV_ROUND_NEAR_INF
+ # This is nice, but only in FFMpeg:
+ # AV_ROUND_PASS_MINMAX
+
+ cdef enum AVColorSpace:
+ AVCOL_SPC_RGB
+ AVCOL_SPC_BT709
+ AVCOL_SPC_UNSPECIFIED
+ AVCOL_SPC_RESERVED
+ AVCOL_SPC_FCC
+ AVCOL_SPC_BT470BG
+ AVCOL_SPC_SMPTE170M
+ AVCOL_SPC_SMPTE240M
+ AVCOL_SPC_YCOCG
+ AVCOL_SPC_BT2020_NCL
+ AVCOL_SPC_BT2020_CL
+ AVCOL_SPC_NB
+
+ cdef enum AVColorRange:
+ AVCOL_RANGE_UNSPECIFIED
+ AVCOL_RANGE_MPEG
+ AVCOL_RANGE_JPEG
+ AVCOL_RANGE_NB
+
+ cdef enum AVColorPrimaries:
+ AVCOL_PRI_RESERVED0
+ AVCOL_PRI_BT709
+ AVCOL_PRI_UNSPECIFIED
+ AVCOL_PRI_RESERVED
+ AVCOL_PRI_BT470M
+ AVCOL_PRI_BT470BG
+ AVCOL_PRI_SMPTE170M
+ AVCOL_PRI_SMPTE240M
+ AVCOL_PRI_FILM
+ AVCOL_PRI_BT2020
+ AVCOL_PRI_SMPTE428
+ AVCOL_PRI_SMPTEST428_1
+ AVCOL_PRI_SMPTE431
+ AVCOL_PRI_SMPTE432
+ AVCOL_PRI_EBU3213
+ AVCOL_PRI_JEDEC_P22
+ AVCOL_PRI_NB
+
+ cdef enum AVColorTransferCharacteristic:
+ AVCOL_TRC_RESERVED0
+ AVCOL_TRC_BT709
+ AVCOL_TRC_UNSPECIFIED
+ AVCOL_TRC_RESERVED
+ AVCOL_TRC_GAMMA22
+ AVCOL_TRC_GAMMA28
+ AVCOL_TRC_SMPTE170M
+ AVCOL_TRC_SMPTE240M
+ AVCOL_TRC_LINEAR
+ AVCOL_TRC_LOG
+ AVCOL_TRC_LOG_SQRT
+ AVCOL_TRC_IEC61966_2_4
+ AVCOL_TRC_BT1361_ECG
+ AVCOL_TRC_IEC61966_2_1
+ AVCOL_TRC_BT2020_10
+ AVCOL_TRC_BT2020_12
+ AVCOL_TRC_SMPTE2084
+ AVCOL_TRC_SMPTEST2084
+ AVCOL_TRC_SMPTE428
+ AVCOL_TRC_SMPTEST428_1
+ AVCOL_TRC_ARIB_STD_B67
+ AVCOL_TRC_NB
+
+ cdef double M_PI
+
+ cdef void* av_malloc(size_t size)
+ cdef void *av_calloc(size_t nmemb, size_t size)
+ cdef void *av_realloc(void *ptr, size_t size)
+
+ cdef void av_freep(void *ptr)
+
+ cdef int av_get_bytes_per_sample(AVSampleFormat sample_fmt)
+
+ cdef int av_samples_get_buffer_size(
+ int *linesize,
+ int nb_channels,
+ int nb_samples,
+ AVSampleFormat sample_fmt,
+ int align
+ )
+
+ # See: http://ffmpeg.org/doxygen/trunk/structAVRational.html
+ ctypedef struct AVRational:
+ int num
+ int den
+
+ cdef AVRational AV_TIME_BASE_Q
+
+ # Rescales from one time base to another
+ cdef int64_t av_rescale_q(
+ int64_t a, # time stamp
+ AVRational bq, # source time base
+ AVRational cq # target time base
+ )
+
+ # Rescale a 64-bit integer with specified rounding.
+ # A simple a*b/c isn't possible as it can overflow
+ cdef int64_t av_rescale_rnd(
+ int64_t a,
+ int64_t b,
+ int64_t c,
+ int r # should be AVRounding, but then we can't use bitwise logic.
+ )
+
+ cdef int64_t av_rescale_q_rnd(
+ int64_t a,
+ AVRational bq,
+ AVRational cq,
+ int r # should be AVRounding, but then we can't use bitwise logic.
+ )
+
+ cdef int64_t av_rescale(
+ int64_t a,
+ int64_t b,
+ int64_t c
+ )
+
+ cdef char* av_strdup(char *s)
+
+ cdef int av_opt_set_int(
+ void *obj,
+ char *name,
+ int64_t value,
+ int search_flags
+ )
+
+ cdef const char* av_get_media_type_string(AVMediaType media_type)
+
+cdef extern from "libavutil/pixdesc.h" nogil:
+
+ # See: http://ffmpeg.org/doxygen/trunk/structAVComponentDescriptor.html
+ cdef struct AVComponentDescriptor:
+ unsigned int plane
+ unsigned int step
+ unsigned int offset
+ unsigned int shift
+ unsigned int depth
+
+ cdef enum AVPixFmtFlags:
+ AV_PIX_FMT_FLAG_BE
+ AV_PIX_FMT_FLAG_PAL
+ AV_PIX_FMT_FLAG_BITSTREAM
+ AV_PIX_FMT_FLAG_HWACCEL
+ AV_PIX_FMT_FLAG_PLANAR
+ AV_PIX_FMT_FLAG_RGB
+ AV_PIX_FMT_FLAG_PSEUDOPAL
+ AV_PIX_FMT_FLAG_ALPHA
+ AV_PIX_FMT_FLAG_BAYER
+ AV_PIX_FMT_FLAG_FLOAT
+
+ # See: http://ffmpeg.org/doxygen/trunk/structAVPixFmtDescriptor.html
+ cdef struct AVPixFmtDescriptor:
+ const char *name
+ uint8_t nb_components
+ uint8_t log2_chroma_w
+ uint8_t log2_chroma_h
+ uint8_t flags
+ AVComponentDescriptor comp[4]
+
+ cdef AVPixFmtDescriptor* av_pix_fmt_desc_get(AVPixelFormat pix_fmt)
+ cdef AVPixFmtDescriptor* av_pix_fmt_desc_next(AVPixFmtDescriptor *prev)
+
+ cdef char * av_get_pix_fmt_name(AVPixelFormat pix_fmt)
+ cdef AVPixelFormat av_get_pix_fmt(char* name)
+
+ int av_get_bits_per_pixel(AVPixFmtDescriptor *pixdesc)
+ int av_get_padded_bits_per_pixel(AVPixFmtDescriptor *pixdesc)
+
+
+cdef extern from "libavutil/channel_layout.h" nogil:
+
+ # Layouts.
+ cdef uint64_t av_get_channel_layout(char* name)
+ cdef int av_get_channel_layout_nb_channels(uint64_t channel_layout)
+ cdef int64_t av_get_default_channel_layout(int nb_channels)
+
+ # Channels.
+ cdef uint64_t av_channel_layout_extract_channel(uint64_t layout, int index)
+ cdef char* av_get_channel_name(uint64_t channel)
+ cdef char* av_get_channel_description(uint64_t channel)
+
+
+cdef extern from "libavutil/audio_fifo.h" nogil:
+
+ cdef struct AVAudioFifo:
+ pass
+
+ cdef void av_audio_fifo_free(AVAudioFifo *af)
+
+ cdef AVAudioFifo* av_audio_fifo_alloc(
+ AVSampleFormat sample_fmt,
+ int channels,
+ int nb_samples
+ )
+
+ cdef int av_audio_fifo_write(
+ AVAudioFifo *af,
+ void **data,
+ int nb_samples
+ )
+
+ cdef int av_audio_fifo_read(
+ AVAudioFifo *af,
+ void **data,
+ int nb_samples
+ )
+
+ cdef int av_audio_fifo_size(AVAudioFifo *af)
+ cdef int av_audio_fifo_space (AVAudioFifo *af)
+
+
+cdef extern from "stdarg.h" nogil:
+ # For logging. Should really be in another PXD.
+ ctypedef struct va_list:
+ pass
+
+
+cdef extern from "Python.h" nogil:
+ # For logging. See av/logging.pyx for an explanation.
+ cdef int Py_AddPendingCall(void *, void *)
+ void PyErr_PrintEx(int set_sys_last_vars)
+ int Py_IsInitialized()
+ void PyErr_Display(object, object, object)
+
+
+cdef extern from "libavutil/opt.h" nogil:
+ cdef enum AVOptionType:
+ AV_OPT_TYPE_FLAGS
+ AV_OPT_TYPE_INT
+ AV_OPT_TYPE_INT64
+ AV_OPT_TYPE_DOUBLE
+ AV_OPT_TYPE_FLOAT
+ AV_OPT_TYPE_STRING
+ AV_OPT_TYPE_RATIONAL
+ AV_OPT_TYPE_BINARY
+ AV_OPT_TYPE_DICT
+ AV_OPT_TYPE_UINT64
+ AV_OPT_TYPE_CONST
+ AV_OPT_TYPE_IMAGE_SIZE
+ AV_OPT_TYPE_PIXEL_FMT
+ AV_OPT_TYPE_SAMPLE_FMT
+ AV_OPT_TYPE_VIDEO_RATE
+ AV_OPT_TYPE_DURATION
+ AV_OPT_TYPE_COLOR
+ AV_OPT_TYPE_CHLAYOUT
+ AV_OPT_TYPE_BOOL
+
+ cdef struct AVOption_default_val:
+ int64_t i64
+ double dbl
+ const char *str
+ AVRational q
+
+ cdef enum:
+ AV_OPT_FLAG_ENCODING_PARAM
+ AV_OPT_FLAG_DECODING_PARAM
+ AV_OPT_FLAG_AUDIO_PARAM
+ AV_OPT_FLAG_VIDEO_PARAM
+ AV_OPT_FLAG_SUBTITLE_PARAM
+ AV_OPT_FLAG_EXPORT
+ AV_OPT_FLAG_READONLY
+ AV_OPT_FLAG_FILTERING_PARAM
+
+ cdef struct AVOption:
+
+ const char *name
+ const char *help
+ AVOptionType type
+ int offset
+
+ AVOption_default_val default_val
+
+ double min
+ double max
+ int flags
+ const char *unit
+
+
+cdef extern from "libavutil/imgutils.h" nogil:
+
+ cdef int av_image_alloc(
+ uint8_t *pointers[4],
+ int linesizes[4],
+ int width,
+ int height,
+ AVPixelFormat pix_fmt,
+ int align
+ )
+ cdef int av_image_fill_pointers(
+ uint8_t *pointers[4],
+ AVPixelFormat pix_fmt,
+ int height,
+ uint8_t *ptr,
+ const int linesizes[4]
+ )
+ cdef int av_image_fill_linesizes(
+ int linesizes[4],
+ AVPixelFormat pix_fmt,
+ int width,
+ )
+
+
+cdef extern from "libavutil/log.h" nogil:
+
+ cdef enum AVClassCategory:
+ AV_CLASS_CATEGORY_NA
+ AV_CLASS_CATEGORY_INPUT
+ AV_CLASS_CATEGORY_OUTPUT
+ AV_CLASS_CATEGORY_MUXER
+ AV_CLASS_CATEGORY_DEMUXER
+ AV_CLASS_CATEGORY_ENCODER
+ AV_CLASS_CATEGORY_DECODER
+ AV_CLASS_CATEGORY_FILTER
+ AV_CLASS_CATEGORY_BITSTREAM_FILTER
+ AV_CLASS_CATEGORY_SWSCALER
+ AV_CLASS_CATEGORY_SWRESAMPLER
+ AV_CLASS_CATEGORY_NB
+
+ cdef struct AVClass:
+
+ const char *class_name
+ const char *(*item_name)(void*) nogil
+
+ AVClassCategory category
+ int parent_log_context_offset
+
+ const AVOption *option
+
+ cdef enum:
+ AV_LOG_QUIET
+ AV_LOG_PANIC
+ AV_LOG_FATAL
+ AV_LOG_ERROR
+ AV_LOG_WARNING
+ AV_LOG_INFO
+ AV_LOG_VERBOSE
+ AV_LOG_DEBUG
+ AV_LOG_TRACE
+ AV_LOG_MAX_OFFSET
+
+ # Send a log.
+ void av_log(void *ptr, int level, const char *fmt, ...)
+
+ # Get the logs.
+ ctypedef void(*av_log_callback)(void *, int, const char *, va_list)
+ void av_log_default_callback(void *, int, const char *, va_list)
+ void av_log_set_callback (av_log_callback callback)
+ void av_log_set_level(int level)
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/buffer.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/buffer.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..daf86105b347c3a47b94f48be5fc63dcecabaa18
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/buffer.pxd
@@ -0,0 +1,9 @@
+from libc.stdint cimport uint8_t
+
+cdef extern from "libavutil/buffer.h" nogil:
+
+ AVBufferRef *av_buffer_create(uint8_t *data, size_t size, void (*free)(void *opaque, uint8_t *data), void *opaque, int flags)
+ void av_buffer_unref(AVBufferRef **buf)
+
+ cdef struct AVBufferRef:
+ uint8_t *data
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/channel_layout.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/channel_layout.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..1459fbd22b64319e060e9035f5739b65c893bc0e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/channel_layout.pxd
@@ -0,0 +1,11 @@
+cdef extern from "libavutil/channel_layout.h" nogil:
+
+ # This is not a comprehensive list.
+ cdef uint64_t AV_CH_LAYOUT_MONO
+ cdef uint64_t AV_CH_LAYOUT_STEREO
+ cdef uint64_t AV_CH_LAYOUT_2POINT1
+ cdef uint64_t AV_CH_LAYOUT_4POINT0
+ cdef uint64_t AV_CH_LAYOUT_5POINT0_BACK
+ cdef uint64_t AV_CH_LAYOUT_5POINT1_BACK
+ cdef uint64_t AV_CH_LAYOUT_6POINT1
+ cdef uint64_t AV_CH_LAYOUT_7POINT1
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/dict.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/dict.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..c2b141101d93446dc4f121955363c85cc33f6341
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/dict.pxd
@@ -0,0 +1,38 @@
+cdef extern from "libavutil/dict.h" nogil:
+
+ # See: http://ffmpeg.org/doxygen/trunk/structAVDictionary.html
+ ctypedef struct AVDictionary:
+ pass
+
+ cdef void av_dict_free(AVDictionary **)
+
+ # See: http://ffmpeg.org/doxygen/trunk/structAVDictionaryEntry.html
+ ctypedef struct AVDictionaryEntry:
+ char *key
+ char *value
+
+ cdef int AV_DICT_IGNORE_SUFFIX
+
+ cdef AVDictionaryEntry* av_dict_get(
+ AVDictionary *dict,
+ char *key,
+ AVDictionaryEntry *prev,
+ int flags,
+ )
+
+ cdef int av_dict_set(
+ AVDictionary **pm,
+ const char *key,
+ const char *value,
+ int flags
+ )
+
+ cdef int av_dict_count(
+ AVDictionary *m
+ )
+
+ cdef int av_dict_copy(
+ AVDictionary **dst,
+ AVDictionary *src,
+ int flags
+ )
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/error.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/error.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..2122772e316d4beb94b6a4672e8cb3d27d250064
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/error.pxd
@@ -0,0 +1,43 @@
+cdef extern from "libavutil/error.h" nogil:
+
+ # Not actually from here, but whatever.
+ cdef int ENOMEM
+ cdef int EAGAIN
+
+ cdef int AVERROR_BSF_NOT_FOUND
+ cdef int AVERROR_BUG
+ cdef int AVERROR_BUFFER_TOO_SMALL
+ cdef int AVERROR_DECODER_NOT_FOUND
+ cdef int AVERROR_DEMUXER_NOT_FOUND
+ cdef int AVERROR_ENCODER_NOT_FOUND
+ cdef int AVERROR_EOF
+ cdef int AVERROR_EXIT
+ cdef int AVERROR_EXTERNAL
+ cdef int AVERROR_FILTER_NOT_FOUND
+ cdef int AVERROR_INVALIDDATA
+ cdef int AVERROR_MUXER_NOT_FOUND
+ cdef int AVERROR_OPTION_NOT_FOUND
+ cdef int AVERROR_PATCHWELCOME
+ cdef int AVERROR_PROTOCOL_NOT_FOUND
+ cdef int AVERROR_UNKNOWN
+ cdef int AVERROR_EXPERIMENTAL
+ cdef int AVERROR_INPUT_CHANGED
+ cdef int AVERROR_OUTPUT_CHANGED
+
+ cdef int AVERROR_HTTP_BAD_REQUEST
+ cdef int AVERROR_HTTP_UNAUTHORIZED
+ cdef int AVERROR_HTTP_FORBIDDEN
+ cdef int AVERROR_HTTP_NOT_FOUND
+ cdef int AVERROR_HTTP_OTHER_4XX
+ cdef int AVERROR_HTTP_SERVER_ERROR
+
+ cdef int AVERROR_NOMEM "AVERROR(ENOMEM)"
+
+ # cdef int FFERRTAG(int, int, int, int)
+
+ cdef int AVERROR(int error)
+
+ cdef int AV_ERROR_MAX_STRING_SIZE
+
+ cdef int av_strerror(int errno, char *output, size_t output_size)
+ cdef char* av_err2str(int errnum)
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/frame.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/frame.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..aa8dc3a006aac923c0c815850d3d9781dca0f942
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/frame.pxd
@@ -0,0 +1,14 @@
+cdef extern from "libavutil/frame.h" nogil:
+
+ cdef AVFrame* av_frame_alloc()
+ cdef void av_frame_free(AVFrame**)
+ cdef int av_frame_ref(AVFrame *dst, const AVFrame *src)
+ cdef AVFrame* av_frame_clone(const AVFrame *src)
+ cdef void av_frame_unref(AVFrame *frame)
+ cdef void av_frame_move_ref(AVFrame *dst, AVFrame *src)
+ cdef int av_frame_get_buffer(AVFrame *frame, int align)
+ cdef int av_frame_is_writable(AVFrame *frame)
+ cdef int av_frame_make_writable(AVFrame *frame)
+ cdef int av_frame_copy(AVFrame *dst, const AVFrame *src)
+ cdef int av_frame_copy_props(AVFrame *dst, const AVFrame *src)
+ cdef AVFrameSideData* av_frame_get_side_data(AVFrame *frame, AVFrameSideDataType type)
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/motion_vector.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/motion_vector.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..457d7149e686435feae490cc5a24b7dba2d1beb4
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/motion_vector.pxd
@@ -0,0 +1,17 @@
+from libc.stdint cimport int16_t, int32_t, uint8_t, uint16_t, uint64_t
+
+
+cdef extern from "libavutil/motion_vector.h" nogil:
+
+ cdef struct AVMotionVector:
+ int32_t source
+ uint8_t w
+ uint8_t h
+ int16_t src_x
+ int16_t src_y
+ int16_t dst_x
+ int16_t dst_y
+ uint64_t flags
+ int32_t motion_x
+ int32_t motion_y
+ uint16_t motion_scale
diff --git a/venv/lib/python3.11/site-packages/av/include/libavutil/samplefmt.pxd b/venv/lib/python3.11/site-packages/av/include/libavutil/samplefmt.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..a26c6ecfdc8c9e47c936902a56f2a817dab39ff0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libavutil/samplefmt.pxd
@@ -0,0 +1,62 @@
+cdef extern from "libavutil/samplefmt.h" nogil:
+
+ cdef enum AVSampleFormat:
+ AV_SAMPLE_FMT_NONE
+ AV_SAMPLE_FMT_U8
+ AV_SAMPLE_FMT_S16
+ AV_SAMPLE_FMT_S32
+ AV_SAMPLE_FMT_FLT
+ AV_SAMPLE_FMT_DBL
+ AV_SAMPLE_FMT_U8P
+ AV_SAMPLE_FMT_S16P
+ AV_SAMPLE_FMT_S32P
+ AV_SAMPLE_FMT_FLTP
+ AV_SAMPLE_FMT_DBLP
+ AV_SAMPLE_FMT_NB # Number.
+
+ # Find by name.
+ cdef AVSampleFormat av_get_sample_fmt(char* name)
+
+ # Inspection.
+ cdef char * av_get_sample_fmt_name(AVSampleFormat sample_fmt)
+ cdef int av_get_bytes_per_sample(AVSampleFormat sample_fmt)
+ cdef int av_sample_fmt_is_planar(AVSampleFormat sample_fmt)
+
+ # Alternative forms.
+ cdef AVSampleFormat av_get_packed_sample_fmt(AVSampleFormat sample_fmt)
+ cdef AVSampleFormat av_get_planar_sample_fmt(AVSampleFormat sample_fmt)
+
+ cdef int av_samples_alloc(
+ uint8_t** audio_data,
+ int* linesize,
+ int nb_channels,
+ int nb_samples,
+ AVSampleFormat sample_fmt,
+ int align
+ )
+
+ cdef int av_samples_get_buffer_size(
+ int *linesize,
+ int nb_channels,
+ int nb_samples,
+ AVSampleFormat sample_fmt,
+ int align
+ )
+
+ cdef int av_samples_fill_arrays(
+ uint8_t **audio_data,
+ int *linesize,
+ const uint8_t *buf,
+ int nb_channels,
+ int nb_samples,
+ AVSampleFormat sample_fmt,
+ int align
+ )
+
+ cdef int av_samples_set_silence(
+ uint8_t **audio_data,
+ int offset,
+ int nb_samples,
+ int nb_channels,
+ AVSampleFormat sample_fmt
+ )
diff --git a/venv/lib/python3.11/site-packages/av/include/libswresample/swresample.pxd b/venv/lib/python3.11/site-packages/av/include/libswresample/swresample.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..65b8314df537f78443aa04fa75aab6ac57fe987e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libswresample/swresample.pxd
@@ -0,0 +1,39 @@
+from libc.stdint cimport int64_t, uint8_t
+
+
+cdef extern from "libswresample/swresample.h" nogil:
+
+ cdef int swresample_version()
+ cdef char* swresample_configuration()
+ cdef char* swresample_license()
+
+ cdef struct SwrContext:
+ pass
+
+ cdef SwrContext* swr_alloc_set_opts(
+ SwrContext *ctx,
+ int64_t out_ch_layout,
+ AVSampleFormat out_sample_fmt,
+ int out_sample_rate,
+ int64_t in_ch_layout,
+ AVSampleFormat in_sample_fmt,
+ int in_sample_rate,
+ int log_offset,
+ void *log_ctx # logging context, can be NULL
+ )
+
+ cdef int swr_convert(
+ SwrContext *ctx,
+ uint8_t ** out_buffer,
+ int out_count,
+ uint8_t **in_buffer,
+ int in_count
+ )
+ # Gets the delay the next input sample will
+ # experience relative to the next output sample.
+ cdef int64_t swr_get_delay(SwrContext *s, int64_t base)
+
+ cdef SwrContext* swr_alloc()
+ cdef int swr_init(SwrContext* ctx)
+ cdef void swr_free(SwrContext **ctx)
+ cdef void swr_close(SwrContext *ctx)
diff --git a/venv/lib/python3.11/site-packages/av/include/libswscale/swscale.pxd b/venv/lib/python3.11/site-packages/av/include/libswscale/swscale.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..af7f7d134843b6cc25c4bc6b2cf6e18439634fe7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/include/libswscale/swscale.pxd
@@ -0,0 +1,98 @@
+
+cdef extern from "libswscale/swscale.h" nogil:
+
+ cdef int swscale_version()
+ cdef char* swscale_configuration()
+ cdef char* swscale_license()
+
+ # See: http://ffmpeg.org/doxygen/trunk/structSwsContext.html
+ cdef struct SwsContext:
+ pass
+
+ # See: http://ffmpeg.org/doxygen/trunk/structSwsFilter.html
+ cdef struct SwsFilter:
+ pass
+
+ # Flags.
+ cdef int SWS_FAST_BILINEAR
+ cdef int SWS_BILINEAR
+ cdef int SWS_BICUBIC
+ cdef int SWS_X
+ cdef int SWS_POINT
+ cdef int SWS_AREA
+ cdef int SWS_BICUBLIN
+ cdef int SWS_GAUSS
+ cdef int SWS_SINC
+ cdef int SWS_LANCZOS
+ cdef int SWS_SPLINE
+
+ cdef int SWS_CS_ITU709
+ cdef int SWS_CS_FCC
+ cdef int SWS_CS_ITU601
+ cdef int SWS_CS_ITU624
+ cdef int SWS_CS_SMPTE170M
+ cdef int SWS_CS_SMPTE240M
+ cdef int SWS_CS_DEFAULT
+
+ cdef SwsContext* sws_getContext(
+ int src_width,
+ int src_height,
+ AVPixelFormat src_format,
+ int dst_width,
+ int dst_height,
+ AVPixelFormat dst_format,
+ int flags,
+ SwsFilter *src_filter,
+ SwsFilter *dst_filter,
+ double *param,
+ )
+
+ cdef int sws_scale(
+ SwsContext *ctx,
+ unsigned char **src_slice,
+ int *src_stride,
+ int src_slice_y,
+ int src_slice_h,
+ unsigned char **dst_slice,
+ int *dst_stride,
+ )
+
+ cdef void sws_freeContext(SwsContext *ctx)
+
+ cdef SwsContext *sws_getCachedContext(
+ SwsContext *context,
+ int src_width,
+ int src_height,
+ AVPixelFormat src_format,
+ int dst_width,
+ int dst_height,
+ AVPixelFormat dst_format,
+ int flags,
+ SwsFilter *src_filter,
+ SwsFilter *dst_filter,
+ double *param,
+ )
+
+ cdef int* sws_getCoefficients(int colorspace)
+
+ cdef int sws_getColorspaceDetails(
+ SwsContext *context,
+ int **inv_table,
+ int *srcRange,
+ int **table,
+ int *dstRange,
+ int *brightness,
+ int *contrast,
+ int *saturation
+ )
+
+ cdef int sws_setColorspaceDetails(
+ SwsContext *context,
+ const int inv_table[4],
+ int srcRange,
+ const int table[4],
+ int dstRange,
+ int brightness,
+ int contrast,
+ int saturation
+ )
diff --git a/venv/lib/python3.11/site-packages/av/logging.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/logging.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..940e81229e2e365cca9ce293e7e8ea09f0a1510e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/logging.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:60a6e2c1e8b9c913f809fe7cf5034113a9acbab301e242b2a9e7dbe28c23ebff
+size 1044313
diff --git a/venv/lib/python3.11/site-packages/av/logging.pxd b/venv/lib/python3.11/site-packages/av/logging.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..a886a0f206e07928b36ff0b4c3799409446e27bd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/logging.pxd
@@ -0,0 +1,2 @@
+
+cpdef get_last_error()
diff --git a/venv/lib/python3.11/site-packages/av/logging.pyi b/venv/lib/python3.11/site-packages/av/logging.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..8c32de77d21c75ca4f38fbd78ecb5d7c2871c2d6
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/logging.pyi
@@ -0,0 +1,33 @@
+from typing import Any, Callable
+
+PANIC: int
+FATAL: int
+ERROR: int
+WARNING: int
+INFO: int
+VERBOSE: int
+DEBUG: int
+TRACE: int
+CRITICAL: int
+
+def adapt_level(level: int) -> int: ...
+def get_level() -> int | None: ...
+def set_level(level: int | None) -> None: ...
+def set_libav_level(level: int) -> None: ...
+def restore_default_callback() -> None: ...
+def get_skip_repeated() -> bool: ...
+def set_skip_repeated(v: bool) -> None: ...
+def get_last_error() -> tuple[int, tuple[int, str, str] | None]: ...
+def log(level: int, name: str, message: str) -> None: ...
+
+class Capture:
+ logs: list[tuple[int, str, str]]
+
+ def __init__(self, local: bool = True) -> None: ...
+ def __enter__(self) -> list[tuple[int, str, str]]: ...
+ def __exit__(
+ self,
+ type_: type | None,
+ value: Exception | None,
+ traceback: Callable[..., Any] | None,
+ ) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/opaque.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/opaque.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..2309e76646d7c28be239330c1795ab78ae662efc
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/opaque.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:e8226eb0d5ee9069e14ac47d6e97d8cda960de6820b945197fee0a8a5fc195ea
+size 347297
diff --git a/venv/lib/python3.11/site-packages/av/opaque.pxd b/venv/lib/python3.11/site-packages/av/opaque.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..f5c38d7fab9deb147f679b66755902bba63a9603
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/opaque.pxd
@@ -0,0 +1,12 @@
+cimport libav as lib
+
+
+cdef class OpaqueContainer:
+ cdef dict _by_name
+
+ cdef lib.AVBufferRef *add(self, object v)
+ cdef object get(self, bytes name)
+ cdef object pop(self, bytes name)
+
+
+cdef OpaqueContainer opaque_container
diff --git a/venv/lib/python3.11/site-packages/av/option.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/option.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..95ab2888c89773ee5b98552f01ca201d73dd7c77
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/option.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:c9049194622c519ef0cb2e5ea8e5a2fd114b4abeaa423f82a610a5536635449d
+size 642377
diff --git a/venv/lib/python3.11/site-packages/av/option.pxd b/venv/lib/python3.11/site-packages/av/option.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..9087b811c06e3dae2a25ce8557d37e6f7a67c7d3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/option.pxd
@@ -0,0 +1,21 @@
+cimport libav as lib
+
+
+cdef class BaseOption:
+
+ cdef const lib.AVOption *ptr
+
+
+cdef class Option(BaseOption):
+
+ cdef readonly tuple choices
+
+
+cdef class OptionChoice(BaseOption):
+
+ cdef readonly bint is_default
+
+
+cdef Option wrap_option(tuple choices, const lib.AVOption *ptr)
+
+cdef OptionChoice wrap_option_choice(const lib.AVOption *ptr, bint is_default)
diff --git a/venv/lib/python3.11/site-packages/av/option.pyi b/venv/lib/python3.11/site-packages/av/option.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..3132f4a02f10fae20b45dbbc1bd716f7ca00642e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/option.pyi
@@ -0,0 +1,54 @@
+from enum import Enum, Flag
+
+class OptionType(Enum):
+ FLAGS: int
+ INT: int
+ INT64: int
+ DOUBLE: int
+ FLOAT: int
+ STRING: int
+ RATIONAL: int
+ BINARY: int
+ DICT: int
+ CONST: int
+ IMAGE_SIZE: int
+ PIXEL_FMT: int
+ SAMPLE_FMT: int
+ VIDEO_RATE: int
+ DURATION: int
+ COLOR: int
+ CHANNEL_LAYOUT: int
+ BOOL: int
+
+class OptionFlags(Flag):
+ ENCODING_PARAM: int
+ DECODING_PARAM: int
+ AUDIO_PARAM: int
+ VIDEO_PARAM: int
+ SUBTITLE_PARAM: int
+ EXPORT: int
+ READONLY: int
+ FILTERING_PARAM: int
+
+class BaseOption:
+ name: str
+ help: str
+ flags: int
+ is_encoding_param: bool
+ is_decoding_param: bool
+ is_audio_param: bool
+ is_video_param: bool
+ is_subtitle_param: bool
+ is_export: bool
+ is_readonly: bool
+ is_filtering_param: bool
+
+class Option(BaseOption):
+ type: OptionType
+ offset: int
+ default: int
+ min: int
+ max: int
+
+class OptionChoice(BaseOption):
+ value: int
diff --git a/venv/lib/python3.11/site-packages/av/packet.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/packet.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..ec6d0c8f2842347d1b6e5f3925b80b0b69029804
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/packet.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:dd3fe2c2ca3c9198098b631aa37e581ae8713b2188555eec8109805ca3769ba4
+size 597449
diff --git a/venv/lib/python3.11/site-packages/av/packet.pxd b/venv/lib/python3.11/site-packages/av/packet.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..ca21e6b76dcf50c86e9c4a79a4ddf4e64410e04d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/packet.pxd
@@ -0,0 +1,21 @@
+cimport libav as lib
+
+from av.buffer cimport Buffer
+from av.bytesource cimport ByteSource
+from av.stream cimport Stream
+
+
+cdef class Packet(Buffer):
+
+ cdef lib.AVPacket* ptr
+
+ cdef Stream _stream
+
+ # We track our own time.
+ cdef lib.AVRational _time_base
+ cdef _rebase_time(self, lib.AVRational)
+
+ # Hold onto the original reference.
+ cdef ByteSource source
+ cdef size_t _buffer_size(self)
+ cdef void* _buffer_ptr(self)
diff --git a/venv/lib/python3.11/site-packages/av/packet.pyi b/venv/lib/python3.11/site-packages/av/packet.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..9bdbb8c6236faae4fc926941d02c6c5e1b3fd6c4
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/packet.pyi
@@ -0,0 +1,25 @@
+from fractions import Fraction
+
+from av.subtitles.subtitle import SubtitleSet
+
+from .buffer import Buffer
+from .stream import Stream
+
+class Packet(Buffer):
+ stream: Stream
+ stream_index: int
+ time_base: Fraction
+ pts: int | None
+ dts: int
+ pos: int | None
+ size: int
+ duration: int | None
+ opaque: object
+ is_keyframe: bool
+ is_corrupt: bool
+ is_discard: bool
+ is_trusted: bool
+ is_disposable: bool
+
+ def __init__(self, input: int | bytes | None = None) -> None: ...
+ def decode(self) -> list[SubtitleSet]: ...
diff --git a/venv/lib/python3.11/site-packages/av/plane.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/plane.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..de4a33c029a5bf39dcdb35b14bef52a07a84ec12
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/plane.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:c6dbd698b62ce9db756b95e5a28150d82e823b753f26ae95f27af09385c69575
+size 425121
diff --git a/venv/lib/python3.11/site-packages/av/plane.pxd b/venv/lib/python3.11/site-packages/av/plane.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..df3847d7b2423a549e11c508427fbb6a0132b8b3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/plane.pxd
@@ -0,0 +1,11 @@
+from av.buffer cimport Buffer
+from av.frame cimport Frame
+
+
+cdef class Plane(Buffer):
+
+ cdef Frame frame
+ cdef int index
+
+ cdef size_t _buffer_size(self)
+ cdef void* _buffer_ptr(self)
diff --git a/venv/lib/python3.11/site-packages/av/plane.pyi b/venv/lib/python3.11/site-packages/av/plane.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..99594d11a23d021fb0181b6bfe5ecc1bcbc43c1b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/plane.pyi
@@ -0,0 +1,8 @@
+from .buffer import Buffer
+from .frame import Frame
+
+class Plane(Buffer):
+ frame: Frame
+ index: int
+
+ def __init__(self, frame: Frame, index: int) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/py.typed b/venv/lib/python3.11/site-packages/av/py.typed
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/__init__.pxd b/venv/lib/python3.11/site-packages/av/sidedata/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/__init__.py b/venv/lib/python3.11/site-packages/av/sidedata/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..64c6eed8961994cc4afbe3313abe4dd246d3b76c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:521fbfd5cb43288258c17099efc8096d1e965c075c4ac3ea333446b526d9ea86
+size 703953
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.pxd b/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..993c5e283f4b3031c16d989f719ab217d060563b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.pxd
@@ -0,0 +1,16 @@
+cimport libav as lib
+
+from av.frame cimport Frame
+from av.sidedata.sidedata cimport SideData
+
+
+cdef class _MotionVectors(SideData):
+
+ cdef dict _vectors
+ cdef int _len
+
+
+cdef class MotionVector:
+
+ cdef _MotionVectors parent
+ cdef lib.AVMotionVector *ptr
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.pyi b/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..eb514eb7055f4af559803e7981f33ab26c280737
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/sidedata/motionvectors.pyi
@@ -0,0 +1,27 @@
+from typing import Any, Sequence, overload
+
+import numpy as np
+
+from .sidedata import SideData
+
+class MotionVectors(SideData, Sequence[MotionVector]):
+ @overload
+ def __getitem__(self, index: int): ...
+ @overload
+ def __getitem__(self, index: slice): ...
+ @overload
+ def __getitem__(self, index: int | slice): ...
+ def __len__(self) -> int: ...
+ def to_ndarray(self) -> np.ndarray[Any, Any]: ...
+
+class MotionVector:
+ source: int
+ w: int
+ h: int
+ src_x: int
+ src_y: int
+ dst_x: int
+ dst_y: int
+ motion_x: int
+ motion_y: int
+ motion_scale: int
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/sidedata.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/sidedata/sidedata.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..3e77286781dc88c31e1f6120033508e9737a972b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/sidedata/sidedata.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:b42b7e6ef843e0fdc71869029572ef00c6ec725546037df31793658543789c15
+size 831089
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/sidedata.pxd b/venv/lib/python3.11/site-packages/av/sidedata/sidedata.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..8a2f6d07c6c32b2aeba63a4c181ca592b15297c7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/sidedata/sidedata.pxd
@@ -0,0 +1,21 @@
+
+cimport libav as lib
+
+from av.buffer cimport Buffer
+from av.dictionary cimport _Dictionary, wrap_dictionary
+from av.frame cimport Frame
+
+
+cdef class SideData(Buffer):
+ cdef Frame frame
+ cdef lib.AVFrameSideData *ptr
+ cdef _Dictionary metadata
+
+
+cdef SideData wrap_side_data(Frame frame, int index)
+
+cdef class _SideDataContainer:
+ cdef Frame frame
+
+ cdef list _by_index
+ cdef dict _by_type
diff --git a/venv/lib/python3.11/site-packages/av/sidedata/sidedata.pyi b/venv/lib/python3.11/site-packages/av/sidedata/sidedata.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..d165513ab8863dc9a84a07e4a40277e72d759c04
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/sidedata/sidedata.pyi
@@ -0,0 +1,52 @@
+from collections.abc import Mapping
+from enum import Enum
+from typing import ClassVar, Iterator, Sequence, overload
+
+from av.buffer import Buffer
+from av.frame import Frame
+
+class Type(Enum):
+ PANSCAN: ClassVar[Type]
+ A53_CC: ClassVar[Type]
+ STEREO3D: ClassVar[Type]
+ MATRIXENCODING: ClassVar[Type]
+ DOWNMIX_INFO: ClassVar[Type]
+ REPLAYGAIN: ClassVar[Type]
+ DISPLAYMATRIX: ClassVar[Type]
+ AFD: ClassVar[Type]
+ MOTION_VECTORS: ClassVar[Type]
+ SKIP_SAMPLES: ClassVar[Type]
+ AUDIO_SERVICE_TYPE: ClassVar[Type]
+ MASTERING_DISPLAY_METADATA: ClassVar[Type]
+ GOP_TIMECODE: ClassVar[Type]
+ SPHERICAL: ClassVar[Type]
+ CONTENT_LIGHT_LEVEL: ClassVar[Type]
+ ICC_PROFILE: ClassVar[Type]
+ S12M_TIMECODE: ClassVar[Type]
+ DYNAMIC_HDR_PLUS: ClassVar[Type]
+ REGIONS_OF_INTEREST: ClassVar[Type]
+ VIDEO_ENC_PARAMS: ClassVar[Type]
+ SEI_UNREGISTERED: ClassVar[Type]
+ FILM_GRAIN_PARAMS: ClassVar[Type]
+ DETECTION_BBOXES: ClassVar[Type]
+ DOVI_RPU_BUFFER: ClassVar[Type]
+ DOVI_METADATA: ClassVar[Type]
+ DYNAMIC_HDR_VIVID: ClassVar[Type]
+ AMBIENT_VIEWING_ENVIRONMENT: ClassVar[Type]
+ VIDEO_HINT: ClassVar[Type]
+
+class SideData(Buffer):
+ type: Type
+
+class SideDataContainer(Mapping):
+ frame: Frame
+ def __len__(self) -> int: ...
+ def __iter__(self) -> Iterator[SideData]: ...
+ @overload
+ def __getitem__(self, key: str | int | Type) -> SideData: ...
+ @overload
+ def __getitem__(self, key: slice) -> Sequence[SideData]: ...
+ @overload
+ def __getitem__(
+ self, key: str | int | Type | slice
+ ) -> SideData | Sequence[SideData]: ...
diff --git a/venv/lib/python3.11/site-packages/av/stream.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/stream.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..8c74a26e5f3abad8468c7c36ffec4d1e511174ac
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/stream.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:9dd66aa57a857f25fb4de9eca92c5922a522076c328dd6860b02541f6902f22e
+size 609681
diff --git a/venv/lib/python3.11/site-packages/av/stream.pxd b/venv/lib/python3.11/site-packages/av/stream.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..c847f641e8a3602ffd0fc65d3440a964dd53f381
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/stream.pxd
@@ -0,0 +1,26 @@
+cimport libav as lib
+
+from av.codec.context cimport CodecContext
+from av.container.core cimport Container
+from av.frame cimport Frame
+from av.packet cimport Packet
+
+
+cdef class Stream:
+ cdef lib.AVStream *ptr
+
+ # Stream attributes.
+ cdef readonly Container container
+ cdef readonly dict metadata
+
+ # CodecContext attributes.
+ cdef readonly CodecContext codec_context
+
+ # Private API.
+ cdef _init(self, Container, lib.AVStream*, CodecContext)
+ cdef _finalize_for_output(self)
+ cdef _set_time_base(self, value)
+ cdef _set_id(self, value)
+
+
+cdef Stream wrap_stream(Container, lib.AVStream*, CodecContext)
diff --git a/venv/lib/python3.11/site-packages/av/stream.pyi b/venv/lib/python3.11/site-packages/av/stream.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..82bb672b283ecc3a4b59285123e5b6c548c98cc4
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/stream.pyi
@@ -0,0 +1,25 @@
+from fractions import Fraction
+from typing import Literal
+
+from .codec import Codec, CodecContext
+from .container import Container
+
+class Stream:
+ name: str | None
+ container: Container
+ codec: Codec
+ codec_context: CodecContext
+ metadata: dict[str, str]
+ id: int
+ profiles: list[str]
+ profile: str | None
+ index: int
+ time_base: Fraction | None
+ average_rate: Fraction | None
+ base_rate: Fraction | None
+ guessed_rate: Fraction | None
+ start_time: int | None
+ duration: int | None
+ frames: int
+ language: str | None
+ type: Literal["video", "audio", "data", "subtitle", "attachment"]
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/__init__.pxd b/venv/lib/python3.11/site-packages/av/subtitles/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/__init__.py b/venv/lib/python3.11/site-packages/av/subtitles/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..a1f3b076b68f7d047ee545f6885a80953030ccd0
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:3a36d50a41c2cffeb3a6492faea6ec7c8d194384ce5e6e54e91b65787de9c696
+size 371697
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.pxd b/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..42141aa4f80e00ad81dd88841cd1b9457d49c231
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.pxd
@@ -0,0 +1,5 @@
+from av.codec.context cimport CodecContext
+
+
+cdef class SubtitleCodecContext(CodecContext):
+ pass
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.pyi b/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..0762c19f0602551e2b0fa4e218354d3b69ffdaf5
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/codeccontext.pyi
@@ -0,0 +1,6 @@
+from typing import Literal
+
+from av.codec.context import CodecContext
+
+class SubtitleCodecContext(CodecContext):
+ type: Literal["subtitle"]
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/stream.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/subtitles/stream.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..737d24ada8159437024b65cb9afd3edfbaa6a452
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/stream.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:fa2bc31d89f8e32d0aa08455bcbeb2cd7739494303d40fdb58be87a14c4f19b3
+size 400337
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/stream.pxd b/venv/lib/python3.11/site-packages/av/subtitles/stream.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..745032af956c812f0d72f2f3dda9c888b23b1b7f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/stream.pxd
@@ -0,0 +1,6 @@
+from av.packet cimport Packet
+from av.stream cimport Stream
+
+
+cdef class SubtitleStream(Stream):
+ cpdef decode(self, Packet packet=?)
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/stream.pyi b/venv/lib/python3.11/site-packages/av/subtitles/stream.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..cb1ac34a25edbe387b1c1e2c390f14b0cb9fd19c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/stream.pyi
@@ -0,0 +1,6 @@
+from av.packet import Packet
+from av.stream import Stream
+from av.subtitles.subtitle import SubtitleSet
+
+class SubtitleStream(Stream):
+ def decode(self, packet: Packet | None = None) -> list[SubtitleSet]: ...
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/subtitle.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/subtitles/subtitle.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..8d9c9dec9c3774f8325df308dff5393de24600d3
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/subtitle.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:2d69dcfd82663023308a92f97a5aa547a7594b4cee86edfff27dfa75585b933d
+size 1011585
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/subtitle.pxd b/venv/lib/python3.11/site-packages/av/subtitles/subtitle.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..508eb903480a4a0b112f6580eb66ba9ae9820491
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/subtitle.pxd
@@ -0,0 +1,31 @@
+cimport libav as lib
+
+
+cdef class SubtitleProxy:
+ cdef lib.AVSubtitle struct
+
+
+cdef class SubtitleSet:
+ cdef SubtitleProxy proxy
+ cdef readonly tuple rects
+
+
+cdef class Subtitle:
+ cdef SubtitleProxy proxy
+ cdef lib.AVSubtitleRect *ptr
+ cdef readonly bytes type
+
+cdef class TextSubtitle(Subtitle):
+ pass
+
+cdef class ASSSubtitle(Subtitle):
+ pass
+
+cdef class BitmapSubtitle(Subtitle):
+ cdef readonly planes
+
+cdef class BitmapSubtitlePlane:
+ cdef readonly BitmapSubtitle subtitle
+ cdef readonly int index
+ cdef readonly long buffer_size
+ cdef void *_buffer
diff --git a/venv/lib/python3.11/site-packages/av/subtitles/subtitle.pyi b/venv/lib/python3.11/site-packages/av/subtitles/subtitle.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..2a35d0a5546694b1030c53cae6e8914e237ff62d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/subtitles/subtitle.pyi
@@ -0,0 +1,37 @@
+from typing import Iterator, Literal
+
+class SubtitleSet:
+ format: int
+ start_display_time: int
+ end_display_time: int
+ pts: int
+ rects: tuple[Subtitle]
+
+ def __len__(self) -> int: ...
+ def __iter__(self) -> Iterator[Subtitle]: ...
+ def __getitem__(self, i: int) -> Subtitle: ...
+
+class Subtitle: ...
+
+class BitmapSubtitle(Subtitle):
+ type: Literal[b"bitmap"]
+ x: int
+ y: int
+ width: int
+ height: int
+ nb_colors: int
+ planes: tuple[BitmapSubtitlePlane, ...]
+
+class BitmapSubtitlePlane:
+ subtitle: BitmapSubtitle
+ index: int
+ buffer_size: int
+
+class AssSubtitle(Subtitle):
+ type: Literal[b"ass", b"text"]
+ @property
+ def ass(self) -> bytes: ...
+ @property
+ def dialogue(self) -> bytes: ...
+ @property
+ def text(self) -> bytes: ...
diff --git a/venv/lib/python3.11/site-packages/av/utils.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/utils.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..f24b9c25b9851ec785dd309c74bba45a87f6f9c7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/utils.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:0d7359d59c1648e6238c5f3030360aad20831534f4a7315c41c3d756b4a1c654
+size 289681
diff --git a/venv/lib/python3.11/site-packages/av/utils.pxd b/venv/lib/python3.11/site-packages/av/utils.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..9aeb4a2fb742b08fea34c973dbda07a7794d3880
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/utils.pxd
@@ -0,0 +1,12 @@
+cimport libav as lib
+from libc.stdint cimport uint64_t
+
+
+cdef dict avdict_to_dict(lib.AVDictionary *input, str encoding, str errors)
+cdef dict_to_avdict(lib.AVDictionary **dst, dict src, str encoding, str errors)
+
+cdef object avrational_to_fraction(const lib.AVRational *input)
+cdef void to_avrational(object frac, lib.AVRational *input)
+
+cdef check_ndarray(object array, object dtype, int ndim)
+cdef flag_in_bitfield(uint64_t bitfield, uint64_t flag)
diff --git a/venv/lib/python3.11/site-packages/av/video/__init__.pxd b/venv/lib/python3.11/site-packages/av/video/__init__.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/av/video/__init__.py b/venv/lib/python3.11/site-packages/av/video/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..4a25d88376662b42451db9fe11a704f642f94744
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/__init__.py
@@ -0,0 +1,2 @@
+from .frame import VideoFrame
+from .stream import VideoStream
diff --git a/venv/lib/python3.11/site-packages/av/video/__init__.pyi b/venv/lib/python3.11/site-packages/av/video/__init__.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..8fa8fe7e5db5c3aa04eb2185f56d187316e9a18d
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/__init__.pyi
@@ -0,0 +1,4 @@
+from .frame import VideoFrame
+from .stream import VideoStream
+
+__all__ = ("VideoFrame", "VideoStream")
diff --git a/venv/lib/python3.11/site-packages/av/video/codeccontext.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/video/codeccontext.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..61bfb4bd3043ccdb93da100b7c249e4789364497
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/codeccontext.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:4a28bef0ef92beb62cdc114b53880d697d68fece4edc9f4260f8ba97c7311456
+size 761129
diff --git a/venv/lib/python3.11/site-packages/av/video/codeccontext.pxd b/venv/lib/python3.11/site-packages/av/video/codeccontext.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..9693caa9bb91e111c0f0e1f44ec485665648a7b2
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/codeccontext.pxd
@@ -0,0 +1,21 @@
+
+from av.codec.context cimport CodecContext
+from av.video.format cimport VideoFormat
+from av.video.frame cimport VideoFrame
+from av.video.reformatter cimport VideoReformatter
+
+
+cdef class VideoCodecContext(CodecContext):
+
+ cdef VideoFormat _format
+ cdef _build_format(self)
+
+ cdef int last_w
+ cdef int last_h
+ cdef readonly VideoReformatter reformatter
+
+ # For encoding.
+ cdef readonly int encoded_frame_count
+
+ # For decoding.
+ cdef VideoFrame next_frame
diff --git a/venv/lib/python3.11/site-packages/av/video/codeccontext.pyi b/venv/lib/python3.11/site-packages/av/video/codeccontext.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..da72053c4ac8340304fe1e471d405ee484bc46ab
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/codeccontext.pyi
@@ -0,0 +1,35 @@
+from fractions import Fraction
+from typing import Iterator, Literal
+
+from av.codec.context import CodecContext
+from av.packet import Packet
+
+from .format import VideoFormat
+from .frame import VideoFrame
+
+class VideoCodecContext(CodecContext):
+ format: VideoFormat | None
+ width: int
+ height: int
+ bits_per_coded_sample: int
+ pix_fmt: str | None
+ framerate: Fraction
+ rate: Fraction
+ gop_size: int
+ sample_aspect_ratio: Fraction | None
+ display_aspect_ratio: Fraction | None
+ has_b_frames: bool
+ max_b_frames: int
+ coded_width: int
+ coded_height: int
+ color_range: int
+ color_primaries: int
+ color_trc: int
+ colorspace: int
+ qmin: int
+ qmax: int
+ type: Literal["video"]
+
+ def encode(self, frame: VideoFrame | None = None) -> list[Packet]: ...
+ def encode_lazy(self, frame: VideoFrame | None = None) -> Iterator[Packet]: ...
+ def decode(self, packet: Packet | None = None) -> list[VideoFrame]: ...
diff --git a/venv/lib/python3.11/site-packages/av/video/format.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/video/format.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..a970139d1bbf2876e6328c7c06d79628482cbfa9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/format.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:da9b80f7a5b4a73a660f8c83d12c2e0b95c781f874144dc5d1ba6b082ad10017
+size 892793
diff --git a/venv/lib/python3.11/site-packages/av/video/format.pxd b/venv/lib/python3.11/site-packages/av/video/format.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..a2efa9d1d26976debb5d4cca1d2da3883c9fb281
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/format.pxd
@@ -0,0 +1,27 @@
+cimport libav as lib
+
+
+cdef class VideoFormat:
+
+ cdef lib.AVPixelFormat pix_fmt
+ cdef const lib.AVPixFmtDescriptor *ptr
+ cdef readonly unsigned int width, height
+
+ cdef readonly tuple components
+
+ cdef _init(self, lib.AVPixelFormat pix_fmt, unsigned int width, unsigned int height)
+
+ cpdef chroma_width(self, int luma_width=?)
+ cpdef chroma_height(self, int luma_height=?)
+
+
+cdef class VideoFormatComponent:
+
+ cdef VideoFormat format
+ cdef readonly unsigned int index
+ cdef const lib.AVComponentDescriptor *ptr
+
+
+cdef VideoFormat get_video_format(lib.AVPixelFormat c_format, unsigned int width, unsigned int height)
+
+cdef lib.AVPixelFormat get_pix_fmt(const char *name) except lib.AV_PIX_FMT_NONE
diff --git a/venv/lib/python3.11/site-packages/av/video/format.pyi b/venv/lib/python3.11/site-packages/av/video/format.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..ee16b85b88070b695be1cb1b317f994d895258bc
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/format.pyi
@@ -0,0 +1,27 @@
+class VideoFormat:
+ name: str
+ bits_per_pixel: int
+ padded_bits_per_pixel: int
+ is_big_endian: bool
+ has_palette: bool
+ is_bit_stream: bool
+ is_planar: bool
+ is_rgb: bool
+ width: int
+ height: int
+ components: tuple[VideoFormatComponent, ...]
+
+ def __init__(self, name: str, width: int = 0, height: int = 0) -> None: ...
+ def chroma_width(self, luma_width: int = 0) -> int: ...
+ def chroma_height(self, luma_height: int = 0) -> int: ...
+
+class VideoFormatComponent:
+ plane: int
+ bits: int
+ is_alpha: bool
+ is_luma: bool
+ is_chroma: bool
+ width: int
+ height: int
+
+ def __init__(self, format: VideoFormat, index: int) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/video/frame.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/video/frame.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..0d1ce405b3c1ffbe585dfc59cc59a5bf624418fd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/frame.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:e5e37f37913069518057f8ff18c15bac9bb3020f36fa7a4f5894d9e4082f2387
+size 3289633
diff --git a/venv/lib/python3.11/site-packages/av/video/frame.pxd b/venv/lib/python3.11/site-packages/av/video/frame.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..779b239771f49b2d697e00df87a0297aff983b82
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/frame.pxd
@@ -0,0 +1,21 @@
+cimport libav as lib
+from libc.stdint cimport uint8_t
+
+from av.frame cimport Frame
+from av.video.format cimport VideoFormat
+from av.video.reformatter cimport VideoReformatter
+
+
+cdef class VideoFrame(Frame):
+ # This is the buffer that is used to back everything in the AVFrame.
+ # We don't ever actually access it directly.
+ cdef uint8_t *_buffer
+ cdef object _np_buffer
+
+ cdef VideoReformatter reformatter
+ cdef readonly VideoFormat format
+
+ cdef _init(self, lib.AVPixelFormat format, unsigned int width, unsigned int height)
+ cdef _init_user_attributes(self)
+
+cdef VideoFrame alloc_video_frame()
diff --git a/venv/lib/python3.11/site-packages/av/video/frame.pyi b/venv/lib/python3.11/site-packages/av/video/frame.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..0739010c1e9731a01e21cfcd18b7e8fc564929c5
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/frame.pyi
@@ -0,0 +1,77 @@
+from enum import IntEnum
+from typing import Any, ClassVar, Union
+
+import numpy as np
+from PIL import Image
+
+from av.frame import Frame
+
+from .format import VideoFormat
+from .plane import VideoPlane
+
+_SupportedNDarray = Union[
+ np.ndarray[Any, np.dtype[np.uint8]],
+ np.ndarray[Any, np.dtype[np.uint16]],
+ np.ndarray[Any, np.dtype[np.float32]],
+]
+
+class PictureType(IntEnum):
+ NONE: int
+ I: int
+ P: int
+ B: int
+ S: int
+ SI: int
+ SP: int
+ BI: int
+
+class VideoFrame(Frame):
+ format: VideoFormat
+ pts: int
+ planes: tuple[VideoPlane, ...]
+ pict_type: int
+ colorspace: int
+ color_range: int
+
+ @property
+ def time(self) -> float: ...
+ @property
+ def width(self) -> int: ...
+ @property
+ def height(self) -> int: ...
+ @property
+ def interlaced_frame(self) -> bool: ...
+ def __init__(
+ self, width: int = 0, height: int = 0, format: str = "yuv420p"
+ ) -> None: ...
+ def reformat(
+ self,
+ width: int | None = None,
+ height: int | None = None,
+ format: str | None = None,
+ src_colorspace: str | int | None = None,
+ dst_colorspace: str | int | None = None,
+ interpolation: int | str | None = None,
+ src_color_range: int | str | None = None,
+ dst_color_range: int | str | None = None,
+ ) -> VideoFrame: ...
+ def to_rgb(self, **kwargs: Any) -> VideoFrame: ...
+ def to_image(self, **kwargs: Any) -> Image.Image: ...
+ def to_ndarray(self, **kwargs: Any) -> _SupportedNDarray: ...
+ @staticmethod
+ def from_image(img: Image.Image) -> VideoFrame: ...
+ @staticmethod
+ def from_numpy_buffer(
+ array: _SupportedNDarray, format: str = "rgb24", width: int = 0
+ ) -> VideoFrame: ...
+ @staticmethod
+ def from_ndarray(array: _SupportedNDarray, format: str = "rgb24") -> VideoFrame: ...
+ @staticmethod
+ def from_bytes(
+ data: bytes,
+ width: int,
+ height: int,
+ format: str = "rgba",
+ flip_horizontal: bool = False,
+ flip_vertical: bool = False,
+ ) -> VideoFrame: ...
diff --git a/venv/lib/python3.11/site-packages/av/video/plane.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/video/plane.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..d82fbb69c39cd4e3485449434ee88d8197bde385
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/plane.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:f780265b3c03ba83521aecb81a9abbdf0d7bd8451646465b096b914426df7bc9
+size 470233
diff --git a/venv/lib/python3.11/site-packages/av/video/plane.pxd b/venv/lib/python3.11/site-packages/av/video/plane.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..f9abf22b61866d29f45e69a3651b394d8af00a08
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/plane.pxd
@@ -0,0 +1,8 @@
+from av.plane cimport Plane
+from av.video.format cimport VideoFormatComponent
+
+
+cdef class VideoPlane(Plane):
+
+ cdef readonly size_t buffer_size
+ cdef readonly unsigned int width, height
diff --git a/venv/lib/python3.11/site-packages/av/video/plane.pyi b/venv/lib/python3.11/site-packages/av/video/plane.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..e4a0a206cd1135120ba813677523a15bffdbf59a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/plane.pyi
@@ -0,0 +1,11 @@
+from av.plane import Plane
+
+from .frame import VideoFrame
+
+class VideoPlane(Plane):
+ line_size: int
+ width: int
+ height: int
+ buffer_size: int
+
+ def __init__(self, frame: VideoFrame, index: int) -> None: ...
diff --git a/venv/lib/python3.11/site-packages/av/video/reformatter.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/video/reformatter.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..52c5abbb99a231ce5bd9dca62f4f5ab16b86c006
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/reformatter.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:6fcfd2b7f0a7702a59030bcf4d8b5f1603ebeb0d82fbfb0730eb42515b72a9b7
+size 810697
diff --git a/venv/lib/python3.11/site-packages/av/video/reformatter.pxd b/venv/lib/python3.11/site-packages/av/video/reformatter.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..7682fab6d48355ff82932f9f204ff5f05b4392b2
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/reformatter.pxd
@@ -0,0 +1,13 @@
+cimport libav as lib
+
+from av.video.frame cimport VideoFrame
+
+
+cdef class VideoReformatter:
+
+ cdef lib.SwsContext *ptr
+
+ cdef _reformat(self, VideoFrame frame, int width, int height,
+ lib.AVPixelFormat format, int src_colorspace,
+ int dst_colorspace, int interpolation,
+ int src_color_range, int dst_color_range)
diff --git a/venv/lib/python3.11/site-packages/av/video/reformatter.pyi b/venv/lib/python3.11/site-packages/av/video/reformatter.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..fd5dbd053d30e7c1650e94cdf012858a57b9123f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/reformatter.pyi
@@ -0,0 +1,52 @@
+from enum import IntEnum
+
+from .frame import VideoFrame
+
+class Interpolation(IntEnum):
+ FAST_BILINEAER: int
+ BILINEAR: int
+ BICUBIC: int
+ X: int
+ POINT: int
+ AREA: int
+ BICUBLIN: int
+ GAUSS: int
+ SINC: int
+ LANCZOS: int
+ SPLINE: int
+
+class Colorspace(IntEnum):
+ ITU709: int
+ FCC: int
+ ITU601: int
+ ITU624: int
+ SMPTE170M: int
+ SMPTE240M: int
+ DEFAULT: int
+ itu709: int
+ fcc: int
+ itu601: int
+ itu624: int
+ smpte170m: int
+ smpte240m: int
+ default: int
+
+class ColorRange(IntEnum):
+ UNSPECIFIED: int
+ MPEG: int
+ JPEG: int
+ NB: int
+
+class VideoReformatter:
+ def reformat(
+ self,
+ frame: VideoFrame,
+ width: int | None = None,
+ height: int | None = None,
+ format: str | None = None,
+ src_colorspace: int | None = None,
+ dst_colorspace: int | None = None,
+ interpolation: int | str | None = None,
+ src_color_range: int | str | None = None,
+ dst_color_range: int | str | None = None,
+ ) -> VideoFrame: ...
diff --git a/venv/lib/python3.11/site-packages/av/video/stream.cpython-311-x86_64-linux-gnu.so b/venv/lib/python3.11/site-packages/av/video/stream.cpython-311-x86_64-linux-gnu.so
new file mode 100644
index 0000000000000000000000000000000000000000..4e8a27c4b7464906d30c4d1841863c9d98f37a95
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/stream.cpython-311-x86_64-linux-gnu.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:482f0d1180fbd645777744eff9ad1c83a45d38202e5279034d8bad645e2a1cf4
+size 601377
diff --git a/venv/lib/python3.11/site-packages/av/video/stream.pxd b/venv/lib/python3.11/site-packages/av/video/stream.pxd
new file mode 100644
index 0000000000000000000000000000000000000000..f0dcfb9b2ca64283b36675e4470f39fd2e1c514b
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/stream.pxd
@@ -0,0 +1,9 @@
+from av.packet cimport Packet
+from av.stream cimport Stream
+
+from .frame cimport VideoFrame
+
+
+cdef class VideoStream(Stream):
+ cpdef encode(self, VideoFrame frame=?)
+ cpdef decode(self, Packet packet=?)
diff --git a/venv/lib/python3.11/site-packages/av/video/stream.pyi b/venv/lib/python3.11/site-packages/av/video/stream.pyi
new file mode 100644
index 0000000000000000000000000000000000000000..dd670d3cf11c5db9989d85dec7ee55be0b5f1b9e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/av/video/stream.pyi
@@ -0,0 +1,43 @@
+from fractions import Fraction
+from typing import Iterator, Literal
+
+from av.codec.context import ThreadType
+from av.packet import Packet
+from av.stream import Stream
+
+from .codeccontext import VideoCodecContext
+from .format import VideoFormat
+from .frame import VideoFrame
+
+class VideoStream(Stream):
+ bit_rate: int | None
+ max_bit_rate: int | None
+ bit_rate_tolerance: int
+ sample_aspect_ratio: Fraction | None
+ display_aspect_ratio: Fraction | None
+ codec_context: VideoCodecContext
+
+ def encode(self, frame: VideoFrame | None = None) -> list[Packet]: ...
+ def encode_lazy(self, frame: VideoFrame | None = None) -> Iterator[Packet]: ...
+ def decode(self, packet: Packet | None = None) -> list[VideoFrame]: ...
+
+ # from codec context
+ format: VideoFormat
+ thread_count: int
+ thread_type: ThreadType
+ width: int
+ height: int
+ bits_per_coded_sample: int
+ pix_fmt: str | None
+ framerate: Fraction
+ rate: Fraction
+ gop_size: int
+ has_b_frames: bool
+ max_b_frames: int
+ coded_width: int
+ coded_height: int
+ color_range: int
+ color_primaries: int
+ color_trc: int
+ colorspace: int
+ type: Literal["video"]
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..3060237e5ac61763a9ef36a5b6ebf925c17d22b6
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/__init__.py
@@ -0,0 +1,78 @@
+# Copyright (c) Facebook, Inc. and its affiliates.
+#
+# This source code is licensed under the MIT license found in the
+# LICENSE file in the root directory of this source tree.
+
+
+import importlib
+import sys
+
+import torch
+
+from . import _ops, utils
+from .autograd._functions import (
+ MatmulLtState,
+ matmul,
+ matmul_4bit,
+)
+from .backends.cpu import ops as cpu_ops
+from .backends.default import ops as default_ops
+from .nn import modules
+from .optim import adam
+
+# This is a signal for integrations with transformers/diffusers.
+# Eventually we may remove this but it is currently required for compatibility.
+features = {"multi_backend"}
+supported_torch_devices = {
+ "cpu",
+ "cuda", # NVIDIA/AMD GPU
+ "xpu", # Intel GPU
+ "hpu", # Intel Gaudi
+ "npu", # Ascend NPU
+ "mps", # Apple Silicon
+}
+
+if torch.cuda.is_available():
+ from .backends.cuda import ops as cuda_ops
+
+if hasattr(torch, "xpu") and torch.xpu.is_available():
+ from .backends.xpu import ops as xpu_ops
+
+if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
+ from .backends.mps import ops as mps_ops
+
+if importlib.util.find_spec("habana_frameworks") and importlib.util.find_spec("habana_frameworks.torch"):
+ # In case not automatically imported
+ import habana_frameworks.torch
+
+ if hasattr(torch, "hpu") and torch.hpu.is_available():
+ from .backends.hpu import ops as hpu_ops
+
+
+def _import_backends():
+ """
+ Discover and autoload all available backends installed as separate packages.
+ Packages with an entrypoint for "bitsandbytes.backends" will be loaded.
+ Inspired by PyTorch implementation: https://pytorch.org/tutorials/prototype/python_extension_autoload.html
+ """
+ from importlib.metadata import entry_points
+
+ extensions = entry_points(group="bitsandbytes.backends")
+
+ for ext in extensions:
+ try:
+ entry = ext.load()
+ entry()
+ except Exception as e:
+ raise RuntimeError(f"bitsandbytes: failed to load backend {ext.name}: {e}") from e
+
+
+_import_backends()
+
+__pdoc__ = {
+ "libbitsandbytes": False,
+ "optim.optimizer.Optimizer8bit": False,
+ "optim.optimizer.MockArgs": False,
+}
+
+__version__ = "0.50.0"
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/__main__.py b/venv/lib/python3.11/site-packages/bitsandbytes/__main__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e716b6f3ff54cb847682d68fd4cb0e6893dba3fe
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/__main__.py
@@ -0,0 +1,4 @@
+if __name__ == "__main__":
+ from bitsandbytes.diagnostics.main import main
+
+ main()
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/_ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/_ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..43efd860999f38d8d0b20a0398dd73ad41e5ad2e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/_ops.py
@@ -0,0 +1,510 @@
+from collections.abc import Sequence
+from typing import Optional
+
+import torch
+
+register_fake = torch.library.register_fake
+register_kernel = torch.library.register_kernel
+
+# Int8 mixed precision matmul + dequant + bias
+torch.library.define(
+ "bitsandbytes::int8_mixed_scaled_mm",
+ "(Tensor A, Tensor CA, Tensor CB, Tensor SCA, Tensor SCB, Tensor? outlier_cols=None, Tensor? bias=None) -> (Tensor, Tensor?)",
+)
+
+
+@register_fake("bitsandbytes::int8_mixed_scaled_mm")
+def _(
+ A: torch.Tensor,
+ CA: torch.Tensor,
+ CB: torch.Tensor,
+ SCA: torch.Tensor,
+ SCB: torch.Tensor,
+ outlier_cols: Optional[torch.Tensor] = None,
+ bias: Optional[torch.Tensor] = None,
+) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
+ shapeC = (*CA.shape[:-1], CB.shape[0])
+
+ out = torch.empty(shapeC, device=A.device, dtype=A.dtype)
+
+ outlier_cols = torch.library.get_ctx().new_dynamic_size()
+ subA = A.new_empty(outlier_cols, dtype=torch.int64)
+
+ return out, subA
+
+
+# Higher level op: int8 matmul + dequant + bias
+torch.library.define(
+ "bitsandbytes::int8_scaled_mm",
+ "(Tensor A, Tensor B, Tensor row_stats, Tensor col_stats, Tensor? bias=None, ScalarType? dtype=None) -> Tensor",
+)
+
+
+@register_fake("bitsandbytes::int8_scaled_mm")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ row_stats: torch.Tensor,
+ col_stats: torch.Tensor,
+ bias: Optional[torch.Tensor] = None,
+ dtype: Optional[torch.dtype] = None,
+) -> torch.Tensor:
+ shapeC = (*A.shape[:-1], B.shape[0])
+ return torch.empty(shapeC, device=A.device, dtype=dtype or torch.float16)
+
+
+torch.library.define(
+ "bitsandbytes::int8_linear_matmul",
+ "(Tensor A, Tensor B) -> Tensor",
+)
+
+
+@register_fake("bitsandbytes::int8_linear_matmul")
+def _(A: torch.Tensor, B: torch.Tensor):
+ torch._check(A.dtype == torch.int8, lambda: "A must be int8")
+ torch._check(B.dtype == torch.int8, lambda: "B must be int8")
+ shapeC = (*A.shape[:-1], B.shape[0])
+ return torch.empty(shapeC, device=A.device, dtype=torch.int32)
+
+
+# More info on `out` overloads:
+# https://github.com/pytorch/pytorch/issues/125044
+torch.library.define(
+ "bitsandbytes::int8_linear_matmul.out",
+ "(Tensor A, Tensor B, Tensor! out) -> ()",
+)
+
+
+@register_fake("bitsandbytes::int8_linear_matmul.out")
+def _(A: torch.Tensor, B: torch.Tensor, out: torch.Tensor):
+ shapeC = (*A.shape[:-1], B.shape[0])
+
+ torch._check(A.dtype == torch.int8, lambda: "A must be int8")
+ torch._check(B.dtype == torch.int8, lambda: "B must be int8")
+ torch._check(out.shape == shapeC, lambda: f"Expected out.shape == {shapeC}, got {out.shape}")
+ torch._check(out.device == A.device, lambda: f"Expected out.device == {A.device}, got {out.device}")
+ torch._check(out.dtype == torch.int32, lambda: f"Expected out.dtype == int32, got {out.dtype}")
+
+
+torch.library.define(
+ "bitsandbytes::int8_vectorwise_quant",
+ "(Tensor A, float threshold=0.0) -> (Tensor, Tensor, Tensor?)",
+)
+
+
+@register_fake("bitsandbytes::int8_vectorwise_quant")
+def _(A: torch.Tensor, threshold=0.0):
+ out_row = torch.empty(A.shape, device=A.device, dtype=torch.int8)
+ row_stats = torch.empty(A.numel() // A.shape[-1], device=A.device, dtype=torch.float32)
+
+ if threshold == 0.0:
+ return out_row, row_stats, None
+
+ outlier_cols = torch.library.get_ctx().new_dynamic_size()
+
+ return out_row, row_stats, A.new_empty(outlier_cols, dtype=torch.int64)
+
+
+torch.library.define("bitsandbytes::int8_vectorwise_dequant", "(Tensor A, Tensor stats) -> Tensor")
+
+
+@register_fake("bitsandbytes::int8_vectorwise_dequant")
+def _(A: torch.Tensor, stats: torch.Tensor) -> torch.Tensor:
+ torch._check(A.dtype == torch.int8, lambda: "A must be int8")
+ return torch.empty_like(A, dtype=torch.float32)
+
+
+# Default PyTorch-native implementation
+@register_kernel("bitsandbytes::int8_vectorwise_dequant", "default")
+def _(A: torch.Tensor, stats: torch.Tensor):
+ # To dequantize we divide by 127, or multiply by the reciprocal.
+ return A * stats.view(-1, 1) * 7.874015718698502e-3
+
+
+torch.library.define(
+ "bitsandbytes::int8_mm_dequant",
+ "(Tensor A, Tensor row_stats, Tensor col_stats, ScalarType? dtype=None, Tensor? bias=None) -> Tensor",
+)
+
+
+@register_fake("bitsandbytes::int8_mm_dequant")
+def _(
+ A: torch.Tensor,
+ row_stats: torch.Tensor,
+ col_stats: torch.Tensor,
+ dtype: Optional[torch.dtype] = None,
+ bias: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ torch._check(A.dtype == torch.int32, lambda: "A must be int32")
+ return torch.empty_like(A, dtype=dtype or torch.float16)
+
+
+torch.library.define(
+ "bitsandbytes::int8_double_quant",
+ "(Tensor A, float threshold=0.0) -> (Tensor, Tensor, Tensor, Tensor, Tensor?)",
+)
+
+
+@register_fake("bitsandbytes::int8_double_quant")
+def _(
+ A: torch.Tensor,
+ threshold=0.0,
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
+ out_row = torch.empty_like(A, dtype=torch.int8)
+ out_col = torch.empty_like(A, dtype=torch.int8)
+ row_stats = torch.empty(A.numel() // A.shape[-1], device=A.device, dtype=torch.float32)
+ col_stats = torch.empty(A.shape[-1], device=A.device, dtype=torch.float32)
+ outlier_n = torch.library.get_ctx().new_dynamic_size()
+ outlier_cols = A.new_empty(outlier_n, dtype=torch.int64)
+ return out_row, out_col, row_stats, col_stats, outlier_cols
+
+
+torch.library.define(
+ "bitsandbytes::dequantize_4bit",
+ "(Tensor A, Tensor absmax, int blocksize, str quant_type, int[] shape, ScalarType dtype) -> Tensor",
+)
+
+
+@register_fake("bitsandbytes::dequantize_4bit")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ torch._check(blocksize in (32, 64, 128, 256, 512, 1024, 2048, 4096), lambda: f"invalid blocksize {blocksize}")
+ torch._check(quant_type in ("nf4", "fp4"), lambda: f"quant_type must be 'nf4' or 'fp4', got {quant_type!r}")
+ torch._check(absmax.dtype == torch.float32, lambda: f"absmax must be float32, got {absmax.dtype}")
+ torch._check(
+ dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"Blockwise 4bit dequantization only supports 16/32-bit floats, but got {dtype}",
+ )
+ return torch.empty(shape, dtype=dtype, device=A.device)
+
+
+torch.library.define(
+ "bitsandbytes::dequantize_4bit.out",
+ "(Tensor A, Tensor absmax, int blocksize, str quant_type, int[] shape, ScalarType dtype, Tensor! out) -> ()",
+)
+
+
+@register_fake("bitsandbytes::dequantize_4bit.out")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ torch._check(blocksize in (32, 64, 128, 256, 512, 1024, 2048, 4096), lambda: f"invalid blocksize {blocksize}")
+ torch._check(quant_type in ("nf4", "fp4"), lambda: f"quant_type must be 'nf4' or 'fp4', got {quant_type!r}")
+ torch._check(absmax.dtype == torch.float32, lambda: f"absmax must be float32, got {absmax.dtype}")
+ torch._check(
+ dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"Blockwise 4bit dequantization only supports 16/32-bit floats, but got {dtype}",
+ )
+ torch._check(out.shape == shape, lambda: f"Expected out.shape == {shape}, got {out.shape}")
+ torch._check(out.device == A.device, lambda: f"Expected out.device == {A.device}, got {out.device}")
+ torch._check(out.dtype == dtype, lambda: f"Expected out.dtype == {dtype}, got {out.dtype}")
+
+
+torch.library.define(
+ "bitsandbytes::quantize_4bit",
+ "(Tensor A, int blocksize, str quant_type, ScalarType quant_storage) -> (Tensor, Tensor)",
+)
+
+
+@register_fake("bitsandbytes::quantize_4bit")
+def _(
+ A: torch.Tensor, blocksize: int, quant_type: str, quant_storage: torch.dtype
+) -> tuple[torch.Tensor, torch.Tensor]:
+ torch._check(blocksize in (32, 64, 128, 256, 512, 1024, 2048, 4096), lambda: f"invalid blocksize {blocksize}")
+ torch._check(quant_type in ("nf4", "fp4"), lambda: f"quant_type must be 'nf4' or 'fp4', got {quant_type!r}")
+ torch._check(
+ A.dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"Blockwise 4bit quantization only supports 16/32-bit floats, but got {A.dtype}",
+ )
+
+ n = A.numel()
+ blocks = -(n // -blocksize)
+ absmax = torch.empty((blocks,), device=A.device, dtype=torch.float32)
+ out = torch.empty(((n + 1) // (quant_storage.itemsize * 2), 1), device=A.device, dtype=quant_storage)
+ return out, absmax
+
+
+torch.library.define(
+ "bitsandbytes::gemm_4bit",
+ "(Tensor A, Tensor B, int[] shapeB, Tensor absmax, int blocksize, str quant_type, "
+ "Tensor? bias=None, Tensor? absmax_8bit=None, Tensor? absmax_code=None, Tensor? absmax_offset=None) -> Tensor",
+)
+
+
+@register_fake("bitsandbytes::gemm_4bit")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ bias: Optional[torch.Tensor] = None,
+ absmax_8bit: Optional[torch.Tensor] = None,
+ absmax_code: Optional[torch.Tensor] = None,
+ absmax_offset: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ torch._check(len(shapeB) == 2, lambda: f"shapeB must be 2D [N, K], got {list(shapeB)}")
+ torch._check(A.shape[-1] == shapeB[1], lambda: f"A inner dim ({A.shape[-1]}) must match shapeB ({shapeB[1]})")
+ torch._check(
+ A.dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"A must be float16, bfloat16, or float32, got {A.dtype}",
+ )
+ torch._check(
+ B.dtype in (torch.uint8, torch.bfloat16, torch.float16, torch.float32),
+ lambda: f"B must be backed by storage of type uint8, bfloat16, float16, or float32, got {B.dtype}",
+ )
+ torch._check(blocksize in (32, 64, 128, 256, 512, 1024, 2048, 4096), lambda: f"invalid blocksize {blocksize}")
+ torch._check(quant_type in ("nf4", "fp4"), lambda: f"quant_type must be 'nf4' or 'fp4', got {quant_type!r}")
+ torch._check(absmax.dtype == torch.float32, lambda: f"absmax must be float32, got {absmax.dtype}")
+ if absmax_8bit is not None:
+ torch._check(absmax_8bit.ndim == 1, lambda: f"absmax_8bit must be 1D, got {absmax_8bit.ndim}D")
+ torch._check(absmax_8bit.dtype == torch.uint8, lambda: f"absmax_8bit must be uint8, got {absmax_8bit.dtype}")
+ torch._check(absmax_code is not None, lambda: "absmax_code required when absmax_8bit is provided")
+ torch._check(absmax_code.ndim == 1, lambda: f"absmax_code must be 1D, got {absmax_code.ndim}D")
+ torch._check(
+ absmax_code.shape[0] == 256, lambda: f"absmax_code must have 256 entries, got {absmax_code.shape[0]}"
+ )
+ torch._check(
+ absmax_code.dtype == torch.float32, lambda: f"absmax_code must be float32, got {absmax_code.dtype}"
+ )
+ torch._check(absmax_offset is not None, lambda: "absmax_offset required when absmax_8bit is provided")
+ torch._check(
+ absmax_offset.ndim == 0, lambda: f"absmax_offset must be a scalar (0-dim), got {absmax_offset.ndim}D"
+ )
+ torch._check(
+ absmax_offset.dtype == torch.float32, lambda: f"absmax_offset must be float32, got {absmax_offset.dtype}"
+ )
+ if bias is not None:
+ torch._check(bias.ndim == 1, lambda: f"bias must be 1D, got {bias.ndim}D")
+ torch._check(bias.shape[0] == shapeB[0], lambda: f"bias length ({bias.shape[0]}) must match N ({shapeB[0]})")
+ torch._check(bias.dtype == A.dtype, lambda: f"bias dtype ({bias.dtype}) must match A dtype ({A.dtype})")
+ N = shapeB[0]
+ return torch.empty((*A.shape[:-1], N), dtype=A.dtype, device=A.device)
+
+
+torch.library.define(
+ "bitsandbytes::dequantize_blockwise",
+ "(Tensor A, Tensor absmax, Tensor code, int blocksize, ScalarType dtype) -> Tensor",
+)
+
+
+@register_fake("bitsandbytes::dequantize_blockwise")
+def _(A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype) -> torch.Tensor:
+ torch._check(blocksize > 0, lambda: f"blocksize must be positive, got {blocksize}")
+ torch._check(A.dtype == torch.uint8, lambda: f"A must be uint8, got {A.dtype}")
+ torch._check(
+ dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"Blockwise dequantization only supports 16/32-bit floats, but got {dtype}",
+ )
+ return torch.empty_like(A, dtype=dtype)
+
+
+torch.library.define(
+ "bitsandbytes::dequantize_blockwise.out",
+ "(Tensor A, Tensor absmax, Tensor code, int blocksize, ScalarType dtype, Tensor! out) -> ()",
+)
+
+
+@register_fake("bitsandbytes::dequantize_blockwise.out")
+def _(
+ A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype, out: torch.Tensor
+):
+ torch._check(blocksize > 0, lambda: f"blocksize must be positive, got {blocksize}")
+ torch._check(A.dtype == torch.uint8, lambda: f"A must be uint8, got {A.dtype}")
+ torch._check(
+ dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"Blockwise dequantization only supports 16/32-bit floats, but got {dtype}",
+ )
+ torch._check(out.shape == A.shape, lambda: f"Expected out.shape == {A.shape}, got {out.shape}")
+ torch._check(out.device == A.device, lambda: f"Expected out.device == {A.device}, got {out.device}")
+ torch._check(out.dtype == dtype, lambda: f"Expected out.dtype == {dtype}, got {out.dtype}")
+
+
+torch.library.define("bitsandbytes::quantize_blockwise", "(Tensor A, Tensor code, int blocksize) -> (Tensor, Tensor)")
+
+
+@register_fake("bitsandbytes::quantize_blockwise")
+def _(A: torch.Tensor, code: torch.Tensor, blocksize: int) -> tuple[torch.Tensor, torch.Tensor]:
+ torch._check(blocksize > 0, lambda: f"blocksize must be positive, got {blocksize}")
+ torch._check(
+ A.dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"Blockwise quantization only supports 16/32-bit floats, but got {A.dtype}",
+ )
+ n = A.numel()
+ blocks = -(n // -blocksize)
+ absmax = torch.empty((blocks,), device=A.device, dtype=torch.float32)
+ out = torch.empty_like(A, dtype=torch.uint8)
+ return out, absmax
+
+
+torch.library.define(
+ "bitsandbytes::gemv_4bit",
+ "(Tensor A, Tensor B, int[] shapeB, Tensor absmax, Tensor code, int blocksize) -> Tensor",
+)
+
+
+@register_fake("bitsandbytes::gemv_4bit")
+def _(
+ A: torch.Tensor, B: torch.Tensor, shapeB: Sequence[int], absmax: torch.Tensor, code: torch.Tensor, blocksize: int
+) -> torch.Tensor:
+ torch._check(blocksize in (32, 64, 128, 256, 512, 1024, 2048, 4096), lambda: f"invalid blocksize {blocksize}")
+ torch._check(
+ A.dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"A must be float16, bfloat16, or float32, got {A.dtype}",
+ )
+ torch._check(
+ B.dtype in (torch.uint8, torch.bfloat16, torch.float16, torch.float32),
+ lambda: f"B must be backed by storage of type uint8, bfloat16, float16, or float32, got {B.dtype}",
+ )
+ shape = (*A.shape[:-1], shapeB[0])
+ return torch.empty(shape, device=A.device, dtype=A.dtype)
+
+
+torch.library.define(
+ "bitsandbytes::gemv_4bit.out",
+ "(Tensor A, Tensor B, int[] shapeB, Tensor absmax, Tensor code, int blocksize, Tensor! out) -> ()",
+)
+
+
+@register_fake("bitsandbytes::gemv_4bit.out")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ out: torch.Tensor,
+) -> None:
+ torch._check(blocksize in (32, 64, 128, 256, 512, 1024, 2048, 4096), lambda: f"invalid blocksize {blocksize}")
+ torch._check(
+ A.dtype in (torch.float16, torch.bfloat16, torch.float32),
+ lambda: f"A must be float16, bfloat16, or float32, got {A.dtype}",
+ )
+ torch._check(
+ B.dtype in (torch.uint8, torch.bfloat16, torch.float16, torch.float32),
+ lambda: f"B must be backed by storage of type uint8, bfloat16, float16, or float32, got {B.dtype}",
+ )
+ torch._check(
+ out.shape == (*A.shape[:-1], shapeB[0]),
+ lambda: f"Expected out.shape == {(*A.shape[:-1], shapeB[0])}, got {out.shape}",
+ )
+ torch._check(out.device == A.device, lambda: f"Expected out.device == {A.device}, got {out.device}")
+ torch._check(out.dtype == A.dtype, lambda: f"Expected out.dtype == {A.dtype}, got {out.dtype}")
+
+
+torch.library.define(
+ "bitsandbytes::optimizer_update_32bit",
+ "(str optimizer_name, Tensor(a0!) g, Tensor(a1!) p, Tensor(a2!) state1, Tensor(a3!)? state2, Tensor(a4!)? unorm_vec, float max_unorm, float param_norm, float beta1, float beta2, float beta3, float alpha, float eps, float weight_decay, int step, float lr, float gnorm_scale, bool skip_zeros=False) -> ()",
+)
+
+
+@register_fake("bitsandbytes::optimizer_update_32bit")
+def _(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float,
+ skip_zeros=False,
+) -> None:
+ torch._check(
+ g.numel() == p.numel(),
+ lambda: f"g and p must have the same number of elements, got {g.numel()} and {p.numel()}",
+ )
+ compute_dtypes = [torch.float16, torch.bfloat16, torch.float32]
+
+ torch._check(
+ g.dtype in compute_dtypes,
+ lambda: f"g must be bfloat16, float16, or float32, got {g.dtype}",
+ )
+ torch._check(
+ g.dtype == p.dtype,
+ lambda: f"Expected all tensors to have the same dtype, got g.dtype={g.dtype}, p.dtype={p.dtype}",
+ )
+
+
+torch.library.define(
+ "bitsandbytes::optimizer_update_8bit_blockwise",
+ "(str optimizer_name, Tensor(a0!) g, Tensor(a1!) p, Tensor(a2!) state1, Tensor(a3!)? state2, float beta1, float beta2, float beta3, float alpha, float eps, int step, float lr, Tensor(a4!) qmap1, Tensor(a5!)? qmap2, Tensor(a6!) absmax1, Tensor(a7!)? absmax2, float weight_decay, float gnorm_scale, bool skip_zeros=False) -> ()",
+)
+
+
+@register_fake("bitsandbytes::optimizer_update_8bit_blockwise")
+def _(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: torch.Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: torch.Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float,
+ gnorm_scale: float,
+ skip_zeros=False,
+) -> None:
+ torch._check(
+ g.numel() == p.numel(),
+ lambda: f"g and p must have the same number of elements, got {g.numel()} and {p.numel()}",
+ )
+ compute_dtypes = [torch.float16, torch.bfloat16, torch.float32]
+
+ torch._check(
+ g.dtype in compute_dtypes,
+ lambda: f"g must be bfloat16, float16, or float32, got {g.dtype}",
+ )
+ torch._check(
+ g.dtype == p.dtype,
+ lambda: f"Expected all tensors to have the same dtype, got g.dtype={g.dtype}, p.dtype={p.dtype}",
+ )
+ torch._check(
+ state1.dtype == torch.uint8,
+ lambda: f"state1 must be uint8, got {state1.dtype}",
+ )
+ torch._check(
+ qmap1.dtype == absmax1.dtype == torch.float32,
+ lambda: f"Expected qmap1 and absmax1 to be float32, got qmap1.dtype={qmap1.dtype}, absmax1.dtype={absmax1.dtype}",
+ )
+ if state2 is not None:
+ torch._check(
+ state2.dtype == torch.uint8,
+ lambda: f"state2 must be uint8, got {state2.dtype}",
+ )
+ torch._check(
+ qmap2.dtype == absmax2.dtype == torch.float32,
+ lambda: f"Expected qmap2 and absmax2 to be float32, got qmap2.dtype={qmap2.dtype}, absmax2.dtype={absmax2.dtype}",
+ )
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/autograd/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/autograd/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/autograd/_functions.py b/venv/lib/python3.11/site-packages/bitsandbytes/autograd/_functions.py
new file mode 100644
index 0000000000000000000000000000000000000000..8a069bd101a089afc59b7d4d3d48e5dd399bc967
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/autograd/_functions.py
@@ -0,0 +1,491 @@
+from dataclasses import dataclass
+import logging
+from math import prod
+from typing import Optional
+import warnings
+from warnings import warn
+
+import torch
+
+import bitsandbytes.functional as F
+
+logger = logging.getLogger(__name__)
+
+# The inverse transformation for the colTuring and colAmpere format were contributed by Alex Borzunov:
+# https://github.com/bigscience-workshop/petals/blob/main/src/petals/utils/linear8bitlt_patch.py
+
+
+"""
+ This class pools outlier dimensions across layers.
+ This is particularly important for small models where outlier features
+ are less systematic and occur with low frequency.
+"""
+
+
+class GlobalOutlierPooler:
+ _instance = None
+
+ def __init__(self):
+ raise RuntimeError("Call get_instance() instead")
+
+ def initialize(self):
+ self.outliers = set()
+ self.model_dim = None
+
+ @classmethod
+ def get_instance(cls):
+ if cls._instance is None:
+ cls._instance = cls.__new__(cls)
+ cls._instance.initialize()
+ return cls._instance
+
+ def add_outliers(self, outlier_idx, feature_dim):
+ if self.model_dim is None:
+ self.model_dim = feature_dim
+ if feature_dim != self.model_dim:
+ return # we do not encode outliers for the 2nd FFN layer
+
+ self.outliers.update(outlier_idx.tolist())
+
+ def get_current_outlier_idx(self):
+ return torch.Tensor(list(self.outliers)).to(torch.int64)
+
+
+_is_compiling = torch.compiler.is_compiling
+
+
+@dataclass
+class MatmulLtState:
+ force_no_igemmlt: bool = False
+
+ CB: Optional[torch.Tensor] = None
+ SB: Optional[torch.Tensor] = None
+ SCB: Optional[torch.Tensor] = None
+
+ SBt: Optional[torch.Tensor] = None
+ CBt: Optional[torch.Tensor] = None
+
+ subB: Optional[torch.Tensor] = None
+
+ outlier_pool: Optional[GlobalOutlierPooler] = None
+ has_accumulated_gradients = False
+ threshold = 0.0
+ idx: Optional[torch.Tensor] = None
+ is_training = True
+ has_fp16_weights = True
+ use_pool = False
+
+ # Deprecated attributes kept for downstream compatibility (TGI, vLLM).
+ # These are always None and will be fully removed in the next release.
+ _deprecated_fields = frozenset({"CxB", "CxBt", "formatB", "_tile_indices"})
+
+ def __getattr__(self, name):
+ if name in MatmulLtState._deprecated_fields:
+ warnings.warn(
+ f"MatmulLtState.{name} is deprecated and will be removed in the next bitsandbytes release.",
+ FutureWarning,
+ stacklevel=2,
+ )
+ return None
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
+
+ def reset_grads(self):
+ self.CB = None
+ self.SB = None
+ self.SCB = None
+
+ self.SBt = None
+ self.CBt = None
+
+
+class MatMul8bitLt(torch.autograd.Function):
+ @staticmethod
+ def forward(
+ ctx: torch.autograd.function.FunctionCtx,
+ A: torch.Tensor,
+ B: torch.Tensor,
+ out: Optional[torch.Tensor] = None,
+ bias: Optional[torch.Tensor] = None,
+ state: Optional[MatmulLtState] = None,
+ ):
+ state = state or MatmulLtState()
+
+ # default of pytorch behavior if inputs are empty
+ ctx.is_empty = False
+ if prod(A.shape) == 0:
+ ctx.is_empty = True
+ ctx.A = A
+ ctx.B = B
+ ctx.bias = bias
+ if A.shape[-1] == B.shape[0]:
+ return torch.empty(A.shape[:-1] + B.shape[1:], dtype=A.dtype, device=A.device)
+ else:
+ return torch.empty(A.shape[:-1] + B.shape[:1], dtype=A.dtype, device=A.device)
+
+ input_shape = A.shape
+
+ # Cast A to fp16
+ if A.dtype != torch.float16 and not _is_compiling():
+ logger.warning("MatMul8bitLt: inputs will be cast from %s to float16 during quantization", A.dtype)
+
+ if len(A.shape) == 3:
+ A = A.reshape(-1, A.shape[-1])
+
+ # 1. Quantize A. Note that as a side-effect, outliers are suppressed in CA/CAt.
+ if ctx.needs_input_grad[1]:
+ # Slower path
+ CA, CAt, SCA, SCAt, outlier_cols = F.int8_double_quant(A.to(torch.float16), threshold=state.threshold)
+ else:
+ # Fast path
+ CA, SCA, outlier_cols = F.int8_vectorwise_quant(A.to(torch.float16), threshold=state.threshold)
+ CAt = SCAt = None
+
+ has_grad = False
+
+ if state.has_fp16_weights or state.CB is None:
+ has_grad = getattr(B, "grad", None) is not None
+ is_transposed = not B.is_contiguous() and B.shape[0] == B.stride(1)
+ if is_transposed:
+ B = B.contiguous()
+
+ if (state.is_training and not has_grad) or state.CB is None or state.SCB is None:
+ state.reset_grads()
+
+ # 2. Quantize B
+ state.CB, state.SCB, _ = F.int8_vectorwise_quant(B.to(torch.float16))
+
+ # Handle sparse decomposition
+ if state.threshold > 0.0:
+ state.idx = outlier_cols
+
+ # Mixed Int8 Matmul + Dequant + Bias
+ output, subA = torch.ops.bitsandbytes.int8_mixed_scaled_mm(
+ A,
+ CA,
+ state.CB,
+ SCA,
+ state.SCB,
+ outlier_cols,
+ bias,
+ )
+
+ else:
+ # Int8 Matmul + Dequant + Bias
+ output = torch.ops.bitsandbytes.int8_scaled_mm.default(
+ CA, state.CB, SCA, state.SCB, bias=bias, dtype=A.dtype
+ )
+ subA = None
+
+ # 5. Save state
+ ctx.state = state
+
+ ctx.grad_shape = input_shape
+ ctx.dtype_A = A.dtype
+ ctx.dtype_bias = None if bias is None else bias.dtype
+
+ if any(ctx.needs_input_grad[:2]):
+ ctx.tensors = (CAt, subA, A)
+ ctx.tensor_states = (SCAt, state.idx)
+ else:
+ ctx.tensors = [None, None, None]
+ ctx.tensor_states = (None, None)
+ ctx.save_for_backward(None, None)
+
+ output_shape = (*input_shape[:-1], state.CB.shape[0])
+
+ if len(input_shape) == 3:
+ return output.reshape(output_shape)
+
+ return output
+
+ @staticmethod
+ def backward(ctx: torch.autograd.function.FunctionCtx, grad_output: torch.Tensor):
+ if ctx.is_empty:
+ bias_grad = None if ctx.bias is None else torch.zeros_like(ctx.bias)
+ return torch.zeros_like(ctx.A), torch.zeros_like(ctx.B), None, bias_grad, None
+
+ req_gradA, req_gradB, _, req_gradBias, _ = ctx.needs_input_grad
+ CAt, subA, _A = ctx.tensors
+ SCAt, idx = ctx.tensor_states
+ state: MatmulLtState = ctx.state
+ grad_A = grad_B = grad_bias = None
+
+ if req_gradBias:
+ # compute grad_bias first before changing grad_output dtype
+ grad_bias = grad_output.sum(0, dtype=ctx.dtype_bias)
+
+ # Cast grad_output to fp16
+ if len(grad_output.shape) == 3:
+ grad_output = grad_output.reshape(-1, grad_output.shape[-1]).contiguous()
+
+ if req_gradB:
+ Cgrad, _, _, SCgradt, _ = F.int8_double_quant(grad_output.to(torch.float16))
+
+ grad_B = torch.ops.bitsandbytes.int8_scaled_mm.default(
+ Cgrad.t().contiguous(),
+ CAt.t(),
+ SCgradt,
+ SCAt,
+ dtype=torch.float16,
+ )
+
+ if state.threshold > 0.0 and subA is not None and subA.numel() > 0:
+ grad_B[:, idx] += torch.matmul(grad_output.t(), subA)
+
+ if req_gradA:
+ if state.CB is not None:
+ CB = state.CB.to(ctx.dtype_A, copy=True).mul_(state.SCB.unsqueeze(1).mul(1.0 / 127.0))
+ grad_A = torch.matmul(grad_output.to(ctx.dtype_A), CB).view(ctx.grad_shape)
+ else:
+ raise Exception("State must contain CB matrix for backward")
+
+ return grad_A, grad_B, None, grad_bias, None
+
+
+class MatMul8bitFp(torch.autograd.Function):
+ # For Intel CPU and XPU MatMul8bitFp is much faster (~3x) than MatMul8bitLt in finetune.
+ # Because the MatMul8bitLt has more mechanisms in computing grad.
+ # We don't have fast kernel for quant/dequant 8bit in CPU/XPU, so it's very slow.
+ # We'd like to use dequant + matmul to run finetune with good performance.
+
+ @staticmethod
+ def forward(ctx, A, B, out=None, bias=None, state=MatmulLtState):
+ if state.has_fp16_weights or state.CB is None:
+ has_grad = getattr(B, "grad", None) is not None
+ is_transposed = not B.is_contiguous() and B.shape[0] == B.stride(1)
+ if is_transposed:
+ B = B.contiguous()
+
+ if (state.is_training and not has_grad) or state.CB is None or state.SCB is None:
+ state.reset_grads()
+ state.CB, state.SCB, _ = F.int8_vectorwise_quant(B.to(torch.float16))
+ B = state.CB
+
+ CB = state.CB.data.to(A.dtype).mul_(state.SCB.unsqueeze(1).mul(1.0 / 127.0))
+ output = torch.nn.functional.linear(A, CB, bias)
+ ctx.state = state
+ ctx.dtype_A = A.dtype
+ ctx.grad_shape = A.shape
+ ctx.A = A
+ ctx.dtype_bias = None if bias is None else bias.dtype
+ return output
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ req_gradA, req_gradB, _, req_gradBias, _ = ctx.needs_input_grad
+ A = ctx.A
+ state = ctx.state
+ grad_A = grad_B = grad_bias = None
+ if req_gradBias:
+ # compute grad_bias first before changing grad_output dtype
+ grad_bias = grad_output.sum(0, dtype=ctx.dtype_bias)
+
+ # Cast grad_output to fp16
+ if len(grad_output.shape) == 3:
+ grad_output = grad_output.reshape(-1, grad_output.shape[-1]).contiguous()
+
+ if req_gradB:
+ grad_B = torch.matmul(A.t(), grad_output).t()
+
+ if req_gradA:
+ if state.CB is not None:
+ CB = state.CB.to(ctx.dtype_A, copy=True).mul_(state.SCB.unsqueeze(1).mul(1.0 / 127.0))
+ grad_A = torch.matmul(grad_output.to(ctx.dtype_A), CB).view(ctx.grad_shape)
+ else:
+ raise Exception("State must contain CB matrix for backward")
+
+ return grad_A, grad_B, None, grad_bias, None
+
+
+class MatMul4Bit(torch.autograd.Function):
+ # forward is the same, but we added the fallback for pre-turing GPUs
+
+ @staticmethod
+ def forward(ctx, A, B, out=None, bias=None, quant_state: Optional[F.QuantState] = None):
+ # default of pytorch behavior if inputs are empty
+ ctx.is_empty = False
+ if A.numel() == 0:
+ ctx.is_empty = True
+ ctx.A = A
+ ctx.B = B
+ ctx.bias = bias
+ B_shape = quant_state.shape
+ if A.shape[-1] == B_shape[0]:
+ return torch.empty(A.shape[:-1] + B_shape[1:], dtype=A.dtype, device=A.device)
+ else:
+ return torch.empty(A.shape[:-1] + B_shape[:1], dtype=A.dtype, device=A.device)
+
+ # Normalize to canonical [(N*K+1)//2, 1]. Packed weights are always contiguous
+ # in this orientation (B.t() callers get strides [1,1], still compatible).
+ # quant_state.shape is the source of truth for N and K.
+ B = B.view(-1, 1)
+
+ if not quant_state.nested:
+ output = torch.ops.bitsandbytes.gemm_4bit.default(
+ A,
+ B,
+ quant_state.shape,
+ quant_state.absmax,
+ quant_state.blocksize,
+ quant_state.quant_type,
+ bias=bias,
+ )
+ elif quant_state.state2.blocksize == 256:
+ output = torch.ops.bitsandbytes.gemm_4bit.default(
+ A,
+ B,
+ quant_state.shape,
+ quant_state.state2.absmax,
+ quant_state.blocksize,
+ quant_state.quant_type,
+ bias=bias,
+ absmax_8bit=quant_state.absmax,
+ absmax_code=quant_state.state2.code,
+ absmax_offset=quant_state.offset,
+ )
+ else:
+ raise NotImplementedError("nested quantization with state2.blocksize != 256 is not supported")
+
+ if out is not None:
+ out.copy_(output)
+ output = out
+
+ # 3. Save state
+ ctx.state = quant_state
+ ctx.dtype_A, ctx.dtype_B, ctx.dtype_bias = A.dtype, B.dtype, None if bias is None else bias.dtype
+
+ if any(ctx.needs_input_grad[:2]):
+ ctx.tensors = (None, B)
+ else:
+ ctx.tensors = (None, None)
+
+ return output
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ if ctx.is_empty:
+ bias_grad = None if ctx.bias is None else torch.zeros_like(ctx.bias)
+ return torch.zeros_like(ctx.A), torch.zeros_like(ctx.B), None, bias_grad, None
+
+ req_gradA, _, _, req_gradBias, _ = ctx.needs_input_grad
+ _, B = ctx.tensors
+
+ grad_A, grad_B, grad_bias = None, None, None
+
+ if req_gradBias:
+ # compute grad_bias first before changing grad_output dtype
+ grad_bias = grad_output.sum(0, dtype=ctx.dtype_bias)
+
+ # not supported by PyTorch. TODO: create work-around
+ # if req_gradB: grad_B = torch.matmul(grad_output.t(), A)
+ if req_gradA:
+ # B in ctx.tensors is already in canonical [(N*K+1)//2, 1] form (normalized in forward).
+ # dequantize returns [N, K]; matmul(grad_output[M,N], [N,K]) = grad_A[M,K].
+ grad_A = torch.matmul(grad_output, F.dequantize_4bit(B, ctx.state).to(grad_output.dtype))
+
+ return grad_A, grad_B, None, grad_bias, None
+
+
+def matmul(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ out: Optional[torch.Tensor] = None,
+ state: Optional[MatmulLtState] = None,
+ threshold=0.0,
+ bias: Optional[torch.Tensor] = None,
+):
+ state = state or MatmulLtState()
+ if threshold > 0.0:
+ state.threshold = threshold
+ # MatMul8bitLt is slower because no fast kernel for quant/dequant 8bit in CPU/XPU
+ if state.is_training:
+ if A.device.type in ("cpu", "xpu"):
+ return MatMul8bitFp.apply(A, B, out, bias, state)
+ return MatMul8bitLt.apply(A, B, out, bias, state)
+
+
+def matmul_4bit(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ quant_state: F.QuantState,
+ out: Optional[torch.Tensor] = None,
+ bias: Optional[torch.Tensor] = None,
+):
+ if quant_state is None:
+ raise ValueError("quant_state is required")
+ if len(quant_state.shape) != 2:
+ raise ValueError("matmul_4bit: quant_state.shape must be 2D [N, K]")
+
+ # packing_format_for_cpu uses a different memory layout optimized for AVX512BF16.
+ # This flag is only set for inference (weight conversion happens at eval time).
+ # The underlying kernel supports any M via tiled GEMM despite the gemv name.
+ if A.device.type == "cpu" and getattr(quant_state, "packing_format_for_cpu", False):
+ result = F.gemv_4bit(A, B, out=out, state=quant_state)
+ if bias is not None:
+ result += bias
+ return result
+
+ # Normalize B to canonical [(N*K+1)//2, 1]. Packed weights are always contiguous
+ # in this orientation (B.t() callers get strides [1,1], still compatible).
+ # quant_state.shape is the source of truth for N and K.
+ B = B.view(-1, 1)
+
+ K = A.shape[-1]
+
+ # Weight is in [K, N] orientation when A's inner dim matches shape[0] not shape[1].
+ # Square weights (K==N) are ambiguous and treated as [N, K].
+ if K == quant_state.shape[0] and K != quant_state.shape[1]:
+ if not _is_compiling():
+ warn(
+ f"matmul_4bit: weight was quantized from a [K, N] tensor (quant_state.shape={list(quant_state.shape)}). "
+ "Re-quantize from the weight in [N, K] (out_features, in_features) orientation. "
+ "This will be an error in a future version.",
+ DeprecationWarning,
+ stacklevel=2,
+ )
+ B_dq = F.dequantize_4bit(B, quant_state).to(A.dtype)
+ result = torch.nn.functional.linear(A, B_dq.t(), bias)
+ if out is not None:
+ out.copy_(result)
+ return out
+ return result
+
+ needs_grad = torch.is_grad_enabled() and (A.requires_grad or (bias is not None and bias.requires_grad))
+ if not needs_grad:
+ A_numel = A.numel()
+ if A_numel == 0:
+ if out is not None:
+ return out
+ return torch.empty((*A.shape[:-1], quant_state.shape[0]), dtype=A.dtype, device=A.device)
+
+ if not quant_state.nested:
+ result = torch.ops.bitsandbytes.gemm_4bit.default(
+ A,
+ B,
+ quant_state.shape,
+ quant_state.absmax,
+ quant_state.blocksize,
+ quant_state.quant_type,
+ bias=bias,
+ )
+ elif quant_state.state2.blocksize == 256:
+ result = torch.ops.bitsandbytes.gemm_4bit.default(
+ A,
+ B,
+ quant_state.shape,
+ quant_state.state2.absmax,
+ quant_state.blocksize,
+ quant_state.quant_type,
+ bias=bias,
+ absmax_8bit=quant_state.absmax,
+ absmax_code=quant_state.state2.code,
+ absmax_offset=quant_state.offset,
+ )
+ else:
+ raise NotImplementedError("nested quantization with state2.blocksize != 256 is not supported")
+ if out is not None:
+ out.copy_(result)
+ return out
+ return result
+
+ return MatMul4Bit.apply(A, B, out, bias, quant_state)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/cpu/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/cpu/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/cpu/ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/cpu/ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..44fb5dcebc794d7cd8e5a509f5144016d775fbee
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/cpu/ops.py
@@ -0,0 +1,580 @@
+from collections.abc import Sequence
+import ctypes as ct
+import logging
+import math
+from math import prod
+from typing import Optional
+
+import torch
+
+from bitsandbytes.functional import get_ptr, has_avx512bf16
+
+from ..._ops import register_kernel
+from ...cextension import ErrorHandlerMockBNBNativeLibrary, lib
+
+logger = logging.getLogger(__name__)
+
+_has_avx512 = torch.backends.cpu.get_cpu_capability() == "AVX512"
+
+# torch._int_mm for s8@s8->s32 is supported on CPU from torch 2.4+.
+# However, we can overflow if we use this without AVX512_VNNI support.
+# This is fixed in torch 2.6+, so we set this as the minimum to be safe.
+# For more information: https://github.com/pytorch/pytorch/pull/136942
+#
+# Without AVX-512 (including aarch64), torch._int_mm uses a scalar fallback
+# that is much slower than fp32 matmul. Only use it when AVX-512 is available.
+if torch.__version__ >= (2, 6) and _has_avx512:
+
+ @register_kernel("bitsandbytes::int8_linear_matmul", "cpu")
+ def _(A: torch.Tensor, B: torch.Tensor):
+ return torch._int_mm(
+ A.reshape(-1, A.shape[-1]),
+ B.t(),
+ ).reshape(*A.shape[:-1], B.shape[0])
+
+
+if not isinstance(lib, ErrorHandlerMockBNBNativeLibrary):
+
+ @register_kernel("bitsandbytes::quantize_blockwise", "cpu")
+ def _(A: torch.Tensor, code: torch.Tensor, blocksize: int) -> tuple[torch.Tensor, torch.Tensor]:
+ A = A.contiguous()
+ n = A.numel()
+ blocks = -(n // -blocksize)
+
+ absmax = torch.empty((blocks,), device=A.device, dtype=torch.float32)
+ out = torch.empty(A.shape, device=A.device, dtype=torch.uint8)
+
+ if A.dtype == torch.float32:
+ lib.cquantize_blockwise_cpu_fp32(
+ get_ptr(code),
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(n),
+ )
+ elif A.dtype == torch.bfloat16:
+ lib.cquantize_blockwise_cpu_bf16(
+ get_ptr(code),
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(n),
+ )
+ elif A.dtype == torch.float16:
+ lib.cquantize_blockwise_cpu_fp16(
+ get_ptr(code),
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(n),
+ )
+ else:
+ # Generic fallback for other dtypes
+ A_flat = A.reshape(n).float()
+ rem = n % blocksize
+ has_rem = rem > 0
+ A_com = A_flat[: n - rem]
+ A_com_reshaped = A_com.reshape(n // blocksize, blocksize)
+ absmax[: blocks - has_rem] = torch.abs(A_com_reshaped).max(dim=-1)[0]
+ scaled_A = torch.clamp(A_com_reshaped * (1 / absmax[: blocks - has_rem].view(-1, 1)), -1, 1)
+ scaled_A = scaled_A.reshape(-1)
+ if has_rem:
+ absmax[-1] = torch.abs(A_flat[n - rem :]).max()
+ scaled_A_rem = torch.clamp(A_flat[n - rem :] * (1 / absmax[-1]), -1, 1)
+ scaled_A = torch.cat([scaled_A, scaled_A_rem], dim=0)
+
+ diff = torch.abs(scaled_A.unsqueeze(-1) - code.to(scaled_A.device))
+ out = torch.argmin(diff, dim=-1).to(torch.uint8).to(scaled_A.device).reshape(A.shape)
+
+ return out, absmax
+
+ @register_kernel("bitsandbytes::dequantize_blockwise", "cpu")
+ def _(
+ A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype
+ ) -> torch.Tensor:
+ A = A.contiguous()
+ out = torch.empty_like(A, dtype=dtype)
+ if dtype == torch.float32:
+ lib.cdequantize_blockwise_cpu_fp32(
+ get_ptr(code),
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(A.numel()),
+ )
+ elif dtype == torch.bfloat16:
+ lib.cdequantize_blockwise_cpu_bf16(
+ get_ptr(code),
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(A.numel()),
+ )
+ elif dtype == torch.float16:
+ lib.cdequantize_blockwise_cpu_fp16(
+ get_ptr(code),
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(A.numel()),
+ )
+ else:
+ out = code[A.reshape(-1).int()]
+ blocks = out.shape[-1] // blocksize
+ res = out.shape[-1] % blocksize
+ if res != 0:
+ out = torch.nn.functional.pad(out, (0, blocksize - res), mode="constant", value=0)
+ out = (out.view(-1, blocksize) * absmax.view(-1, 1)).to(dtype).reshape(-1)
+ out = out[: blocks * blocksize + res]
+ out = out.reshape(A.shape)
+
+ return out
+
+ @register_kernel("bitsandbytes::dequantize_4bit", "cpu")
+ def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+ ) -> torch.Tensor:
+ # Fallback as AVX512 implementation has accuracy issues with blocksize >= 2048.
+ # Note: this is not a common use case.
+ avx512_fallback = _has_avx512 and blocksize >= 2048
+
+ # Odd shape is not supported by this kernel; fallback to generic implementation
+ shape_fallback = shape[-1] % 2 != 0
+
+ if avx512_fallback or shape_fallback:
+ from ..default.ops import _dequantize_4bit_compute
+ from ..utils import _get_4bit_code
+
+ if A.dtype != torch.uint8:
+ A = A.view(torch.uint8)
+ code = _get_4bit_code(quant_type, A.device)
+ return _dequantize_4bit_compute(A.reshape(-1), absmax, code, blocksize, shape, dtype)
+
+ # Enable non uint8 dtype
+ if A.dtype != torch.uint8:
+ A = A.view(torch.uint8)
+
+ # TODO: support half precision absmax
+ if absmax.dtype != torch.float32:
+ absmax = absmax.float()
+
+ if len(shape) == 1:
+ shape = (1, shape[0])
+
+ m = prod(shape[:-1])
+ n = shape[-1]
+
+ A = A.reshape(m, n // 2)
+ out = torch.empty(shape, dtype=dtype, device=A.device)
+
+ if quant_type == "fp4":
+ if dtype == torch.float32:
+ lib.cdequantize_blockwise_cpu_fp4_fp32(
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(m),
+ ct.c_longlong(n),
+ )
+ elif dtype == torch.bfloat16:
+ lib.cdequantize_blockwise_cpu_fp4_bf16(
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(m),
+ ct.c_longlong(n),
+ )
+ elif dtype == torch.float16:
+ lib.cdequantize_blockwise_cpu_fp4_fp16(
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(m),
+ ct.c_longlong(n),
+ )
+ elif quant_type == "nf4":
+ if dtype == torch.float32:
+ lib.cdequantize_blockwise_cpu_nf4_fp32(
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(m),
+ ct.c_longlong(n),
+ )
+ elif dtype == torch.bfloat16:
+ lib.cdequantize_blockwise_cpu_nf4_bf16(
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(m),
+ ct.c_longlong(n),
+ )
+ elif dtype == torch.float16:
+ lib.cdequantize_blockwise_cpu_nf4_fp16(
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_longlong(blocksize),
+ ct.c_longlong(m),
+ ct.c_longlong(n),
+ )
+ else:
+ raise ValueError
+
+ return out
+
+ if has_avx512bf16():
+ gemm_4bit_forward_kernel = None
+ try:
+ from kernels import get_kernel
+
+ gemm_4bit_forward_kernel = get_kernel(
+ "kernels-community/quantization-bitsandbytes", version=1
+ ).gemm_4bit_forward
+ except Exception as exc: # pragma: no cover - best effort fallback
+ gemm_4bit_forward_kernel = None
+ logger.warning(
+ "Failed to load CPU gemm_4bit_forward from kernels-community: %s. Please make sure you already `pip install kernels` and the kernels >= 0.11.1",
+ exc,
+ )
+
+ @register_kernel("bitsandbytes::gemv_4bit", "cpu")
+ def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ ) -> torch.Tensor:
+ if B.dtype != torch.uint8:
+ B = B.contiguous().view(torch.uint8)
+ dtype = A.dtype
+ quant_type = "fp4" if code[1] > 0 else "nf4"
+ # cpu fused op only support bf16 for now.
+ if dtype != torch.bfloat16:
+ A = A.to(torch.bfloat16)
+ if absmax.dtype != torch.bfloat16:
+ absmax = absmax.to(torch.bfloat16)
+
+ final_out_shape = (*A.shape[:-1], shapeB[0])
+ A = A.reshape(-1, A.shape[-1])
+ out_shape = (*A.shape[:-1], shapeB[0])
+ if gemm_4bit_forward_kernel is not None:
+ quant_type_num = 1 if quant_type == "fp4" else 0
+ # C++ kernel expects weight shape (N, K_packed), ensure 2D contiguous
+ B_2d = B.reshape(shapeB[0], -1).contiguous()
+ out = gemm_4bit_forward_kernel(A, B_2d, absmax, blocksize, quant_type_num)
+ else:
+ out = torch.empty(out_shape, dtype=A.dtype, device=A.device)
+ M = A.shape[0]
+ N = shapeB[0]
+ K = A.shape[1]
+ x_strideM = A.stride(0)
+ out_strideM = out.stride(0)
+ if quant_type == "fp4":
+ lib.gemv_4bit_inference_cpu_fp4_bf16(
+ ct.c_int64(M),
+ ct.c_int64(N),
+ ct.c_int64(K),
+ get_ptr(A),
+ get_ptr(B),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_int64(blocksize),
+ ct.c_int64(x_strideM),
+ ct.c_int64(out_strideM),
+ )
+ elif quant_type == "nf4":
+ lib.gemv_4bit_inference_cpu_nf4_bf16(
+ ct.c_int64(M),
+ ct.c_int64(N),
+ ct.c_int64(K),
+ get_ptr(A),
+ get_ptr(B),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_int64(blocksize),
+ ct.c_int64(x_strideM),
+ ct.c_int64(out_strideM),
+ )
+
+ if dtype != torch.bfloat16:
+ out = out.to(dtype)
+
+ return out.reshape(final_out_shape)
+
+
+# ==================== CPU Optimizer Kernels ====================
+
+
+def _compute_update_norm_and_scale(
+ update: torch.Tensor,
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+) -> float:
+ """Compute trust-ratio scaling factor for LAMB/LARS and store update norm."""
+ if max_unorm <= 0.0:
+ return 1.0
+ unorm = torch.norm(update).item()
+ if unorm_vec is not None:
+ unorm_vec.fill_(unorm)
+ if unorm > max_unorm * param_norm:
+ return (max_unorm * param_norm) / unorm
+ return 1.0
+
+
+@torch.no_grad()
+def _optimizer_update_32bit_cpu(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float,
+ skip_zeros: bool = False,
+) -> None:
+ g_float = g.float() * gnorm_scale
+ p_float = p.data.float()
+
+ if optimizer_name in ("adam", "lamb"):
+ # Adam / LAMB (2-state): m and v
+ state1.mul_(beta1).add_(g_float, alpha=1.0 - beta1)
+ state2.mul_(beta2).addcmul_(g_float, g_float, value=1.0 - beta2)
+
+ correction1 = 1.0 - beta1**step
+ correction2 = math.sqrt(1.0 - beta2**step)
+ step_size = -lr * correction2 / correction1
+
+ if weight_decay > 0.0:
+ p_float.mul_(1.0 - lr * weight_decay)
+
+ update = state1 / (state2.sqrt() + eps * correction2)
+
+ update_scale = _compute_update_norm_and_scale(update, unorm_vec, max_unorm, param_norm)
+ p_float.add_(update, alpha=step_size * update_scale)
+
+ elif optimizer_name == "ademamix":
+ # AdEMAMix (2-state): state1 shape is (2, *p.shape), state1[0]=m1, state1[1]=m2
+ m1 = state1[0]
+ m2 = state1[1]
+ nu = state2
+
+ m1.mul_(beta1).add_(g_float, alpha=1.0 - beta1)
+ m2.mul_(beta3).add_(g_float, alpha=1.0 - beta3)
+ nu.mul_(beta2).addcmul_(g_float, g_float, value=1.0 - beta2)
+
+ correction1 = 1.0 - beta1**step
+ correction2 = math.sqrt(1.0 - beta2**step)
+
+ if weight_decay > 0.0:
+ p_float.mul_(1.0 - lr * weight_decay)
+
+ mixed_momentum = (m1 / correction1) + (alpha * m2)
+ adaptive_term = (nu.sqrt() / correction2) + eps
+ p_float.add_(mixed_momentum / adaptive_term, alpha=-lr)
+
+ elif optimizer_name in ("momentum", "lars"):
+ # SGD with momentum / LARS (1-state)
+ g_wd = g_float.add(p_float, alpha=weight_decay) if weight_decay > 0.0 else g_float
+
+ if step == 1:
+ state1.copy_(g_wd)
+ else:
+ state1.mul_(beta1).add_(g_wd)
+
+ update_scale = _compute_update_norm_and_scale(state1, unorm_vec, max_unorm, param_norm)
+ p_float.add_(state1, alpha=-lr * update_scale)
+
+ elif optimizer_name == "lion":
+ # Lion (2-state sign update)
+ if weight_decay > 0.0:
+ p_float.mul_(1.0 - lr * weight_decay)
+
+ update = state1.mul(beta1).add(g_float, alpha=1.0 - beta1)
+ p_float.add_(update.sign(), alpha=-lr)
+
+ state1.mul_(beta2).add_(g_float, alpha=1.0 - beta2)
+
+ elif optimizer_name == "rmsprop":
+ # RMSprop (1-state)
+ g_wd = g_float.add(p_float, alpha=weight_decay) if weight_decay > 0.0 else g_float
+ state1.mul_(beta1).addcmul_(g_wd, g_wd, value=1.0 - beta1)
+
+ update = g_wd / (state1.sqrt() + eps)
+ update_scale = _compute_update_norm_and_scale(update, unorm_vec, max_unorm, param_norm)
+ p_float.add_(update, alpha=-lr * update_scale)
+
+ elif optimizer_name == "adagrad":
+ # Adagrad (1-state)
+ g_wd = g_float.add(p_float, alpha=weight_decay) if weight_decay > 0.0 else g_float
+ state1.addcmul_(g_wd, g_wd, value=1.0)
+
+ update = g_wd / (state1.sqrt() + eps)
+ p_float.add_(update, alpha=-lr)
+
+ else:
+ raise ValueError(f"Unsupported optimizer for CPU: {optimizer_name}")
+
+ # Write back to original precision
+ p.data.copy_(p_float)
+
+
+register_kernel("bitsandbytes::optimizer_update_32bit", "cpu")(_optimizer_update_32bit_cpu)
+
+
+@torch.no_grad()
+def _dequant_blockwise_fp32_direct(
+ A_uint8: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int
+) -> torch.Tensor:
+ return torch.ops.bitsandbytes.dequantize_blockwise(A_uint8, absmax, code, blocksize, torch.float32)
+
+
+def _quant_blockwise_fp32_direct(
+ A_fp32: torch.Tensor, code: torch.Tensor, absmax_out: torch.Tensor, out_uint8: torch.Tensor, blocksize: int
+) -> None:
+ out, absmax = torch.ops.bitsandbytes.quantize_blockwise(A_fp32, code, blocksize)
+ out_uint8.copy_(out)
+ absmax_out.copy_(absmax)
+
+
+def _optimizer_update_8bit_blockwise_cpu(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: torch.Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: torch.Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float,
+ gnorm_scale: float,
+ skip_zeros: bool = False,
+) -> None:
+ blocksize = 256
+
+ # Dequantize states
+ if optimizer_name == "ademamix" and absmax1.ndim == 2:
+ s1_1 = _dequant_blockwise_fp32_direct(state1[0], absmax1[0], qmap1, blocksize)
+ s1_2 = _dequant_blockwise_fp32_direct(state1[1], absmax1[1], qmap1, blocksize)
+ state1_fp32 = torch.stack([s1_1, s1_2])
+ else:
+ state1_fp32 = _dequant_blockwise_fp32_direct(state1, absmax1, qmap1, blocksize)
+
+ state2_fp32 = None
+ if state2 is not None and qmap2 is not None and absmax2 is not None:
+ state2_fp32 = _dequant_blockwise_fp32_direct(state2, absmax2, qmap2, blocksize)
+
+ grad = g.float() * gnorm_scale
+ p_fp32 = p.data.float()
+
+ if optimizer_name in ("adam", "lamb"):
+ state1_fp32.mul_(beta1).add_(grad, alpha=1.0 - beta1)
+ state2_fp32.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
+
+ correction1 = 1.0 - beta1**step
+ correction2 = math.sqrt(1.0 - beta2**step)
+
+ denom = (state2_fp32.sqrt() / correction2).add_(eps)
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+ p_fp32.addcdiv_(state1_fp32, denom, value=-lr / correction1)
+
+ elif optimizer_name == "ademamix":
+ m1_fp32, m2_fp32 = state1_fp32[0], state1_fp32[1]
+ nu_fp32 = state2_fp32
+
+ m1_fp32.mul_(beta1).add_(grad, alpha=1.0 - beta1)
+ m2_fp32.mul_(beta3).add_(grad, alpha=1.0 - beta3)
+ nu_fp32.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
+
+ correction1 = 1.0 - beta1**step
+ correction2 = math.sqrt(1.0 - beta2**step)
+
+ update = (m1_fp32 / correction1 + alpha * m2_fp32) / (nu_fp32.sqrt() / correction2 + eps)
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+ p_fp32.add_(update, alpha=-lr)
+
+ state1_fp32 = torch.stack([m1_fp32, m2_fp32])
+
+ elif optimizer_name in ("momentum", "lars"):
+ grad.add_(p_fp32, alpha=weight_decay)
+ if step == 1:
+ state1_fp32.copy_(grad)
+ else:
+ state1_fp32.mul_(beta1).add_(grad)
+ p_fp32.add_(state1_fp32, alpha=-lr)
+
+ elif optimizer_name == "lion":
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+
+ update_dir = torch.sign(state1_fp32.mul(beta1) + grad.mul(1.0 - beta1))
+ p_fp32.add_(update_dir, alpha=-lr)
+
+ state1_fp32.mul_(beta2).add_(grad, alpha=1.0 - beta2)
+
+ elif optimizer_name == "rmsprop":
+ grad.add_(p_fp32, alpha=weight_decay)
+ state1_fp32.mul_(beta1).addcmul_(grad, grad, value=1.0 - beta1)
+ p_fp32.addcdiv_(grad, state1_fp32.sqrt().add_(eps), value=-lr)
+
+ elif optimizer_name == "adagrad":
+ grad.add_(p_fp32, alpha=weight_decay)
+ state1_fp32.addcmul_(grad, grad, value=1.0)
+ p_fp32.addcdiv_(grad, state1_fp32.sqrt().add_(eps), value=-lr)
+
+ else:
+ raise ValueError(f"Unsupported optimizer for CPU 8-bit: {optimizer_name}")
+
+ p.data.copy_(p_fp32)
+
+ # Re-quantize states
+ if optimizer_name == "ademamix":
+ _quant_blockwise_fp32_direct(state1_fp32[0], qmap1, absmax1[0], state1[0], blocksize)
+ _quant_blockwise_fp32_direct(state1_fp32[1], qmap1, absmax1[1], state1[1], blocksize)
+ _quant_blockwise_fp32_direct(state2_fp32, qmap2, absmax2, state2, blocksize)
+ else:
+ _quant_blockwise_fp32_direct(state1_fp32, qmap1, absmax1, state1, blocksize)
+ if state2_fp32 is not None:
+ _quant_blockwise_fp32_direct(state2_fp32, qmap2, absmax2, state2, blocksize)
+
+
+register_kernel("bitsandbytes::optimizer_update_8bit_blockwise", "cpu")(_optimizer_update_8bit_blockwise_cpu)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/cuda/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/cuda/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/cuda/ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/cuda/ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..0d288d82b641cb866e886932fde1241223d3b6dc
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/cuda/ops.py
@@ -0,0 +1,1199 @@
+from collections.abc import Sequence
+import ctypes as ct
+import functools
+from math import prod
+from typing import Optional
+from warnings import warn
+
+import torch
+
+from bitsandbytes.functional import CUBLAS_Context, _cuda_device_of, get_ptr
+
+from ..._ops import register_kernel
+from ...cextension import lib
+
+
+def _setup_ctypes(names, argtypes, restype=None):
+ for name in names:
+ fn = getattr(lib, name)
+ fn.argtypes = argtypes
+ fn.restype = restype
+
+
+# 4-bit/8-bit dequantize: (code, A, absmax, out, blocksize, numel, stream)
+_setup_ctypes(
+ [f"cdequantize_blockwise_{d}_{q}" for d in ("fp32", "bf16", "fp16") for q in ("nf4", "fp4")]
+ + [f"cdequantize_blockwise_{d}" for d in ("fp32", "bf16", "fp16")],
+ [ct.c_void_p] * 4 + [ct.c_int32, ct.c_int32, ct.c_void_p],
+)
+
+# 4-bit GEMM: (A, B, absmax, absmax_8bit, absmax_code, absmax_offset, out, bias, M, N, K, blocksize, quant_type, stream)
+_setup_ctypes(
+ [f"cgemm_4bit_{d}" for d in ("bf16", "fp16", "fp32")],
+ [ct.c_void_p] * 8 + [ct.c_int32, ct.c_int32, ct.c_int32, ct.c_int32, ct.c_int32, ct.c_void_p],
+)
+
+# 4-bit GEMV: (m, n, k, A, B, absmax, code, out, lda, ldb, ldc, blocksize, stream)
+_setup_ctypes(
+ [f"cgemm_4bit_inference_naive_{d}" for d in ("bf16", "fp16", "fp32")],
+ [ct.c_int32] * 3 + [ct.c_void_p] * 5 + [ct.c_int32] * 3 + [ct.c_int32, ct.c_void_p],
+)
+
+# int8 igemm: (ctx, m, n, k, A, B, C, rowscale, lda, ldb, ldc, stream) -> int32
+_setup_ctypes(
+ ["cigemmlt_32"],
+ [ct.c_void_p] + [ct.c_int32] * 3 + [ct.c_void_p] * 4 + [ct.c_int32] * 3 + [ct.c_void_p],
+ restype=ct.c_int32,
+)
+
+# int8 mm dequant: (A, row_stats, col_stats, out, bias, numRows, numCols, stream)
+_setup_ctypes(
+ ["cdequant_mm_int32_fp16"],
+ [ct.c_void_p] * 5 + [ct.c_int32, ct.c_int32, ct.c_void_p],
+)
+
+# int8 vectorwise quant: (A, out, row_stats, threshold, rows, cols, stream)
+_setup_ctypes(
+ ["cint8_vector_quant"],
+ [ct.c_void_p] * 3 + [ct.c_float, ct.c_int32, ct.c_int32, ct.c_void_p],
+)
+
+# 4-bit/8-bit blockwise quantize: (code, A, absmax, out, blocksize, n)
+_setup_ctypes(
+ [f"cquantize_blockwise_{d}_{q}" for d in ("fp32", "bf16", "fp16") for q in ("nf4", "fp4")]
+ + [f"cquantize_blockwise_{d}" for d in ("fp32", "bf16", "fp16")],
+ [ct.c_void_p] * 4 + [ct.c_int32, ct.c_int32],
+)
+
+
+_get_raw_stream = torch._C._cuda_getCurrentRawStream
+
+
+@functools.cache
+def _gpu_dispatch_props(device_index):
+ props = torch.cuda.get_device_properties(device_index)
+ return props.multi_processor_count, props.major, props.minor
+
+
+@register_kernel("bitsandbytes::int8_linear_matmul", "cuda")
+def _(A: torch.Tensor, B: torch.Tensor):
+ out = torch.empty((*A.shape[:-1], B.shape[0]), device=A.device, dtype=torch.int32)
+ return _int8_linear_matmul_impl(A, B, out)
+
+
+@register_kernel("bitsandbytes::int8_linear_matmul.out", "cuda")
+def _(A: torch.Tensor, B: torch.Tensor, out: torch.Tensor):
+ _int8_linear_matmul_impl(A, B, out)
+
+
+def _int8_linear_matmul_impl(A: torch.Tensor, B: torch.Tensor, out: torch.Tensor):
+ A, B = B, A
+
+ shapeA = A.shape
+ shapeB = B.shape
+
+ if A.dtype != torch.int8:
+ raise ValueError("B must be int8")
+ if B.dtype != torch.int8:
+ raise ValueError("A must be int8")
+ if A.ndim != 2:
+ raise ValueError("Only two dimensional matrices are supported for argument B")
+ if B.ndim not in (2, 3):
+ raise ValueError("Only two or three dimensional matrices are supported for argument A")
+ if prod(shapeB) <= 0:
+ raise ValueError(f"Input tensor dimensions need to be > 0: {shapeB}")
+ if out.dtype != torch.int32:
+ raise ValueError(f"out must be int32, got {out.dtype}")
+
+ shapeC = (*shapeB[:-1], shapeA[0])
+ if out.shape != shapeC:
+ raise ValueError(f"Output shape {out.shape} does not match expected shape {shapeC}")
+
+ k, m = shapeA
+ n = prod(shapeB[:-1])
+ lda = shapeA[-1] # Weights (outputs, inputs)
+ ldb = shapeB[-1] # Activations (batch, tokens, inputs)
+ ldc = shapeC[-1] # Output (batch, tokens, outputs)
+
+ if lda != ldb:
+ raise ValueError(
+ f"int8_linear_matmul only supports B^T @ A. Inner dimensions do not match: B @ A = {shapeB} @ {shapeA}"
+ )
+
+ # cuBLASLt does not support int8 matmul with inner dimensions that are not divisible by 4.
+ # We'll fall back to a slower fp32 calculation in this circumstance.
+ # Fortunately, this should not be very common.
+ if lda % 4 != 0:
+ result = torch.matmul(B.float(), A.float().t()).to(torch.int32)
+ return out.copy_(result)
+
+ with _cuda_device_of(A):
+ ctx = CUBLAS_Context.get_instance().get_context(A.device)
+ has_error = lib.cigemmlt_32(
+ ctx,
+ m,
+ n,
+ k,
+ A.data_ptr(),
+ B.data_ptr(),
+ out.data_ptr(),
+ None,
+ lda,
+ ldb,
+ ldc,
+ _get_raw_stream(A.device.index),
+ )
+
+ if has_error:
+ if has_error == 100:
+ # `ERR_NOT_IMPLEMENTED` is defined as 100 in `ops.cu`. The HIP backend
+ # also returns this when no usable hipBLASLt algo exists for the shape
+ # (seen on MI300X for some small-n int8 gemms). Fall back to fp32 — same
+ # path used for the `lda % 4 != 0` case above.
+ import warnings
+
+ warnings.warn(
+ f"int8_linear_matmul has no usable (hip|cu)blasLt algo for shape "
+ f"{shapeA=} {shapeB=}; falling back to fp32 matmul.",
+ RuntimeWarning,
+ stacklevel=2,
+ )
+ result = torch.matmul(B.float(), A.float().t()).to(torch.int32)
+ return out.copy_(result)
+ else:
+ raise RuntimeError(
+ f"cublasLt ran into an error!\n\t{shapeA=}, {shapeB=}, {shapeC=}\n\t{(lda, ldb, ldc)=}\n\t{(m, n, k)=}"
+ )
+
+ return out
+
+
+@register_kernel("bitsandbytes::int8_mm_dequant", "cuda")
+def _(
+ A: torch.Tensor,
+ row_stats: torch.Tensor,
+ col_stats: torch.Tensor,
+ dtype: Optional[torch.dtype] = None,
+ bias: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ if A.dtype != torch.int32:
+ raise ValueError(f"A must be int32, got {A.dtype}")
+ if row_stats.dtype != torch.float32:
+ raise ValueError(f"row_stats must be float32, got {row_stats.dtype}")
+ if col_stats.dtype != torch.float32:
+ raise ValueError(f"col_stats must be float32, got {col_stats.dtype}")
+
+ # Note: cuda kernel only currently supports fp16 output.
+ # We'll later cast to desired dtype if needed.
+ out = torch.empty_like(A, dtype=torch.float16)
+
+ # Note: fused bias in the kernel is only supported for fp16
+ # TODO(matthewdouglas): Consider supporting bf16 fused bias
+ bias_ptr = bias.data_ptr() if bias is not None and bias.dtype == torch.float16 else None
+
+ with _cuda_device_of(A):
+ lib.cdequant_mm_int32_fp16(
+ A.data_ptr(),
+ row_stats.data_ptr(),
+ col_stats.data_ptr(),
+ out.data_ptr(),
+ bias_ptr,
+ A.numel() // A.shape[-1],
+ A.shape[-1],
+ _get_raw_stream(A.device.index),
+ )
+
+ # Add bias separately if not fused in kernel
+ if bias is not None and bias.dtype != torch.float16:
+ out.add_(bias)
+
+ return out.to(dtype or torch.float16)
+
+
+@register_kernel("bitsandbytes::int8_vectorwise_quant", "cuda")
+def _(A: torch.Tensor, threshold=0.0):
+ if A.dtype != torch.float16:
+ raise ValueError(f"A must be float16, got {A.dtype}")
+ if threshold < 0.0:
+ raise ValueError("threshold must be non-negative")
+
+ rows = A.numel() // A.shape[-1]
+ cols = A.shape[-1]
+
+ row_stats = torch.empty(rows, device=A.device, dtype=torch.float32)
+ out_row = torch.empty(A.shape, device=A.device, dtype=torch.int8)
+
+ outlier_cols = None
+
+ if threshold > 0.0:
+ # TODO we could improve perf of this
+ outliers = A.abs() >= threshold
+
+ if outliers.any():
+ outlier_cols = torch.argwhere(outliers.any(dim=0)).view(-1)
+ else:
+ # Needed for torch.compile support.
+ outlier_cols = torch.empty(0, device=A.device, dtype=torch.int64)
+
+ with _cuda_device_of(A):
+ lib.cint8_vector_quant(
+ A.data_ptr(),
+ out_row.data_ptr(),
+ row_stats.data_ptr(),
+ threshold,
+ rows,
+ cols,
+ _get_raw_stream(A.device.index),
+ )
+
+ # Zero out values from outlier columns across all rows.
+ # The kernel will handle this for outliers themselves, so we can optimize for rows=1.
+ if rows > 1 and outlier_cols is not None:
+ out_row[:, outlier_cols] = 0
+
+ return out_row, row_stats, outlier_cols
+
+
+@register_kernel("bitsandbytes::int8_double_quant", "cuda")
+def _(
+ A: torch.Tensor,
+ threshold=0.0,
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
+ # Use CUDA kernel for rowwise quant and outlier column detection
+ quant_row, row_stats, outlier_cols = torch.ops.bitsandbytes.int8_vectorwise_quant.default(
+ A,
+ threshold=threshold,
+ )
+
+ # PyTorch impl for colwise
+ col_stats, outlier_mask = _get_col_absmax(A, threshold=threshold)
+ if threshold > 0.0 and outlier_mask is not None:
+ A = A.masked_fill(outlier_mask, 0.0)
+ quant_col = torch.round(A.mul(127.0) / col_stats.unsqueeze(0)).to(torch.int8)
+
+ return quant_row, quant_col, row_stats, col_stats.flatten().float(), outlier_cols
+
+
+def _get_col_absmax(
+ A: torch.Tensor,
+ threshold=0.0,
+) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
+ if not A.is_floating_point():
+ raise ValueError(f"A must be a floating point tensor, got {A.dtype}")
+
+ outlier_mask = None
+
+ absA = A.abs().view(-1, A.shape[-1])
+
+ if threshold > 0.0:
+ # Filter outliers from stats when enabled
+ outlier_mask = absA >= threshold
+ absA.masked_fill_(outlier_mask, 0.0)
+
+ # shape [cols]; unsqueeze(0) gives [1,cols]
+ col_stats = absA.amax(dim=0, keepdim=False).float()
+
+ return col_stats, outlier_mask
+
+
+@register_kernel("bitsandbytes::quantize_blockwise", "cuda")
+def _(A: torch.Tensor, code: torch.Tensor, blocksize: int) -> tuple[torch.Tensor, torch.Tensor]:
+ A = A.contiguous()
+
+ if code.dtype != torch.float32:
+ raise ValueError(f"code must be float32, got {code.dtype}")
+ if blocksize not in (64, 128, 256, 512, 1024, 2048, 4096):
+ raise ValueError(f"invalid blocksize {blocksize}")
+
+ n = A.numel()
+ blocks = -(n // -blocksize)
+ absmax = torch.empty((blocks,), device=A.device, dtype=torch.float32)
+ out = torch.empty_like(A, dtype=torch.uint8)
+
+ if A.dtype == torch.float32:
+ fn = lib.cquantize_blockwise_fp32
+ elif A.dtype == torch.float16:
+ fn = lib.cquantize_blockwise_fp16
+ elif A.dtype == torch.bfloat16:
+ fn = lib.cquantize_blockwise_bf16
+ else:
+ raise ValueError(f"Blockwise quantization only supports 16/32-bit floats, but got {A.dtype}")
+
+ with _cuda_device_of(A):
+ fn(
+ code.data_ptr(),
+ A.data_ptr(),
+ absmax.data_ptr(),
+ out.data_ptr(),
+ blocksize,
+ n,
+ )
+
+ return out, absmax
+
+
+@register_kernel("bitsandbytes::dequantize_blockwise", "cuda")
+def _(A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype) -> torch.Tensor:
+ out = torch.empty_like(A, dtype=dtype)
+ _dequantize_blockwise_impl(A, absmax, code, blocksize, dtype, out=out)
+ return out
+
+
+@register_kernel("bitsandbytes::dequantize_blockwise.out", "cuda")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ if out.dtype != dtype:
+ raise ValueError(f"Expected out.dtype == {dtype}, got {out.dtype}")
+ if out.shape != A.shape:
+ raise ValueError(f"Expected out.shape == {A.shape}, got {out.shape}")
+ _dequantize_blockwise_impl(A, absmax, code, blocksize, dtype, out=out)
+
+
+def _dequantize_blockwise_impl(
+ A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype, out: torch.Tensor
+) -> None:
+ A = A.contiguous()
+
+ if dtype == torch.float32:
+ fn = lib.cdequantize_blockwise_fp32
+ elif dtype == torch.float16:
+ fn = lib.cdequantize_blockwise_fp16
+ elif dtype == torch.bfloat16:
+ fn = lib.cdequantize_blockwise_bf16
+ else:
+ raise ValueError(f"Blockwise dequantization only supports 16/32-bit floats, but got {dtype}")
+
+ with _cuda_device_of(A):
+ fn(
+ code.data_ptr(),
+ A.data_ptr(),
+ absmax.data_ptr(),
+ out.data_ptr(),
+ blocksize,
+ A.numel(),
+ _get_raw_stream(A.device.index),
+ )
+
+
+@register_kernel("bitsandbytes::quantize_4bit", "cuda")
+def _(
+ A: torch.Tensor, blocksize: int, quant_type: str, quant_storage: torch.dtype
+) -> tuple[torch.Tensor, torch.Tensor]:
+ A = A.contiguous()
+ n = A.numel()
+ blocks = -(n // -blocksize)
+ absmax = torch.empty((blocks,), device=A.device, dtype=torch.float32)
+ out = torch.empty(((n + 1) // (quant_storage.itemsize * 2), 1), device=A.device, dtype=quant_storage)
+
+ if A.dtype == torch.bfloat16:
+ if quant_type == "fp4":
+ fn = lib.cquantize_blockwise_bf16_fp4
+ else:
+ fn = lib.cquantize_blockwise_bf16_nf4
+ elif A.dtype == torch.float16:
+ if quant_type == "fp4":
+ fn = lib.cquantize_blockwise_fp16_fp4
+ else:
+ fn = lib.cquantize_blockwise_fp16_nf4
+ elif A.dtype == torch.float32:
+ if quant_type == "fp4":
+ fn = lib.cquantize_blockwise_fp32_fp4
+ else:
+ fn = lib.cquantize_blockwise_fp32_nf4
+
+ with _cuda_device_of(A):
+ fn(
+ None,
+ A.data_ptr(),
+ absmax.data_ptr(),
+ out.data_ptr(),
+ blocksize,
+ n,
+ )
+
+ return out, absmax
+
+
+@register_kernel("bitsandbytes::dequantize_4bit", "cuda")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ out = torch.empty(shape, dtype=dtype, device=A.device)
+ _dequantize_4bit_impl(A, absmax, blocksize, quant_type, dtype, out=out)
+ return out
+
+
+@register_kernel("bitsandbytes::dequantize_4bit.out", "cuda")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ if out.shape != tuple(shape):
+ raise ValueError(f"Expected out.shape == {shape}, got {out.shape}")
+ if out.dtype != dtype:
+ raise ValueError(f"Expected out.dtype == {dtype}, got {out.dtype}")
+ _dequantize_4bit_impl(A, absmax, blocksize, quant_type, dtype, out=out)
+
+
+def _dequantize_4bit_impl(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ A = A.contiguous()
+
+ if dtype == torch.bfloat16:
+ if quant_type == "fp4":
+ fn = lib.cdequantize_blockwise_bf16_fp4
+ else:
+ fn = lib.cdequantize_blockwise_bf16_nf4
+ elif dtype == torch.float16:
+ if quant_type == "fp4":
+ fn = lib.cdequantize_blockwise_fp16_fp4
+ else:
+ fn = lib.cdequantize_blockwise_fp16_nf4
+ elif dtype == torch.float32:
+ if quant_type == "fp4":
+ fn = lib.cdequantize_blockwise_fp32_fp4
+ else:
+ fn = lib.cdequantize_blockwise_fp32_nf4
+ else:
+ raise ValueError(f"Blockwise 4bit dequantization only supports 16/32-bit floats, but got {dtype}")
+
+ with _cuda_device_of(A):
+ fn(
+ None,
+ A.data_ptr(),
+ absmax.data_ptr(),
+ out.data_ptr(),
+ blocksize,
+ out.numel(),
+ _get_raw_stream(A.device.index),
+ )
+
+
+@register_kernel("bitsandbytes::gemv_4bit", "cuda")
+def _(
+ A: torch.Tensor, B: torch.Tensor, shapeB: Sequence[int], absmax: torch.Tensor, code: torch.Tensor, blocksize: int
+) -> torch.Tensor:
+ shape = (*A.shape[:-1], shapeB[0])
+ out = torch.empty(shape, device=A.device, dtype=A.dtype)
+ _gemv_4bit_impl(A, B, shapeB, absmax, code, blocksize, out=out)
+ return out
+
+
+@register_kernel("bitsandbytes::gemv_4bit.out", "cuda")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ out: torch.Tensor,
+) -> None:
+ expected_shape = (*A.shape[:-1], shapeB[0])
+ if out.shape != expected_shape:
+ raise ValueError(f"Expected out.shape == {expected_shape}, got {out.shape}")
+ if out.dtype != A.dtype:
+ raise ValueError(f"Expected out.dtype == {A.dtype}, got {out.dtype}")
+ _gemv_4bit_impl(A, B, shapeB, absmax, code, blocksize, out=out)
+
+
+def _gemv_4bit_impl(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ out: torch.Tensor,
+) -> None:
+ if blocksize not in (32, 64, 128, 256, 512, 1024, 2048, 4096):
+ raise ValueError(f"invalid blocksize {blocksize}")
+
+ # Note: these checks are not strictly necessary, and cost more than they are worth, so they are commented out for now.
+ # torch._check(
+ # A.numel() == A.size(-1),
+ # lambda: f"A must be a vector with leading dimensions of 1, got {A.shape}",
+ # )
+ # torch._check(
+ # A.dtype in [torch.float16, torch.bfloat16, torch.float32],
+ # lambda: f"A must be float16, bfloat16, or float32, got {A.dtype}",
+ # )
+ # torch._check(
+ # B.dtype in [torch.uint8, torch.bfloat16, torch.float16, torch.float32],
+ # lambda: f"B must be backed by storage of type uint8, bfloat16, float16, or float32, got {B.dtype}",
+ # )
+ # torch._check(absmax.dtype == torch.float32, lambda: f"absmax must be float32, got {absmax.dtype}")
+ # torch._check(code.dtype == torch.float32, lambda: f"code must be float32, got {code.dtype}")
+
+ m = shapeB[0]
+ n = 1
+ k = shapeB[1]
+
+ lda = m
+ ldb = (A.shape[-1] + 1) // 2
+ ldc = m
+
+ if A.dtype == torch.float16:
+ fn = lib.cgemm_4bit_inference_naive_fp16
+ elif A.dtype == torch.bfloat16:
+ fn = lib.cgemm_4bit_inference_naive_bf16
+ elif A.dtype == torch.float32:
+ fn = lib.cgemm_4bit_inference_naive_fp32
+
+ with _cuda_device_of(A):
+ fn(
+ m,
+ n,
+ k,
+ A.data_ptr(),
+ B.data_ptr(),
+ absmax.data_ptr(),
+ code.data_ptr(),
+ out.data_ptr(),
+ lda,
+ ldb,
+ ldc,
+ blocksize,
+ _get_raw_stream(A.device.index),
+ )
+
+
+@functools.cache
+def _gemm_4bit_use_custom_cuda(device_index, dtype, M, N, K):
+ """Custom kernel vs dequant+F.linear heuristic for M in [5, 1536].
+
+ Per-arch notes (bf16/fp16, M >= 8, large weight):
+ sm75 (T4, ~300 GB/s GDDR6): fp16 MMA only; GDDR makes dequant expensive.
+ sm80 (A100, ~2 TB/s HBM2e): mma.sync; HBM thresholds; K-heavy shapes handled explicitly.
+ sm86 (A10, ~600 GB/s GDDR6): dedicated block; wider M caps than sm89 at medium N.
+ sm89 (4090, L40S, GDDR6X): default fallback; tall-K and large-N get higher M caps.
+ sm90 (H100/H200, HBM3/HBM3e): dequant+linear is much faster; thresholds are tight.
+ sm100 (B200/B300, HBM3e): exits early at top of function.
+ sm120 (RTX 5000, GDDR7): dedicated block; medium-N tiers differ from sm89.
+ """
+ if M <= _GEMM_4BIT_CUSTOM_FLOOR_M:
+ return True
+
+ num_sms, major, minor = _gpu_dispatch_props(device_index)
+ n_blocks = (N + 63) // 64
+
+ # fp32 has no MMA kernel; pre-sm75 has no MMA kernel; sm75 has fp16 MMA only.
+ # For all of these, custom only wins in the SIMT range (M<8).
+ if dtype == torch.float32 or major < 7:
+ return M < 8
+ if major == 7 and (minor < 5 or dtype != torch.float16):
+ return M < 8
+
+ # sm87 and sm110: no calibration data, conservative fallback.
+ if (major == 8 and minor == 7) or major == 11:
+ return False
+
+ # sm100 (B200/B300): dequant+F.linear is significantly faster than our mma.sync kernel.
+ if major == 10:
+ if n_blocks >= num_sms * 3:
+ return M <= 32
+ if n_blocks >= num_sms:
+ return False if K >= N else M <= 8
+ return False
+
+ is_sm75 = major == 7 and minor == 5
+ is_sm80 = major == 8 and minor == 0
+ is_sm86 = major == 8 and minor == 6
+ is_sm90 = major == 9
+ is_sm120 = major == 12 and minor == 0
+ is_hbm = is_sm80 or is_sm90 # sm100 already returned above
+ tall_k_2xn = K > N * 2
+
+ # Small-weight path (N*K < 4MB): dequant overhead dominates.
+ if N * K < 4 * 1024 * 1024:
+ if K * 2 < N:
+ # Very short K (K < N/2): latency-dominated, custom 3-9x cheaper.
+ if is_hbm:
+ # Calibrated on A100: custom wins to M=1536 (low wave), M=512 (high wave).
+ # Calibrated on H100/H200: custom wins to M=512 (low wave), M=320 (high wave).
+ low_wave = n_blocks * 3 < num_sms
+ if is_sm80:
+ return M <= (1536 if low_wave else 512)
+ return M <= (512 if low_wave else 320)
+ if is_sm75:
+ # T4: wins require >=3 waves; M cap scales with K depth.
+ if n_blocks >= num_sms * 3:
+ return M <= 320
+ if K >= 1024:
+ return M <= 64
+ if K >= 704:
+ return M <= 96
+ return M <= 320
+ # sm86/sm89/sm120: well-subscribed wins to M=320; undersubscribed tighter.
+ if n_blocks >= num_sms:
+ return M <= 320
+ return M <= 192 if n_blocks * K > num_sms * 320 else M <= 320
+ # K*2 >= N: arch-specific handling at low occupancy.
+ quarter_wave = n_blocks * 4 <= num_sms
+ if is_sm80 and quarter_wave:
+ # A100 <1/4 wave: K>=N loses earlier (K-tiling efficient on HBM2e).
+ if K >= N:
+ return M <= (32 if n_blocks * 8 <= num_sms else 128)
+ return M <= 384
+ # T4 <1 wave non-short-K: M>8 routes through occupancy caps below.
+ if is_sm75 and n_blocks < num_sms and M > 8:
+ return M <= 64
+ # General tiers (sm90, sm86, sm89, sm120):
+ # GDDR tall-K (K>=N) at <1/4 wave: K-tiling in default impl wins above M=23.
+ if quarter_wave:
+ return M <= (32 if (K < N or is_hbm) else 23)
+ if n_blocks * 2 <= num_sms:
+ return M <= 16
+ return False # >=1/2 wave: no validated wins for remaining small-weight shapes
+
+ # Non-small-weight: custom wins up to M=512; dequant+F.linear wins above that.
+ if M > 512:
+ return False
+
+ # M=5-7: custom SIMT generally wins because dequant cost dominates.
+ # Exceptions where K-tiling efficiency or MMA occupancy favors dequant+F.linear:
+ # HBM at M=6-7: tall-K (K>N) at ~3/4 MMA wave.
+ # sm90 square (K==N) at specific occupancy bands: arch-specific crossover.
+ if M < 8:
+ hbm_m67_thresh = 36 if is_sm90 else 48
+ if is_hbm and M >= 6 and n_blocks >= hbm_m67_thresh:
+ lt_75pct_wave = n_blocks * 4 < num_sms * 3
+ lt_60pct_wave = n_blocks * 5 < num_sms * 3
+ # Tall-K: K-tiling in default impl wins when under-subscribed.
+ if K > N and lt_75pct_wave:
+ return False
+ # Square: arch-specific crossover around 0.6 wave.
+ # A100 (HBM2e): loses below 0.6 wave. H100/H200 (HBM3/3e): loses above.
+ if K == N:
+ if is_sm80 and lt_60pct_wave:
+ return False
+ if is_sm90 and lt_75pct_wave and not lt_60pct_wave:
+ return False
+ return True
+
+ # M in [8, 512]: per-arch tier ladders.
+
+ if is_sm75:
+ # fp16 MMA (m16n8k8). GDDR bandwidth makes dequant relatively expensive.
+ if n_blocks >= num_sms * 3:
+ return M <= (128 if K < N else 64)
+ if n_blocks >= num_sms // 2:
+ return M <= 64
+ return M <= 32
+
+ if is_sm80:
+ # mma.sync (m16n8k16). HBM2e thresholds; K-heavy shapes handled explicitly.
+ if n_blocks >= num_sms * 3:
+ return M <= 128
+ if n_blocks >= num_sms:
+ return M <= (64 if K < N else 32)
+ # Very tall-K (K>=3N) at >1/4 wave: K-tiling in default impl wins at all M.
+ # Uses >= to catch K==3N (e.g. N=4096,K=12288 M=9-16: measured regression on A100).
+ if K >= N * 3 and n_blocks * 4 > num_sms:
+ return False
+ # Square (K==N) at 0.5-1 wave: K-tiling wins at ~0.6 wave.
+ # n_blocks>=48 excludes small N where SIMT still wins.
+ if K == N and n_blocks >= 48 and n_blocks * 5 < num_sms * 3:
+ return False
+ # <0.5 wave: K<=N custom wins to M=128; K>N default wins above wave threshold.
+ if n_blocks * 2 < num_sms:
+ if K <= N:
+ return M <= 128
+ if n_blocks * 3 >= num_sms:
+ return False
+ # 0.5-1 wave K= num_sms // 2 and K < N:
+ return M <= 128
+ return M <= 16
+
+ if is_sm86:
+ # ~600-940 GB/s GDDR6/GDDR6X. Dedicated block: sm89 fallback tiers are too
+ # loose for 600 GB/s bandwidth and cause regressions at medium N (~N=4096).
+ if n_blocks >= num_sms:
+ return M <= 128
+ if n_blocks >= num_sms // 2:
+ return M <= 64
+ return M <= 16
+
+ if is_sm90:
+ # HBM3/HBM3e. dequant+F.linear (WGMMA path) is significantly faster than our
+ # mma.sync kernel; thresholds are calibrated conservatively (H100/H200 share path).
+ if n_blocks >= num_sms * 3:
+ return M <= 64
+ if n_blocks >= num_sms * 2:
+ return M <= 48
+ if n_blocks >= num_sms:
+ return M <= 32
+ if n_blocks >= num_sms // 2:
+ # Square/tall-K at <3/4 wave: K-tiling too efficient on HBM3e.
+ if K >= N and n_blocks * 4 < num_sms * 3:
+ return False
+ return M <= 16
+ return False
+
+ if is_sm120:
+ # GDDR7 (~1-1.8 TB/s). Medium-N threshold tiers differ from sm89.
+ # sm121 (DGX Spark) has a different bandwidth/SM profile; uses sm89
+ # fallback below until validated.
+ if n_blocks >= num_sms * 3:
+ return M <= 256
+ if n_blocks >= num_sms * 2:
+ return M <= 128
+ # Short-K (K= num_sms * 4:
+ return M <= (96 if K >= N else 64)
+ if n_blocks >= num_sms:
+ return M <= 64
+ if n_blocks >= num_sms // 2:
+ # Large-N (n_blocks>=128, N>=8192) with K>=N/2: calibrated on RTX Pro 6000 to M=64.
+ return M <= (64 if (K * 2 >= N and n_blocks >= 128) else 8)
+ if tall_k_2xn and n_blocks > 64:
+ return M <= 16
+ return M <= 8
+
+ # Fallback: sm89 (4090, L40S, L4), sm121 (DGX Spark), unrecognized arches.
+ # GDDR bandwidth makes dequant relatively expensive so custom wins at higher M.
+ if n_blocks >= num_sms * 3:
+ return M <= 256
+ if n_blocks >= num_sms * 2:
+ return M <= 128
+ # Near-wave (~0.8x): tall-K and very large N (n_blocks>=200, N>=14336) raise cap to M=128.
+ # N=10240 (n_blocks=160) deliberately excluded to avoid regressions there.
+ if n_blocks * 5 >= num_sms * 4:
+ if tall_k_2xn or n_blocks >= 200:
+ return M <= 128
+ # Square/tall-K: >=60 SMs wins to M=128; <60 SMs default wins earlier.
+ if K >= N:
+ return M <= (128 if num_sms >= 60 else 32)
+ return M <= 64
+ if n_blocks >= num_sms // 2:
+ if tall_k_2xn:
+ return M <= 64
+ if n_blocks >= 64:
+ return M <= 8
+ return M <= 32
+ # Tall-K (K>N) at narrow N (n_blocks<=48): M-driven crossover.
+ # K>=3N (e.g. N=2560,K=10240): SIMT wins to M=12. Moderate K>N: M=10.
+ if K > N and n_blocks <= 48:
+ return M <= (12 if K >= N * 3 else 10)
+ return M <= (16 if (tall_k_2xn or n_blocks < 48) else 8)
+
+
+@functools.cache
+def _gemm_4bit_use_custom_rocm(device_index, dtype, M, N, K):
+ """
+ Fused SIMT kernel vs dequant+F.linear heuristic for ROCm.
+
+ RDNA3/RDNA4 calibration keeps the SIMT kernel through ~M=8.
+ CDNA/gfx9 is calibrated on MI308X (gfx942): bf16/fp16 win through M<=4
+ after the SIMT math-path tuning, while fp32 only has a broad win through M<=2.
+
+ TODO: revisit once WMMA/MFMA kernels land.
+ """
+ if M <= _GEMM_4BIT_CUSTOM_FLOOR_M and dtype != torch.float32:
+ return True
+
+ arch = _rocm_gfx_arch(device_index)
+ if arch.startswith("gfx11") or arch.startswith("gfx12"): # RDNA3 / RDNA4
+ return M <= 8
+ if arch.startswith("gfx9"): # CDNA / MI-series
+ return M <= (2 if dtype == torch.float32 else 4)
+ return M <= 4 # unknown ROCm arch: conservative tiny-batch floor
+
+
+@functools.cache
+def _rocm_gfx_arch(device_index):
+ """gfx arch string (e.g. 'gfx1100') for a ROCm device, feature flags stripped."""
+ name = getattr(torch.cuda.get_device_properties(device_index), "gcnArchName", "") or ""
+ return name.split(":")[0]
+
+
+def _gemm_4bit_kernel_impl(
+ A, B, shapeB, absmax, blocksize, quant_type, bias=None, absmax_8bit=None, absmax_code=None, absmax_offset=None
+):
+ """Invoke the fused cgemm_4bit_* kernel (shared by the CUDA and ROCm dispatch; the
+ C dispatch in gemm_4bit.cu picks SIMT vs MMA per arch/shape). A is made contiguous
+ because the kernel reads it as row-major (stride K)."""
+ K = A.shape[-1]
+ M = A.numel() // K
+ N = shapeB[0]
+
+ if K != shapeB[1]:
+ raise RuntimeError(f"A inner dim ({K}) does not match weight ({shapeB[1]})")
+ if absmax.dtype != torch.float32:
+ raise RuntimeError(f"absmax must be float32, got {absmax.dtype}")
+ if bias is not None:
+ if bias.ndim != 1:
+ raise RuntimeError(f"bias must be 1D, got {bias.ndim}D")
+ if bias.dtype != A.dtype:
+ raise RuntimeError(f"bias dtype ({bias.dtype}) must match A dtype ({A.dtype})")
+
+ A = A.contiguous()
+ quant_type_int = 1 if quant_type == "fp4" else 2
+ out = torch.empty((*A.shape[:-1], N), dtype=A.dtype, device=A.device)
+ stream = _get_raw_stream(A.device.index)
+
+ if A.dtype == torch.bfloat16:
+ fn = lib.cgemm_4bit_bf16
+ elif A.dtype == torch.float16:
+ fn = lib.cgemm_4bit_fp16
+ elif A.dtype == torch.float32:
+ fn = lib.cgemm_4bit_fp32
+ else:
+ raise RuntimeError(f"unsupported dtype {A.dtype}")
+
+ # Offset is expected to be a float32 tensor.
+ absmax_offset_f32 = absmax_offset.to(dtype=torch.float32) if absmax_offset is not None else None
+
+ with _cuda_device_of(A):
+ fn(
+ A.data_ptr(),
+ B.data_ptr(),
+ absmax.data_ptr(),
+ absmax_8bit.data_ptr() if absmax_8bit is not None else None,
+ absmax_code.data_ptr() if absmax_code is not None else None,
+ absmax_offset_f32.data_ptr() if absmax_offset_f32 is not None else None,
+ out.data_ptr(),
+ bias.data_ptr() if bias is not None else None,
+ M,
+ N,
+ K,
+ blocksize,
+ quant_type_int,
+ stream,
+ )
+
+ return out
+
+
+def _dequant_linear_fallback(
+ A, B, shapeB, absmax, blocksize, quant_type, bias=None, absmax_8bit=None, absmax_code=None, absmax_offset=None
+):
+ """Unfused fallback shared by CUDA and ROCm: reconstruct the (optionally nested)
+ absmax, dequantize the 4-bit weight via the backend dequant impls (reusing
+ preallocated buffers), then F.linear."""
+ if absmax_8bit is not None:
+ absmax_dq = torch.empty_like(absmax_8bit, dtype=torch.float32)
+ _dequantize_blockwise_impl(absmax_8bit, absmax, absmax_code, 256, torch.float32, out=absmax_dq)
+ absmax = absmax_dq + absmax_offset
+ B_dq = torch.empty(shapeB, dtype=A.dtype, device=A.device)
+ _dequantize_4bit_impl(B, absmax, blocksize, quant_type, A.dtype, out=B_dq)
+ return torch.nn.functional.linear(A, B_dq, bias)
+
+
+# Unified CUDA/ROCm dispatch for bitsandbytes::gemm_4bit. The choice *among* custom
+# kernels (CUDA SIMT vs MMA; ROCm SIMT) is made in the C dispatch (csrc/gemm_4bit.cu).
+_GEMM_4BIT_CUSTOM_FLOOR_M = 4
+if torch.version.hip is None:
+ _gemm_4bit_use_custom_fn = _gemm_4bit_use_custom_cuda
+ # CUDA: dequant+F.linear wins past M=1536 (dequant savings negligible at very
+ # large batch).
+ _gemm_4bit_custom_max_m = 1536
+else:
+ _gemm_4bit_use_custom_fn = _gemm_4bit_use_custom_rocm
+ # ROCm: the custom path is SIMT-only today; the per-arch heuristic above owns
+ # RDNA/CDNA thresholds. Keep a hard upper cap while WMMA/MFMA paths are absent.
+ _gemm_4bit_custom_max_m = 256
+
+
+@register_kernel("bitsandbytes::gemm_4bit", "cuda")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ bias: Optional[torch.Tensor] = None,
+ absmax_8bit: Optional[torch.Tensor] = None,
+ absmax_code: Optional[torch.Tensor] = None,
+ absmax_offset: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ K = A.shape[-1]
+ M = A.numel() // K
+ N = shapeB[0]
+
+ # The backend-specific heuristic owns tiny-M floors and per-arch thresholds.
+ # Past custom_max_m (or for blocksize-misaligned K), use the dequant+F.linear
+ # fallback.
+ if M > _gemm_4bit_custom_max_m:
+ use_custom = False
+ elif K % blocksize != 0:
+ warn(
+ f"inner dimension ({K}) is not aligned for fast kernel "
+ f"with blocksize={blocksize}, falling back to slower implementation.",
+ UserWarning,
+ )
+ use_custom = False
+ else:
+ use_custom = _gemm_4bit_use_custom_fn(A.device.index, A.dtype, M, N, K)
+
+ if not use_custom:
+ return _dequant_linear_fallback(
+ A,
+ B,
+ shapeB,
+ absmax,
+ blocksize,
+ quant_type,
+ bias,
+ absmax_8bit=absmax_8bit,
+ absmax_code=absmax_code,
+ absmax_offset=absmax_offset,
+ )
+
+ return _gemm_4bit_kernel_impl(
+ A, B, shapeB, absmax, blocksize, quant_type, bias, absmax_8bit, absmax_code, absmax_offset
+ )
+
+
+"""C FUNCTIONS FOR OPTIMIZERS"""
+str2optimizer32bit = {
+ "adam": (
+ lib.cadam32bit_grad_fp32,
+ lib.cadam32bit_grad_fp16,
+ lib.cadam32bit_grad_bf16,
+ ),
+ "momentum": (
+ lib.cmomentum32bit_grad_32,
+ lib.cmomentum32bit_grad_16,
+ ),
+ "rmsprop": (
+ lib.crmsprop32bit_grad_32,
+ lib.crmsprop32bit_grad_16,
+ ),
+ "lion": (
+ lib.clion32bit_grad_fp32,
+ lib.clion32bit_grad_fp16,
+ lib.clion32bit_grad_bf16,
+ ),
+ "adagrad": (
+ lib.cadagrad32bit_grad_32,
+ lib.cadagrad32bit_grad_16,
+ ),
+ "lamb": (
+ lib.cadam32bit_grad_fp32,
+ lib.cadam32bit_grad_fp16,
+ lib.cadam32bit_grad_bf16,
+ ),
+ "ademamix": (
+ lib.cademamix32bit_grad_fp32,
+ lib.cademamix32bit_grad_fp16,
+ lib.cademamix32bit_grad_bf16,
+ ),
+ "lars": (
+ lib.cmomentum32bit_grad_32,
+ lib.cmomentum32bit_grad_16,
+ ),
+}
+
+str2optimizer8bit_blockwise = {
+ "adam": (
+ lib.cadam_8bit_blockwise_grad_fp32,
+ lib.cadam_8bit_blockwise_grad_fp16,
+ lib.cadam_8bit_blockwise_grad_bf16,
+ ),
+ "momentum": (
+ lib.cmomentum_8bit_blockwise_grad_fp32,
+ lib.cmomentum_8bit_blockwise_grad_fp16,
+ lib.cmomentum_8bit_blockwise_grad_bf16,
+ ),
+ "rmsprop": (
+ lib.crmsprop_8bit_blockwise_grad_fp32,
+ lib.crmsprop_8bit_blockwise_grad_fp16,
+ lib.crmsprop_8bit_blockwise_grad_bf16,
+ ),
+ "lion": (
+ lib.clion_8bit_blockwise_grad_fp32,
+ lib.clion_8bit_blockwise_grad_fp16,
+ lib.clion_8bit_blockwise_grad_bf16,
+ ),
+ "adagrad": (
+ lib.cadagrad_8bit_blockwise_grad_fp32,
+ lib.cadagrad_8bit_blockwise_grad_fp16,
+ lib.cadagrad_8bit_blockwise_grad_bf16,
+ ),
+ "ademamix": (
+ lib.cademamix_8bit_blockwise_grad_fp32,
+ lib.cademamix_8bit_blockwise_grad_fp16,
+ lib.cademamix_8bit_blockwise_grad_bf16,
+ ),
+}
+
+
+def _optimizer_update_32bit_impl(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float,
+ skip_zeros=False,
+) -> None:
+ optim_fns = str2optimizer32bit.get(optimizer_name, None)
+ if optim_fns is None:
+ raise ValueError(
+ f"Unsupported optimizer name: {optimizer_name}. Supported optimizers: {list(str2optimizer32bit.keys())}"
+ )
+ if g.dtype == torch.float32:
+ optim_func = optim_fns[0]
+ elif g.dtype == torch.float16:
+ optim_func = optim_fns[1]
+ elif g.dtype == torch.bfloat16 and len(optim_fns) == 3:
+ optim_func = optim_fns[2]
+ else:
+ raise ValueError(
+ f"Gradient+optimizer bit data type combination not supported: grad {g.dtype}, optimizer {state1.dtype}",
+ )
+
+ with _cuda_device_of(g):
+ optim_func(
+ get_ptr(g),
+ get_ptr(p),
+ get_ptr(state1),
+ get_ptr(state2),
+ get_ptr(unorm_vec),
+ ct.c_float(max_unorm),
+ ct.c_float(param_norm),
+ ct.c_float(beta1),
+ ct.c_float(beta2),
+ ct.c_float(beta3),
+ ct.c_float(alpha),
+ ct.c_float(eps),
+ ct.c_float(weight_decay),
+ ct.c_int32(step),
+ ct.c_float(lr),
+ ct.c_float(gnorm_scale),
+ ct.c_bool(skip_zeros),
+ ct.c_int32(g.numel()),
+ )
+
+
+def _optimizer_update_8bit_blockwise_impl(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: torch.Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: torch.Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float,
+ gnorm_scale: float,
+ skip_zeros=False,
+) -> None:
+ # torch._check(
+ # g.numel() == p.numel(),
+ # lambda: f"g and p must have the same number of elements, got {g.numel()} and {p.numel()}",
+ # )
+ # compute_dtypes = [torch.float16, torch.bfloat16, torch.float32]
+
+ # torch._check(
+ # g.dtype in compute_dtypes,
+ # lambda: f"g must be bfloat16, float16, or float32, got {g.dtype}",
+ # )
+ # torch._check(
+ # g.dtype == p.dtype,
+ # lambda: f"Expected all tensors to have the same dtype, got g.dtype={g.dtype}, p.dtype={p.dtype}",
+ # )
+ # torch._check(
+ # state1.dtype == torch.uint8,
+ # lambda: f"state1 must be uint8, got {state1.dtype}",
+ # )
+ # torch._check(
+ # qmap1.dtype == absmax1.dtype == torch.float32,
+ # lambda: f"Expected qmap1 and absmax1 to be float32, got qmap1.dtype={qmap1.dtype}, absmax1.dtype={absmax1.dtype}",
+ # )
+ # if state2 is not None:
+ # torch._check(
+ # state2.dtype == torch.uint8,
+ # lambda: f"state2 must be uint8, got {state2.dtype}",
+ # )
+ # torch._check(
+ # qmap2.dtype == absmax2.dtype == torch.float32,
+ # lambda: f"Expected qmap2 and absmax2 to be float32, got qmap2.dtype={qmap2.dtype}, absmax2.dtype={absmax2.dtype}",
+ # )
+ optimizer_fns = str2optimizer8bit_blockwise.get(optimizer_name)
+ if optimizer_fns is None:
+ raise ValueError(
+ f"Unsupported optimizer name: {optimizer_name}. Supported optimizers: {list(str2optimizer8bit_blockwise.keys())}"
+ )
+
+ if g.dtype == torch.float32:
+ optimizer_fn = optimizer_fns[0]
+ elif g.dtype == torch.float16:
+ optimizer_fn = optimizer_fns[1]
+ elif g.dtype == torch.bfloat16:
+ optimizer_fn = optimizer_fns[2]
+ else:
+ raise ValueError(
+ f"Unsupported gradient dtype: {g.dtype}. Supported dtypes: torch.float32, torch.float16, torch.bfloat16"
+ )
+
+ with _cuda_device_of(g):
+ optimizer_fn(
+ get_ptr(p),
+ get_ptr(g),
+ get_ptr(state1),
+ get_ptr(state2),
+ ct.c_float(beta1),
+ ct.c_float(beta2),
+ ct.c_float(beta3),
+ ct.c_float(alpha),
+ ct.c_float(eps),
+ ct.c_int32(step),
+ ct.c_float(lr),
+ get_ptr(qmap1),
+ get_ptr(qmap2),
+ get_ptr(absmax1),
+ get_ptr(absmax2),
+ ct.c_float(weight_decay),
+ ct.c_float(gnorm_scale),
+ ct.c_bool(skip_zeros),
+ ct.c_int32(g.numel()),
+ )
+
+
+register_kernel("bitsandbytes::optimizer_update_8bit_blockwise", "cuda")(_optimizer_update_8bit_blockwise_impl)
+register_kernel("bitsandbytes::optimizer_update_32bit", "cuda")(_optimizer_update_32bit_impl)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/default/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/default/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/default/ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/default/ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..521802922810fd6f6432564ef1f71233936d2508
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/default/ops.py
@@ -0,0 +1,632 @@
+from collections.abc import Sequence
+from functools import cache, wraps
+from math import prod, sqrt
+from typing import Optional
+
+import torch
+
+from ..._ops import register_kernel
+from ..utils import _get_4bit_code
+
+
+def _try_torch_compile(func=None, **compile_kwargs):
+ """
+ Wrapper around torch.compile that falls back to the original function if compilation fails.
+ """
+
+ def decorator(fn):
+ try:
+ compiled_fn = torch.compile(fn, **compile_kwargs)
+
+ @wraps(fn)
+ def wrapper(*args, **kwargs):
+ try:
+ return compiled_fn(*args, **kwargs)
+ except Exception:
+ return fn(*args, **kwargs)
+
+ return wrapper
+ except Exception:
+ return fn
+
+ if func is None:
+ return decorator
+ else:
+ return decorator(func)
+
+
+@register_kernel("bitsandbytes::int8_mm_dequant", "default")
+def _(
+ A: torch.Tensor,
+ row_stats: torch.Tensor,
+ col_stats: torch.Tensor,
+ dtype: Optional[torch.dtype] = None,
+ bias: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ if A.dtype != torch.int32:
+ raise ValueError(f"A must be int32, got {A.dtype}")
+ if row_stats.dtype != torch.float32:
+ raise ValueError(f"row_stats must be float32, got {row_stats.dtype}")
+ if col_stats.dtype != torch.float32:
+ raise ValueError(f"col_stats must be float32, got {col_stats.dtype}")
+
+ A_calc = A.view(-1, A.shape[-1])
+ row_stats = row_stats.reshape(-1).unsqueeze(-1)
+ col_stats = col_stats.reshape(-1).unsqueeze(0)
+
+ out = A_calc * (row_stats * col_stats) * 6.200124e-05
+ if bias is not None:
+ out += bias
+
+ return out.to(dtype or torch.float16)
+
+
+@register_kernel("bitsandbytes::int8_mixed_scaled_mm", "default")
+def _(
+ A: torch.Tensor,
+ CA: torch.Tensor,
+ CB: torch.Tensor,
+ SCA: torch.Tensor,
+ SCB: torch.Tensor,
+ outlier_cols: Optional[torch.Tensor] = None,
+ bias: Optional[torch.Tensor] = None,
+) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
+ subB = None
+
+ if outlier_cols is not None and outlier_cols.numel():
+ # Extract the inputs with outliers in original precision
+ subA = A[:, outlier_cols].contiguous()
+
+ # Dequantize the corresponding weight columns
+ subB = (
+ torch.ops.bitsandbytes.int8_vectorwise_dequant.default(CB[:, outlier_cols].contiguous(), SCB)
+ .to(A.dtype)
+ .t()
+ )
+
+ # TODO: if state.has_fp16_weights: subB = B[:, outlier_cols].t()
+
+ else:
+ # Needed for torch.compile when there are no outliers.
+ subA = torch.empty(0, device=A.device, dtype=A.dtype)
+
+ # Int8 Matmul + Dequant + Bias
+ output = torch.ops.bitsandbytes.int8_scaled_mm.default(CA, CB, SCA, SCB, bias=bias, dtype=A.dtype)
+
+ if subB is not None:
+ # Add the outlier columns back to the output
+ output = output.addmm(subA, subB)
+
+ return output, subA
+
+
+@register_kernel("bitsandbytes::int8_scaled_mm", "default")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ row_stats: torch.Tensor,
+ col_stats: torch.Tensor,
+ bias: Optional[torch.Tensor] = None,
+ dtype: Optional[torch.dtype] = None,
+) -> torch.Tensor:
+ out_i32 = torch.ops.bitsandbytes.int8_linear_matmul.default(A, B)
+ return torch.ops.bitsandbytes.int8_mm_dequant.default(
+ out_i32,
+ row_stats,
+ col_stats,
+ dtype=dtype or torch.float16,
+ bias=bias,
+ )
+
+
+@register_kernel("bitsandbytes::int8_linear_matmul", "default")
+def _(A: torch.Tensor, B: torch.Tensor):
+ return _int8_linear_matmul_impl(A, B)
+
+
+@register_kernel("bitsandbytes::int8_linear_matmul.out", "default")
+def _(A: torch.Tensor, B: torch.Tensor, out: torch.Tensor):
+ if out.dtype != torch.int32:
+ raise ValueError(f"out must be int32, got {out.dtype}")
+ _int8_linear_matmul_impl(A, B, out)
+
+
+def _int8_linear_matmul_impl(A: torch.Tensor, B: torch.Tensor, out: Optional[torch.Tensor] = None):
+ # Naive implementation: perform matmul in fp32
+ result = torch.matmul(A.float(), B.float().t()).to(torch.int32)
+ if out is not None:
+ result = out.copy_(result)
+ return result
+
+
+@register_kernel("bitsandbytes::int8_vectorwise_quant", "default")
+def _(A: torch.Tensor, threshold=0.0):
+ rows = A.numel() // A.shape[-1]
+ outlier_cols = None
+
+ outlier_restore = None
+
+ if threshold > 0.0:
+ outliers = A.abs() >= threshold
+
+ if outliers.any():
+ # Determine which columns contain outliers, and zero out the
+ # outliers ahead of quantization. We need to keep a backup of these
+ # outliers to restore them after quantization.
+ outlier_cols = torch.argwhere(outliers.any(dim=0)).view(-1)
+ outlier_restore = A[outliers].clone()
+ A[outliers] = 0
+ else:
+ # Needed for torch.compile support.
+ outlier_cols = torch.empty(0, device=A.device, dtype=torch.int64)
+
+ # Get absmax for each row.
+ row_stats = torch.max(A.abs(), dim=1).values.float()
+
+ # Quantize row-wise to int8.
+ out_row = torch.round(A * (127.0 / row_stats.unsqueeze(-1))).to(torch.int8)
+
+ # Zero out values from outlier columns across all rows.
+ if rows > 1 and outlier_cols is not None:
+ out_row[:, outlier_cols] = 0
+
+ # Restore outliers.
+ if outlier_restore is not None:
+ A[outliers] = outlier_restore
+
+ return out_row, row_stats, outlier_cols
+
+
+@register_kernel("bitsandbytes::quantize_blockwise", "default")
+def _(A: torch.Tensor, code: torch.Tensor, blocksize: int) -> tuple[torch.Tensor, torch.Tensor]:
+ A_flat = A.reshape(-1).float()
+ n = A_flat.numel()
+ rem = n % blocksize
+ full = n - rem
+ blocks = full // blocksize
+ A_com = A_flat[:full].reshape(blocks, blocksize)
+ absmax = A_com.abs().max(dim=-1)[0]
+ scaled = torch.clamp(A_com * (1.0 / absmax.clamp(min=1e-38).view(-1, 1)), -1, 1).reshape(-1)
+ if rem:
+ am = A_flat[full:].abs().max().clamp(min=1e-38)
+ absmax = torch.cat([absmax, am.unsqueeze(0)])
+ scaled = torch.cat([scaled, torch.clamp(A_flat[full:] / am, -1, 1)])
+ bounds = (code[:-1] + code[1:]) / 2 # code is always sorted (same assumption as CUDA kernel)
+ q = torch.bucketize(scaled, bounds, out_int32=True).to(torch.uint8)
+ return q.reshape(A.shape), absmax
+
+
+@_try_torch_compile(dynamic=False)
+def _dequantize_blockwise_compute(
+ A_flat: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype
+):
+ n = A_flat.numel()
+ out = code[A_flat.to(torch.int64)]
+ rem = n % blocksize
+ if rem == 0:
+ out = (out.reshape(-1, blocksize) * absmax.view(-1, 1)).reshape(n)
+ else:
+ full = n - rem
+ blocks = full // blocksize
+ out = torch.cat(
+ [
+ (out[:full].reshape(blocks, blocksize) * absmax[:blocks].view(-1, 1)).reshape(full),
+ out[full:] * absmax[blocks],
+ ]
+ )
+ return out.to(dtype)
+
+
+@register_kernel("bitsandbytes::dequantize_blockwise", "default")
+def _(A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype) -> torch.Tensor:
+ return _dequantize_blockwise_compute(A.reshape(-1), absmax, code, blocksize, dtype).reshape(A.shape)
+
+
+@cache
+def _get_4bit_quantize_bounds(quant_type: str, device: torch.device):
+ code = _get_4bit_code(quant_type, device)
+ order = torch.argsort(code)
+ midpoints = (code[order[:-1]] + code[order[1:]]) / 2
+ return midpoints, order # NF4 order is identity (sorted); FP4 needs remap
+
+
+@register_kernel("bitsandbytes::quantize_4bit", "default")
+def _(
+ A: torch.Tensor, blocksize: int, quant_type: str, quant_storage: torch.dtype
+) -> tuple[torch.Tensor, torch.Tensor]:
+ bounds, order = _get_4bit_quantize_bounds(quant_type, A.device)
+ A_flat = A.reshape(-1).float()
+ n = A_flat.numel()
+ rem = n % blocksize
+ full = n - rem
+ blocks = full // blocksize
+ A_com = A_flat[:full].reshape(blocks, blocksize)
+ absmax = A_com.abs().max(dim=-1)[0]
+ scaled = torch.clamp(A_com * (1.0 / absmax.clamp(min=1e-38).view(-1, 1)), -1, 1).reshape(-1)
+ if rem:
+ am = A_flat[full:].abs().max().clamp(min=1e-38)
+ absmax = torch.cat([absmax, am.unsqueeze(0)])
+ scaled = torch.cat([scaled, torch.clamp(A_flat[full:] / am, -1, 1)])
+ if scaled.numel() % 2:
+ scaled = torch.nn.functional.pad(scaled, (0, 1))
+ q = torch.bucketize(scaled, bounds, out_int32=True)
+ if quant_type != "nf4":
+ q = order[q]
+ q8 = q.to(torch.uint8)
+ packed = ((q8[::2] << 4) | q8[1::2]).unsqueeze(1)
+ if quant_storage != torch.uint8:
+ packed = packed.squeeze().view(quant_storage).unsqueeze(1)
+ return packed, absmax
+
+
+@_try_torch_compile(dynamic=False)
+def _dequantize_4bit_compute(
+ A_flat: torch.Tensor,
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+):
+ n = prod(shape)
+ out_dq = torch.empty(A_flat.size(0) * 2, dtype=torch.int32, device=A_flat.device)
+ out_dq[1::2] = A_flat & 0xF
+ out_dq[::2] = A_flat >> 4
+ out_dq = code[out_dq][:n] # stays fp32, matches C++ / CUDA behavior
+ rem = n % blocksize
+ if rem:
+ full = n - rem
+ blocks = full // blocksize
+ out = torch.empty(n, dtype=torch.float32, device=A_flat.device)
+ out[:full] = (out_dq[:full].view(-1, blocksize) * absmax[:blocks].view(-1, 1)).reshape(full)
+ out[full:] = out_dq[full:] * absmax[blocks]
+ else:
+ out = (out_dq.view(-1, blocksize) * absmax.view(-1, 1)).reshape(n)
+ return out.reshape(-1, *shape[1:]).to(dtype)
+
+
+@register_kernel("bitsandbytes::dequantize_4bit", "default")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ if A.dtype != torch.uint8:
+ A = A.view(torch.uint8)
+ code = _get_4bit_code(quant_type, A.device)
+ return _dequantize_4bit_compute(A.reshape(-1), absmax, code, blocksize, shape, dtype)
+
+
+@register_kernel("bitsandbytes::gemv_4bit", "default")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+) -> torch.Tensor:
+ # Applied from dequantize_4bit
+ quant_type = "fp4" if code[1] > 0 else "nf4"
+ B_dq = torch.ops.bitsandbytes.dequantize_4bit.default(B, absmax, blocksize, quant_type, shapeB, A.dtype)
+
+ return torch.nn.functional.linear(
+ A,
+ B_dq,
+ bias=None,
+ )
+
+
+def _gemm_4bit_default_impl(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ bias: Optional[torch.Tensor] = None,
+ absmax_8bit: Optional[torch.Tensor] = None,
+ absmax_code: Optional[torch.Tensor] = None,
+ absmax_offset: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ # When nested, per-block scale = absmax_code[absmax_8bit[i]] * absmax[i // 256] + absmax_offset
+ if absmax_8bit is not None:
+ absmax = (
+ torch.ops.bitsandbytes.dequantize_blockwise.default(absmax_8bit, absmax, absmax_code, 256, torch.float32)
+ + absmax_offset
+ )
+ B_dq = torch.ops.bitsandbytes.dequantize_4bit.default(B, absmax, blocksize, quant_type, shapeB, A.dtype)
+ return torch.nn.functional.linear(A, B_dq, bias)
+
+
+register_kernel("bitsandbytes::gemm_4bit", "default")(_gemm_4bit_default_impl)
+
+
+MOMENTUM = 0
+RMSPROP = 1
+ADAGRAD = 2
+ADAM = 3
+# LION should be larger than MOMENTUM, RMSPROP, ADAGRAD due to comparison in kernels
+LION = 4
+ADEMAMIX = 5
+
+name2optimizer_id = {
+ "momentum": MOMENTUM,
+ "lars": MOMENTUM,
+ "rmsprop": RMSPROP,
+ "adagrad": ADAGRAD,
+ "adam": ADAM,
+ "lamb": ADAM,
+ "lion": LION,
+ "ademamix": ADEMAMIX,
+}
+
+
+@_try_torch_compile
+def _optimizer_precondition_32bit(
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: torch.Tensor,
+ beta1: float,
+ beta2: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float,
+ optimizer_id: int,
+):
+ """Preprocessing optimizer, computing update norm"""
+
+ g_vals = gnorm_scale * g
+
+ if optimizer_id == 3: # ADAM
+ correction1 = 1.0 / (1.0 - beta1**step)
+ correction2 = 1.0 / (1.0 - beta2**step)
+
+ s1_vals = state1 * beta1 + (1.0 - beta1) * g_vals
+ s2_vals = state2 * beta2 + (1.0 - beta2) * g_vals * g_vals
+
+ s1_vals = s1_vals * correction1
+ s2_vals = s2_vals * correction2
+
+ update_vals = s1_vals / (torch.sqrt(s2_vals) + eps)
+ update_norm = update_vals * update_vals
+
+ elif optimizer_id == 5: # ADEMAMIX
+ update_norm = state1
+
+ elif optimizer_id == 0: # MOMENTUM
+ if step == 1:
+ s1_vals = g_vals
+ else:
+ s1_vals = state1 * beta1 + g_vals
+ update_norm = s1_vals * s1_vals
+
+ elif optimizer_id == 4: # LION
+ s1_vals = state1 * beta2 + (1.0 - beta2) * g_vals
+ update_norm = s1_vals
+
+ elif optimizer_id == 1: # RMSPROP
+ s1_vals = state1 * beta1 + (1.0 - beta1) * g_vals * g_vals
+ update_vals = g_vals / (torch.sqrt(s1_vals) + eps)
+ update_norm = update_vals * update_vals
+
+ elif optimizer_id == 2: # ADAGRAD
+ s1_vals = state1 + g_vals * g_vals
+ update_vals = g_vals / (torch.sqrt(s1_vals) + eps)
+ update_norm = update_vals * update_vals
+
+ total_norm = torch.sum(update_norm)
+ unorm_vec.add_(total_norm)
+
+
+@_try_torch_compile
+def _optimizer_update_32bit(
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float,
+ optimizer_id: int,
+):
+ """Unified optimizer update kernel"""
+
+ p_vals = p.float()
+ g_vals = (gnorm_scale * g).float()
+ # Coupled (L2) weight decay: fold wd into the gradient. This is correct for
+ # MOMENTUM/RMSPROP/ADAGRAD, but NOT for LION (id 4), which uses *decoupled*
+ # (AdamW-style) weight decay applied to the param directly (see the LION branch
+ # below and Chen et al. 2023). LION is intentionally excluded here.
+ if optimizer_id in [0, 1, 2] and weight_decay > 0.0:
+ g_vals = g_vals + p_vals * weight_decay
+
+ update_scale = 1.0
+ if max_unorm > 0.0:
+ current_unorm = torch.sqrt(unorm_vec)
+ if optimizer_id in [0, 1, 2, 4]: # 1-state optimizers
+ if current_unorm > max_unorm * param_norm + eps:
+ update_scale = (max_unorm * param_norm + eps) / current_unorm
+ else: # 2-state optimizers
+ if current_unorm > max_unorm * param_norm:
+ update_scale = (max_unorm * param_norm) / current_unorm
+
+ if optimizer_id == 3: # ADAM
+ s1_vals = state1 * beta1 + (1.0 - beta1) * g_vals
+ s2_vals = state2 * beta2 + (1.0 - beta2) * g_vals * g_vals
+
+ correction1 = 1.0 - beta1**step
+ correction2 = sqrt(1.0 - beta2**step)
+ step_size = -lr * correction2 / correction1
+
+ if weight_decay > 0.0:
+ p_vals = p_vals * (1.0 - lr * weight_decay)
+
+ update_val = update_scale * step_size * (s1_vals / (torch.sqrt(s2_vals) + eps * correction2))
+ p_vals = p_vals + update_val
+
+ state1.copy_(s1_vals)
+ state2.copy_(s2_vals)
+
+ elif optimizer_id == 5: # ADEMAMIX
+ s1_vals = state1[0]
+ s3_vals = state1[1]
+ s2_vals = state2
+
+ m1 = s1_vals * beta1 + (1.0 - beta1) * g_vals
+ m2 = s3_vals * beta3 + (1.0 - beta3) * g_vals
+ nu = s2_vals * beta2 + (1.0 - beta2) * g_vals * g_vals
+
+ correction1 = 1.0 - beta1**step
+ correction2 = sqrt(1.0 - beta2**step)
+
+ if weight_decay > 0.0:
+ p_vals = p_vals * (1.0 - lr * weight_decay)
+
+ mixed_momentum = (m1 / correction1) + (alpha * m2)
+ adaptive_term = (torch.sqrt(nu) / correction2) + eps
+ p_vals = p_vals - lr * (mixed_momentum / adaptive_term)
+
+ state1[0].copy_(m1)
+ state1[1].copy_(m2)
+ state2.copy_(nu)
+
+ elif optimizer_id == 0: # MOMENTUM
+ if step == 1:
+ s1_vals = g_vals
+ else:
+ s1_vals = state1 * beta1 + g_vals
+
+ update_val = update_scale * (-lr * s1_vals)
+ p_vals = p_vals + update_val
+
+ state1.copy_(s1_vals)
+
+ elif optimizer_id == 4: # LION
+ # Lion uses decoupled weight decay: shrink the param directly (p *= 1 - lr*wd)
+ # rather than folding wd into the gradient. Matches the cpu backend, the CUDA
+ # 8-bit blockwise kernel, and the Lion paper (Chen et al. 2023).
+ if weight_decay > 0.0:
+ p_vals = p_vals * (1.0 - lr * weight_decay)
+
+ momentum_update = state1 * beta1 + (1.0 - beta1) * g_vals
+ update_val = update_scale * lr * torch.sign(momentum_update)
+ p_vals = p_vals - update_val
+
+ s1_vals = state1 * beta2 + (1.0 - beta2) * g_vals
+ state1.copy_(s1_vals)
+
+ elif optimizer_id == 1: # RMSPROP
+ s1_vals = state1 * beta1 + (1.0 - beta1) * g_vals * g_vals
+ update_val = update_scale * lr * g_vals / (torch.sqrt(s1_vals) + eps)
+ p_vals = p_vals - update_val
+
+ state1.copy_(s1_vals)
+
+ elif optimizer_id == 2: # ADAGRAD
+ s1_vals = state1 + g_vals * g_vals
+ update_val = lr * g_vals / (torch.sqrt(s1_vals) + eps)
+ p_vals = p_vals - update_val
+
+ state1.copy_(s1_vals)
+
+ p.copy_(p_vals)
+
+
+@register_kernel("bitsandbytes::optimizer_update_32bit", "default")
+def _(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float = 1.0,
+ skip_zeros=False,
+) -> None:
+ """
+ 32-bit optimizer implemented by PyTorch with @torch.compile
+ """
+ if skip_zeros:
+ raise NotImplementedError("skip_zeros is not supported yet")
+
+ optimizer_id = name2optimizer_id[optimizer_name]
+
+ if optimizer_name == "lion":
+ _optimizer_update_32bit(
+ g,
+ p,
+ state1,
+ state2,
+ unorm_vec,
+ max_unorm,
+ param_norm,
+ beta1,
+ beta2,
+ beta3,
+ alpha,
+ eps,
+ weight_decay,
+ step,
+ lr,
+ gnorm_scale,
+ optimizer_id,
+ )
+
+ if max_unorm > 0.0:
+ unorm_vec.zero_()
+ _optimizer_precondition_32bit(
+ g, p, state1, state2, unorm_vec, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale, optimizer_id
+ )
+ else:
+ if max_unorm > 0.0:
+ unorm_vec.zero_()
+ _optimizer_precondition_32bit(
+ g, p, state1, state2, unorm_vec, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale, optimizer_id
+ )
+
+ _optimizer_update_32bit(
+ g,
+ p,
+ state1,
+ state2,
+ unorm_vec,
+ max_unorm,
+ param_norm,
+ beta1,
+ beta2,
+ beta3,
+ alpha,
+ eps,
+ weight_decay,
+ step,
+ lr,
+ gnorm_scale,
+ optimizer_id,
+ )
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/hpu/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/hpu/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/hpu/ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/hpu/ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..645687598c23523319450b0613cea42a319156d1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/hpu/ops.py
@@ -0,0 +1,53 @@
+from collections.abc import Sequence
+import math
+
+import torch
+
+from ..._ops import register_kernel
+from ..utils import GAUDI_SW_VER
+
+
+# convert btw standard 4-bit compression format and ipex compression format
+# needed for backward compatibility with older versions of gaudi sw
+def _reverse_4bit_compress_format(weight: torch.Tensor):
+ out_1 = (weight & 0xF0) >> 4
+ out_2 = (weight & 0xF) << 4
+ out = out_1 | out_2
+ return out
+
+
+@register_kernel("bitsandbytes::dequantize_4bit", "hpu")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ if quant_type != "nf4":
+ raise ValueError(f"HPU backend only supports quant_type 'nf4', got {quant_type!r}")
+ if A.dtype not in (torch.bfloat16, torch.uint8):
+ raise ValueError(f"HPU backend only supports uint8 or bfloat16 storage, got {A.dtype}")
+
+ # Enable non uint8 dtype
+ if A.dtype != torch.uint8:
+ A = A.view(torch.uint8)
+
+ A = A.reshape(-1)
+
+ if GAUDI_SW_VER and (GAUDI_SW_VER.major < 1 or GAUDI_SW_VER.minor < 22):
+ A = _reverse_4bit_compress_format(A)
+
+ # HPU dequantization function for NF4 quantized tensors.
+ out_dq = torch.ops.hpu.dequantize_nf4(
+ A,
+ absmax.to(dtype),
+ blocksize,
+ out_shape=(math.prod(shape),),
+ out_dtype=dtype,
+ )
+
+ output = out_dq.reshape(shape)
+
+ return output
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/mps/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/mps/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/mps/ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/mps/ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..04c3d4fda9d070014b0dccde9f001cebfac178f7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/mps/ops.py
@@ -0,0 +1,277 @@
+"""MPS backend for bitsandbytes quantization ops.
+
+Hub kernels (kernels-community/bitsandbytes-mps) are attempted lazily on
+macOS 26+. On older macOS the hub kernel path is skipped entirely and
+default fallbacks or MPS-specific pure PyTorch fallbacks are used for all ops.
+
+Note: not all ops have implementations on the Hub kernels. Those that do not are also
+implemented using pure PyTorch fallbacks.
+"""
+
+from collections.abc import Sequence
+from math import prod
+import platform
+from typing import Optional
+
+import torch
+
+from ..._ops import register_kernel
+from ..default.ops import (
+ _dequantize_4bit_compute,
+ _get_4bit_quantize_bounds,
+ _try_torch_compile,
+)
+from ..utils import _get_4bit_code
+
+_QUANT_MAP = {"fp4": 1, "nf4": 2}
+
+_kernel = None
+
+_macos_major = int(platform.mac_ver()[0].split(".")[0]) if platform.mac_ver()[0] else 0
+
+# Pre-set to True on macOS < 26 so _get_kernel() never attempts the import.
+_kernel_load_failed = _macos_major < 26
+
+
+def _get_kernel():
+ global _kernel, _kernel_load_failed
+ if _kernel_load_failed:
+ return None
+ if _kernel is not None:
+ return _kernel
+ try:
+ from kernels import get_kernel
+
+ _kernel = get_kernel("kernels-community/bitsandbytes-mps", version=1)
+ except Exception:
+ _kernel_load_failed = True
+ return None
+ return _kernel
+
+
+@_try_torch_compile(dynamic=True)
+def _quantize_blockwise_compute(
+ A_flat: torch.Tensor, code: torch.Tensor, blocksize: int
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """
+ On torch <= 2.12, torch.bucketize does not perform well.
+ Implements blockwise quantization using a binary search instead of using the default.
+ """
+ n = A_flat.numel()
+ rem = n % blocksize
+ full = n - rem
+ blocks = full // blocksize
+ A_com = A_flat[:full].reshape(blocks, blocksize)
+ absmax = A_com.abs().max(dim=-1)[0]
+ scaled = torch.clamp(A_com * (1.0 / absmax.clamp(min=1e-38).view(-1, 1)), -1, 1).reshape(-1)
+ if rem:
+ am = A_flat[full:].abs().max().clamp(min=1e-38)
+ absmax = torch.cat([absmax, am.unsqueeze(0)])
+ scaled = torch.cat([scaled, torch.clamp(A_flat[full:] / am, -1, 1)])
+ bounds = (code[:-1] + code[1:]) / 2
+ n_bounds = bounds.shape[0]
+ n_iters = n_bounds.bit_length()
+ lo = torch.zeros(scaled.shape, dtype=torch.int16, device=scaled.device)
+ hi = torch.full(scaled.shape, n_bounds, dtype=torch.int16, device=scaled.device)
+ for _ in range(n_iters):
+ mid = (lo + hi) >> 1
+ val = bounds[mid.to(torch.int64)]
+ lo = torch.where(val < scaled, (mid + 1).to(torch.int16), lo)
+ hi = torch.where(val >= scaled, mid, hi)
+ return lo.to(torch.uint8), absmax
+
+
+@register_kernel("bitsandbytes::quantize_blockwise", "mps")
+def _(A: torch.Tensor, code: torch.Tensor, blocksize: int) -> tuple[torch.Tensor, torch.Tensor]:
+ q, absmax = _quantize_blockwise_compute(A.reshape(-1).float(), code.float(), blocksize)
+ return q.reshape(A.shape), absmax
+
+
+@_try_torch_compile(dynamic=True)
+def _quantize_4bit_compute(
+ A_flat: torch.Tensor,
+ blocksize: int,
+ bounds: torch.Tensor,
+ order: torch.Tensor,
+ nf4: bool,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ n = A_flat.numel()
+ rem = n % blocksize
+ full = n - rem
+ blocks = full // blocksize
+ A_com = A_flat[:full].reshape(blocks, blocksize)
+ absmax = A_com.abs().max(dim=-1)[0]
+ scaled = torch.clamp(A_com * (1.0 / absmax.clamp(min=1e-38).view(-1, 1)), -1, 1).reshape(-1)
+ if rem:
+ am = A_flat[full:].abs().max().clamp(min=1e-38)
+ absmax = torch.cat([absmax, am.unsqueeze(0)])
+ scaled = torch.cat([scaled, torch.clamp(A_flat[full:] / am, -1, 1)])
+ if scaled.numel() % 2:
+ scaled = torch.nn.functional.pad(scaled, (0, 1))
+ idx = torch.zeros(scaled.shape, dtype=torch.int8, device=scaled.device)
+ for b in bounds:
+ idx = idx + (scaled > b).to(torch.int8)
+ if not nf4:
+ idx = order[idx.to(torch.int32)]
+ q8 = idx.to(torch.uint8)
+ return (q8[::2] << 4) | q8[1::2], absmax
+
+
+def _quantize_4bit_fallback(
+ A: torch.Tensor, blocksize: int, quant_type: str, quant_storage: torch.dtype
+) -> tuple[torch.Tensor, torch.Tensor]:
+ bounds, order = _get_4bit_quantize_bounds(quant_type, A.device)
+ packed, absmax = _quantize_4bit_compute(A.reshape(-1).float(), blocksize, bounds, order, quant_type == "nf4")
+ packed = packed.unsqueeze(1)
+ if quant_storage != torch.uint8:
+ packed = packed.squeeze().view(quant_storage).unsqueeze(1)
+ return packed, absmax
+
+
+@register_kernel("bitsandbytes::quantize_4bit", "mps")
+def _(
+ A: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ quant_storage: torch.dtype,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ if blocksize in (64, 128, 256, 512) and (k := _get_kernel()) is not None:
+ packed, absmax = k.quantize_4bit(A.contiguous(), blocksize, _QUANT_MAP[quant_type])
+ packed = packed.view(quant_storage).unsqueeze(1)
+ return packed, absmax
+ return _quantize_4bit_fallback(A, blocksize, quant_type, quant_storage)
+
+
+def _dequantize_4bit_impl(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ if A.dtype != torch.uint8:
+ A = A.view(torch.uint8)
+
+ # Use HF Hub kernel when supported.
+ if blocksize in (64, 128, 256, 512) and (k := _get_kernel()) is not None:
+ numel = prod(shape)
+ out = k.dequantize_4bit(A, absmax, blocksize, _QUANT_MAP[quant_type], numel, dtype)
+ return out.reshape(shape)
+
+ # Fallback to implementation from default backend.
+ code = _get_4bit_code(quant_type, A.device)
+ return _dequantize_4bit_compute(A.reshape(-1), absmax, code, blocksize, shape, dtype)
+
+
+@register_kernel("bitsandbytes::dequantize_4bit", "mps")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ return _dequantize_4bit_impl(A, absmax, blocksize, quant_type, shape, dtype)
+
+
+@register_kernel("bitsandbytes::dequantize_4bit.out", "mps")
+def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ result = _dequantize_4bit_impl(A, absmax, blocksize, quant_type, shape, dtype)
+ out.copy_(result)
+
+
+def _gemv_4bit_impl(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+) -> torch.Tensor:
+ if blocksize in (64, 128, 256) and (k := _get_kernel()) is not None:
+ if B.dtype != torch.uint8:
+ B = B.view(torch.uint8)
+
+ output_features = shapeB[0]
+ quant_type_int = _QUANT_MAP["fp4"] if code[1] > 0 else _QUANT_MAP["nf4"]
+
+ return k.gemv_4bit(A, B, absmax, output_features, blocksize, quant_type_int)
+
+ quant_type = "fp4" if code[1] > 0 else "nf4"
+ B_dq = _dequantize_4bit_impl(B, absmax, blocksize, quant_type, shapeB, A.dtype)
+ return torch.nn.functional.linear(A, B_dq)
+
+
+@register_kernel("bitsandbytes::gemv_4bit", "mps")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+) -> torch.Tensor:
+ return _gemv_4bit_impl(A, B, shapeB, absmax, code, blocksize)
+
+
+@register_kernel("bitsandbytes::gemv_4bit.out", "mps")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ out: torch.Tensor,
+) -> None:
+ result = _gemv_4bit_impl(A, B, shapeB, absmax, code, blocksize)
+ out.copy_(result)
+
+
+@register_kernel("bitsandbytes::gemm_4bit", "mps")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ bias: Optional[torch.Tensor] = None,
+ absmax_8bit: Optional[torch.Tensor] = None,
+ absmax_code: Optional[torch.Tensor] = None,
+ absmax_offset: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ K = A.shape[-1]
+ M = A.numel() // K
+ N = shapeB[0]
+
+ # For nested absmax, we don't have a fused implementation yet.
+ # Dequantize the absmax values first.
+ if absmax_8bit is not None:
+ absmax = (
+ torch.ops.bitsandbytes.dequantize_blockwise.default(absmax_8bit, absmax, absmax_code, 256, torch.float32)
+ + absmax_offset
+ )
+
+ # Use HF Hub kernel when supported for GEMV.
+ if M == 1 and blocksize in (64, 128, 256) and (k := _get_kernel()) is not None:
+ if B.dtype != torch.uint8:
+ B = B.view(torch.uint8)
+ result = k.gemv_4bit(A, B, absmax.view(N, -1), N, blocksize, _QUANT_MAP[quant_type])
+ if bias is not None:
+ result = result + bias
+ return result
+
+ # Fallback: dequantize + linear.
+ B_dq = _dequantize_4bit_impl(B, absmax, blocksize, quant_type, shapeB, A.dtype)
+ return torch.nn.functional.linear(A, B_dq, bias)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_4bit.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_4bit.py
new file mode 100644
index 0000000000000000000000000000000000000000..bdd59fad2eaf5d991f39d85627133964dbcaefa9
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_4bit.py
@@ -0,0 +1,577 @@
+import torch
+
+import triton
+import triton.language as tl
+
+
+# Triton implementation of similar CUDA kernel to avoid loading code from csrc/kernels.cu::dQuantizeFP4
+# @triton.autotune(
+# configs=[
+# triton.Config({"SPLIT_NUM_BLOCKS": 1, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# triton.Config({"SPLIT_NUM_BLOCKS": 2, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# triton.Config({"SPLIT_NUM_BLOCKS": 1}),
+# triton.Config({"SPLIT_NUM_BLOCKS": 2}),
+# triton.Config({"SPLIT_NUM_BLOCKS": 4}),
+# triton.Config({"SPLIT_NUM_BLOCKS": 8}),
+# ],
+# key=["n_elements"],
+# )
+@triton.jit
+def quantize_fp4_blockwise_kernel(
+ A_ptr,
+ absmax_ptr,
+ out_ptr,
+ n_elements,
+ BLOCK_SIZE: tl.constexpr,
+ SPLIT_NUM_BLOCKS: tl.constexpr,
+):
+ PAIRED_SPLIT_NUM_BLOCKS: tl.constexpr = SPLIT_NUM_BLOCKS * 2
+ block_start_idx = tl.program_id(0) * PAIRED_SPLIT_NUM_BLOCKS
+ thread_idx = tl.arange(0, PAIRED_SPLIT_NUM_BLOCKS * BLOCK_SIZE)
+
+ offsets = block_start_idx * BLOCK_SIZE + thread_idx
+ mask = offsets < n_elements
+
+ A = tl.load(A_ptr + offsets, mask=mask, other=0.0)
+
+ # To be able process several blocks -> (PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE)
+ A_reshaped = tl.reshape(A, (PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE))
+
+ # Calculating absamax for each block
+ absmax = tl.max(tl.abs(A_reshaped), axis=1)
+ tl.store(absmax_ptr + block_start_idx + tl.arange(0, PAIRED_SPLIT_NUM_BLOCKS), absmax)
+
+ A_normalized = A_reshaped / absmax[:, None]
+ A_normalized = tl.clamp(A_normalized, -1.0, 1.0)
+
+ sign = tl.where(A_normalized < 0, 0b1000, 0b0000)
+ A_absf = tl.abs(A_normalized)
+
+ result = tl.where(
+ A_absf > 0.29166667,
+ tl.where(
+ A_absf > 0.583333, tl.where(A_absf > 0.8333333, 0b011, 0b010), tl.where(A_absf > 0.4166667, 0b101, 0b100)
+ ),
+ tl.where(
+ A_absf > 0.0859375,
+ tl.where(A_absf > 0.20833333, 0b0111, 0b0110),
+ tl.where(A_absf > 0.00260417, 0b0001, 0b0000),
+ ),
+ )
+ quantized = (result ^ sign).to(tl.uint8)
+
+ quantized = quantized.reshape((PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE // 2, 2))
+ left, right = quantized.split()
+ packed = left << 4 | (right & 0xF)
+
+ packed_flat = tl.reshape(packed, (BLOCK_SIZE * SPLIT_NUM_BLOCKS,))
+ out_offsets = block_start_idx * BLOCK_SIZE // 2 + tl.arange(0, SPLIT_NUM_BLOCKS * BLOCK_SIZE)
+ # Use n - n//2 instead of (n+1)//2 to avoid integer overflow for large n
+ out_mask = out_offsets < (n_elements - n_elements // 2)
+ tl.store(out_ptr + out_offsets, packed_flat, mask=out_mask)
+
+
+# Triton implementation of similar CUDA kernel to avoid loading code from csrc/kernels.cu::dQuantizeNF4
+# @triton.autotune(
+# configs=[
+# triton.Config({"SPLIT_NUM_BLOCKS": 1, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# triton.Config({"SPLIT_NUM_BLOCKS": 2, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# triton.Config({"SPLIT_NUM_BLOCKS": 1}),
+# triton.Config({"SPLIT_NUM_BLOCKS": 2}),
+# triton.Config({"SPLIT_NUM_BLOCKS": 4}),
+# triton.Config({"SPLIT_NUM_BLOCKS": 8}),
+# ],
+# key=["n_elements"],
+# )
+@triton.jit
+def quantize_nf4_blockwise_kernel(
+ A_ptr,
+ absmax_ptr,
+ out_ptr,
+ n_elements,
+ BLOCK_SIZE: tl.constexpr,
+ SPLIT_NUM_BLOCKS: tl.constexpr,
+):
+ PAIRED_SPLIT_NUM_BLOCKS: tl.constexpr = SPLIT_NUM_BLOCKS * 2
+ block_start_idx = tl.program_id(0) * PAIRED_SPLIT_NUM_BLOCKS
+ thread_idx = tl.arange(0, PAIRED_SPLIT_NUM_BLOCKS * BLOCK_SIZE)
+
+ offsets = block_start_idx * BLOCK_SIZE + thread_idx
+ mask = offsets < n_elements
+
+ A = tl.load(A_ptr + offsets, mask=mask, other=0.0)
+
+ # To be able process several blocks -> (PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE)
+ A_reshaped = tl.reshape(A, (PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE))
+
+ # Calculating absamax for each block
+ absmax = tl.max(tl.abs(A_reshaped), axis=1)
+ tl.store(absmax_ptr + block_start_idx + tl.arange(0, PAIRED_SPLIT_NUM_BLOCKS), absmax)
+
+ A_normalized = A_reshaped / absmax[:, None]
+ A_normalized = tl.clamp(A_normalized, -1.0, 1.0)
+
+ result = tl.where(
+ A_normalized > 0.03979014977812767,
+ tl.where(
+ A_normalized > 0.3893125355243683,
+ tl.where(
+ A_normalized > 0.6427869200706482,
+ tl.where(A_normalized > 0.8614784181118011, 0b1111, 0b1110),
+ tl.where(A_normalized > 0.5016634166240692, 0b1101, 0b1100),
+ ),
+ tl.where(
+ A_normalized > 0.2035212516784668,
+ tl.where(A_normalized > 0.2920137718319893, 0b1011, 0b1010),
+ tl.where(A_normalized > 0.1202552504837513, 0b1001, 0b1000),
+ ),
+ ),
+ tl.where(
+ A_normalized > -0.33967943489551544,
+ tl.where(
+ A_normalized > -0.13791173323988914,
+ tl.where(A_normalized > -0.045525018125772476, 0b0111, 0b0110),
+ tl.where(A_normalized > -0.23460740596055984, 0b0101, 0b0100),
+ ),
+ tl.where(
+ A_normalized > -0.6106329262256622,
+ tl.where(A_normalized > -0.4599952697753906, 0b0011, 0b0010),
+ tl.where(A_normalized > -0.8480964004993439, 0b0001, 0b0000),
+ ),
+ ),
+ )
+ quantized = result.to(tl.uint8)
+
+ quantized = quantized.reshape((PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE // 2, 2))
+
+ left, right = quantized.split()
+ packed = left << 4 | (right & 0xF)
+
+ packed_flat = tl.reshape(packed, (BLOCK_SIZE * SPLIT_NUM_BLOCKS,))
+ out_offsets = block_start_idx * BLOCK_SIZE // 2 + tl.arange(0, SPLIT_NUM_BLOCKS * BLOCK_SIZE)
+ # Use n - n//2 instead of (n+1)//2 to avoid integer overflow for large n
+ out_mask = out_offsets < (n_elements - n_elements // 2)
+ tl.store(out_ptr + out_offsets, packed_flat, mask=out_mask)
+
+
+def quantize_4bit_blockwise_triton(A, blocksize, quant_type, blocks, absmax, num_elements, quantized_out):
+ # grid = lambda META: (triton.cdiv(blocks, META["SPLIT_NUM_BLOCKS"]),)
+ split_num_blocks = 4
+ grid = (triton.cdiv(blocks, split_num_blocks),)
+ if quant_type == "fp4":
+ quantize_fp4_blockwise_kernel[grid](
+ A_ptr=A,
+ absmax_ptr=absmax,
+ out_ptr=quantized_out,
+ n_elements=num_elements,
+ BLOCK_SIZE=blocksize,
+ SPLIT_NUM_BLOCKS=split_num_blocks,
+ )
+ else:
+ quantize_nf4_blockwise_kernel[grid](
+ A_ptr=A,
+ absmax_ptr=absmax,
+ out_ptr=quantized_out,
+ n_elements=num_elements,
+ BLOCK_SIZE=blocksize,
+ SPLIT_NUM_BLOCKS=split_num_blocks,
+ )
+ return quantized_out, absmax
+
+
+@triton.jit
+def dequant_4bit_body_util(a, offsets, quant_ptr, absmax_ptr, n_elems, QUANT_BLOCK: tl.constexpr):
+ PAIRED_QUANT_BLOCK: tl.constexpr = QUANT_BLOCK // 2
+ mask = offsets < n_elems
+ higher = a & 0xF
+ # lower 4bits
+ lower = a >> 4
+
+ abs_offsets = offsets // PAIRED_QUANT_BLOCK
+ absmax = tl.load(absmax_ptr + abs_offsets, mask=mask, other=1.0, eviction_policy="evict_last")
+
+ # apply conversion
+ lower_4 = tl.load(quant_ptr + lower, eviction_policy="evict_last")
+ higher_4 = tl.load(quant_ptr + higher, eviction_policy="evict_last")
+
+ mul_high = higher_4 * absmax
+ mul_low = lower_4 * absmax
+ out_dq = tl.interleave(mul_low, mul_high)
+ return out_dq
+
+
+# Triton implementation of similar CUDA kernel to avoid loading code from csrc/kernels.cu::dDequantizeFP4Tree
+@triton.jit
+def dequantize_fp4_tree(val, absmax):
+ # val: tl.tensor (uint8)
+ # absmax: tl.tensor (float32/float16)
+ # 00001100 00001011 00001001 00001111
+ sign = tl.where((val & 0b1000) == 0b1000, -1.0, 1.0) # -1
+ third_bit = (val & 0b0100) == 0b0100 # True
+ second_bit = (val & 0b0010) == 0b0010 # False
+ first_bit = (val & 0b0001) == 0b0001 # False
+
+ branch1 = tl.where(
+ second_bit,
+ tl.where(first_bit, 0.25, 0.16666667), # 1111, 1110
+ tl.where(first_bit, 0.5, 0.33333333), # 1101, 1100
+ )
+ branch2 = tl.where(
+ second_bit,
+ tl.where(first_bit, 1.0, 0.66666667), # 1011, 1010
+ tl.where(first_bit, 0.00520833, 0.0), # 1001, 1000
+ )
+ out = tl.where(third_bit, branch1, branch2)
+ return out * sign * absmax
+
+
+@triton.jit
+def dequant_fp4_body_util(a, offsets, absmax_ptr, n_elems, QUANT_BLOCK: tl.constexpr):
+ PAIRED_QUANT_BLOCK: tl.constexpr = QUANT_BLOCK // 2
+ mask = offsets < n_elems
+ higher = a & 0xF
+ lower = a >> 4
+
+ abs_offsets = offsets // PAIRED_QUANT_BLOCK
+ absmax = tl.load(absmax_ptr + abs_offsets, mask=mask, other=1.0, eviction_policy="evict_last")
+ mul_high = dequantize_fp4_tree(higher, absmax)
+ mul_low = dequantize_fp4_tree(lower, absmax)
+ out_dq = tl.interleave(mul_low, mul_high)
+ return out_dq
+
+
+# Triton implementation of similar CUDA kernel to avoid loading code from csrc/kernels.cu::dDequantizeNF4
+@triton.jit
+def dequantize_nf4_tree(val):
+ # val: tl.tensor (uint8)
+ cond0 = (val & 0b1000) == 0b1000
+ cond1 = (val & 0b0100) == 0b0100
+ cond2 = (val & 0b0010) == 0b0010
+ cond3 = (val & 0b0001) == 0b0001
+
+ # Positive branch (val & 0b1000) == 8
+ branch_pos = tl.where(
+ cond1,
+ tl.where(
+ cond2,
+ tl.where(cond3, 1.0, 0.7229568362236023), # 1111, 1110
+ tl.where(cond3, 0.5626170039176941, 0.44070982933044434), # 1101, 1100
+ ),
+ tl.where(
+ cond2,
+ tl.where(cond3, 0.33791524171829224, 0.24611230194568634), # 1011, 1010
+ tl.where(cond3, 0.16093020141124725, 0.07958029955625534), # 1001, 1000
+ ),
+ )
+
+ # Negative branch (val & 0b1000) == 0
+ branch_neg = tl.where(
+ cond1,
+ tl.where(
+ cond2,
+ tl.where(cond3, 0.0, -0.09105003625154495), # 0111, 0110
+ tl.where(cond3, -0.18477343022823334, -0.28444138169288635), # 0101, 0100
+ ),
+ tl.where(
+ cond2,
+ tl.where(cond3, -0.39491748809814453, -0.5250730514526367), # 0011, 0010
+ tl.where(cond3, -0.6961928009986877, -1.0), # 0001, 0000
+ ),
+ )
+ return tl.where(cond0, branch_pos, branch_neg)
+
+
+@triton.jit
+def dequant_nf4_body_util(a, offsets, absmax_ptr, n_elems, QUANT_BLOCK: tl.constexpr):
+ PAIRED_QUANT_BLOCK: tl.constexpr = QUANT_BLOCK // 2
+ mask = offsets < n_elems
+ higher = a & 0xF
+ # lower 4bits
+ lower = a >> 4
+
+ abs_offsets = offsets // PAIRED_QUANT_BLOCK
+ absmax = tl.load(absmax_ptr + abs_offsets, mask=mask, other=1.0, eviction_policy="evict_last")
+ mul_high = dequantize_nf4_tree(higher) * absmax
+ mul_low = dequantize_nf4_tree(lower) * absmax
+ out_dq = tl.interleave(mul_low, mul_high)
+ return out_dq
+
+
+# All such kernels are similar, so maybe code can be generalised.
+# @triton.autotune(
+# configs=[
+# # # triton.Config({'SPLIT_SIZE': 64}),
+# # # # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'auto'}, num_stages=2, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'large'}, num_stages=4, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'auto'}, num_stages=4, num_warps=32),
+# triton.Config({'SPLIT_SIZE': 128}),
+# triton.Config({'SPLIT_SIZE': 128}, num_warps = 32, num_stages = 2),
+# # # triton.Config({'SPLIT_SIZE': 128}, num_warps = 4, num_stages = 4),
+# # # # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'auto'}, num_stages=2, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'large'}, num_stages=4, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'auto'}, num_stages=4, num_warps=32),
+# triton.Config({'SPLIT_SIZE': 256}),
+# triton.Config({'SPLIT_SIZE': 256}, num_warps = 32, num_stages = 2),
+# # triton.Config({'SPLIT_SIZE': 256}, num_warps = 4, num_stages = 4),
+# triton.Config({'SPLIT_SIZE': 512}),
+# triton.Config({'SPLIT_SIZE': 512}, num_warps = 32, num_stages = 2),
+# # triton.Config({'SPLIT_SIZE': 512}, num_warps = 4, num_stages = 4),
+# # # # triton.Config({'SPLIT_SIZE': 512, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 512, 'grf_mode': 'auto'}, num_stages=2, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 512, 'grf_mode': 'large'}, num_stages=4, num_warps=32),
+# # # # triton.Config({'SPLIT_SIZE': 512, 'grf_mode': 'auto'}, num_stages=4, num_warps=32),
+# # # triton.Config({'SPLIT_SIZE': 1024}),
+# # # # triton.Config({'SPLIT_SIZE': 2048}),
+# # # # triton.Config({'SPLIT_SIZE': 4096}),
+# # # # triton.Config({'SPLIT_SIZE': 8192}),
+# # # # triton.Config({'SPLIT_SIZE': 16384}),
+# ],
+# key=['num_paired_elements'],
+# )
+@triton.jit
+def dequant_4bit_kernel(
+ a_ptr,
+ c_ptr,
+ quant_ptr,
+ absmax_ptr,
+ num_paired_elements,
+ num_output_elements,
+ QUANT_BLOCK: tl.constexpr,
+ SPLIT_SIZE: tl.constexpr,
+):
+ pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0.
+ block_start = pid * SPLIT_SIZE
+ offsets = block_start + tl.arange(0, SPLIT_SIZE)
+ mask = offsets < num_paired_elements
+
+ a = tl.load(a_ptr + offsets, mask, eviction_policy="evict_first")
+
+ out_dq = dequant_4bit_body_util(
+ a=a,
+ offsets=offsets,
+ quant_ptr=quant_ptr,
+ absmax_ptr=absmax_ptr,
+ n_elems=num_paired_elements,
+ QUANT_BLOCK=QUANT_BLOCK,
+ )
+
+ out_block_start = pid * SPLIT_SIZE * 2
+ offs = out_block_start + tl.arange(0, SPLIT_SIZE * 2)
+ mask = offs < num_output_elements
+ tl.store(c_ptr + offs, out_dq, mask)
+
+
+# @triton.autotune(
+# configs=[
+# triton.Config({'SPLIT_SIZE': 128}, num_warps = 32, num_stages = 2),
+# triton.Config({'SPLIT_SIZE': 256}),
+# triton.Config({'SPLIT_SIZE': 256}, num_warps = 32, num_stages = 2),
+# triton.Config({'SPLIT_SIZE': 512}),
+# triton.Config({'SPLIT_SIZE': 512}, num_warps = 32, num_stages = 2),
+# triton.Config({'SPLIT_SIZE': 1024}, num_warps = 32, num_stages = 2),
+# ],
+# key=['num_paired_elements'],
+# )
+@triton.jit
+def dequant_fp4_kernel(
+ a_ptr,
+ c_ptr,
+ absmax_ptr,
+ num_paired_elements,
+ num_output_elements,
+ QUANT_BLOCK: tl.constexpr,
+ SPLIT_SIZE: tl.constexpr,
+):
+ pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0.
+ block_start = pid * SPLIT_SIZE
+ offsets = block_start + tl.arange(0, SPLIT_SIZE)
+ mask = offsets < num_paired_elements
+
+ a = tl.load(a_ptr + offsets, mask, eviction_policy="evict_first")
+
+ out_dq = dequant_fp4_body_util(
+ a=a,
+ offsets=offsets,
+ absmax_ptr=absmax_ptr,
+ n_elems=num_paired_elements,
+ QUANT_BLOCK=QUANT_BLOCK,
+ )
+
+ out_block_start = pid * SPLIT_SIZE * 2
+ offs = out_block_start + tl.arange(0, SPLIT_SIZE * 2)
+ mask = offs < num_output_elements
+ tl.store(c_ptr + offs, out_dq, mask)
+
+
+# @triton.autotune(
+# configs=[
+# triton.Config({'SPLIT_SIZE': 128}, num_warps = 32, num_stages = 2),
+# triton.Config({'SPLIT_SIZE': 256}),
+# triton.Config({'SPLIT_SIZE': 256}, num_warps = 32, num_stages = 2),
+# triton.Config({'SPLIT_SIZE': 512}),
+# triton.Config({'SPLIT_SIZE': 512}, num_warps = 32, num_stages = 2),
+# triton.Config({'SPLIT_SIZE': 1024}, num_warps = 32, num_stages = 2),
+# ],
+# key=['num_paired_elements'],
+# )
+@triton.jit
+def dequant_nf4_kernel(
+ a_ptr,
+ c_ptr,
+ absmax_ptr,
+ num_paired_elements,
+ num_output_elements,
+ QUANT_BLOCK: tl.constexpr,
+ SPLIT_SIZE: tl.constexpr,
+):
+ pid = tl.program_id(axis=0) # We use a 1D launch grid so axis is 0.
+ block_start = pid * SPLIT_SIZE
+ offsets = block_start + tl.arange(0, SPLIT_SIZE)
+ mask = offsets < num_paired_elements
+
+ a = tl.load(a_ptr + offsets, mask, eviction_policy="evict_first")
+
+ out_dq = dequant_nf4_body_util(
+ a=a,
+ offsets=offsets,
+ absmax_ptr=absmax_ptr,
+ n_elems=num_paired_elements,
+ QUANT_BLOCK=QUANT_BLOCK,
+ )
+
+ out_block_start = pid * SPLIT_SIZE * 2
+ offs = out_block_start + tl.arange(0, SPLIT_SIZE * 2)
+ mask = offs < num_output_elements
+ tl.store(c_ptr + offs, out_dq, mask)
+
+
+def dequantize_4bit_impl(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ # It's will be processed as an array, so
+ # actual length is row * col
+ # Elements are in uint8 format, so interleaved
+ # so total amount of data is 2 * elem_count
+ number_of_paired_elements = A.numel()
+ num_output_elements = out.numel()
+ # we assume that split_size > quant_blocksize
+
+ SPLIT_SIZE = 256
+ # grid = lambda META: (triton.cdiv(number_of_paired_elements, META['SPLIT_SIZE']), )
+ grid = (triton.cdiv(number_of_paired_elements, SPLIT_SIZE),)
+ if quant_type == "fp4":
+ dequant_fp4_kernel[grid](A, out, absmax, number_of_paired_elements, num_output_elements, blocksize, SPLIT_SIZE)
+ else:
+ dequant_nf4_kernel[grid](A, out, absmax, number_of_paired_elements, num_output_elements, blocksize, SPLIT_SIZE)
+
+
+def dequantize_4bit_impl_passing_code(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ code: torch.Tensor,
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ number_of_paired_elements = A.numel()
+ num_output_elements = out.numel()
+ # we assume that split_size > quant_blocksize
+
+ SPLIT_SIZE = 256
+ # grid = lambda META: (triton.cdiv(number_of_paired_elements, META['SPLIT_SIZE']), )
+ grid = (triton.cdiv(number_of_paired_elements, SPLIT_SIZE),)
+ dequant_4bit_kernel[grid](
+ A, out, code, absmax, number_of_paired_elements, num_output_elements, blocksize, SPLIT_SIZE
+ )
+
+
+######################### Fallback dequantization functions #########################
+## for debug ##
+
+
+# @triton.autotune(
+# configs=[
+# # triton.Config({'SPLIT_NUM_BLOCKS': 1, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# # triton.Config({'SPLIT_NUM_BLOCKS': 1, 'grf_mode': 'auto'}, num_stages=2, num_warps=32),
+# # triton.Config({'SPLIT_NUM_BLOCKS': 1, 'grf_mode': 'large'}, num_stages=4, num_warps=32),
+# # #
+# # triton.Config({"SPLIT_NUM_BLOCKS": 1, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# #
+# triton.Config({"SPLIT_NUM_BLOCKS": 2}),
+# # triton.Config({"SPLIT_NUM_BLOCKS": 2, "grf_mode": "large"}, num_stages=2, num_warps=32),
+# # # triton.Config({'SPLIT_NUM_BLOCKS': 2, 'grf_mode': 'large'}, num_stages=4, num_warps=32),
+# # triton.Config({"SPLIT_NUM_BLOCKS": 2, "grf_mode": "auto"}, num_stages=2, num_warps=32),
+# # triton.Config({"SPLIT_NUM_BLOCKS": 2, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# # triton.Config({"SPLIT_NUM_BLOCKS": 4, "grf_mode": "large"}, num_stages=2, num_warps=32),
+# # triton.Config({"SPLIT_NUM_BLOCKS": 4, "grf_mode": "large"}, num_stages=4, num_warps=32),
+# # triton.Config({'SPLIT_NUM_BLOCKS': 8, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# ],
+# key=["n_elements", "BLOCK_SIZE"],
+# )
+@triton.jit
+def quantize_4bit_blockwise_kernel(
+ A_ptr,
+ code_ptr,
+ absmax_ptr,
+ out_ptr,
+ n_elements,
+ BLOCK_SIZE: tl.constexpr,
+ CODE_SIZE: tl.constexpr,
+ SPLIT_NUM_BLOCKS: tl.constexpr,
+):
+ PAIRED_SPLIT_NUM_BLOCKS: tl.constexpr = SPLIT_NUM_BLOCKS * 2
+ block_start_idx = tl.program_id(0) * PAIRED_SPLIT_NUM_BLOCKS
+ thread_idx = tl.arange(0, PAIRED_SPLIT_NUM_BLOCKS * BLOCK_SIZE)
+
+ offsets = block_start_idx * BLOCK_SIZE + thread_idx
+ mask = offsets < n_elements
+
+ A = tl.load(A_ptr + offsets, mask=mask, other=0.0)
+
+ # To be able process several blocks -> (PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE)
+ A_reshaped = tl.reshape(A, (PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE))
+
+ # Calculating absamax for each block
+ absmax = tl.max(tl.abs(A_reshaped), axis=1)
+ tl.store(absmax_ptr + block_start_idx + tl.arange(0, PAIRED_SPLIT_NUM_BLOCKS), absmax)
+
+ A_normalized = A_reshaped / absmax[:, None]
+ A_normalized = tl.clamp(A_normalized, -1.0, 1.0)
+
+ lower_pivot = tl.zeros((PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE), dtype=tl.int32)
+ upper_pivot = tl.full((PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE), CODE_SIZE - 1, dtype=tl.int32)
+
+ for _ in range(4): # ceil(log2(code_size)) = 4, actually, in general case should be input parameter
+ pivot = (lower_pivot + upper_pivot) // 2
+ val = tl.load(code_ptr + pivot)
+ is_higher = A_normalized > val # code[pivot]
+ lower_pivot = tl.where(is_higher, pivot, lower_pivot)
+ upper_pivot = tl.where(is_higher, upper_pivot, pivot)
+
+ # Choose closest level
+ lower_val = tl.load(code_ptr + lower_pivot)
+ upper_val = tl.load(code_ptr + upper_pivot)
+ lower_dist = tl.abs(A_normalized - lower_val)
+ upper_dist = tl.abs(A_normalized - upper_val)
+ quantized = tl.where(lower_dist <= upper_dist, lower_pivot, upper_pivot).to(tl.uint8)
+
+ quantized = quantized.reshape((PAIRED_SPLIT_NUM_BLOCKS, BLOCK_SIZE // 2, 2))
+ quantized = quantized.to(tl.uint8, bitcast=True)
+ left, right = quantized.split()
+ packed = left << 4 | (right & 0xF)
+
+ # Reduce don't guarantee the order of the elements passed to unite_2_int4
+ # packed = tl.reduce(quantized, axis=2, combine_fn=unite_2_int4)
+ # packed = packed.to(tl.uint8, bitcast=True)
+
+ packed_flat = tl.reshape(packed, (BLOCK_SIZE * SPLIT_NUM_BLOCKS,))
+ out_offsets = block_start_idx * BLOCK_SIZE // 2 + tl.arange(0, SPLIT_NUM_BLOCKS * BLOCK_SIZE)
+ out_mask = out_offsets < n_elements // 2
+ tl.store(out_ptr + out_offsets, packed_flat, mask=out_mask)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_8bit_quant.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_8bit_quant.py
new file mode 100644
index 0000000000000000000000000000000000000000..c0a5a21efff7507e1409fc2d021c9814bed22c1e
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_8bit_quant.py
@@ -0,0 +1,195 @@
+import torch
+
+import triton
+import triton.language as tl
+
+
+# @triton.autotune(
+# configs=[
+# # triton.Config({'SPLIT_SIZE': 64}),
+# # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'auto'}, num_stages=2, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'large'}, num_stages=4, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 64, 'grf_mode': 'auto'}, num_stages=4, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 128}),
+# # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'auto'}, num_stages=2, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'large'}, num_stages=4, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 128, 'grf_mode': 'auto'}, num_stages=4, num_warps=32),
+# triton.Config({"SPLIT_SIZE": 256}),
+# # triton.Config({'SPLIT_SIZE': 256, 'grf_mode': 'large'}, num_stages=2, num_warps=32),
+# # triton.Config({'SPLIT_SIZE': 256, 'grf_mode': 'auto'}, num_stages=2, num_warps=32),
+# triton.Config({"SPLIT_SIZE": 512}),
+# # triton.Config({'SPLIT_SIZE': 1024}),
+# ],
+# key=["num_paired_elements", "QUANT_BLOCK"],
+# )
+@triton.jit
+def dequant_8bit_kernel(
+ a_ptr,
+ out_ptr,
+ code_ptr,
+ absmax_ptr,
+ n,
+ QUANT_BLOCK: tl.constexpr,
+ SPLIT_SIZE: tl.constexpr,
+):
+ pid = tl.program_id(axis=0)
+ block_start = pid * SPLIT_SIZE
+ offsets = block_start + tl.arange(0, SPLIT_SIZE)
+ mask = offsets < n
+ out_dq = dequant_8bit_blockwise_kernel_util(a_ptr, offsets, code_ptr, absmax_ptr, mask, QUANT_BLOCK)
+ tl.store(out_ptr + offsets, out_dq, mask)
+
+
+def dequant_8bit_blockwise(
+ a: torch.Tensor,
+ absmax: torch.Tensor,
+ quant_state_code: torch.Tensor,
+ quant_blocksize: int = 64,
+ dtype: torch.dtype = None,
+ out: torch.Tensor = None,
+):
+ n = a.numel()
+ if out is None:
+ if dtype is None:
+ raise ValueError("If out is None, dtype must be specified")
+ out = torch.empty_like(a, dtype=dtype, device=a.device)
+
+ SPLIT_SIZE = 256
+ # grid = lambda META: (triton.cdiv(number_of_paired_elements, META["SPLIT_SIZE"]),)
+ grid = (triton.cdiv(n, SPLIT_SIZE),)
+ dequant_8bit_kernel[grid](
+ a,
+ out,
+ quant_state_code,
+ absmax,
+ n,
+ quant_blocksize,
+ SPLIT_SIZE,
+ )
+ return out
+
+
+# @triton.autotune(
+# configs=[
+# triton.Config({"SPLIT_NUM_BLOCKS": 1, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# triton.Config({"SPLIT_NUM_BLOCKS": 2, "grf_mode": "auto"}, num_stages=4, num_warps=32),
+# triton.Config({"SPLIT_NUM_BLOCKS": 1}),
+# triton.Config({"SPLIT_NUM_BLOCKS": 2}),
+# ],
+# key=["n_elements"],
+# )
+@triton.jit
+def quantize_8bit_blockwise_kernel(
+ A_ptr,
+ code_ptr,
+ absmax_ptr,
+ out_ptr,
+ n_elements,
+ BLOCK_SIZE: tl.constexpr,
+ CODE_SIZE: tl.constexpr,
+ SPLIT_NUM_BLOCKS: tl.constexpr,
+):
+ block_start_idx = tl.program_id(0) * SPLIT_NUM_BLOCKS
+ thread_idx = tl.arange(0, SPLIT_NUM_BLOCKS * BLOCK_SIZE)
+
+ offsets = block_start_idx * BLOCK_SIZE + thread_idx
+ mask = offsets < n_elements
+
+ A = tl.load(A_ptr + offsets, mask=mask, other=0.0)
+
+ quantized, absmax = quantize_8bit_blockwise_kernel_util(A, code_ptr, CODE_SIZE, BLOCK_SIZE, SPLIT_NUM_BLOCKS)
+ tl.store(absmax_ptr + block_start_idx + tl.arange(0, SPLIT_NUM_BLOCKS), absmax)
+ tl.store(out_ptr + offsets, quantized, mask=mask)
+
+
+def quantize_blockwise_triton(A, code, blocksize, absmax=None, out=None):
+ n = A.numel()
+ blocks = -(n // -blocksize)
+
+ if absmax is None:
+ absmax = torch.empty((blocks,), device=A.device, dtype=A.dtype)
+ if out is None:
+ out = torch.empty_like(A.flatten(), dtype=torch.uint8)
+
+ split_num_blocks = 1
+ grid = (triton.cdiv(blocks, split_num_blocks),)
+ # grid = lambda META: (triton.cdiv(blocks, META["SPLIT_NUM_BLOCKS"]),)
+ quantize_8bit_blockwise_kernel[grid](
+ A_ptr=A,
+ code_ptr=code,
+ absmax_ptr=absmax,
+ out_ptr=out,
+ n_elements=n,
+ BLOCK_SIZE=blocksize,
+ CODE_SIZE=code.numel(),
+ SPLIT_NUM_BLOCKS=split_num_blocks,
+ # num_warps=1,
+ # num_stages=2,
+ )
+ out = out.reshape(A.shape)
+
+ return out, absmax
+
+
+@triton.jit
+def quantize_8bit_blockwise_kernel_util(
+ a,
+ code_ptr,
+ CODE_SIZE: tl.constexpr,
+ BLOCK_SIZE: tl.constexpr,
+ N_PER_TH: tl.constexpr,
+):
+ # To be able process several blocks -> (BLOCK_SIZE, SPLIT_NUM_BLOCKS)
+ a_reshaped = tl.reshape(a, (N_PER_TH, BLOCK_SIZE))
+
+ # Calculating absmax for each block
+ absmax = tl.max(tl.abs(a_reshaped), axis=1)
+
+ a_normalized = a_reshaped / absmax[:, None]
+ a_normalized = tl.clamp(a_normalized, -1.0, 1.0)
+
+ lower_pivot = tl.zeros((N_PER_TH, BLOCK_SIZE), dtype=tl.int32)
+ upper_pivot = tl.full((N_PER_TH, BLOCK_SIZE), CODE_SIZE - 1, dtype=tl.int32)
+
+ # ceil(log2(code_size)) = 8, actually, in general case should be input parameter
+ for _ in range(8):
+ pivot = (lower_pivot + upper_pivot) // 2
+ val = tl.load(code_ptr + pivot)
+ is_higher = a_normalized > val # code[pivot]
+ lower_pivot = tl.where(is_higher, pivot, lower_pivot)
+ upper_pivot = tl.where(is_higher, upper_pivot, pivot)
+
+ # Choose closest level
+ lower_val = tl.load(code_ptr + lower_pivot)
+ upper_val = tl.load(code_ptr + upper_pivot)
+ lower_dist = tl.abs(a_normalized - lower_val)
+ upper_dist = tl.abs(a_normalized - upper_val)
+ quantized = tl.where(lower_dist <= upper_dist, lower_pivot, upper_pivot).to(tl.uint8)
+
+ # too slow approach
+ # diff = tl.abs(A_normalized[:, :, None] - code[None, None, :])
+ # quantized = tl.argmin(diff, axis=2).to(tl.uint8)
+
+ quantized_flat = tl.reshape(quantized, (BLOCK_SIZE * N_PER_TH,))
+ return quantized_flat, absmax
+
+
+@triton.jit
+def dequant_8bit_blockwise_kernel_util(
+ a_ptr,
+ offsets,
+ code_ptr,
+ absmax_ptr,
+ mask,
+ BLOCK_SIZE: tl.constexpr,
+):
+ a = tl.load(a_ptr + offsets, mask, other=0).to(tl.uint8)
+ scaled_int8 = tl.load(code_ptr + a, mask)
+ # Load scales
+ absmax_offsets = offsets // BLOCK_SIZE
+ absmax = tl.load(absmax_ptr + absmax_offsets, mask=mask, other=0.0, eviction_policy="evict_last")
+ # Apply scales
+ out_dq = scaled_int8 * absmax
+ return out_dq
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_optim.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_optim.py
new file mode 100644
index 0000000000000000000000000000000000000000..f7eb2e213be9f4a5a5ace6393b11ad616dc33d57
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/kernels_optim.py
@@ -0,0 +1,1177 @@
+import math
+from typing import Optional
+
+import torch
+
+import triton
+import triton.language as tl
+
+# from triton.language.extra import libdevice
+from .kernels_8bit_quant import (
+ dequant_8bit_blockwise,
+ dequant_8bit_blockwise_kernel_util,
+ quantize_8bit_blockwise_kernel_util,
+ quantize_blockwise_triton,
+)
+
+MOMENTUM = 0
+RMSPROP = 1
+ADAGRAD = 2
+ADAM = 3
+# LION should be larger than MOMENTUM, RMSPROP, ADAGRAD due to comparison in kernels
+LION = 4
+ADEMAMIX = 5
+
+name2optimizer_id = {
+ "momentum": MOMENTUM,
+ "lars": MOMENTUM,
+ "rmsprop": RMSPROP,
+ "adagrad": ADAGRAD,
+ "adam": ADAM,
+ "lamb": ADAM,
+ "lion": LION,
+ "ademamix": ADEMAMIX,
+}
+
+
+@triton.jit
+def _optimizer_precondition_2state_32bit(
+ g_ptr,
+ p_ptr,
+ state1_ptr,
+ state2_ptr,
+ unorm_ptr,
+ beta1: tl.constexpr,
+ beta2: tl.constexpr,
+ eps: tl.constexpr,
+ weight_decay: tl.constexpr,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale: tl.constexpr,
+ n_elements,
+ OPTIMIZER_ID: tl.constexpr,
+ BLOCK_SIZE: tl.constexpr,
+ N_PER_TH: tl.constexpr,
+):
+ """Preprocessing optimizer, computing update norm (2-state optimizer)"""
+ pid = tl.program_id(axis=0)
+ block_start_idx = pid * N_PER_TH
+ offsets = block_start_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE * N_PER_TH)
+ mask = offsets < n_elements
+
+ g_vals = tl.load(g_ptr + offsets, mask=mask, other=0.0)
+ s1_vals = tl.load(state1_ptr + offsets, mask=mask, other=0.0)
+ s2_vals = tl.load(state2_ptr + offsets, mask=mask, other=0.0)
+
+ g_vals = gnorm_scale * g_vals
+
+ correction1 = 1.0 / (1.0 - beta1_step)
+ correction2 = 1.0 / (1.0 - beta2_step)
+
+ if OPTIMIZER_ID == 3: # ADAM
+ s1_vals = s1_vals * beta1 + (1.0 - beta1) * g_vals
+ s2_vals = s2_vals * beta2 + (1.0 - beta2) * g_vals * g_vals
+
+ s1_vals = s1_vals * correction1
+ s2_vals = s2_vals * correction2
+
+ update_vals = s1_vals / (tl.sqrt(s2_vals) + eps)
+
+ update_norm = update_vals * update_vals
+
+ elif OPTIMIZER_ID == 5: # ADEMAMIX
+ update_norm = s1_vals
+
+ total_norm = tl.sum(tl.where(mask, update_norm, 0.0))
+
+ tl.atomic_add(unorm_ptr, total_norm)
+
+
+@triton.jit
+def _optimizer_precondition_1state_32bit(
+ g_ptr,
+ p_ptr,
+ state1_ptr,
+ state2_ptr,
+ unorm_ptr,
+ beta1: tl.constexpr,
+ beta2: tl.constexpr,
+ eps: tl.constexpr,
+ weight_decay,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale: tl.constexpr,
+ n_elements,
+ OPTIMIZER_ID: tl.constexpr,
+ BLOCK_SIZE: tl.constexpr,
+ N_PER_TH: tl.constexpr,
+):
+ """Preprocessing optimizer, computing update norm (1-state optimizer)"""
+ pid = tl.program_id(axis=0)
+ block_start_idx = pid * N_PER_TH
+ offsets = block_start_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE * N_PER_TH)
+ mask = offsets < n_elements
+
+ g_vals = tl.load(g_ptr + offsets, mask=mask, other=0.0)
+ s1_vals = tl.load(state1_ptr + offsets, mask=mask, other=0.0)
+
+ g_vals = gnorm_scale * g_vals
+
+ if OPTIMIZER_ID == 0: # MOMENTUM
+ if step == 1:
+ # Cast to fp32 to avoid type mismatch: s1_vals is fp32 but g_vals may be fp16.
+ s1_vals = g_vals.to(tl.float32)
+ else:
+ s1_vals = s1_vals * beta1 + g_vals
+ update_norm = s1_vals * s1_vals
+
+ elif OPTIMIZER_ID == 4: # LION
+ s1_vals = s1_vals * beta2 + (1.0 - beta2) * g_vals
+ update_norm = s1_vals
+
+ elif OPTIMIZER_ID == 1: # RMSPROP
+ s1_vals = s1_vals * beta1 + (1.0 - beta1) * g_vals * g_vals
+ update_vals = g_vals / (tl.sqrt(s1_vals) + eps)
+ update_norm = update_vals * update_vals
+
+ elif OPTIMIZER_ID == 2: # ADAGRAD
+ s1_vals = s1_vals + g_vals * g_vals
+ update_vals = g_vals / (tl.sqrt(s1_vals) + eps)
+ update_norm = update_vals * update_vals
+
+ total_norm = tl.sum(tl.where(mask, update_norm, 0.0))
+
+ tl.atomic_add(unorm_ptr, total_norm)
+
+
+@triton.jit
+def _optimizer_update_2state_32bit_triton_kernel(
+ g_ptr,
+ p_ptr,
+ state1_ptr,
+ state2_ptr,
+ unorm_ptr,
+ max_unorm: tl.constexpr,
+ param_norm,
+ beta1: tl.constexpr,
+ beta2: tl.constexpr,
+ beta3,
+ alpha,
+ eps: tl.constexpr,
+ weight_decay: tl.constexpr,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale: tl.constexpr,
+ skip_zeros,
+ n_elements,
+ OPTIMIZER_ID: tl.constexpr,
+ BLOCK_SIZE: tl.constexpr,
+ N_PER_TH: tl.constexpr,
+):
+ """2-state optimizer kernel"""
+ pid = tl.program_id(axis=0)
+ block_start_idx = pid * N_PER_TH
+ offsets = block_start_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE * N_PER_TH)
+ mask = offsets < n_elements
+
+ g_vals = tl.load(g_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
+ p_vals = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
+ s1_vals = tl.load(state1_ptr + offsets, mask=mask, other=0.0)
+ s2_vals = tl.load(state2_ptr + offsets, mask=mask, other=0.0)
+
+ if OPTIMIZER_ID == 5: # ADEMAMIX
+ s3_vals = tl.load(state1_ptr + n_elements + offsets, mask=mask, other=0.0)
+
+ g_vals = gnorm_scale * g_vals
+
+ update_scale = 1.0
+ if max_unorm > 0.0:
+ current_unorm = tl.sqrt(tl.load(unorm_ptr))
+ if current_unorm > max_unorm * param_norm:
+ update_scale = (max_unorm * param_norm) / current_unorm
+
+ if OPTIMIZER_ID == 3: # ADAM
+ s1_vals = s1_vals * beta1 + (1.0 - beta1) * g_vals
+ s2_vals = s2_vals * beta2 + (1.0 - beta2) * g_vals * g_vals
+
+ correction1 = 1.0 - beta1_step
+ correction2 = tl.sqrt(1.0 - beta2_step)
+ step_size = -lr * correction2 / correction1
+
+ if weight_decay > 0.0:
+ p_vals = p_vals * (1.0 - lr * weight_decay)
+
+ update_val = update_scale * step_size * (s1_vals / (tl.sqrt(s2_vals) + eps * correction2))
+ p_vals = p_vals + update_val
+
+ elif OPTIMIZER_ID == 5: # ADEMAMIX
+ s1_vals = s1_vals * beta1 + (1.0 - beta1) * g_vals # m1
+ s3_vals = s3_vals * beta3 + (1.0 - beta3) * g_vals # m2
+ s2_vals = s2_vals * beta2 + (1.0 - beta2) * g_vals * g_vals # nu
+
+ correction1 = 1.0 - beta1_step
+ correction2 = tl.sqrt(1.0 - beta2_step)
+
+ if weight_decay > 0.0:
+ p_vals = p_vals * (1.0 - lr * weight_decay)
+
+ mixed_momentum = (s1_vals / correction1) + (alpha * s3_vals)
+ adaptive_term = (tl.sqrt(s2_vals) / correction2) + eps
+ p_vals = p_vals - lr * (mixed_momentum / adaptive_term)
+
+ tl.store(p_ptr + offsets, p_vals, mask=mask)
+ tl.store(state1_ptr + offsets, s1_vals, mask=mask)
+ tl.store(state2_ptr + offsets, s2_vals, mask=mask)
+
+ if OPTIMIZER_ID == 5: # ADEMAMIX
+ tl.store(state1_ptr + n_elements + offsets, s3_vals, mask=mask)
+
+
+@triton.jit
+def _optimizer_update_1state_32bit_triton_kernel(
+ g_ptr,
+ p_ptr,
+ state1_ptr,
+ state2_ptr,
+ unorm_ptr,
+ max_unorm: tl.constexpr,
+ param_norm,
+ beta1: tl.constexpr,
+ beta2: tl.constexpr,
+ beta3,
+ alpha,
+ eps: tl.constexpr,
+ weight_decay: tl.constexpr,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale: tl.constexpr,
+ skip_zeros,
+ n_elements,
+ OPTIMIZER_ID: tl.constexpr,
+ BLOCK_SIZE: tl.constexpr,
+ N_PER_TH: tl.constexpr,
+):
+ """1-state optimizer kernel"""
+ pid = tl.program_id(axis=0)
+ block_start_idx = pid * N_PER_TH
+ offsets = block_start_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE * N_PER_TH)
+ mask = offsets < n_elements
+
+ g_vals = tl.load(g_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
+ p_vals = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
+ s1_vals = tl.load(state1_ptr + offsets, mask=mask, other=0.0)
+
+ g_vals = gnorm_scale * g_vals
+ # Coupled (L2) weight decay: fold wd into the gradient. This is correct for
+ # MOMENTUM/RMSPROP/ADAGRAD, but NOT for LION (id 4), which uses *decoupled*
+ # (AdamW-style) weight decay applied to the param directly (see the LION branch
+ # below and Chen et al. 2023). LION is intentionally excluded here.
+ if OPTIMIZER_ID != 4 and weight_decay > 0.0:
+ g_vals = g_vals + p_vals * weight_decay
+
+ update_scale = 1.0
+ if max_unorm > 0.0:
+ current_unorm = tl.sqrt(tl.load(unorm_ptr))
+ if current_unorm > max_unorm * param_norm + eps:
+ update_scale = (max_unorm * param_norm + eps) / current_unorm
+
+ if OPTIMIZER_ID == 0: # MOMENTUM
+ if step == 1:
+ s1_vals = g_vals
+ else:
+ s1_vals = s1_vals * beta1 + g_vals
+
+ update_val = update_scale * (-lr * s1_vals)
+ p_vals = p_vals + update_val
+
+ elif OPTIMIZER_ID == 4: # LION
+ # Lion uses decoupled weight decay: shrink the param directly (p *= 1 - lr*wd)
+ # rather than folding wd into the gradient. Matches the 8-bit blockwise kernel,
+ # the default/cpu backends, and the Lion paper (Chen et al. 2023).
+ if weight_decay > 0.0:
+ p_vals = p_vals * (1.0 - lr * weight_decay)
+
+ momentum_update = s1_vals * beta1 + (1.0 - beta1) * g_vals
+ update_val = update_scale * lr * tl.where(momentum_update > 0, 1.0, tl.where(momentum_update < 0, -1.0, 0.0))
+ p_vals = p_vals - update_val
+
+ s1_vals = s1_vals * beta2 + (1.0 - beta2) * g_vals
+
+ elif OPTIMIZER_ID == 1: # RMSPROP
+ s1_vals = s1_vals * beta1 + (1.0 - beta1) * g_vals * g_vals
+
+ update_val = update_scale * lr * g_vals / (tl.sqrt(s1_vals) + eps)
+ p_vals = p_vals - update_val
+
+ elif OPTIMIZER_ID == 2: # ADAGRAD
+ s1_vals = s1_vals + g_vals * g_vals
+
+ update_val = lr * g_vals / (tl.sqrt(s1_vals) + eps)
+ p_vals = p_vals - update_val
+
+ tl.store(p_ptr + offsets, p_vals, mask=mask)
+ tl.store(state1_ptr + offsets, s1_vals, mask=mask)
+
+
+name2optimizer_32bit_fn = {
+ "adam": {
+ "preprocess": _optimizer_precondition_2state_32bit,
+ "update": _optimizer_update_2state_32bit_triton_kernel,
+ },
+ "lamb": {
+ "preprocess": _optimizer_precondition_2state_32bit,
+ "update": _optimizer_update_2state_32bit_triton_kernel,
+ },
+ "ademamix": {
+ "preprocess": _optimizer_precondition_2state_32bit,
+ "update": _optimizer_update_2state_32bit_triton_kernel,
+ },
+ "momentum": {
+ "preprocess": _optimizer_precondition_1state_32bit,
+ "update": _optimizer_update_1state_32bit_triton_kernel,
+ },
+ "lars": {
+ "preprocess": _optimizer_precondition_1state_32bit,
+ "update": _optimizer_update_1state_32bit_triton_kernel,
+ },
+ "rmsprop": {
+ "preprocess": _optimizer_precondition_1state_32bit,
+ "update": _optimizer_update_1state_32bit_triton_kernel,
+ },
+ "adagrad": {
+ "preprocess": _optimizer_precondition_1state_32bit,
+ "update": _optimizer_update_1state_32bit_triton_kernel,
+ },
+ "lion": {
+ "preprocess": _optimizer_precondition_1state_32bit,
+ "update": _optimizer_update_1state_32bit_triton_kernel,
+ },
+}
+
+
+def optimizer_update_32bit_impl(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float = 1.0,
+ skip_zeros=False,
+) -> None:
+ """
+ 32-bit optimizer implemented by Triton
+ """
+ if skip_zeros:
+ raise NotImplementedError("skip_zeros is not supported on XPU yet")
+
+ BLOCK_SIZE = 256
+ N_PER_TH = 1 # Number of blocks processed per thread.
+ grid = (triton.cdiv(p.numel(), BLOCK_SIZE * N_PER_TH),)
+ optimizer_id = name2optimizer_id[optimizer_name]
+ fn_preprocess = name2optimizer_32bit_fn[optimizer_name]["preprocess"]
+ fn_update = name2optimizer_32bit_fn[optimizer_name]["update"]
+
+ # In torch=2.7 on XPU there is an issue with libdevice.pow, leading to an error.
+ # For backwards compatibility we precompute the bias correction factors.
+ beta1_step = beta1**step
+ beta2_step = beta2**step
+
+ if optimizer_name == "lion":
+ fn_update[grid](
+ g,
+ p,
+ state1,
+ state2,
+ unorm_vec,
+ max_unorm,
+ param_norm,
+ beta1,
+ beta2,
+ beta3,
+ alpha,
+ eps,
+ weight_decay,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale,
+ skip_zeros,
+ p.numel(),
+ optimizer_id,
+ BLOCK_SIZE,
+ N_PER_TH,
+ num_warps=2,
+ )
+
+ if max_unorm > 0.0:
+ unorm_vec.zero_()
+ fn_preprocess[grid](
+ g,
+ p,
+ state1,
+ state2,
+ unorm_vec,
+ beta1,
+ beta2,
+ eps,
+ weight_decay,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale,
+ p.numel(),
+ optimizer_id,
+ BLOCK_SIZE,
+ N_PER_TH,
+ num_warps=2,
+ )
+
+ else:
+ if max_unorm > 0.0:
+ unorm_vec.zero_()
+ fn_preprocess[grid](
+ g,
+ p,
+ state1,
+ state2,
+ unorm_vec,
+ beta1,
+ beta2,
+ eps,
+ weight_decay,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale,
+ p.numel(),
+ optimizer_id,
+ BLOCK_SIZE,
+ N_PER_TH,
+ num_warps=2,
+ )
+
+ fn_update[grid](
+ g,
+ p,
+ state1,
+ state2,
+ unorm_vec,
+ max_unorm,
+ param_norm,
+ beta1,
+ beta2,
+ beta3,
+ alpha,
+ eps,
+ weight_decay,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ gnorm_scale,
+ skip_zeros,
+ p.numel(),
+ optimizer_id,
+ BLOCK_SIZE,
+ N_PER_TH,
+ num_warps=2,
+ )
+
+
+###########################################
+# Pure torch implementation for reference #
+###########################################
+
+
+@torch.compile
+def _dequantize_blockwise_pytorch(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ """
+ Pure PyTorch reference implementation for block-wise dequantization.
+ """
+ if A.numel() == 0:
+ return torch.empty_like(A, dtype=dtype)
+
+ A_flat = A.flatten()
+ num_elements = A_flat.numel()
+
+ dequantized_flat = code.to(A.device)[A_flat.long()].to(dtype)
+
+ num_blocks = math.ceil(num_elements / blocksize)
+ pad_len = num_blocks * blocksize - num_elements
+ if pad_len > 0:
+ dequantized_flat = torch.nn.functional.pad(dequantized_flat, (0, pad_len))
+
+ dequantized_blocks = dequantized_flat.reshape(num_blocks, blocksize)
+
+ rescaled_blocks = dequantized_blocks * absmax.unsqueeze(1).to(dtype)
+
+ rescaled_flat = rescaled_blocks.flatten()
+ if pad_len > 0:
+ rescaled_flat = rescaled_flat[:-pad_len]
+
+ return rescaled_flat.reshape(A.shape)
+
+
+@torch.compile
+def _quantize_blockwise_pytorch(
+ A: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """
+ Pure PyTorch reference implementation for block-wise quantization.
+ """
+ if A.numel() == 0:
+ return torch.empty_like(A, dtype=torch.uint8), torch.empty(0, dtype=torch.float32, device=A.device)
+
+ A_flat = A.flatten()
+ num_elements = A_flat.numel()
+
+ num_blocks = math.ceil(num_elements / blocksize)
+
+ pad_len = num_blocks * blocksize - num_elements
+ if pad_len > 0:
+ A_flat = torch.nn.functional.pad(A_flat, (0, pad_len))
+
+ A_blocks = A_flat.reshape(num_blocks, blocksize)
+
+ absmax = torch.max(torch.abs(A_blocks), dim=1, keepdim=True)[0]
+ absmax[absmax == 0] = 1.0
+
+ scaled_blocks = A_blocks / absmax
+
+ # Inefficient but straightforward quantization, takes a lot of memory
+ diff = torch.abs(scaled_blocks.unsqueeze(2) - code.to(A.device))
+ quantized_indices = torch.argmin(diff, dim=2).to(torch.uint8)
+
+ quantized_flat = quantized_indices.flatten()
+ if pad_len > 0:
+ quantized_flat = quantized_flat[:-pad_len]
+
+ return quantized_flat.reshape(A.shape), absmax.flatten()
+
+
+# Main updated function
+def optimizer_update_8bit_blockwise_pytorch(
+ p: torch.Tensor,
+ g: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float, # ADEMIX
+ alpha: float, # ADEMIX
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: torch.Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: torch.Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float,
+ gnorm_scale: float,
+ skip_zeros: bool,
+ # ADEMIX
+ *,
+ optimizer_name: str,
+) -> None:
+ """
+ Pure PyTorch implementation of the 8-bit block-wise optimizer update step.
+ This version ensures high-precision updates for float16 parameters.
+ """
+ if skip_zeros:
+ raise ValueError("skip_zeros is not supported on XPU yet.")
+
+ blocksize = 256
+
+ with torch.no_grad():
+ # Dequantize states to perform updates in 32-bit precision
+ if optimizer_name == "ademamix" and absmax1.ndim == 2:
+ # For AdEMAMix, state1 holds two EMAs, so absmax1 is stacked.
+ s1_1_fp32 = _dequantize_blockwise_pytorch(state1[0], absmax1[0], qmap1, blocksize, torch.float32)
+ s1_2_fp32 = _dequantize_blockwise_pytorch(state1[1], absmax1[1], qmap1, blocksize, torch.float32)
+ state1_fp32 = torch.stack([s1_1_fp32, s1_2_fp32])
+ else:
+ state1_fp32 = _dequantize_blockwise_pytorch(state1, absmax1, qmap1, blocksize, torch.float32)
+
+ state2_fp32 = None
+ if state2 is not None:
+ state2_fp32 = _dequantize_blockwise_pytorch(state2, absmax2, qmap2, blocksize, torch.float32)
+
+ grad = g.float() * gnorm_scale
+
+ # Create a 32-bit copy of the parameter for high-precision updates
+ p_fp32 = p.data.float()
+
+ if optimizer_name == "adam":
+ state1_fp32.mul_(beta1).add_(grad, alpha=1.0 - beta1)
+ state2_fp32.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
+
+ bias_correction1 = 1.0 - beta1**step
+ bias_correction2 = 1.0 - beta2**step
+
+ denom = (state2_fp32.sqrt() / math.sqrt(bias_correction2)).add_(eps)
+
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+ p_fp32.addcdiv_(state1_fp32, denom, value=-lr / bias_correction1)
+
+ elif optimizer_name == "ademamix":
+ m1_fp32, m2_fp32 = state1_fp32[0], state1_fp32[1]
+ nu_fp32 = state2_fp32
+
+ m1_fp32.mul_(beta1).add_(grad, alpha=1.0 - beta1)
+ m2_fp32.mul_(beta3).add_(grad, alpha=1.0 - beta3)
+ nu_fp32.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
+
+ bias_correction1 = 1.0 - beta1**step
+ bias_correction2 = math.sqrt(1.0 - beta2**step)
+
+ update = (m1_fp32 / bias_correction1 + alpha * m2_fp32) / (nu_fp32.sqrt() / bias_correction2 + eps)
+
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+
+ p_fp32.add_(update, alpha=-lr)
+ state1_fp32 = torch.stack([m1_fp32, m2_fp32])
+
+ elif optimizer_name == "momentum":
+ grad.add_(p_fp32, alpha=weight_decay)
+ if step == 1:
+ state1_fp32.copy_(grad)
+ else:
+ state1_fp32.mul_(beta1).add_(grad)
+ p_fp32.add_(state1_fp32, alpha=-lr)
+
+ elif optimizer_name == "rmsprop":
+ grad.add_(p_fp32, alpha=weight_decay)
+ state1_fp32.mul_(beta1).addcmul_(grad, grad, value=1.0 - beta1)
+ p_fp32.addcdiv_(grad, state1_fp32.sqrt().add_(eps), value=-lr)
+
+ elif optimizer_name == "lion":
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+
+ update_dir = torch.sign(state1_fp32.mul(beta1) + grad.mul(1.0 - beta1))
+ p_fp32.add_(update_dir, alpha=-lr)
+
+ state1_fp32.mul_(beta2).add_(grad, alpha=1.0 - beta2)
+
+ elif optimizer_name == "adagrad":
+ grad.add_(p_fp32, alpha=weight_decay)
+ state1_fp32.addcmul_(grad, grad, value=1.0)
+ p_fp32.addcdiv_(grad, state1_fp32.sqrt().add_(eps), value=-lr)
+
+ else:
+ raise NotImplementedError(
+ f"Pure PyTorch implementation for optimizer '{optimizer_name}' is not available."
+ )
+
+ # Copy the updated 32-bit parameter back to the original tensor
+ p.data.copy_(p_fp32)
+
+ # Re-quantize states and update state tensors in-place
+ if optimizer_name == "ademamix":
+ new_m1_8bit, new_absmax_m1 = _quantize_blockwise_pytorch(state1_fp32[0], qmap1, blocksize)
+ new_m2_8bit, new_absmax_m2 = _quantize_blockwise_pytorch(state1_fp32[1], qmap1, blocksize)
+ state1[0].copy_(new_m1_8bit)
+ state1[1].copy_(new_m2_8bit)
+ absmax1[0].copy_(new_absmax_m1)
+ absmax1[1].copy_(new_absmax_m2)
+
+ new_state2_8bit, new_absmax2 = _quantize_blockwise_pytorch(state2_fp32, qmap2, blocksize)
+ state2.copy_(new_state2_8bit)
+ absmax2.copy_(new_absmax2)
+ else:
+ new_state1_8bit, new_absmax1 = _quantize_blockwise_pytorch(state1_fp32, qmap1, blocksize)
+ state1.copy_(new_state1_8bit)
+ absmax1.copy_(new_absmax1)
+
+ if state2_fp32 is not None:
+ new_state2_8bit, new_absmax2 = _quantize_blockwise_pytorch(state2_fp32, qmap2, blocksize)
+ state2.copy_(new_state2_8bit)
+ absmax2.copy_(new_absmax2)
+
+
+#######################################
+# Mixed torch + triton implementation #
+#######################################
+
+
+# Much more memory efficient due to using triton for quantization/dequantization
+def optimizer_update_8bit_blockwise_triton_quant(
+ p: torch.Tensor,
+ g: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float, # ADEMIX
+ alpha: float, # ADEMIX
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: torch.Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: torch.Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float,
+ gnorm_scale: float,
+ skip_zeros: bool,
+ # ADEMIX
+ *,
+ optimizer_name: str,
+) -> None:
+ """
+ Pure PyTorch implementation of the 8-bit block-wise optimizer update step.
+ This version ensures high-precision updates for float16 parameters.
+ """
+ if skip_zeros and not torch.any(g):
+ return
+
+ blocksize = 256
+ grad = g.float() * gnorm_scale
+
+ with torch.no_grad():
+ # Create a 32-bit copy of the parameter for high-precision updates
+ p_fp32 = p.data.float()
+
+ # Dequantize states to perform updates in 32-bit precision
+ if optimizer_name == "ademamix" and absmax1.ndim == 2:
+ # For AdEMAMix, state1 holds two EMAs, so absmax1 is stacked.
+ s1_1_fp32 = dequant_8bit_blockwise(state1[0], absmax1[0], qmap1, blocksize, dtype=torch.float32)
+ s1_2_fp32 = dequant_8bit_blockwise(state1[1], absmax1[1], qmap1, blocksize, dtype=torch.float32)
+ state1_fp32 = torch.stack([s1_1_fp32, s1_2_fp32])
+ else:
+ state1_fp32 = dequant_8bit_blockwise(state1, absmax1, qmap1, blocksize, dtype=torch.float32)
+
+ state2_fp32 = None
+ if state2 is not None:
+ state2_fp32 = dequant_8bit_blockwise(state2, absmax2, qmap2, blocksize, dtype=torch.float32)
+
+ # Apply optimizer-specific update logic
+ if optimizer_name == "adam":
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+
+ state1_fp32.mul_(beta1).add_(grad, alpha=1.0 - beta1)
+ state2_fp32.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
+
+ bias_correction1 = 1.0 - beta1**step
+ bias_correction2 = 1.0 - beta2**step
+
+ denom = (state2_fp32.sqrt() / math.sqrt(bias_correction2)).add_(eps)
+ p_fp32.addcdiv_(state1_fp32, denom, value=-lr / bias_correction1)
+
+ elif optimizer_name == "ademamix":
+ m1_fp32, m2_fp32 = state1_fp32[0], state1_fp32[1]
+ nu_fp32 = state2_fp32
+
+ m1_fp32.mul_(beta1).add_(grad, alpha=1.0 - beta1)
+ m2_fp32.mul_(beta3).add_(grad, alpha=1.0 - beta3)
+ nu_fp32.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
+
+ bias_correction1 = 1.0 - beta1**step
+ bias_correction2 = math.sqrt(1.0 - beta2**step)
+
+ update = (m1_fp32 / bias_correction1 + alpha * m2_fp32) / (nu_fp32.sqrt() / bias_correction2 + eps)
+
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+
+ p_fp32.add_(update, alpha=-lr)
+ state1_fp32 = torch.stack([m1_fp32, m2_fp32])
+
+ elif optimizer_name == "momentum":
+ grad.add_(p_fp32, alpha=weight_decay)
+ if step == 1:
+ state1_fp32.copy_(grad)
+ else:
+ state1_fp32.mul_(beta1).add_(grad)
+ p_fp32.add_(state1_fp32, alpha=-lr)
+
+ elif optimizer_name == "rmsprop":
+ grad.add_(p_fp32, alpha=weight_decay)
+ state1_fp32.mul_(beta1).addcmul_(grad, grad, value=1.0 - beta1)
+ p_fp32.addcdiv_(grad, state1_fp32.sqrt().add_(eps), value=-lr)
+
+ elif optimizer_name == "lion":
+ if weight_decay > 0.0:
+ p_fp32.mul_(1.0 - lr * weight_decay)
+
+ update_dir = torch.sign(state1_fp32.mul(beta1) + grad.mul(1.0 - beta1))
+ p_fp32.add_(update_dir, alpha=-lr)
+
+ state1_fp32.mul_(beta2).add_(grad, alpha=1.0 - beta2)
+
+ elif optimizer_name == "adagrad":
+ grad.add_(p_fp32, alpha=weight_decay)
+ state1_fp32.addcmul_(grad, grad, value=1.0)
+ p_fp32.addcdiv_(grad, state1_fp32.sqrt().add_(eps), value=-lr)
+
+ else:
+ raise NotImplementedError(
+ f"Pure PyTorch implementation for optimizer '{optimizer_name}' is not available."
+ )
+
+ # Copy the updated 32-bit parameter back to the original tensor
+ p.data.copy_(p_fp32)
+
+ # Re-quantize states and update state tensors in-place
+ if optimizer_name == "ademamix":
+ new_m1_8bit, new_absmax_m1 = quantize_blockwise_triton(state1_fp32[0], qmap1, blocksize)
+ new_m2_8bit, new_absmax_m2 = quantize_blockwise_triton(state1_fp32[1], qmap1, blocksize)
+ state1[0].copy_(new_m1_8bit)
+ state1[1].copy_(new_m2_8bit)
+ absmax1[0].copy_(new_absmax_m1)
+ absmax1[1].copy_(new_absmax_m2)
+
+ new_state2_8bit, new_absmax2 = quantize_blockwise_triton(state2_fp32, qmap2, blocksize)
+ state2.copy_(new_state2_8bit)
+ absmax2.copy_(new_absmax2)
+ else:
+ new_state1_8bit, new_absmax1 = quantize_blockwise_triton(state1_fp32, qmap1, blocksize)
+ state1.copy_(new_state1_8bit)
+ absmax1.copy_(new_absmax1)
+
+ if state2_fp32 is not None:
+ new_state2_8bit, new_absmax2 = quantize_blockwise_triton(state2_fp32, qmap2, blocksize)
+ state2.copy_(new_state2_8bit)
+ absmax2.copy_(new_absmax2)
+
+
+#########################
+# Triton implementation #
+#########################
+
+
+@triton.jit
+def _optimizer_update_1state_8bit_blockwise_triton_kernel(
+ # Tensors
+ p_ptr,
+ g_ptr,
+ state1_ptr,
+ state2_ptr,
+ beta1: tl.constexpr,
+ beta2: tl.constexpr,
+ beta3,
+ alpha,
+ eps: tl.constexpr,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ qmap1_ptr,
+ qmap2_ptr,
+ absmax1_ptr,
+ absmax2_ptr,
+ weight_decay,
+ gnorm_scale,
+ # Meta-parameters
+ n_elements,
+ BLOCK_SIZE_N: tl.constexpr,
+ N_PER_TH: tl.constexpr,
+ OPTIMIZER_ID: tl.constexpr,
+):
+ """
+ Triton kernel for 8-bit optimizers that use one momentum state.
+ Supports: Momentum, RMSprop, Adagrad, Lion.
+ """
+ # 1. Boilerplate: pid, offsets, mask
+ pid = tl.program_id(axis=0)
+ block_start_idx = pid * N_PER_TH
+ offsets = block_start_idx * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N * N_PER_TH)
+ mask = offsets < n_elements
+
+ # 2. Load and dequantize tensors
+ g = tl.load(g_ptr + offsets, mask=mask, other=0.0).to(tl.float32) * gnorm_scale
+ p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
+ s1 = dequant_8bit_blockwise_kernel_util(state1_ptr, offsets, qmap1_ptr, absmax1_ptr, mask, BLOCK_SIZE_N)
+
+ # 3. Optimizer-specific updates
+ # LION
+ if weight_decay > 0.0 and OPTIMIZER_ID == 2:
+ p *= 1.0 - lr * weight_decay
+ # Apply weight decay for momentum, rmsprop, adagrad
+ elif weight_decay > 0.0:
+ g += p * weight_decay
+
+ # Momentum update
+ if OPTIMIZER_ID == 0: # MOMENTUM
+ if step == 1:
+ s1 = g
+ else:
+ s1 = s1 * beta1 + g
+ p -= lr * s1
+
+ # RMSprop update
+ elif OPTIMIZER_ID == 1: # RMSPROP
+ s1 = s1 * beta1 + (1.0 - beta1) * g * g
+ p -= lr * (g / (tl.sqrt(s1) + eps))
+
+ # Adagrad update
+ elif OPTIMIZER_ID == 2: # ADAGRAD
+ s1 += g * g
+ p -= lr * (g / (tl.sqrt(s1) + eps))
+
+ # Lion update
+ elif OPTIMIZER_ID == 4: # LION
+ val = s1 * beta1 + (1.0 - beta1) * g
+ update = tl.where(val > 0.0, 1.0, tl.where(val < 0.0, -1.0, 0.0))
+ p -= lr * update
+ s1 = s1 * beta2 + (1.0 - beta2) * g
+
+ # 4. Store updated parameter and requantized state
+ tl.store(p_ptr + offsets, p.to(p_ptr.dtype.element_ty), mask=mask)
+ s1_codes, new_absmax1 = quantize_8bit_blockwise_kernel_util(s1, qmap1_ptr, 256, BLOCK_SIZE_N, N_PER_TH)
+ tl.store(state1_ptr + offsets, s1_codes, mask=mask)
+ tl.store(absmax1_ptr + block_start_idx + tl.arange(0, N_PER_TH), new_absmax1)
+
+
+@triton.jit
+def _optimizer_update_2state_8bit_blockwise_triton_kernel(
+ # Tensors
+ p_ptr,
+ g_ptr,
+ state1_ptr,
+ state2_ptr,
+ beta1: tl.constexpr,
+ beta2: tl.constexpr,
+ # ademamix changes alpha and beta3
+ beta3,
+ # ademamix changes alpha and beta3
+ alpha,
+ eps: tl.constexpr,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ qmap1_ptr,
+ qmap2_ptr,
+ absmax1_ptr,
+ absmax2_ptr,
+ weight_decay: tl.constexpr,
+ gnorm_scale: tl.constexpr,
+ # Meta-parameters
+ n_elements,
+ BLOCK_SIZE_N: tl.constexpr,
+ N_PER_TH: tl.constexpr,
+ OPTIMIZER_ID: tl.constexpr,
+):
+ """
+ Triton kernel for 8-bit optimizers that use two momentum states.
+ Supports: Adam, AdEMAMix.
+ """
+ # 1. Boilerplate: pid, offsets, mask
+ pid = tl.program_id(axis=0)
+ block_start_idx = pid * N_PER_TH
+ offsets = block_start_idx * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N * N_PER_TH)
+ mask = offsets < n_elements
+
+ # 2. Load and dequantize tensors
+ g = tl.load(g_ptr + offsets, mask=mask, other=0.0).to(tl.float32) * gnorm_scale
+ p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
+
+ # 3. Optimizer-specific updates
+ if OPTIMIZER_ID == 3: # ADAM
+ s1 = dequant_8bit_blockwise_kernel_util(state1_ptr, offsets, qmap1_ptr, absmax1_ptr, mask, BLOCK_SIZE_N)
+ s2 = dequant_8bit_blockwise_kernel_util(state2_ptr, offsets, qmap2_ptr, absmax2_ptr, mask, BLOCK_SIZE_N)
+
+ s1 = s1 * beta1 + (1.0 - beta1) * g
+ s2 = s2 * beta2 + (1.0 - beta2) * g * g
+
+ # In torch=2.7 on XPU there is an issue with libdevice.pow, leading to an error.
+ # For backwards compatibility we precompute the bias correction factors.
+ # bias_correction1 = 1.0 - libdevice.pow(beta1, step)
+ # bias_correction2 = 1.0 - libdevice.pow(beta2, step)
+ bias_correction1 = 1.0 - beta1_step
+ bias_correction2 = 1.0 - beta2_step
+
+ if weight_decay > 0.0:
+ p *= 1.0 - lr * weight_decay
+
+ denom = tl.sqrt(s2) / tl.sqrt(bias_correction2) + eps
+ p -= (lr / bias_correction1) * (s1 / denom)
+
+ # Store updated parameter
+ tl.store(p_ptr + offsets, p.to(p_ptr.dtype.element_ty), mask=mask)
+
+ # Requantize and store states
+ s1_codes, new_absmax1 = quantize_8bit_blockwise_kernel_util(s1, qmap1_ptr, 256, BLOCK_SIZE_N, N_PER_TH)
+ tl.store(state1_ptr + offsets, s1_codes, mask=mask)
+ tl.store(absmax1_ptr + block_start_idx + tl.arange(0, N_PER_TH), new_absmax1)
+
+ s2_codes, new_absmax2 = quantize_8bit_blockwise_kernel_util(s2, qmap2_ptr, 256, BLOCK_SIZE_N, N_PER_TH)
+ tl.store(state2_ptr + offsets, s2_codes, mask=mask)
+ tl.store(absmax2_ptr + block_start_idx + tl.arange(0, N_PER_TH), new_absmax2)
+
+ elif OPTIMIZER_ID == 5: # ADEMAMIX
+ # AdEMAMix has a stacked state1 (m1, m2) and state2 (nu)
+ m1 = dequant_8bit_blockwise_kernel_util(state1_ptr, offsets, qmap1_ptr, absmax1_ptr, mask, BLOCK_SIZE_N)
+ m2 = dequant_8bit_blockwise_kernel_util(
+ state1_ptr + n_elements,
+ offsets,
+ qmap1_ptr,
+ absmax1_ptr + n_elements // BLOCK_SIZE_N,
+ mask,
+ BLOCK_SIZE_N,
+ )
+ nu = dequant_8bit_blockwise_kernel_util(state2_ptr, offsets, qmap2_ptr, absmax2_ptr, mask, BLOCK_SIZE_N)
+
+ m1 = m1 * beta1 + (1.0 - beta1) * g
+ m2 = m2 * beta3 + (1.0 - beta3) * g
+ nu = nu * beta2 + (1.0 - beta2) * g * g
+
+ # In torch=2.7 on XPU there is an issue with libdevice.pow, leading to an error.
+ # For backwards compatibility we precompute the bias correction factors.
+ # bias_correction1 = 1.0 - libdevice.pow(beta1, step)
+ # bias_correction2 = tl.sqrt(1.0 - libdevice.pow(beta2, step))
+ bias_correction1 = 1.0 - beta1_step
+ bias_correction2 = tl.sqrt(1.0 - beta2_step)
+
+ update = (m1 / bias_correction1 + alpha * m2) / (tl.sqrt(nu) / bias_correction2 + eps)
+
+ if weight_decay > 0.0:
+ p *= 1.0 - lr * weight_decay
+
+ p -= lr * update
+
+ # Store updated parameter
+ tl.store(p_ptr + offsets, p.to(p_ptr.dtype.element_ty), mask=mask)
+
+ # Requantize and store all three states
+ m1_codes, new_absmax_m1 = quantize_8bit_blockwise_kernel_util(m1, qmap1_ptr, 256, BLOCK_SIZE_N, N_PER_TH)
+ tl.store(state1_ptr + offsets, m1_codes, mask=mask)
+ tl.store(absmax1_ptr + block_start_idx + tl.arange(0, N_PER_TH), new_absmax_m1)
+
+ m2_codes, new_absmax_m2 = quantize_8bit_blockwise_kernel_util(m2, qmap1_ptr, 256, BLOCK_SIZE_N, N_PER_TH)
+ tl.store(state1_ptr + n_elements + offsets, m2_codes, mask=mask)
+ tl.store(
+ absmax1_ptr + block_start_idx + tl.arange(0, N_PER_TH) + n_elements // BLOCK_SIZE_N,
+ new_absmax_m2,
+ )
+
+ nu_codes, new_absmax_nu = quantize_8bit_blockwise_kernel_util(nu, qmap2_ptr, 256, BLOCK_SIZE_N, N_PER_TH)
+ tl.store(state2_ptr + offsets, nu_codes, mask=mask)
+ tl.store(absmax2_ptr + block_start_idx + tl.arange(0, N_PER_TH), new_absmax_nu)
+
+
+name2optimizer_fn = {
+ "momentum": _optimizer_update_1state_8bit_blockwise_triton_kernel,
+ "lars": _optimizer_update_1state_8bit_blockwise_triton_kernel,
+ "rmsprop": _optimizer_update_1state_8bit_blockwise_triton_kernel,
+ "adagrad": _optimizer_update_1state_8bit_blockwise_triton_kernel,
+ "adam": _optimizer_update_2state_8bit_blockwise_triton_kernel,
+ "lamb": _optimizer_update_2state_8bit_blockwise_triton_kernel,
+ "lion": _optimizer_update_1state_8bit_blockwise_triton_kernel,
+ "ademamix": _optimizer_update_2state_8bit_blockwise_triton_kernel,
+}
+
+
+def optimizer_update_8bit_blockwise_impl(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: torch.Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: torch.Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float = 0.0,
+ gnorm_scale: float = 1.0,
+ skip_zeros=False,
+) -> None:
+ if skip_zeros:
+ raise NotImplementedError("skip_zeros is not supported on XPU yet")
+
+ if optimizer_name == "ademamix":
+ # Handle AdEMAMIX's stacked state tensors
+ if state1.dim() < 2 or state1.shape[0] != 2:
+ raise ValueError(
+ f"For ademamix, state1 must be a stacked tensor of shape (2, ...), but got {state1.shape}"
+ )
+ if absmax1.dim() < 2 or absmax1.shape[0] != 2:
+ raise ValueError(
+ f"For ademamix, absmax1 must be a stacked tensor of shape (2, ...), but got {absmax1.shape}"
+ )
+
+ BLOCK_SIZE = 256
+ N_PER_TH = 1 # Number of blocks processed per thread.
+ grid = (triton.cdiv(p.numel(), BLOCK_SIZE * N_PER_TH),)
+ fn = name2optimizer_fn[optimizer_name]
+ optimizer_id = name2optimizer_id[optimizer_name]
+
+ # In torch=2.7 on XPU there is an issue with libdevice.pow, leading to an error.
+ # For backwards compatibility we precompute the bias correction factors.
+ beta1_step = beta1**step
+ beta2_step = beta2**step
+
+ fn[grid](
+ p,
+ g,
+ state1,
+ state2,
+ beta1,
+ beta2,
+ beta3,
+ alpha,
+ eps,
+ step,
+ beta1_step,
+ beta2_step,
+ lr,
+ qmap1,
+ qmap2,
+ absmax1,
+ absmax2,
+ weight_decay,
+ gnorm_scale,
+ p.numel(),
+ BLOCK_SIZE_N=BLOCK_SIZE,
+ N_PER_TH=N_PER_TH,
+ OPTIMIZER_ID=optimizer_id,
+ num_warps=2,
+ )
+
+
+# optimizer_update_8bit_blockwise_impl = optimizer_update_8bit_blockwise_pytorch
+# optimizer_update_8bit_blockwise_impl = torch.compile(optimizer_update_8bit_blockwise_pytorch_impl)
+# optimizer_update_8bit_blockwise_impl = optimizer_update_8bit_blockwise_triton_quant
+# optimizer_update_8bit_blockwise_impl = torch.compile(optimizer_update_8bit_blockwise_triton_quant)
+optimizer_update_8bit_blockwise_impl = optimizer_update_8bit_blockwise_impl
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..2dfe7e758f1a118e7a2d80251e373306bffb107f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/triton/ops.py
@@ -0,0 +1,304 @@
+from collections.abc import Sequence
+from typing import Optional
+
+import torch
+
+from . import kernels_4bit, kernels_8bit_quant, kernels_optim
+
+# currently codes unused, kept for reference
+# Should be the same for quant/dequant
+# from bitsandbytes.functional import get_4bit_type
+# _FP4_QUANT_TABLE = get_4bit_type("fp4", device="xpu")
+# _NF4_QUANT_TABLE = get_4bit_type("nf4", device="xpu")
+device_type = torch.accelerator.current_accelerator().type if hasattr(torch, "accelerator") else "cuda"
+torch_accelerator_module = getattr(torch, device_type, torch.cuda)
+
+
+def quantize_blockwise(A: torch.Tensor, code: torch.Tensor, blocksize: int) -> tuple[torch.Tensor, torch.Tensor]:
+ # torch._check(A.dtype == torch.float32, lambda: f"A must be float32 on xpu, got {A.dtype}")
+ with torch_accelerator_module.device(A.device):
+ out, absmax = kernels_8bit_quant.quantize_blockwise_triton(A.contiguous(), code, blocksize)
+ return out, absmax.float()
+
+
+def dequantize_blockwise(
+ A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype
+) -> torch.Tensor:
+ if A.dtype != torch.uint8:
+ raise ValueError(f"A must be uint8, got {A.dtype}")
+ # torch._check(dtype == torch.float32, lambda: f"dtype must be float32 on xpu, got {dtype}")
+ with torch_accelerator_module.device(A.device):
+ out = kernels_8bit_quant.dequant_8bit_blockwise(
+ A.contiguous(),
+ absmax,
+ code,
+ blocksize,
+ dtype=dtype,
+ )
+ return out
+
+
+def dequantize_blockwise_inplace(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ if A.dtype != torch.uint8:
+ raise ValueError(f"A must be uint8, got {A.dtype}")
+ if out.shape != A.shape:
+ raise ValueError(f"Expected out.shape == {A.shape}, got {out.shape}")
+ if out.device != A.device:
+ raise ValueError(f"Expected out.device == {A.device}, got {out.device}")
+ if out.dtype != dtype:
+ raise ValueError(f"Expected out.dtype == {dtype}, got {out.dtype}")
+
+ with torch_accelerator_module.device(A.device):
+ kernels_8bit_quant.dequant_8bit_blockwise(
+ A,
+ absmax,
+ code,
+ blocksize,
+ dtype=dtype,
+ out=out,
+ )
+
+
+def quantize_4bit(
+ A: torch.Tensor, blocksize: int, quant_type: str, quant_storage: torch.dtype
+) -> tuple[torch.Tensor, torch.Tensor]:
+ # torch._check(quant_type == "nf4", lambda: f"quant_type must be nf4 on CPU, got {quant_type}")
+ if A.dtype not in (torch.bfloat16, torch.float16, torch.float32):
+ raise ValueError(f"Blockwise 4bit quantization only supports 16/32-bit floats, but got {A.dtype}")
+
+ n = A.numel()
+
+ # Pad to next multiple of blocksize so the kernel always processes full blocks
+ remainder = n % blocksize
+ if remainder != 0:
+ padding = blocksize - remainder
+ A = torch.nn.functional.pad(A.view(-1), (0, padding), value=0.0)
+ n = A.numel()
+
+ blocks = -(n // -(blocksize * 2))
+
+ absmax = torch.empty((blocks * 2,), device=A.device, dtype=A.dtype)
+ # Use n - n//2 instead of (n+1)//2 to avoid integer overflow for large n
+ out = torch.empty((n - n // 2, 1), device=A.device, dtype=torch.uint8)
+
+ with torch_accelerator_module.device(A.device):
+ kernels_4bit.quantize_4bit_blockwise_triton(
+ A, blocksize, quant_type, blocks, absmax, num_elements=n, quantized_out=out
+ )
+ packed = out
+
+ if quant_storage != torch.uint8:
+ packed = out.squeeze().view(quant_storage).unsqueeze(1)
+
+ return packed, absmax.float()
+
+
+def dequantize_4bit(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+) -> torch.Tensor:
+ # torch._check(quant_type == "nf4", lambda: f"quant_type must be nf4 on XPU, got {quant_type}")
+ if dtype not in (torch.bfloat16, torch.float16, torch.float32):
+ raise ValueError(f"Blockwise 4bit dequantization only supports 16/32-bit floats, but got {dtype}")
+ # torch._check(
+ # A.dtype == torch.uint8,
+ # lambda: f"Blockwise 4bit dequantization on XPU only supports uint8 storage, got {A.dtype}",
+ # )
+ # Check if this is fine and fast
+ if A.dtype != torch.uint8:
+ A = A.squeeze().view(torch.uint8).unsqueeze(1)
+
+ out = torch.empty(shape, dtype=dtype, device=A.device)
+ with torch_accelerator_module.device(A.device):
+ kernels_4bit.dequantize_4bit_impl(A, absmax, blocksize, quant_type, dtype, out=out)
+
+ return out
+
+
+def dequantize_4bit_inplace(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ if out.shape != tuple(shape):
+ raise ValueError(f"Expected out.shape == {shape}, got {out.shape}")
+ if out.dtype != dtype:
+ raise ValueError(f"Expected out.dtype == {dtype}, got {out.dtype}")
+ with torch_accelerator_module.device(A.device):
+ kernels_4bit.dequantize_4bit_impl(A, absmax, blocksize, quant_type, dtype, out=out)
+
+
+def gemv_4bit(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+) -> torch.Tensor:
+ if B.dtype != torch.uint8:
+ B = B.squeeze().view(torch.uint8).unsqueeze(1)
+
+ B_dq_triton = torch.empty(shapeB, dtype=A.dtype, device=A.device)
+
+ with torch_accelerator_module.device(A.device):
+ kernels_4bit.dequantize_4bit_impl_passing_code(
+ B,
+ absmax,
+ blocksize,
+ code,
+ dtype=A.dtype,
+ out=B_dq_triton,
+ )
+
+ return torch.nn.functional.linear(
+ A,
+ B_dq_triton,
+ bias=None,
+ )
+
+
+# optimizer_update_8bit_blockwise_impl = kernels_optim.optimizer_update_8bit_blockwise_pytorch
+# optimizer_update_8bit_blockwise_impl = torch.compile(kernels_optim.optimizer_update_8bit_blockwise_pytorch) # 60ms
+# optimizer_update_8bit_blockwise_impl = kernels_optim.optimizer_update_8bit_blockwise_triton_quant #2.8ms
+# optimizer_update_8bit_blockwise_impl = torch.compile(kernels_optim.optimizer_update_8bit_blockwise_triton_quant) # 2.3ms
+optimizer_update_8bit_blockwise_impl = kernels_optim.optimizer_update_8bit_blockwise_impl # ~0.95ms for adam
+
+
+def optimizer_update_8bit_blockwise(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: torch.Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: torch.Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float = 0.0,
+ gnorm_scale: float = 1.0,
+ skip_zeros=False,
+) -> None:
+ # torch._check(
+ # g.numel() == p.numel(),
+ # lambda: f"g and p must have the same number of elements, got {g.numel()} and {p.numel()}",
+ # )
+ # compute_dtypes = [torch.float16, torch.bfloat16, torch.float32]
+
+ # torch._check(
+ # g.dtype in compute_dtypes,
+ # lambda: f"g must be bfloat16, float16, or float32, got {g.dtype}",
+ # )
+ # torch._check(
+ # g.dtype == p.dtype,
+ # lambda: f"Expected all tensors to have the same dtype, got g.dtype={g.dtype}, p.dtype={p.dtype}",
+ # )
+ # torch._check(
+ # state1.dtype == torch.uint8,
+ # lambda: f"state1 must be uint8, got {state1.dtype}",
+ # )
+ # torch._check(
+ # qmap1.dtype == absmax1.dtype == torch.float32,
+ # lambda: f"Expected qmap1 and absmax1 to be float32, got qmap1.dtype={qmap1.dtype}, absmax1.dtype={absmax1.dtype}",
+ # )
+ # if state2 is not None:
+ # torch._check(
+ # state2.dtype == torch.uint8,
+ # lambda: f"state2 must be uint8, got {state2.dtype}",
+ # )
+ # torch._check(
+ # qmap2.dtype == absmax2.dtype == torch.float32,
+ # lambda: f"Expected qmap2 and absmax2 to be float32, got qmap2.dtype={qmap2.dtype}, absmax2.dtype={absmax2.dtype}",
+ # )
+
+ # Use g.device for device context: paged state tensors appear as CPU tensors
+ # but are backed by USM shared memory and accessible from the accelerator.
+ with torch_accelerator_module.device(g.device):
+ optimizer_update_8bit_blockwise_impl(
+ optimizer_name=optimizer_name,
+ g=g,
+ p=p,
+ state1=state1,
+ state2=state2,
+ beta1=beta1,
+ beta2=beta2,
+ beta3=beta3,
+ alpha=alpha,
+ eps=eps,
+ step=step,
+ lr=lr,
+ qmap1=qmap1,
+ qmap2=qmap2,
+ absmax1=absmax1,
+ absmax2=absmax2,
+ weight_decay=weight_decay,
+ gnorm_scale=gnorm_scale,
+ skip_zeros=skip_zeros,
+ )
+
+
+def optimizer_update_32bit(
+ optimizer_name: str,
+ g: torch.Tensor,
+ p: torch.Tensor,
+ state1: torch.Tensor,
+ state2: Optional[torch.Tensor],
+ unorm_vec: Optional[torch.Tensor],
+ max_unorm: float,
+ param_norm: float,
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ weight_decay: float,
+ step: int,
+ lr: float,
+ gnorm_scale: float,
+ skip_zeros=False,
+) -> None:
+ # Use g.device for device context: paged state tensors appear as CPU tensors
+ # but are backed by USM shared memory and accessible from the accelerator.
+ with torch_accelerator_module.device(g.device):
+ kernels_optim.optimizer_update_32bit_impl(
+ optimizer_name=optimizer_name,
+ g=g,
+ p=p,
+ state1=state1,
+ state2=state2,
+ unorm_vec=unorm_vec,
+ max_unorm=max_unorm,
+ param_norm=param_norm,
+ beta1=beta1,
+ beta2=beta2,
+ beta3=beta3,
+ alpha=alpha,
+ eps=eps,
+ weight_decay=weight_decay,
+ step=step,
+ lr=lr,
+ gnorm_scale=gnorm_scale,
+ skip_zeros=skip_zeros,
+ )
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/utils.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..a63e59a99609e1ee42a982590cc1e9e67023d33c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/utils.py
@@ -0,0 +1,94 @@
+from importlib.metadata import metadata
+
+from packaging import version
+import torch
+
+try:
+ import triton # noqa: F401
+ import triton.language as tl # noqa: F401
+
+ triton_available = True
+except ImportError:
+ triton_available = False
+
+
+_NF4_QUANT_TABLE = torch.tensor(
+ [
+ -1.0,
+ -0.6961928009986877,
+ -0.5250730514526367,
+ -0.39491748809814453,
+ -0.28444138169288635,
+ -0.18477343022823334,
+ -0.09105003625154495,
+ 0.0,
+ 0.07958029955625534,
+ 0.16093020141124725,
+ 0.24611230194568634,
+ 0.33791524171829224,
+ 0.44070982933044434,
+ 0.5626170039176941,
+ 0.7229568362236023,
+ 1.0,
+ ],
+ dtype=torch.float32,
+ device="xpu"
+ if hasattr(torch, "xpu") and torch.xpu.is_available()
+ else "cpu", # Only cpu/xpu use this table for now.
+)
+_FP4_QUANT_TABLE = torch.tensor(
+ [
+ 0.0000,
+ 0.0052,
+ 0.6667,
+ 1.0000,
+ 0.3333,
+ 0.5000,
+ 0.1667,
+ 0.2500,
+ 0.0000,
+ -0.0052,
+ -0.6667,
+ -1.0000,
+ -0.3333,
+ -0.5000,
+ -0.1667,
+ -0.2500,
+ ],
+ dtype=torch.float32,
+ device="xpu"
+ if hasattr(torch, "xpu") and torch.xpu.is_available()
+ else "cpu", # Only cpu/xpu use this table for now.
+)
+CODE = {"nf4": _NF4_QUANT_TABLE, "fp4": _FP4_QUANT_TABLE}
+
+# Cache 4-bit dequantization code tensors per (quant_type, device).
+_code_4bit_cache: dict[tuple[str, torch.device], torch.Tensor] = {}
+
+
+def _get_4bit_code(quant_type: str, device: torch.device) -> torch.Tensor:
+ key = (quant_type, device)
+ if key not in _code_4bit_cache:
+ from bitsandbytes.functional import get_4bit_type
+
+ _code_4bit_cache[key] = get_4bit_type(quant_type, device=device)
+ return _code_4bit_cache[key]
+
+
+def get_gaudi_sw_version():
+ """
+ Returns the installed version of Gaudi SW.
+ """
+ try:
+ # if we find the spec, examine the installed version
+ plugin_metadata = metadata("habana-torch-plugin")
+ plugin_version = plugin_metadata.get("Version")
+ if plugin_version:
+ gaudi_version = version.parse(plugin_version)
+ except Exception:
+ gaudi_version = None
+
+ return gaudi_version
+
+
+GAUDI_SW_VER = get_gaudi_sw_version()
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/xpu/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/xpu/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/backends/xpu/ops.py b/venv/lib/python3.11/site-packages/bitsandbytes/backends/xpu/ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..731200c534ab1950072daed3d2dd70026a2a64b1
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/backends/xpu/ops.py
@@ -0,0 +1,305 @@
+from collections.abc import Sequence
+import ctypes as ct
+import logging
+from typing import Optional
+from warnings import warn
+
+from packaging import version
+import torch
+
+from bitsandbytes.functional import _get_tensor_stream, get_ptr
+
+from ..._ops import register_kernel
+from ...cextension import ErrorHandlerMockBNBNativeLibrary, lib
+from ..default.ops import _gemm_4bit_default_impl
+from ..utils import _get_4bit_code, triton_available
+
+logger = logging.getLogger(__name__)
+
+# _int_mm is available in torch starting from 2.9 version
+if version.parse(torch.__version__).release >= version.parse("2.9").release:
+
+ @register_kernel("bitsandbytes::int8_linear_matmul", "xpu")
+ def _(A: torch.Tensor, B: torch.Tensor):
+ return torch._int_mm(
+ A.reshape(-1, A.shape[-1]),
+ B.t(),
+ ).reshape(*A.shape[:-1], B.shape[0])
+
+
+def _dequantize_4bit_impl(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ dtype: torch.dtype,
+ out: torch.Tensor,
+) -> None:
+ # XPU SYCL kernels only support contiguous tensors.
+ A = A.contiguous()
+ args = (
+ None,
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_int(blocksize),
+ ct.c_int(out.numel()),
+ _get_tensor_stream(A),
+ )
+ if dtype == torch.bfloat16:
+ if quant_type == "fp4":
+ lib.cdequantize_blockwise_bf16_fp4(*args)
+ else:
+ lib.cdequantize_blockwise_bf16_nf4(*args)
+ elif dtype == torch.float16:
+ if quant_type == "fp4":
+ lib.cdequantize_blockwise_fp16_fp4(*args)
+ else:
+ lib.cdequantize_blockwise_fp16_nf4(*args)
+ elif dtype == torch.float32:
+ if quant_type == "fp4":
+ lib.cdequantize_blockwise_fp32_fp4(*args)
+ else:
+ lib.cdequantize_blockwise_fp32_nf4(*args)
+
+
+def _dequantize_blockwise_impl(
+ A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype, out: torch.Tensor
+) -> None:
+ # XPU SYCL kernels only support contiguous tensors.
+ A = A.contiguous()
+ args = (
+ get_ptr(code),
+ get_ptr(A),
+ get_ptr(absmax),
+ get_ptr(out),
+ ct.c_int(blocksize),
+ ct.c_int(A.numel()),
+ _get_tensor_stream(A),
+ )
+ if dtype == torch.float16:
+ lib.cdequantize_blockwise_fp16(*args)
+ elif dtype == torch.bfloat16:
+ lib.cdequantize_blockwise_bf16(*args)
+ elif dtype == torch.float32:
+ lib.cdequantize_blockwise_fp32(*args)
+
+
+def _gemv_4bit_impl(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ out: torch.Tensor,
+) -> None:
+ m = ct.c_int32(1)
+ n = ct.c_int32(shapeB[0])
+ k = ct.c_int32(shapeB[1])
+
+ lda = m
+ ldb = ct.c_int32((A.shape[-1] + 1) // 2)
+ ldc = m
+
+ stream = _get_tensor_stream(A)
+ if A.dtype == torch.float16:
+ lib.cgemv_4bit_inference_fp16(
+ m,
+ n,
+ k,
+ get_ptr(A),
+ get_ptr(B),
+ get_ptr(absmax),
+ get_ptr(code),
+ get_ptr(out),
+ lda,
+ ldb,
+ ldc,
+ ct.c_int32(blocksize),
+ stream,
+ )
+ elif A.dtype == torch.bfloat16:
+ lib.cgemv_4bit_inference_bf16(
+ m,
+ n,
+ k,
+ get_ptr(A),
+ get_ptr(B),
+ get_ptr(absmax),
+ get_ptr(code),
+ get_ptr(out),
+ lda,
+ ldb,
+ ldc,
+ ct.c_int32(blocksize),
+ stream,
+ )
+ elif A.dtype == torch.float32:
+ lib.cgemv_4bit_inference_fp32(
+ m,
+ n,
+ k,
+ get_ptr(A),
+ get_ptr(B),
+ get_ptr(absmax),
+ get_ptr(code),
+ get_ptr(out),
+ lda,
+ ldb,
+ ldc,
+ ct.c_int32(blocksize),
+ stream,
+ )
+
+
+@register_kernel("bitsandbytes::gemm_4bit", "xpu")
+def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ bias: Optional[torch.Tensor] = None,
+ absmax_8bit: Optional[torch.Tensor] = None,
+ absmax_code: Optional[torch.Tensor] = None,
+ absmax_offset: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ K = A.shape[-1]
+ M = A.numel() // K
+
+ if M == 1:
+ if K % blocksize == 0:
+ if absmax_8bit is not None:
+ absmax = (
+ torch.ops.bitsandbytes.dequantize_blockwise.default(
+ absmax_8bit, absmax, absmax_code, 256, torch.float32
+ )
+ + absmax_offset
+ )
+
+ code = _get_4bit_code(quant_type, A.device)
+ out = torch.ops.bitsandbytes.gemv_4bit.default(A, B, shapeB, absmax, code, blocksize)
+
+ if bias is not None:
+ out = out + bias
+ return out
+
+ warn(
+ f"inner dimension ({K}) is not aligned for fast kernel "
+ f"with blocksize={blocksize}, falling back to slower implementation.",
+ UserWarning,
+ )
+
+ return _gemm_4bit_default_impl(
+ A,
+ B,
+ shapeB,
+ absmax,
+ blocksize,
+ quant_type,
+ bias,
+ absmax_8bit=absmax_8bit,
+ absmax_code=absmax_code,
+ absmax_offset=absmax_offset,
+ )
+
+
+# SYCL should be faster for xpu, so at first checking if it is available.
+if not isinstance(lib, ErrorHandlerMockBNBNativeLibrary):
+ logger.info("Register sycl bitsandbytes kernels for XPU")
+
+ # TODO: Remove the triton register when quantization sycl kernel is ready.
+ if triton_available:
+ from ..triton import ops as triton_ops
+
+ register_kernel("bitsandbytes::quantize_blockwise", "xpu")(triton_ops.quantize_blockwise)
+ register_kernel("bitsandbytes::quantize_4bit", "xpu")(triton_ops.quantize_4bit)
+ register_kernel("bitsandbytes::optimizer_update_8bit_blockwise", "xpu")(
+ triton_ops.optimizer_update_8bit_blockwise
+ )
+ register_kernel("bitsandbytes::optimizer_update_32bit", "xpu")(triton_ops.optimizer_update_32bit)
+
+ @register_kernel("bitsandbytes::dequantize_4bit", "xpu")
+ def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ blocksize: int,
+ quant_type: str,
+ shape: Sequence[int],
+ dtype: torch.dtype,
+ ) -> torch.Tensor:
+ out = torch.empty(shape, dtype=dtype, device=A.device)
+ _dequantize_4bit_impl(A, absmax, blocksize, quant_type, dtype, out=out)
+ return out
+
+ @register_kernel("bitsandbytes::dequantize_blockwise", "xpu")
+ def _(
+ A: torch.Tensor, absmax: torch.Tensor, code: torch.Tensor, blocksize: int, dtype: torch.dtype
+ ) -> torch.Tensor:
+ out = torch.empty_like(A, dtype=dtype)
+ _dequantize_blockwise_impl(A, absmax, code, blocksize, dtype, out=out)
+ return out
+
+ @register_kernel("bitsandbytes::dequantize_blockwise.out", "xpu")
+ def _(
+ A: torch.Tensor,
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ dtype: torch.dtype,
+ out: torch.Tensor,
+ ) -> None:
+ if out.dtype != dtype:
+ raise ValueError(f"Expected out.dtype == {dtype}, got {out.dtype}")
+ if out.shape != A.shape:
+ raise ValueError(f"Expected out.shape == {A.shape}, got {out.shape}")
+ _dequantize_blockwise_impl(A, absmax, code, blocksize, dtype, out=out)
+
+ @register_kernel("bitsandbytes::gemv_4bit", "xpu")
+ def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ ) -> torch.Tensor:
+ shape = (*A.shape[:-1], shapeB[0])
+ out = torch.empty(shape, device=A.device, dtype=A.dtype)
+ _gemv_4bit_impl(A, B, shapeB, absmax, code, blocksize, out=out)
+ return out
+
+ @register_kernel("bitsandbytes::gemv_4bit.out", "xpu")
+ def _(
+ A: torch.Tensor,
+ B: torch.Tensor,
+ shapeB: Sequence[int],
+ absmax: torch.Tensor,
+ code: torch.Tensor,
+ blocksize: int,
+ out: torch.Tensor,
+ ) -> None:
+ expected_shape = (*A.shape[:-1], shapeB[0])
+ if out.shape != expected_shape:
+ raise ValueError(f"Expected out.shape == {expected_shape}, got {out.shape}")
+ if out.dtype != A.dtype:
+ raise ValueError(f"Expected out.dtype == {A.dtype}, got {out.dtype}")
+ _gemv_4bit_impl(A, B, shapeB, absmax, code, blocksize, out=out)
+
+elif triton_available:
+ logger.info("Register triton bitsandbytes kernels for XPU")
+ from ..triton import ops as triton_ops
+
+ register_kernel("bitsandbytes::quantize_blockwise", "xpu")(triton_ops.quantize_blockwise)
+ register_kernel("bitsandbytes::dequantize_blockwise.out", "xpu")(triton_ops.dequantize_blockwise_inplace)
+ register_kernel("bitsandbytes::dequantize_blockwise", "xpu")(triton_ops.dequantize_blockwise)
+ register_kernel("bitsandbytes::quantize_4bit", "xpu")(triton_ops.quantize_4bit)
+ register_kernel("bitsandbytes::dequantize_4bit.out", "xpu")(triton_ops.dequantize_4bit_inplace)
+ register_kernel("bitsandbytes::dequantize_4bit", "xpu")(triton_ops.dequantize_4bit)
+ register_kernel("bitsandbytes::gemv_4bit", "xpu")(triton_ops.gemv_4bit)
+ register_kernel("bitsandbytes::optimizer_update_8bit_blockwise", "xpu")(triton_ops.optimizer_update_8bit_blockwise)
+ register_kernel("bitsandbytes::optimizer_update_32bit", "xpu")(triton_ops.optimizer_update_32bit)
+else:
+ logger.warning("Register pytorch bitsandbytes kernels for XPU because no native library or triton packages found.")
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/cextension.py b/venv/lib/python3.11/site-packages/bitsandbytes/cextension.py
new file mode 100644
index 0000000000000000000000000000000000000000..e234f20d3346c5effcdc15ae426b9e54c16a6412
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/cextension.py
@@ -0,0 +1,405 @@
+import ctypes as ct
+import functools
+import logging
+import os
+from pathlib import Path
+import re
+from typing import Optional
+
+import torch
+
+from bitsandbytes.consts import DYNAMIC_LIBRARY_SUFFIX, PACKAGE_DIR
+from bitsandbytes.cuda_specs import (
+ CUDASpecs,
+ get_cuda_specs,
+ get_cuda_version_tuple,
+ get_rocm_gpu_arch,
+)
+
+logger = logging.getLogger(__name__)
+
+
+def get_cuda_bnb_library_path(cuda_specs: CUDASpecs) -> Path:
+ """
+ Get the path to the best matching CUDA/ROCm BNB native library for the given specs.
+
+ When no override is set, selects from packaged libraries using the following priority:
+ 1. Exact version match.
+ 2. Highest packaged version <= runtime version, same major (e.g. runtime 12.9, ship 12.8).
+ 3. Lowest packaged version > runtime version, same major (e.g. runtime 12.0, ship 12.1).
+ No cross-major fallback: if no same-major library exists, returns the exact non-existent
+ path so the caller raises a clear "not found" error.
+ A warning is logged when falling back. Override env vars bypass selection entirely
+ and load the named version with no fallback. The returned path is not guaranteed to
+ exist when no packaged libs are found, or when an override names an absent version.
+ """
+ is_hip = bool(torch.version.hip)
+ prefix = "rocm" if is_hip else "cuda"
+ override_var = "BNB_ROCM_VERSION" if is_hip else "BNB_CUDA_VERSION"
+
+ override_value = os.environ.get(override_var)
+
+ if override_value is not None:
+ if not override_value.isdigit():
+ raise RuntimeError(f"{override_var}={override_value!r}: value must be digits only (e.g. '124' for 12.4).")
+ library_name = f"libbitsandbytes_{prefix}{override_value}{DYNAMIC_LIBRARY_SUFFIX}"
+ logger.warning(
+ f"WARNING: {override_var}={override_value} environment variable detected; loading {library_name}.\n"
+ f"This overrides automatic {'ROCm' if is_hip else 'CUDA'} version selection.\n"
+ f"If this was unintended clear the variable and retry: unset {override_var}\n",
+ )
+ return PACKAGE_DIR / library_name
+
+ available = _find_cuda_libs(prefix, is_hip)
+ runtime_version = cuda_specs.cuda_version_tuple
+
+ if not available:
+ return PACKAGE_DIR / f"libbitsandbytes_{prefix}{cuda_specs.cuda_version_string}{DYNAMIC_LIBRARY_SUFFIX}"
+
+ if runtime_version in available:
+ return available[runtime_version]
+
+ lower = [v for v in available if v[0] == runtime_version[0] and v < runtime_version]
+ if lower:
+ selected = max(lower)
+ else:
+ higher_same = [v for v in available if v[0] == runtime_version[0] and v > runtime_version]
+ if higher_same:
+ selected = min(higher_same)
+ else:
+ # No same-major library available. Return the non-existent exact path so
+ # get_native_library() raises a clear "not found" error.
+ return PACKAGE_DIR / f"libbitsandbytes_{prefix}{cuda_specs.cuda_version_string}{DYNAMIC_LIBRARY_SUFFIX}"
+
+ logger.warning(
+ f"No prebuilt binary for {'ROCm' if is_hip else 'CUDA'} "
+ f"{runtime_version[0]}.{runtime_version[1]}, loading "
+ f"{'ROCm' if is_hip else 'CUDA'} {selected[0]}.{selected[1]} instead. "
+ f"Set {override_var} to override."
+ )
+ return available[selected]
+
+
+class BNBNativeLibrary:
+ _lib: ct.CDLL
+ compiled_with_cuda = False
+
+ def __init__(self, lib: ct.CDLL):
+ self._lib = lib
+
+ @functools.cache # noqa: B019
+ def __getattr__(self, name):
+ fn = getattr(self._lib, name, None)
+
+ if fn is not None:
+ return fn
+
+ def throw_on_call(*args, **kwargs):
+ raise RuntimeError(
+ f"Method '{name}' not available in CPU-only version of bitsandbytes.\n"
+ "Reinstall with GPU support or use CUDA-enabled hardware."
+ )
+
+ return throw_on_call
+
+ def __getitem__(self, item):
+ return self.__getattr__(item)
+
+
+class CudaBNBNativeLibrary(BNBNativeLibrary):
+ compiled_with_cuda = True
+
+ def __init__(self, lib: ct.CDLL):
+ super().__init__(lib)
+ lib.get_context.restype = ct.c_void_p
+ lib.cget_managed_ptr.restype = ct.c_void_p
+
+
+class XpuBNBNativeLibrary(BNBNativeLibrary):
+ """XPU native library with SYCL USM paged memory support."""
+
+ def __init__(self, lib: ct.CDLL):
+ super().__init__(lib)
+ if hasattr(lib, "cget_managed_ptr"):
+ lib.cget_managed_ptr.restype = ct.c_void_p
+
+
+def _split_cuda_version(compact: str, is_hip: bool) -> tuple[int, int]:
+ """Split a compact CUDA/ROCm version string from a library filename into (major, minor).
+
+ CUDA: major is always 2 digits (11, 12, 13...), e.g. '118' -> (11, 8), '132' -> (13, 2).
+ ROCm: major is always 1 digit for now (6, 7...), e.g. '72' -> (7, 2), '713' -> (7, 13).
+ Note: revisit if ROCm major reaches 10.
+ """
+ if is_hip:
+ return int(compact[:1]), int(compact[1:])
+ return int(compact[:2]), int(compact[2:])
+
+
+def _find_cuda_libs(prefix: str, is_hip: bool) -> dict[tuple[int, int], Path]:
+ """Return a {(major, minor): Path} mapping for all packaged CUDA/ROCm library files."""
+ result = {}
+ for lib in PACKAGE_DIR.glob(f"libbitsandbytes_{prefix}*{DYNAMIC_LIBRARY_SUFFIX}"):
+ match = re.search(rf"{prefix}(\d+)", lib.name)
+ if match:
+ try:
+ result[_split_cuda_version(match.group(1), is_hip)] = lib
+ except (ValueError, IndexError):
+ continue
+ return result
+
+
+def get_available_cuda_binary_versions() -> list[str]:
+ """Get formatted CUDA/ROCm versions from existing library files."""
+ is_hip = bool(torch.version.hip)
+ prefix = "rocm" if is_hip else "cuda"
+ return sorted(f"{major}.{minor}" for major, minor in _find_cuda_libs(prefix, is_hip))
+
+
+def parse_cuda_version(version_str: str) -> str:
+ """Convert a raw version code string (e.g. '118', '713') to a dotted version (e.g. '11.8', '7.13')."""
+ if version_str.isdigit():
+ is_hip = bool(torch.version.hip)
+ try:
+ major, minor = _split_cuda_version(version_str, is_hip)
+ return f"{major}.{minor}"
+ except (ValueError, IndexError):
+ pass
+ return version_str
+
+
+class ErrorHandlerMockBNBNativeLibrary(BNBNativeLibrary):
+ """
+ Mock library handler that defers errors until native methods are called.
+
+ This class serves as a fallback when the native bitsandbytes library fails to load.
+ It captures the original error and generates detailed troubleshooting guidance.
+
+ Key behaviors:
+ - Allows attribute access and method assignment without immediate errors
+ - Throws a RuntimeError with diagnostic information only when a native method is called, as otherwise it would error out on import, breaking backward compatibility
+ - Handles both missing CUDA dependencies and version mismatch scenarios
+
+ Error scenarios covered:
+ 1. Missing shared library dependencies (e.g., libcudart.so not in LD_LIBRARY_PATH or through PyTorch CUDA installation)
+ 2. CUDA version mismatch between PyTorch and available pre-compiled binaries
+ 3. Completely missing pre-compiled binaries when CUDA is detected
+ 4. Custom BNB_CUDA_VERSION or BNB_ROCM_VERSION override but mismatch
+ 5. CPU-only installation attempts when GPU functionality is requested
+
+ """
+
+ def __init__(self, error_msg: str):
+ self.error_msg = error_msg
+ self.available_versions = get_available_cuda_binary_versions()
+ override_value = os.environ.get("BNB_ROCM_VERSION") if HIP_ENVIRONMENT else os.environ.get("BNB_CUDA_VERSION")
+ user_version = get_cuda_version_tuple()
+ user_version_str = f"{user_version[0]}.{user_version[1]}" if user_version else "unknown"
+ self.requested_version = parse_cuda_version(override_value) if override_value else user_version_str
+
+ # Pre-generate the error message based on error type
+ if "cannot open shared object file" in error_msg:
+ self.formatted_error = self._format_dependency_error()
+ else: # lib loading errors
+ self.formatted_error = self._format_lib_error_message(
+ available_versions=self.available_versions,
+ user_cuda_version=user_version_str,
+ original_error=f"Original error: {self.error_msg}\n" if self.error_msg else "",
+ requested_version=self.requested_version,
+ )
+
+ def _format_lib_error_message(
+ self,
+ available_versions: list[str],
+ user_cuda_version: str,
+ original_error: str = "",
+ requested_version: Optional[str] = None,
+ ) -> str:
+ """Format detailed error message for library loading failures"""
+ analysis = ""
+ no_cpu_lib_found = "libbitsandbytes_cpu.so: cannot open" in original_error
+ no_cuda_lib_found = f"{BNB_BACKEND} binary not found" in original_error
+
+ if no_cpu_lib_found:
+ analysis = "\n🚨 Failed to load CPU-only bitsandbytes library 🚨\n\n"
+
+ elif no_cuda_lib_found:
+ version_list_str = "\n - " + "\n - ".join(available_versions) if available_versions else "NONE"
+ analysis = (
+ (
+ f"\n🚨 {BNB_BACKEND} VERSION MISMATCH 🚨\n"
+ f"Requested {BNB_BACKEND} version: {requested_version}\n"
+ f"Detected PyTorch {BNB_BACKEND} version: {user_cuda_version}\n"
+ f"Available pre-compiled versions: {version_list_str}\n\n"
+ "This means:\n"
+ "The version you're trying to use is NOT distributed with this package\n\n"
+ )
+ if available_versions
+ else "\n🚨 Forgot to compile the bitsandbytes library? 🚨\n"
+ "1. You're not using the package but checked-out the source code\n"
+ "2. You MUST compile from source\n\n"
+ )
+
+ base_msg = "Attempted to use bitsandbytes native library functionality but it's not available.\n\n"
+
+ troubleshooting = (
+ (
+ f"This typically happens when:\n"
+ f"1. bitsandbytes doesn't ship with a pre-compiled binary for your {BNB_BACKEND} version\n"
+ f"2. The library wasn't compiled properly during installation from source\n\n"
+ )
+ if no_cuda_lib_found
+ else f"This typically happens when you checked the code out from source and your torch installation doesn't detect {BNB_BACKEND} on your machine.\n\n"
+ )
+
+ note = (
+ (
+ f"bitsandbytes tried to find a compatible {BNB_BACKEND} binary but none could be loaded.\n"
+ f"If your {BNB_BACKEND} version isn't among the available pre-compiled versions above, you must compile from source.\n\n"
+ )
+ if no_cuda_lib_found
+ else ""
+ )
+
+ compile_instructions = (
+ ("COMPILE FROM SOURCE for CPU-only:\n `cmake -DCOMPUTE_BACKEND=cpu -S . && make`\n\n")
+ if not no_cuda_lib_found
+ else (
+ "You have two options:\n"
+ "1. COMPILE FROM SOURCE (required if no binary exists):\n"
+ " https://huggingface.co/docs/bitsandbytes/main/en/installation#cuda-compile\n"
+ "2. Use BNB_CUDA_VERSION to specify a DIFFERENT CUDA version from the detected one, which is installed on your machine and matching an available pre-compiled version listed above\n\n"
+ )
+ if not HIP_ENVIRONMENT
+ else (
+ "You have two options:\n"
+ "1. COMPILE FROM SOURCE as mentioned here:\n"
+ " https://huggingface.co/docs/bitsandbytes/main/en/installation?backend=AMD+ROCm#amd-gpu\n"
+ "2. Use BNB_ROCM_VERSION to specify a DIFFERENT ROCm version from the detected one, matching the version the library was built with.\n\n"
+ )
+ )
+
+ diagnostics = (
+ f"🔍 Run this command for detailed diagnostics:\n"
+ f"python -m bitsandbytes\n\n"
+ f"If you've tried everything and still have issues:\n"
+ f"1. Include ALL version info (operating system, bitsandbytes, pytorch, {BNB_BACKEND.lower()}, python)\n"
+ f"2. Describe what you've tried in detail\n"
+ f"3. Open an issue with this information:\n"
+ f" https://github.com/bitsandbytes-foundation/bitsandbytes/issues\n\n"
+ )
+
+ return f"{analysis}{base_msg}{troubleshooting}{note}{compile_instructions}{original_error}\n{diagnostics}"
+
+ def _format_dependency_error(self) -> str:
+ """Format error message for missing shared libraries"""
+ # Extract missing library name from error
+ error_parts = self.error_msg.split(":")
+ missing_lib = error_parts[0].strip() if len(error_parts) > 0 else "unknown library"
+ cuda_major_version = (
+ self.requested_version.split(".")[0] if "." in self.requested_version else self.requested_version
+ )
+
+ return (
+ f"\n🚨 {BNB_BACKEND} SETUP ERROR: Missing dependency: {missing_lib} 🚨\n\n"
+ f"{BNB_BACKEND} {cuda_major_version}.x runtime libraries were not found in the LD_LIBRARY_PATH.\n\n"
+ f"To fix this, make sure that:\n"
+ f"1. You have installed {BNB_BACKEND} {cuda_major_version}.x toolkit on your system\n"
+ f"2. The {BNB_BACKEND} runtime libraries are in your LD_LIBRARY_PATH\n\n"
+ f"You can add them with (and persist the change by adding the line to your .bashrc):\n"
+ f" export LD_LIBRARY_PATH=$LD_LIBRARY_PATH:/path/to/{BNB_BACKEND.lower()}-{cuda_major_version}.x/"
+ f"{'lib64' if not HIP_ENVIRONMENT else 'lib'}\n\n"
+ f"Original error: {self.error_msg}\n\n"
+ f"🔍 Run this command for detailed diagnostics:\n"
+ f"python -m bitsandbytes\n\n"
+ f"If you've tried everything and still have issues:\n"
+ f"1. Include ALL version info (operating system, bitsandbytes, pytorch, {BNB_BACKEND.lower()}, python)\n"
+ f"2. Describe what you've tried in detail\n"
+ f"3. Open an issue with this information:\n"
+ f" https://github.com/bitsandbytes-foundation/bitsandbytes/issues\n\n"
+ )
+
+ def __getattr__(self, name):
+ """Return a dummy function that throws when called, rather than on attribute access"""
+
+ def throw_on_call(*args, **kwargs):
+ raise RuntimeError(f"{self.formatted_error}Native code method attempted to call: lib.{name}()")
+
+ return throw_on_call
+
+ def __getitem__(self, name):
+ return self.__getattr__(name)
+
+
+def get_xpu_bnb_library_path() -> Path:
+ """Get the path to the XPU native library matching the oneAPI toolchain.
+
+ Prefers the versioned library (e.g. libbitsandbytes_xpu2026) matching the first
+ 4 digits of torch.version.xpu, falling back to an unversioned libbitsandbytes_xpu.
+ """
+ xpu_version = getattr(torch.version, "xpu", None)
+ if xpu_version:
+ versioned = PACKAGE_DIR / f"libbitsandbytes_xpu{xpu_version[:4]}{DYNAMIC_LIBRARY_SUFFIX}"
+ if versioned.exists():
+ return versioned
+ return PACKAGE_DIR / f"libbitsandbytes_xpu{DYNAMIC_LIBRARY_SUFFIX}"
+
+
+def get_native_library() -> BNBNativeLibrary:
+ """
+ Load CUDA library XOR CPU, as the latter contains a subset of symbols of the former.
+ """
+ cuda_specs = get_cuda_specs()
+ binary_path = PACKAGE_DIR / f"libbitsandbytes_cpu{DYNAMIC_LIBRARY_SUFFIX}"
+
+ if cuda_specs:
+ cuda_binary_path = get_cuda_bnb_library_path(cuda_specs)
+
+ if not cuda_binary_path.exists():
+ raise RuntimeError(f"No compatible {BNB_BACKEND} binary found at {cuda_binary_path}")
+
+ binary_path = cuda_binary_path
+
+ if torch._C._has_xpu:
+ binary_path = get_xpu_bnb_library_path()
+
+ logger.debug(f"Loading bitsandbytes native library from: {binary_path}")
+
+ # Try to load the library - any errors will propagate up
+ dll = ct.cdll.LoadLibrary(str(binary_path))
+
+ if hasattr(dll, "get_context"): # only a CUDA-built library exposes this
+ return CudaBNBNativeLibrary(dll)
+
+ if torch._C._has_xpu:
+ return XpuBNBNativeLibrary(dll)
+
+ return BNBNativeLibrary(dll)
+
+
+ROCM_GPU_ARCH = get_rocm_gpu_arch()
+
+HIP_ENVIRONMENT = False
+BNB_BACKEND = "CPU"
+if torch.version.hip:
+ HIP_ENVIRONMENT = True
+ BNB_BACKEND = "ROCm"
+elif torch.cuda.is_available():
+ BNB_BACKEND = "CUDA"
+elif torch._C._has_xpu:
+ BNB_BACKEND = "XPU"
+
+try:
+ lib = get_native_library()
+except Exception as e:
+ if BNB_BACKEND in ("CPU", "XPU"):
+ lib = ErrorHandlerMockBNBNativeLibrary("XPU/CPU can run without native library.")
+ else:
+ error_msg = str(e)
+ logger.error(
+ f"bitsandbytes library load error: {error_msg}",
+ exc_info=True,
+ )
+
+ # create a mock with error messaging as fallback
+ lib = ErrorHandlerMockBNBNativeLibrary(error_msg)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/consts.py b/venv/lib/python3.11/site-packages/bitsandbytes/consts.py
new file mode 100644
index 0000000000000000000000000000000000000000..8242d104e0084179af9ac564af574f63e7403a30
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/consts.py
@@ -0,0 +1,12 @@
+from pathlib import Path
+import platform
+
+DYNAMIC_LIBRARY_SUFFIX = {
+ "Darwin": ".dylib",
+ "Linux": ".so",
+ "Windows": ".dll",
+}.get(platform.system(), ".so")
+
+PACKAGE_DIR = Path(__file__).parent
+PACKAGE_GITHUB_URL = "https://github.com/TimDettmers/bitsandbytes"
+NONPYTORCH_DOC_URL = "https://github.com/TimDettmers/bitsandbytes/blob/main/docs/source/nonpytorchcuda.mdx"
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/cuda_specs.py b/venv/lib/python3.11/site-packages/bitsandbytes/cuda_specs.py
new file mode 100644
index 0000000000000000000000000000000000000000..25ce3cd1efacf417cc20e57a528777f864a48e52
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/cuda_specs.py
@@ -0,0 +1,111 @@
+import dataclasses
+from functools import lru_cache
+import logging
+import platform
+import re
+import subprocess
+from typing import Optional
+
+import torch
+
+
+@dataclasses.dataclass(frozen=True)
+class CUDASpecs:
+ highest_compute_capability: tuple[int, int]
+ cuda_version_string: str
+ cuda_version_tuple: tuple[int, int]
+
+ @property
+ def has_imma(self) -> bool:
+ return torch.version.hip or self.highest_compute_capability >= (7, 5)
+
+
+def get_compute_capabilities() -> list[tuple[int, int]]:
+ return sorted(torch.cuda.get_device_capability(torch.cuda.device(i)) for i in range(torch.cuda.device_count()))
+
+
+@lru_cache(None)
+def get_cuda_version_tuple() -> Optional[tuple[int, int]]:
+ """Get CUDA/HIP version as a tuple of (major, minor)."""
+ try:
+ if torch.version.cuda:
+ version_str = torch.version.cuda
+ elif torch.version.hip:
+ version_str = torch.version.hip
+ else:
+ return None
+
+ parts = version_str.split(".")
+ if len(parts) >= 2:
+ return tuple(map(int, parts[:2]))
+ return None
+ except (AttributeError, ValueError, IndexError):
+ return None
+
+
+def get_cuda_version_string() -> Optional[str]:
+ """Get CUDA/HIP version as a string."""
+ version_tuple = get_cuda_version_tuple()
+ if version_tuple is None:
+ return None
+ major, minor = version_tuple
+ return f"{major}{minor}"
+
+
+def get_cuda_specs() -> Optional[CUDASpecs]:
+ """Get CUDA/HIP specifications."""
+ if not torch.cuda.is_available():
+ return None
+
+ try:
+ compute_capabilities = get_compute_capabilities()
+ if not compute_capabilities:
+ return None
+
+ version_tuple = get_cuda_version_tuple()
+ if version_tuple is None:
+ return None
+
+ version_string = get_cuda_version_string()
+ if version_string is None:
+ return None
+
+ return CUDASpecs(
+ highest_compute_capability=compute_capabilities[-1],
+ cuda_version_string=version_string,
+ cuda_version_tuple=version_tuple,
+ )
+ except Exception:
+ return None
+
+
+def get_rocm_gpu_arch() -> str:
+ """Get ROCm GPU architecture."""
+ logger = logging.getLogger(__name__)
+ try:
+ if torch.version.hip:
+ # On Windows, use hipinfo.exe; on Linux, use rocminfo
+ if platform.system() == "Windows":
+ cmd = ["hipinfo.exe"]
+ arch_pattern = r"gcnArchName:\s+gfx([a-zA-Z\d]+)"
+ else:
+ cmd = ["rocminfo"]
+ arch_pattern = r"Name:\s+gfx([a-zA-Z\d]+)"
+
+ result = subprocess.run(cmd, capture_output=True, text=True)
+ match = re.search(arch_pattern, result.stdout)
+ if match:
+ return "gfx" + match.group(1)
+ else:
+ return "unknown"
+ else:
+ return "unknown"
+ except Exception as e:
+ logger.error(f"Could not detect ROCm GPU architecture: {e}")
+ if torch.cuda.is_available():
+ logger.warning(
+ """
+ROCm GPU architecture detection failed despite ROCm being available.
+ """,
+ )
+ return "unknown"
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/cuda.py b/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/cuda.py
new file mode 100644
index 0000000000000000000000000000000000000000..655da84a05fbf4288398f9d0c3ff3cfb9b38d3d7
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/cuda.py
@@ -0,0 +1,193 @@
+from collections.abc import Iterable, Iterator
+import logging
+import os
+from pathlib import Path
+
+import torch
+
+from bitsandbytes.cextension import HIP_ENVIRONMENT, get_cuda_bnb_library_path
+from bitsandbytes.cuda_specs import CUDASpecs
+from bitsandbytes.diagnostics.utils import print_dedented
+
+CUDART_PATH_PREFERRED_ENVVARS = ("CONDA_PREFIX", "LD_LIBRARY_PATH")
+
+CUDART_PATH_IGNORED_ENVVARS = {
+ "DBUS_SESSION_BUS_ADDRESS", # hardware related
+ "GOOGLE_VM_CONFIG_LOCK_FILE", # GCP: requires elevated permissions, causing problems in VMs and Jupyter notebooks
+ "HOME", # Linux shell default
+ "LESSCLOSE",
+ "LESSOPEN", # related to the `less` command
+ "MAIL", # something related to emails
+ "OLDPWD",
+ "PATH", # this is for finding binaries, not libraries
+ "PWD", # PWD: this is how the shell keeps track of the current working dir
+ "SHELL", # binary for currently invoked shell
+ "SSH_AUTH_SOCK", # SSH stuff, therefore unrelated
+ "SSH_TTY",
+ "TMUX", # Terminal Multiplexer
+ "XDG_DATA_DIRS", # XDG: Desktop environment stuff
+ "XDG_GREETER_DATA_DIR", # XDG: Desktop environment stuff
+ "XDG_RUNTIME_DIR",
+ "_", # current Python interpreter
+}
+
+CUDA_RUNTIME_LIB_PATTERNS = (
+ (
+ "libamdhip64.so*", # Linux
+ "amdhip64*.dll", # Windows
+ )
+ if HIP_ENVIRONMENT
+ else (
+ "cudart64*.dll", # Windows
+ "libcudart*.so*", # libcudart.so, libcudart.so.11.0, libcudart.so.12.0, libcudart.so.12.1, libcudart.so.12.2 etc.
+ "nvcuda*.dll", # Windows
+ )
+)
+
+logger = logging.getLogger(__name__)
+
+
+def find_cuda_libraries_in_path_list(paths_list_candidate: str) -> Iterable[Path]:
+ for dir_string in paths_list_candidate.split(os.pathsep):
+ if not dir_string:
+ continue
+ if os.sep not in dir_string:
+ continue
+ try:
+ dir = Path(dir_string)
+ try:
+ if not dir.exists():
+ logger.warning(f"The directory listed in your path is found to be non-existent: {dir}")
+ continue
+ except OSError: # Assume an esoteric error trying to poke at the directory
+ pass
+ for lib_pattern in CUDA_RUNTIME_LIB_PATTERNS:
+ for pth in dir.glob(lib_pattern):
+ if pth.is_file() and not pth.is_symlink():
+ yield pth
+ except (OSError, PermissionError):
+ pass
+
+
+def is_relevant_candidate_env_var(env_var: str, value: str) -> bool:
+ return (
+ env_var in CUDART_PATH_PREFERRED_ENVVARS # is a preferred location
+ or (
+ os.sep in value # might contain a path
+ and env_var not in CUDART_PATH_IGNORED_ENVVARS # not ignored
+ and "CONDA" not in env_var # not another conda envvar
+ and "BASH_FUNC" not in env_var # not a bash function defined via envvar
+ and "\n" not in value # likely e.g. a script or something?
+ )
+ )
+
+
+def get_potentially_lib_path_containing_env_vars() -> dict[str, str]:
+ return {env_var: value for env_var, value in os.environ.items() if is_relevant_candidate_env_var(env_var, value)}
+
+
+def find_cudart_libraries() -> Iterator[Path]:
+ """
+ Searches for a cuda installations, in the following order of priority:
+ 1. active conda env
+ 2. LD_LIBRARY_PATH
+ 3. any other env vars, while ignoring those that
+ - are known to be unrelated
+ - don't contain the path separator `/`
+
+ If multiple libraries are found in part 3, we optimistically try one,
+ while giving a warning message.
+ """
+ candidate_env_vars = get_potentially_lib_path_containing_env_vars()
+
+ for envvar in CUDART_PATH_PREFERRED_ENVVARS:
+ if envvar in candidate_env_vars:
+ directory = candidate_env_vars[envvar]
+ yield from find_cuda_libraries_in_path_list(directory)
+ candidate_env_vars.pop(envvar)
+
+ for env_var, value in candidate_env_vars.items():
+ yield from find_cuda_libraries_in_path_list(value)
+
+
+def _print_cuda_diagnostics(cuda_specs: CUDASpecs) -> None:
+ print(
+ f"PyTorch settings found: CUDA_VERSION={cuda_specs.cuda_version_string}, "
+ f"Highest Compute Capability: {cuda_specs.highest_compute_capability}.",
+ )
+
+ binary_path = get_cuda_bnb_library_path(cuda_specs)
+ if not binary_path.exists():
+ print_dedented(
+ f"""
+ No compatible CUDA library found (tried: {binary_path.name}). You may need to compile from source:
+ https://huggingface.co/docs/bitsandbytes/main/en/installation#cuda-compile
+ """,
+ )
+
+ # 7.5 is the minimum CC for int8 tensor cores
+ if not cuda_specs.has_imma:
+ print_dedented(
+ """
+ WARNING: Compute capability < 7.5 detected! Only slow 8-bit matmul is supported for your GPU!
+ If you run into issues with 8-bit matmul, you can try 4-bit quantization:
+ https://huggingface.co/blog/4bit-transformers-bitsandbytes
+ """,
+ )
+
+
+def _print_hip_diagnostics(cuda_specs: CUDASpecs) -> None:
+ print(f"PyTorch settings found: ROCM_VERSION={cuda_specs.cuda_version_string}")
+
+ rocm_override = os.environ.get("BNB_ROCM_VERSION")
+ if rocm_override:
+ print(f"BNB_ROCM_VERSION override: {rocm_override}")
+
+ binary_path = get_cuda_bnb_library_path(cuda_specs)
+ if not binary_path.exists():
+ print_dedented(
+ f"""
+ No compatible ROCm library found (tried: {binary_path.name}). You may need to compile from source:
+ https://huggingface.co/docs/bitsandbytes/main/en/installation#rocm-compile
+ Use BNB_ROCM_VERSION to force a specific version if needed.
+ """,
+ )
+
+ hip_major, hip_minor = cuda_specs.cuda_version_tuple
+ if (hip_major, hip_minor) < (6, 1):
+ print_dedented(
+ """
+ WARNING: bitsandbytes is fully supported only from ROCm 6.1.
+ """,
+ )
+
+
+def print_diagnostics(cuda_specs: CUDASpecs) -> None:
+ if HIP_ENVIRONMENT:
+ _print_hip_diagnostics(cuda_specs)
+ else:
+ _print_cuda_diagnostics(cuda_specs)
+
+
+def print_runtime_diagnostics() -> None:
+ backend = "ROCm" if HIP_ENVIRONMENT else "CUDA"
+ runtime_version = torch.version.hip if HIP_ENVIRONMENT else torch.version.cuda
+ override_var = "BNB_ROCM_VERSION" if HIP_ENVIRONMENT else "BNB_CUDA_VERSION"
+ override_example = "72" if HIP_ENVIRONMENT else "122"
+
+ cudart_paths = list(find_cudart_libraries())
+ if not cudart_paths:
+ print(f"{backend} SETUP: WARNING! {backend} runtime files not found in any environmental path.")
+ elif len(cudart_paths) > 1:
+ print_dedented(
+ f"""
+ Found duplicate {backend} runtime files (see below).
+
+ bitsandbytes will use PyTorch's {backend} runtime ({runtime_version}) and auto-select
+ the closest available library version. If you need to force a specific version,
+ set {override_var}, e.g.:
+ export {override_var}={override_example}
+ """,
+ )
+ for pth in cudart_paths:
+ print(f"* Found {backend} runtime at: {pth}")
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/main.py b/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/main.py
new file mode 100644
index 0000000000000000000000000000000000000000..a64925c062cf19571f18cd8a5179874dcb8d2894
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/main.py
@@ -0,0 +1,134 @@
+import importlib
+import platform
+import sys
+import traceback
+
+import torch
+
+from bitsandbytes import __version__ as bnb_version
+from bitsandbytes.consts import PACKAGE_GITHUB_URL
+from bitsandbytes.cuda_specs import get_cuda_specs
+from bitsandbytes.diagnostics.cuda import print_diagnostics
+from bitsandbytes.diagnostics.utils import print_dedented, print_header
+
+_RELATED_PACKAGES = [
+ "accelerate",
+ "diffusers",
+ "numpy",
+ "pip",
+ "peft",
+ "safetensors",
+ "transformers",
+ "triton",
+ "trl",
+]
+
+
+def sanity_check():
+ from bitsandbytes.optim import Adam
+
+ p = torch.nn.Parameter(torch.rand(10, 10).cuda())
+ a = torch.rand(10, 10).cuda()
+ p1 = p.data.sum().item()
+ adam = Adam([p])
+ out = a * p
+ loss = out.sum()
+ loss.backward()
+ adam.step()
+ p2 = p.data.sum().item()
+ assert p1 != p2
+
+
+def get_package_version(name: str) -> str:
+ try:
+ version = importlib.metadata.version(name)
+ except importlib.metadata.PackageNotFoundError:
+ version = "not found"
+ return version
+
+
+def show_environment():
+ """Simple utility to print out environment information."""
+
+ print(f"Platform: {platform.platform()}")
+ if platform.system() == "Linux":
+ print(f" libc: {'-'.join(platform.libc_ver())}")
+
+ print(f"Python: {platform.python_version()}")
+
+ print(f"PyTorch: {torch.__version__}")
+ print(f" CUDA: {torch.version.cuda or 'N/A'}")
+ print(f" HIP: {torch.version.hip or 'N/A'}")
+ print(f" XPU: {getattr(torch.version, 'xpu', 'N/A') or 'N/A'}")
+
+ print("Related packages:")
+ for pkg in _RELATED_PACKAGES:
+ version = get_package_version(pkg)
+ print(f" {pkg}: {version}")
+
+
+def main():
+ print_header(f"bitsandbytes v{bnb_version}")
+ show_environment()
+ print_header("")
+
+ cuda_specs = get_cuda_specs()
+
+ if cuda_specs:
+ print_diagnostics(cuda_specs)
+
+ has_rocm = torch.version.hip is not None
+ has_cuda = not has_rocm and torch.version.cuda is not None and torch.cuda.is_available()
+ has_xpu = hasattr(torch, "xpu") and torch.xpu.is_available()
+
+ from bitsandbytes.cextension import ErrorHandlerMockBNBNativeLibrary, lib
+
+ lib_loaded = not isinstance(lib, ErrorHandlerMockBNBNativeLibrary)
+
+ if not (has_cuda or has_rocm or has_xpu):
+ print(
+ f"No CUDA, ROCm, or XPU detected; CPU library {'loaded successfully' if lib_loaded else 'failed to load'}."
+ )
+ elif has_xpu:
+ from bitsandbytes.backends.utils import triton_available
+
+ if not isinstance(lib, ErrorHandlerMockBNBNativeLibrary):
+ print("XPU native library loaded successfully.")
+ elif triton_available:
+ print("XPU native library not loaded; using triton fallback.")
+ else:
+ print("XPU native library not loaded and triton not available.")
+ else:
+ if not lib_loaded:
+ print_dedented(
+ f"""
+ See above for details on why the library failed to load.
+ Please provide this info when creating an issue via {PACKAGE_GITHUB_URL}/issues/new/choose
+ WARNING: Please be sure to sanitize sensitive info from the output before posting it.
+ """,
+ )
+ sys.exit(1)
+
+ print("Checking that the library is importable and callable...")
+ try:
+ sanity_check()
+ print("SUCCESS!")
+ return
+ except RuntimeError as e:
+ if "not available in CPU-only" in str(e):
+ print("WARNING: bitsandbytes is running as CPU-only!")
+ print("8-bit optimizers and GPU quantization are unavailable.")
+ print("If you think this is an error, please report an issue.")
+ else:
+ raise e
+ except Exception:
+ traceback.print_exc()
+
+ print_dedented(
+ f"""
+ Above we output some debug information.
+ Please provide this info when creating an issue via {PACKAGE_GITHUB_URL}/issues/new/choose
+ WARNING: Please be sure to sanitize sensitive info from the output before posting it.
+ """,
+ )
+ sys.exit(1)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/utils.py b/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..facc58b30af8f3eb3a7895e712588dcf092721ee
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/diagnostics/utils.py
@@ -0,0 +1,12 @@
+import textwrap
+
+HEADER_WIDTH = 60
+
+
+def print_header(txt: str, width: int = HEADER_WIDTH, filler: str = "=") -> None:
+ txt = f" {txt} " if txt else ""
+ print(txt.center(width, filler))
+
+
+def print_dedented(text):
+ print("\n".join(textwrap.dedent(text).strip().split("\n")))
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/functional.py b/venv/lib/python3.11/site-packages/bitsandbytes/functional.py
new file mode 100644
index 0000000000000000000000000000000000000000..33b2cd9ac674f338429fc972669596dbedc88143
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/functional.py
@@ -0,0 +1,1810 @@
+# Copyright (c) Facebook, Inc. and its affiliates.
+#
+# This source code is licensed under the MIT license found in the
+# LICENSE file in the root directory of this source tree.
+from collections.abc import Iterable
+import ctypes as ct
+import itertools
+from math import prod
+from typing import Any, Optional
+
+import numpy as np
+import torch
+from torch import Tensor
+from typing_extensions import deprecated
+
+from bitsandbytes.utils import pack_dict_to_tensor, unpack_tensor_to_dict
+
+from .cextension import lib
+
+name2qmap = {}
+
+"""C FUNCTIONS FOR OPTIMIZERS"""
+
+
+class GlobalPageManager:
+ _instance = None
+
+ def __init__(self):
+ raise RuntimeError("Call get_instance() instead")
+
+ def initialize(self):
+ self.paged_tensors = []
+
+ @classmethod
+ def get_instance(cls):
+ if cls._instance is None:
+ cls._instance = cls.__new__(cls)
+ cls._instance.initialize()
+ return cls._instance
+
+ def prefetch_all(self, to_cpu=False):
+ # assume the first added, will be the
+ # ones that are used first, so swap them in last
+ # in the case they are evicted again
+ for t in self.paged_tensors[::-1]:
+ prefetch_tensor(t, to_cpu)
+
+
+class CUBLAS_Context:
+ _instance = None
+
+ def __init__(self):
+ raise RuntimeError("Call get_instance() instead")
+
+ def initialize(self):
+ self.context = {}
+
+ @classmethod
+ def get_instance(cls):
+ if cls._instance is None:
+ cls._instance = cls.__new__(cls)
+ cls._instance.initialize()
+ return cls._instance
+
+ def get_context(self, device):
+ if device.index not in self.context:
+ prev_device = torch.cuda.current_device()
+ torch.cuda.set_device(device)
+ self.context[device.index] = ct.c_void_p(lib.get_context())
+ torch.cuda.set_device(prev_device)
+ return self.context[device.index]
+
+
+FIRST_CUDA_DEVICE = torch.device("cuda", index=0)
+
+# When multiple GPUs are present, we use a context manager to
+# switch to the correct device of a tensor before invoking our CUDA
+# kernels in the C++ library. However, when there's only one device
+# there is no need to incur the overhead of cudaGetDevice/cudaSetDevice.
+if torch.cuda.device_count() > 1:
+
+ def _cuda_device_of(a: torch.Tensor):
+ return torch.cuda.device_of(a)
+else:
+ import contextlib
+
+ def _cuda_device_of(a: torch.Tensor):
+ return contextlib.nullcontext()
+
+
+def get_paged(*shape, dtype=torch.float32, device=FIRST_CUDA_DEVICE):
+ num_bytes = dtype.itemsize * prod(shape)
+ managed_ptr = lib.cget_managed_ptr(ct.c_size_t(num_bytes))
+ c_ptr = ct.cast(managed_ptr, ct.POINTER(ct.c_int))
+ new_array = np.ctypeslib.as_array(c_ptr, shape=shape)
+ out = torch.frombuffer(new_array, dtype=dtype, count=prod(shape)).view(shape)
+ out.is_paged = True
+ out.page_deviceid = device.index
+ return out
+
+
+def prefetch_tensor(A: torch.Tensor, to_cpu=False):
+ assert A.is_paged, "Only paged tensors can be prefetched!"
+ if to_cpu:
+ deviceid = -1
+ else:
+ deviceid = A.page_deviceid
+
+ lib.cprefetch(get_ptr(A), ct.c_size_t(A.nbytes), ct.c_int32(deviceid))
+
+
+def elementwise_func(func_name, A, B, value, prefetch=True):
+ func = None
+ if A.dtype == torch.float32:
+ func = getattr(lib, f"c{func_name}_fp32", None)
+ cvalue = ct.c_float(value)
+ elif A.dtype == torch.uint8:
+ func = getattr(lib, f"c{func_name}_uint8", None)
+ cvalue = ct.c_uint8(value)
+
+ if func is None:
+ raise NotImplementedError(f"Function not implemented: {func_name}")
+
+ is_managed = getattr(A, "is_managed", False)
+ if is_managed and prefetch:
+ prefetch_tensor(A)
+ if B is not None:
+ prefetch_tensor(B)
+
+ func(get_ptr(A), get_ptr(B), cvalue, ct.c_int64(A.numel()))
+ if A.is_paged or B.is_paged:
+ # paged function are fully asynchronous
+ # if we return from this function, we want to the tensor
+ # to be in the correct state, that is the final state after the
+ # operation occurred. So we synchronize.
+ if torch.cuda.is_available():
+ torch.cuda.synchronize()
+ elif hasattr(torch, "xpu") and torch.xpu.is_available():
+ torch.xpu.synchronize()
+
+
+def fill(A, value, device=None, prefetch=True):
+ elementwise_func("fill", A, None, value)
+
+
+def _mul(A, B, device=None):
+ elementwise_func("_mul", A, B, 0)
+
+
+def create_linear_map(signed=True, total_bits=8, add_zero=True):
+ sign = -1.0 if signed else 0.0
+ total_values = 2**total_bits
+ if add_zero or total_bits < 8:
+ # add a zero
+ # since we simulate less bits by having zeros in the data type, we
+ # we need to center the quantization around zero and as such lose
+ # a single value
+ total_values = 2**total_bits if not signed else 2**total_bits - 1
+
+ values = torch.linspace(sign, 1.0, total_values)
+ gap = 256 - values.numel()
+ if gap == 0:
+ return values
+ else:
+ l = values.numel() // 2 # noqa: E741
+ return torch.Tensor(values[:l].tolist() + [0] * gap + values[l:].tolist())
+
+
+def create_normal_map(offset=0.9677083, use_extra_value=True):
+ """Create the NormalFloat (NF4) quantization map.
+
+ Constructs a lookup table of 16 quantization values (stored in a 256-element tensor for
+ indexing convenience) derived from quantiles of the standard normal distribution N(0, 1).
+ Each bin has approximately equal probability mass under the normal distribution, which is
+ optimal for normally-distributed data like neural network weights.
+
+ Unlike floating-point types (FP4, FP8), NF4 is NOT a float encoding — the 4-bit index is
+ simply a lookup into this table. There is no sign/exponent/mantissa decomposition.
+
+ The values are generated by computing ``scipy.stats.norm.ppf()`` (inverse CDF) at evenly
+ spaced quantile points, then normalizing to [-1, 1].
+
+ For more details, see: QLoRA: Efficient Finetuning of Quantized LLMs
+ (https://arxiv.org/abs/2305.14314)
+
+ Args:
+ offset: The outermost quantile boundary, controlling the range of the normal distribution
+ that is covered. ``norm.ppf(offset)`` gives the largest bin edge in standard deviations.
+ The default (0.9677083) covers up to ~1.845 standard deviations and was empirically
+ optimized to minimize quantization error for typical neural network weight distributions.
+ use_extra_value: If True, creates an asymmetric type with 8 negative and 9 positive values
+ (including zero), for 15 non-zero values total. If False, creates a symmetric type
+ with 7 negative and 7 positive values (14 non-zero values total).
+
+ Returns:
+ A 256-element tensor where the first 16 values are the sorted NF4 quantization levels
+ normalized to [-1, 1], and the remaining values are zero (padding for 8-bit indexing).
+ """
+ try:
+ from scipy.stats import norm
+ except ImportError as ie:
+ raise ImportError(
+ "Scipy is required for `create_normal_map`. Install `bitsandbytes` with the `[test]` extra.",
+ ) from ie
+
+ if use_extra_value:
+ # one more positive value, this is an asymmetric type
+ v1 = norm.ppf(torch.linspace(offset, 0.5, 9)[:-1]).tolist()
+ v2 = [0] * (256 - 15) ## we have 15 non-zero values in this data type
+ v3 = (-norm.ppf(torch.linspace(offset, 0.5, 8)[:-1])).tolist()
+ else:
+ v1 = norm.ppf(torch.linspace(offset, 0.5, 8)[:-1]).tolist()
+ v2 = [0] * (256 - 14) ## we have 14 non-zero values in this data type
+ v3 = (-norm.ppf(torch.linspace(offset, 0.5, 8)[:-1])).tolist()
+
+ v = v1 + v2 + v3
+
+ values = torch.Tensor(v)
+ values = values.sort().values
+ values /= values.max()
+
+ assert values.numel() == 256
+
+ return values
+
+
+def create_fp8_map(signed=True, exponent_bits=5, precision_bits=2, total_bits=8):
+ """Create a floating-point quantization map with configurable bit layout.
+
+ Generates a lookup table for a custom floating-point format following IEEE 754-like encoding
+ with configurable exponent and mantissa (precision) bits. Despite the name, this function
+ handles any total bit width (including FP4 when called with ``total_bits=4``).
+
+ The encoding uses:
+ - Exponent bias: ``2^(exponent_bits - 1)``
+ - Normal values: ``(1 + mantissa) * 2^(exponent - bias - 1)``
+ - Subnormal values (exponent field = 0): ``mantissa * 2^(-bias)``
+
+ Note: The values in the returned tensor are normalized by dividing by the maximum value,
+ so the actual represented range is [-1, 1].
+
+ For the FP4 type used in bitsandbytes (2 exponent bits, 1 mantissa bit, signed):
+ ``create_fp8_map(signed=True, exponent_bits=2, precision_bits=1, total_bits=4)``
+
+ Args:
+ signed: Whether the format includes a sign bit.
+ exponent_bits: Number of bits for the exponent field.
+ precision_bits: Number of bits for the mantissa (precision/fraction) field.
+ total_bits: Total number of bits per value (must equal sign + exponent + precision).
+
+ Returns:
+ A 256-element tensor of sorted quantization levels normalized to [-1, 1].
+ For types with fewer than 8 bits, the remaining entries are zero-padded.
+ """
+ e = exponent_bits
+ p = precision_bits
+ has_sign = 1 if signed else 0
+ assert e + p == total_bits - has_sign
+ # the exponent is biased to 2^(e-1) -1 == 0
+ evalues = []
+ for i, val in enumerate(range(-(2 ** (exponent_bits - has_sign)), 2 ** (exponent_bits - has_sign), 1)):
+ evalues.append(2**val)
+
+ values = []
+ lst = list(itertools.product([0, 1], repeat=precision_bits))
+ # for ev in evalues:
+ bias = 2 ** (exponent_bits - 1)
+ for evalue in range(2 ** (exponent_bits)):
+ for bit_pattern in lst:
+ value = 1 if evalue != 0 else 0
+ for i, pval in enumerate(list(bit_pattern)):
+ value += pval * (2 ** -(i + 1))
+ if evalue == 0:
+ # subnormals
+ value = value * 2**-(bias)
+ else:
+ # normals
+ value = value * 2 ** -(evalue - bias - 1)
+ values.append(value)
+ if signed:
+ values.append(-value)
+
+ assert len(values) == 2**total_bits
+ values.sort()
+ if total_bits < 8:
+ gap = 256 - len(values)
+ for i in range(gap):
+ values.append(0)
+ values.sort()
+ code = torch.tensor(values)
+ code /= code.max()
+
+ return code
+
+
+def create_dynamic_map(signed=True, max_exponent_bits=7, total_bits=8):
+ """
+ Creates the dynamic quantiztion map.
+
+ The dynamic data type is made up of a dynamic exponent and
+ fraction. As the exponent increase from 0 to -7 the number
+ of bits available for the fraction shrinks.
+
+ This is a generalization of the dynamic type where a certain
+ number of the bits and be reserved for the linear quantization
+ region (the fraction). n determines the maximum number of
+ exponent bits.
+
+ For more details see
+ (8-Bit Approximations for Parallelism in Deep Learning)[https://arxiv.org/abs/1511.04561]
+ """
+
+ data = []
+ # these are additional items that come from the case
+ # where all the exponent bits are zero and no
+ # indicator bit is present
+ non_sign_bits = total_bits - 1
+ additional_items = 2 ** (non_sign_bits - max_exponent_bits) - 1
+ for i in range(max_exponent_bits):
+ fraction_items = int(
+ 2 ** (i + non_sign_bits - max_exponent_bits) + 1
+ if signed
+ else 2 ** (i + non_sign_bits - max_exponent_bits + 1) + 1,
+ )
+ boundaries = torch.linspace(0.1, 1, fraction_items, dtype=torch.float32)
+ means = (boundaries[:-1] + boundaries[1:]) / 2.0
+ data += ((10 ** (-(max_exponent_bits - 1) + i)) * means).tolist()
+ if signed:
+ data += (-(10 ** (-(max_exponent_bits - 1) + i)) * means).tolist()
+
+ if additional_items > 0:
+ boundaries = torch.linspace(0.1, 1, additional_items + 1, dtype=torch.float32)
+ means = (boundaries[:-1] + boundaries[1:]) / 2.0
+ data += ((10 ** (-(max_exponent_bits - 1) + i)) * means).tolist()
+ if signed:
+ data += (-(10 ** (-(max_exponent_bits - 1) + i)) * means).tolist()
+
+ data.append(0)
+ data.append(1.0)
+
+ assert len(data) == 2**total_bits
+
+ gap = 256 - len(data)
+ for i in range(gap):
+ data.append(0)
+
+ data.sort()
+ return torch.tensor(data, dtype=torch.float32)
+
+
+def is_on_gpu(tensors: Iterable[Optional[torch.Tensor]]):
+ """Verifies that the input tensors are all on the same device.
+
+ An input tensor may also be marked as `paged`, in which case the device placement is ignored.
+ CPU tensors are allowed and checked for consistency among themselves.
+
+ Args:
+ tensors (`Iterable[Optional[torch.Tensor]]`): A list of tensors to verify.
+
+ Raises:
+ `RuntimeError`: Raised when the verification fails.
+
+ Returns:
+ `Literal[True]`
+ """
+
+ devices = set()
+
+ for t in tensors:
+ # NULL pointers and paged tensors are OK.
+ if t is not None and not getattr(t, "is_paged", False):
+ devices.add((t.device.type, t.device.index))
+
+ # All tensors on CPU is valid
+ if devices == {("cpu", None)}:
+ return True
+
+ # Check that no CPU tensors are mixed with GPU tensors
+ has_cpu = ("cpu", None) in devices
+ if has_cpu and len(devices) > 1:
+ raise RuntimeError(
+ f"Input tensors need to be on the same device, but found the following tensor and device combinations:\n {[(t.shape, t.device) for t in tensors if t is not None]}",
+ )
+
+ # GPU path: all tensors must be on the same single GPU
+ if len(devices) > 1:
+ raise RuntimeError(
+ f"Input tensors need to be on the same GPU, but found the following tensor and device combinations:\n {[(t.shape, t.device) for t in tensors if t is not None]}",
+ )
+ return True
+
+
+def _get_tensor_stream(tensor: Tensor) -> ct.c_void_p:
+ # We use the raw stream for performance reasons.
+ if tensor.device.type == "cuda":
+ return ct.c_void_p(torch._C._cuda_getCurrentRawStream(tensor.device.index))
+ if tensor.device.type == "xpu":
+ return ct.c_void_p(torch._C._xpu_getCurrentRawStream(tensor.device.index))
+ # For CPU tensors (e.g. paged optimizer states), use current device's stream.
+ if hasattr(torch, "xpu") and torch.xpu.is_available():
+ return ct.c_void_p(torch._C._xpu_getCurrentRawStream(torch.xpu.current_device()))
+ return ct.c_void_p(torch._C._cuda_getCurrentRawStream(torch.cuda.current_device()))
+
+
+def get_ptr(A: Optional[Tensor]) -> Optional[ct.c_void_p]:
+ """Gets the memory address of the first element of a tenso
+
+ Args:
+ A (`Optional[Tensor]`): A PyTorch tensor.
+
+ Returns:
+ `Optional[ct.c_void_p]`: A pointer to the underlying tensor data.
+ """
+ if A is None:
+ return None
+
+ return ct.c_void_p(A.data_ptr())
+
+
+class QuantState:
+ """container for quantization state components to work with Params4bit and similar classes"""
+
+ valid_quant_types = ("fp4", "nf4")
+ valid_qs_type_keys = [f"bitsandbytes__{x}" for x in valid_quant_types]
+ valid_qs_keys = [
+ "absmax",
+ "quant_map",
+ "nested_absmax",
+ "nested_quant_map",
+ "quant_state",
+ "quant_type",
+ "blocksize",
+ "dtype",
+ "shape",
+ "nested_blocksize",
+ "nested_dtype",
+ "nested_offset",
+ ]
+
+ def __init__(
+ self,
+ absmax,
+ shape=None,
+ code=None,
+ blocksize=None,
+ quant_type=None,
+ dtype=None,
+ offset=None,
+ state2=None,
+ ):
+ self.absmax = absmax
+ self.shape = shape
+ self.code = code
+ self.dtype = dtype
+ self.blocksize = blocksize
+ self.quant_type = quant_type
+ self.offset = offset
+ self.state2 = state2
+ self.nested = state2 is not None
+
+ def __getattr__(self, name):
+ # Support attribute access for packed state_dict keys like "bitsandbytes__nf4".
+ # PyTorch's FSDP state_dict traversal (_get_fqns) resolves dotted FQN paths via
+ # getattr. The packed key "quant_state.bitsandbytes__nf4" causes it to call
+ # getattr(quant_state_obj, "bitsandbytes__nf4"), which we handle here.
+ if name.startswith("bitsandbytes__"):
+ qs_dict = self.as_dict(packed=True)
+ packed_key = "quant_state." + name
+ if packed_key in qs_dict:
+ return qs_dict[packed_key]
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
+
+ def __getitem__(self, idx):
+ """
+ ensures compatibility with older quant state scheme with nested lists.
+ assumes the following layout:
+ state = [qabsmax, input_shape, A.dtype, blocksize, [offset, state2], quant_type]
+ state2 = [absmax, input_shape, A.dtype, blocksize, None, quant_type]
+ """
+ if self.nested:
+ list_repr = [
+ self.absmax,
+ self.shape,
+ self.dtype,
+ self.blocksize,
+ [self.offset, self.state2],
+ self.quant_type,
+ ]
+ else:
+ list_repr = [self.absmax, self.shape, self.dtype, self.blocksize, None, self.quant_type]
+ return list_repr[idx]
+
+ @classmethod
+ def from_dict(cls, qs_dict: dict[str, Any], device: torch.device) -> "QuantState":
+ """
+ unpacks components of state_dict into QuantState
+ where necessary, convert into strings, torch.dtype, ints, etc.
+
+ qs_dict: based on state_dict, with only relevant keys, striped of prefixes.
+
+ item with key `quant_state.bitsandbytes__[nf4/fp4]` may contain minor and non-tensor quant state items.
+ """
+
+ # unpacking tensor with non-tensor components
+ qs_key = [k for k, v in qs_dict.items() if "quant_state" in k and isinstance(v, torch.Tensor)]
+ if "quant_type" not in qs_dict:
+ if not qs_key:
+ raise ValueError("Expected packed or unpacked quant_state items, found neither")
+ elif len(qs_key) != 1 or qs_key[0].split(".")[-1] not in cls.valid_qs_type_keys:
+ raise ValueError(
+ f"There should be exactly one `quant_state` item with ending from {cls.valid_qs_type_keys}.\nDetected {qs_key}.",
+ )
+
+ # unpacking minor and non-tensor quant state items if necessary
+ if len(qs_key) == 1:
+ first_qs_key = qs_key[0]
+ qs_dict.update(unpack_tensor_to_dict(qs_dict.pop(first_qs_key)))
+
+ qs_dict = {k.split(".")[-1]: v for k, v in qs_dict.items()} # strip prefixes
+ assert set(qs_dict.keys()).issubset(cls.valid_qs_keys)
+
+ if "nested_absmax" in qs_dict:
+ offset = torch.tensor(float(qs_dict["nested_offset"])).to(device)
+ state2 = cls(
+ absmax=qs_dict["nested_absmax"].to(device),
+ blocksize=qs_dict["nested_blocksize"],
+ code=qs_dict["nested_quant_map"].to(device),
+ dtype=getattr(torch, qs_dict["nested_dtype"]),
+ )
+ else:
+ offset, state2 = None, None
+
+ quant_state = cls(
+ quant_type=qs_dict["quant_type"],
+ absmax=qs_dict["absmax"].to(device),
+ blocksize=qs_dict["blocksize"],
+ code=qs_dict["quant_map"].to(device),
+ dtype=getattr(torch, qs_dict["dtype"]),
+ shape=torch.Size(qs_dict["shape"]) if qs_dict["shape"] is not None else None,
+ offset=offset,
+ state2=state2,
+ )
+ return quant_state
+
+ def as_dict(self, packed: bool = False) -> dict[str, Any]:
+ """
+ returns dict of tensors and strings to use in serialization via _save_to_state_dict()
+ param: packed -- returns dict[str, torch.Tensor] for state_dict fit for safetensors saving
+ """
+ qs_dict = {
+ "quant_type": self.quant_type,
+ "absmax": self.absmax,
+ "blocksize": self.blocksize,
+ "quant_map": self.code,
+ "dtype": str(self.dtype).strip("torch."),
+ "shape": tuple(self.shape) if self.shape is not None else None,
+ }
+ if self.nested:
+ qs_dict.update(
+ {
+ "nested_absmax": self.state2.absmax,
+ "nested_blocksize": self.state2.blocksize,
+ "nested_quant_map": self.state2.code.clone(), # un-shared to avoid restoring it after shared tensors are removed by safetensors
+ "nested_dtype": str(self.state2.dtype).strip("torch."),
+ "nested_offset": self.offset.item(),
+ },
+ )
+ if not packed or self.quant_type is None:
+ return qs_dict
+
+ # packed format allows serialization of non-tensor components, critical for saving in safetensors format
+ qs_packed_dict = {k: v for k, v in qs_dict.items() if isinstance(v, torch.Tensor)}
+ non_tensor_dict = {k: v for k, v in qs_dict.items() if not isinstance(v, torch.Tensor)}
+ key = "quant_state.bitsandbytes__"
+ if self.quant_type is not None:
+ key += self.quant_type
+ qs_packed_dict[key] = pack_dict_to_tensor(non_tensor_dict)
+ return qs_packed_dict
+
+ def to(self, device):
+ # make sure the quantization state is on the right device
+ self.code = self.code.to(device)
+ self.absmax = self.absmax.to(device)
+ if self.nested:
+ self.offset = self.offset.to(device)
+ self.state2.absmax = self.state2.absmax.to(device)
+ self.state2.code = self.state2.code.to(device)
+
+ def __eq__(self, other):
+ if not isinstance(other, QuantState):
+ return False
+
+ return (
+ torch.allclose(self.absmax, other.absmax, atol=1e-6)
+ and self.shape == other.shape
+ and torch.allclose(self.code, other.code, atol=1e-6)
+ and self.dtype == other.dtype
+ and self.blocksize == other.blocksize
+ and self.quant_type == other.quant_type
+ and (
+ self.offset == other.offset
+ if self.offset is not None and other.offset is not None
+ else self.offset is other.offset
+ )
+ and (
+ self.state2 == other.state2
+ if self.state2 is not None and other.state2 is not None
+ else self.state2 is other.state2
+ )
+ )
+
+
+def quantize_blockwise(
+ A: torch.Tensor,
+ code: Optional[torch.Tensor] = None,
+ absmax: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize=4096,
+ nested=False,
+) -> tuple[torch.Tensor, QuantState]:
+ """Quantize a tensor in blocks of values.
+
+ The input tensor is quantized by dividing it into blocks of `blocksize` values.
+ The the absolute maximum value within these blocks is calculated for scaling
+ the non-linear quantization.
+
+ Args:
+ A (`torch.Tensor`): The input tensor. Supports `float16`, `bfloat16`, or `float32` datatypes.
+ code (`torch.Tensor`, *optional*):
+ A mapping describing the low-bit data type. Defaults to a signed 8-bit dynamic type.
+ For more details, see (8-Bit Approximations for Parallelism in Deep Learning)[https://arxiv.org/abs/1511.04561].
+ absmax (`torch.Tensor`, *optional*): A tensor to use to store the absmax values.
+ out (`torch.Tensor`, *optional*): A tensor to use to store the result.
+ blocksize (`int`, *optional*):
+ The size of the blocks. Defaults to 4096.
+ Valid values are 64, 128, 256, 512, 1024, 2048, and 4096.
+ nested (`bool`, *optional*): Whether to additionally quantize the absmax values. Defaults to False.
+
+ Raises:
+ ValueError: Raised when the input data type is not supported.
+
+ Returns:
+ `Tuple[torch.Tensor, QuantState]`: A tuple containing the quantization results.
+ - `torch.Tensor`: The quantized tensor.
+ - [`QuantState`]: The state object used to undo the quantization.
+ """
+
+ if blocksize <= 0:
+ raise ValueError(f"blocksize must be positive, got {blocksize}")
+ if A.dtype not in (torch.float32, torch.float16, torch.bfloat16):
+ raise ValueError(f"Blockwise quantization only supports 16/32-bit floats, but got {A.dtype}")
+
+ if code is None:
+ if "dynamic" not in name2qmap:
+ name2qmap["dynamic"] = create_dynamic_map().to(A.device)
+ code = name2qmap["dynamic"]
+
+ _out, _absmax = torch.ops.bitsandbytes.quantize_blockwise.default(
+ A,
+ code.to(A.device),
+ blocksize,
+ )
+
+ if nested:
+ offset = _absmax.mean()
+ _absmax -= offset
+ qabsmax, state2 = quantize_blockwise(_absmax, blocksize=blocksize, nested=False)
+ quant_state = QuantState(
+ absmax=qabsmax,
+ code=code.to(A.device, copy=True),
+ blocksize=blocksize,
+ dtype=A.dtype,
+ offset=offset,
+ state2=state2,
+ )
+ else:
+ quant_state = QuantState(absmax=_absmax, code=code.to(A.device, copy=True), blocksize=blocksize, dtype=A.dtype)
+
+ # TODO(matthewdouglas): Deprecate out kwarg
+ out = out.copy_(_out) if out is not None else _out
+
+ # TODO(matthewdouglas): Deprecate absmax kwarg
+ if absmax is not None:
+ quant_state.absmax = absmax.copy_(quant_state.absmax)
+
+ return out, quant_state
+
+
+def dequantize_blockwise(
+ A: torch.Tensor,
+ quant_state: Optional[QuantState] = None,
+ absmax: Optional[torch.Tensor] = None,
+ code: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize: int = 4096,
+ nested=False,
+) -> torch.Tensor:
+ """Dequantize a tensor in blocks of values.
+
+ The input tensor is dequantized by dividing it into blocks of `blocksize` values.
+ The the absolute maximum value within these blocks is used for scaling
+ the non-linear dequantization.
+
+ Args:
+ A (`torch.Tensor`): The quantized input tensor.
+ quant_state ([`QuantState`], *optional*):
+ The quantization state as returned by [`quantize_blockwise`].
+ Required if `absmax` is not provided.
+ absmax (`torch.Tensor`, *optional*):
+ A tensor containing the scaling values.
+ Required if `quant_state` is not provided and ignored otherwise.
+ code (`torch.Tensor`, *optional*):
+ A mapping describing the low-bit data type. Defaults to a signed 8-bit dynamic type.
+ For more details, see (8-Bit Approximations for Parallelism in Deep Learning)[https://arxiv.org/abs/1511.04561].
+ Ignored when `quant_state` is provided.
+ out (`torch.Tensor`, *optional*): A tensor to use to store the result.
+ blocksize (`int`, *optional*):
+ The size of the blocks. Defaults to 4096.
+ Valid values are 64, 128, 256, 512, 1024, 2048, and 4096.
+ Ignored when `quant_state` is provided.
+
+ Raises:
+ ValueError: Raised when the input data type is not supported.
+
+ Returns:
+ `torch.Tensor`:
+ The dequantized tensor. The datatype is indicated by `quant_state.dtype` and defaults to `torch.float32`.
+ """
+
+ if quant_state is None and absmax is None:
+ raise ValueError("dequantize_blockwise requires either quant_state or absmax")
+ if A.dtype != torch.uint8:
+ raise ValueError(f"A must be uint8, got {A.dtype}")
+ if code is None and quant_state is None:
+ if "dynamic" not in name2qmap:
+ name2qmap["dynamic"] = create_dynamic_map().to(A.device)
+ code = name2qmap["dynamic"]
+
+ if quant_state is None:
+ quant_state = QuantState(absmax=absmax, code=code, blocksize=blocksize, dtype=torch.float32)
+
+ if quant_state.blocksize <= 0:
+ raise ValueError(f"blocksize must be positive, got {quant_state.blocksize}")
+
+ absmax = quant_state.absmax
+ if quant_state.nested:
+ absmax = dequantize_blockwise(quant_state.absmax, quant_state.state2)
+ absmax += quant_state.offset
+ if absmax.dtype != torch.float32:
+ absmax = absmax.float()
+
+ if out is not None:
+ torch.ops.bitsandbytes.dequantize_blockwise.out(
+ A,
+ absmax,
+ quant_state.code.to(A.device),
+ quant_state.blocksize,
+ quant_state.dtype,
+ out=out,
+ )
+ return out
+
+ return torch.ops.bitsandbytes.dequantize_blockwise.default(
+ A,
+ absmax,
+ quant_state.code.to(A.device),
+ quant_state.blocksize,
+ quant_state.dtype,
+ )
+
+
+def get_4bit_type(typename, device=None, blocksize=64):
+ if device is None:
+ device = "cuda"
+ data = None
+ if typename == "nf4":
+ # NF4 (NormalFloat4) quantization type.
+ #
+ # These 16 values are a lookup table derived from quantiles of the standard normal
+ # distribution N(0, 1), where each bin has equal probability mass. The 4-bit index
+ # is just a position in this table — NF4 is NOT a floating-point encoding (no
+ # sign/exponent/mantissa decomposition). This is fundamentally different from FP4.
+ #
+ # Generated by: create_normal_map(offset=0.9677083, use_extra_value=True)
+ # Values are hardcoded to avoid a scipy dependency at runtime.
+ #
+ # For details see: QLoRA (https://arxiv.org/abs/2305.14314)
+ data = [
+ -1.0,
+ -0.6961928009986877,
+ -0.5250730514526367,
+ -0.39491748809814453,
+ -0.28444138169288635,
+ -0.18477343022823334,
+ -0.09105003625154495,
+ 0.0,
+ 0.07958029955625534,
+ 0.16093020141124725,
+ 0.24611230194568634,
+ 0.33791524171829224,
+ 0.44070982933044434,
+ 0.5626170039176941,
+ 0.7229568362236023,
+ 1.0,
+ ]
+ elif typename == "fp4":
+ # FP4 (4-bit floating point) quantization type.
+ #
+ # Unlike NF4, FP4 is an actual floating-point encoding with 1 sign bit, 2 exponent
+ # bits, and 1 mantissa bit. Values below are listed in bit-pattern order (not value
+ # order), where only the 3 non-sign bits are shown:
+ #
+ # 0b000 = 0 (subnormal: zero)
+ # 0b001 = 0.0625 (subnormal: 0.5 * 2^-2)
+ # 0b010 = 8 0b011 = 12 0b100 = 4
+ # 0b101 = 6 0b110 = 2 0b111 = 3
+ #
+ # The exponent bias is 2^(e-1) = 2, which differs from IEEE 754's convention.
+ # These can be regenerated with:
+ # create_fp8_map(signed=True, exponent_bits=2, precision_bits=1, total_bits=4)
+ #
+ # All values are normalized to [-1, 1] after construction (see end of function).
+ data = [0, 0.0625, 8.0, 12.0, 4.0, 6.0, 2.0, 3.0, -0, -0.0625, -8.0, -12.0, -4.0, -6.0, -2.0, -3.0]
+ elif typename == "int4":
+ data = [7, 6, 5, 4, 3, 2, 1, 0, -0, -1, -2, -3, -4, -5, -6, -7]
+ elif typename == "af4":
+ # Taken from: NF4 Isn't Information Theoretically Optimal (and that's Good)
+ # https://arxiv.org/abs/2306.06965
+ if blocksize == 64:
+ data = [
+ -1.0,
+ -0.69441008,
+ -0.51243739,
+ -0.3736951,
+ -0.25607552,
+ -0.14982478,
+ -0.04934812,
+ 0.0,
+ 0.04273164,
+ 0.12934483,
+ 0.21961274,
+ 0.31675666,
+ 0.42563882,
+ 0.55496234,
+ 0.72424863,
+ 1.0,
+ ][::-1]
+ else:
+ raise NotImplementedError("4-bit AbnormalFloats currently only support blocksize 64.")
+
+ if data is None:
+ raise NotImplementedError(f"Typename {typename} not supported")
+
+ data = torch.tensor(data, device=device)
+ data.div_(data.abs().max())
+
+ assert data.numel() == 16
+
+ return data
+
+
+def quantize_fp4(
+ A: torch.Tensor,
+ absmax: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize=None,
+ compress_statistics=False,
+ quant_storage=torch.uint8,
+):
+ return quantize_4bit(A, absmax, out, blocksize, compress_statistics, "fp4", quant_storage)
+
+
+def quantize_nf4(
+ A: torch.Tensor,
+ absmax: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize=None,
+ compress_statistics=False,
+ quant_storage=torch.uint8,
+):
+ return quantize_4bit(A, absmax, out, blocksize, compress_statistics, "nf4", quant_storage)
+
+
+def quantize_4bit(
+ A: torch.Tensor,
+ absmax: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize=None,
+ compress_statistics=False,
+ quant_type="fp4",
+ quant_storage=torch.uint8,
+) -> tuple[torch.Tensor, QuantState]:
+ """Quantize tensor A in blocks of 4-bit values.
+
+ Quantizes tensor A by dividing it into blocks which are independently quantized.
+
+ Args:
+ A (`torch.Tensor`): The input tensor. Supports `float16`, `bfloat16`, or `float32` datatypes.
+ absmax (`torch.Tensor`, *optional*): A tensor to use to store the absmax values.
+ out (`torch.Tensor`, *optional*): A tensor to use to store the result.
+ blocksize (`int`, *optional*):
+ The size of the blocks. Defaults to 64.
+ Valid values are 32, 64, 128, 256, 512, 1024, 2048, and 4096.
+ compress_statistics (`bool`, *optional*): Whether to additionally quantize the absmax values. Defaults to False.
+ quant_type (`str`, *optional*): The data type to use: `nf4` or `fp4`. Defaults to `fp4`.
+ quant_storage (`torch.dtype`, *optional*): The dtype of the tensor used to store the result. Defaults to `torch.uint8`.
+
+ Raises:
+ ValueError: Raised when the input data type is not supported.
+
+ Returns:
+ Tuple[`torch.Tensor`, `QuantState`]: A tuple containing the quantization results.
+ - `torch.Tensor`: The quantized tensor with packed 4-bit values.
+ - [`QuantState`]: The state object used to undo the quantization.
+ """
+
+ if blocksize is None:
+ blocksize = 64
+
+ if blocksize not in (32, 64, 128, 256, 512, 1024, 2048, 4096):
+ raise ValueError(f"invalid blocksize {blocksize}")
+ if quant_type not in ("nf4", "fp4"):
+ raise ValueError(f"quant_type must be 'nf4' or 'fp4', got {quant_type!r}")
+ if A.dtype not in (torch.bfloat16, torch.float16, torch.float32):
+ raise ValueError(f"Blockwise 4bit quantization only supports 16/32-bit floats, but got {A.dtype}")
+
+ input_shape = A.shape
+
+ _out, _absmax = torch.ops.bitsandbytes.quantize_4bit.default(
+ A,
+ blocksize,
+ quant_type,
+ quant_storage,
+ )
+
+ code = get_4bit_type(quant_type, device=A.device)
+
+ if compress_statistics:
+ offset = _absmax.mean()
+ qabsmax, state2 = quantize_blockwise(_absmax - offset, blocksize=256)
+ del _absmax
+ state = QuantState(
+ absmax=qabsmax,
+ shape=input_shape,
+ dtype=A.dtype,
+ blocksize=blocksize,
+ code=code,
+ quant_type=quant_type,
+ offset=offset,
+ state2=state2,
+ )
+ else:
+ state = QuantState(
+ absmax=_absmax,
+ shape=input_shape,
+ dtype=A.dtype,
+ blocksize=blocksize,
+ code=code,
+ quant_type=quant_type,
+ )
+
+ # TODO(matthewdouglas): Deprecate out kwarg
+ out = out.copy_(_out) if out is not None else _out
+
+ # TODO(matthewdouglas): Deprecate absmax kwarg
+ if absmax is not None:
+ state.absmax = absmax.copy_(state.absmax)
+
+ return out, state
+
+
+def dequantize_fp4(
+ A: torch.Tensor,
+ quant_state: Optional[QuantState] = None,
+ absmax: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize: Optional[int] = None,
+) -> torch.Tensor:
+ return dequantize_4bit(A, quant_state, absmax, out, blocksize, "fp4")
+
+
+def dequantize_nf4(
+ A: torch.Tensor,
+ quant_state: Optional[QuantState] = None,
+ absmax: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize: Optional[int] = None,
+) -> torch.Tensor:
+ return dequantize_4bit(A, quant_state, absmax, out, blocksize, "nf4")
+
+
+def dequantize_4bit(
+ A: torch.Tensor,
+ quant_state: Optional[QuantState] = None,
+ absmax: Optional[torch.Tensor] = None,
+ out: Optional[torch.Tensor] = None,
+ blocksize: Optional[int] = None,
+ quant_type="fp4",
+) -> torch.Tensor:
+ """Dequantizes a packed 4-bit quantized tensor.
+
+ The input tensor is dequantized by dividing it into blocks of `blocksize` values.
+ The absolute maximum value within these blocks is used for scaling
+ the non-linear dequantization.
+
+ Args:
+ A (`torch.Tensor`): The quantized input tensor.
+ quant_state ([`QuantState`], *optional*):
+ The quantization state as returned by [`quantize_4bit`].
+ Required if `absmax` is not provided.
+ absmax (`torch.Tensor`, *optional*):
+ A tensor containing the scaling values.
+ Required if `quant_state` is not provided and ignored otherwise.
+ out (`torch.Tensor`, *optional*): A tensor to use to store the result.
+ blocksize (`int`, *optional*):
+ The size of the blocks. Defaults to 64.
+ Valid values are 32, 64, 128, 256, 512, 1024, 2048, and 4096.
+ quant_type (`str`, *optional*): The data type to use: `nf4` or `fp4`. Defaults to `fp4`.
+
+ Raises:
+ ValueError: Raised when the input data type or blocksize is not supported.
+
+ Returns:
+ `torch.Tensor`: The dequantized tensor.
+ """
+
+ if blocksize is None:
+ blocksize = 64
+
+ if quant_state is None:
+ if absmax is None or out is None:
+ raise ValueError("dequantize_4bit requires both absmax and out when quant_state is not provided")
+
+ quant_state = QuantState(
+ absmax=absmax,
+ shape=out.shape,
+ dtype=out.dtype,
+ blocksize=blocksize,
+ quant_type=quant_type,
+ )
+
+ else:
+ absmax = quant_state.absmax
+
+ if quant_state.blocksize not in (32, 64, 128, 256, 512, 1024, 2048, 4096):
+ raise ValueError(f"invalid blocksize {quant_state.blocksize}")
+ if quant_state.quant_type not in ("nf4", "fp4"):
+ raise ValueError(f"quant_type must be 'nf4' or 'fp4', got {quant_state.quant_type!r}")
+ if quant_state.dtype not in (torch.bfloat16, torch.float16, torch.float32):
+ raise ValueError(f"Blockwise 4bit dequantization only supports 16/32-bit floats, but got {quant_state.dtype}")
+
+ if quant_state.nested:
+ absmax = dequantize_blockwise(quant_state.absmax, quant_state.state2)
+ absmax += quant_state.offset
+ if absmax.dtype != torch.float32:
+ absmax = absmax.float()
+
+ if out is not None:
+ torch.ops.bitsandbytes.dequantize_4bit.out(
+ A, absmax, quant_state.blocksize, quant_state.quant_type, quant_state.shape, quant_state.dtype, out=out
+ )
+ else:
+ out = torch.ops.bitsandbytes.dequantize_4bit.default(
+ A,
+ absmax,
+ quant_state.blocksize,
+ quant_state.quant_type,
+ quant_state.shape,
+ quant_state.dtype,
+ )
+
+ # BC shim: callers that pass the packed weight in transposed [1, (N*K+1)//2] form
+ # receive the output transposed back to [K, N]. bnb's own paths no longer trigger
+ # this since B is normalized to [(N*K+1)//2, 1] at the matmul_4bit entry point.
+ if A.shape[0] == 1:
+ return out.t()
+ return out
+
+
+def optimizer_update_32bit(
+ optimizer_name: str,
+ g: Tensor,
+ p: Tensor,
+ state1: Tensor,
+ beta1: float,
+ eps: float,
+ step: int,
+ lr: float,
+ state2: Optional[torch.Tensor] = None,
+ beta2: float = 0.0,
+ beta3: float = 0.0,
+ alpha: float = 0.0,
+ weight_decay: float = 0.0,
+ gnorm_scale: float = 1.0,
+ unorm_vec: Optional[torch.Tensor] = None,
+ max_unorm: float = 0.0,
+ skip_zeros=False,
+) -> None:
+ """
+ Performs an inplace optimizer update with one or two optimizer states.
+
+ Universal optimizer update for 32-bit state and 32/16-bit gradients/weights.
+
+ Parameters
+ ----------
+ optimizer_name : str
+ The name of the optimizer: {adam}.
+ g : torch.Tensor
+ Gradient tensor.
+ p : torch.Tensor
+ Parameter tensor.
+ state1 : torch.Tensor
+ Optimizer state 1.
+ beta1 : float
+ Optimizer beta1.
+ eps : float
+ Optimizer epsilon.
+ weight_decay : float
+ Weight decay.
+ step : int
+ Current optimizer step.
+ lr : float
+ The learning rate.
+ state2 : torch.Tensor
+ Optimizer state 2.
+ beta2 : float
+ Optimizer beta2.
+ beta3 : float
+ Optimizer beta3.
+ alpha : float
+ Optimizer alpha.
+ gnorm_scale : float
+ The factor to rescale the gradient to the max clip value.
+ unorm_vec : torch.Tensor
+ The tensor for the update norm.
+ max_unorm : float
+ The maximum update norm relative to the weight norm.
+ skip_zeros : bool
+ Whether to skip zero-valued gradients or not (default: False).
+ """
+
+ param_norm = 0.0
+ if max_unorm > 0.0:
+ param_norm = torch.norm(p.data.float())
+
+ is_on_gpu([g, p, state1, state2, unorm_vec])
+ torch.ops.bitsandbytes.optimizer_update_32bit(
+ optimizer_name,
+ g,
+ p,
+ state1,
+ state2,
+ unorm_vec,
+ max_unorm,
+ param_norm,
+ beta1,
+ beta2,
+ beta3,
+ alpha,
+ eps,
+ weight_decay,
+ step,
+ lr,
+ gnorm_scale,
+ skip_zeros,
+ )
+
+
+def optimizer_update_8bit_blockwise(
+ optimizer_name: str,
+ g: Tensor,
+ p: Tensor,
+ state1: Tensor,
+ state2: Optional[torch.Tensor],
+ beta1: float,
+ beta2: float,
+ beta3: float,
+ alpha: float,
+ eps: float,
+ step: int,
+ lr: float,
+ qmap1: Tensor,
+ qmap2: Optional[torch.Tensor],
+ absmax1: Tensor,
+ absmax2: Optional[torch.Tensor],
+ weight_decay: float = 0.0,
+ gnorm_scale: float = 1.0,
+ skip_zeros=False,
+) -> None:
+ is_on_gpu([p, g, state1, state2, qmap1, qmap2, absmax1, absmax2])
+
+ torch.ops.bitsandbytes.optimizer_update_8bit_blockwise(
+ optimizer_name,
+ g,
+ p,
+ state1,
+ state2,
+ beta1,
+ beta2,
+ beta3,
+ alpha,
+ eps,
+ step,
+ lr,
+ qmap1,
+ qmap2,
+ absmax1,
+ absmax2,
+ weight_decay,
+ gnorm_scale,
+ skip_zeros,
+ )
+
+
+@deprecated("This function is deprecated and will be removed in a future release.", category=FutureWarning)
+def check_matmul(A, B, out, transposed_A, transposed_B, expected_type=torch.int8):
+ if not torch.cuda.is_initialized():
+ torch.cuda.init()
+ if A.dtype != expected_type or B.dtype != expected_type:
+ raise TypeError(f"Expected torch.int8 input tensors A and B, but got {A.dtype} and {B.dtype}")
+
+ sA = A.shape
+ sB = B.shape
+ tA = transposed_A
+ tB = transposed_B
+
+ correct = True
+
+ if len(sA) == 2 and len(sB) == 2:
+ if not tA and not tB and A.shape[1] != B.shape[0]:
+ correct = False
+ elif tA and not tB and A.shape[0] != B.shape[0]:
+ correct = False
+ elif tA and tB and A.shape[0] != B.shape[1]:
+ correct = False
+ elif not tA and tB and A.shape[1] != B.shape[1]:
+ correct = False
+ elif len(sA) == 3 and len(sB) == 2:
+ if not tA and not tB and A.shape[2] != B.shape[0]:
+ correct = False
+ elif tA and not tB and A.shape[1] != B.shape[0]:
+ correct = False
+ elif tA and tB and A.shape[1] != B.shape[1]:
+ correct = False
+ elif not tA and tB and A.shape[2] != B.shape[1]:
+ correct = False
+ elif len(sA) == 3 and len(sB) == 3:
+ if not tA and not tB and A.shape[2] != B.shape[1]:
+ correct = False
+ elif tA and not tB and A.shape[1] != B.shape[1]:
+ correct = False
+ elif tA and tB and A.shape[1] != B.shape[2]:
+ correct = False
+ elif not tA and tB and A.shape[2] != B.shape[2]:
+ correct = False
+
+ if out is not None:
+ sout = out.shape
+ # special case common in backprop
+ if not correct and len(sA) == 3 and len(sB) == 3:
+ if sout[0] == sA[2] and sout[1] == sB[2] and sA[0] == sB[0] and sA[1] == sB[1]:
+ correct = True
+ else:
+ if len(sA) == 2 and len(sB) == 2:
+ if not tA and not tB:
+ sout = (sA[0], sB[1])
+ elif tA and tB:
+ sout = (sA[1], sB[0])
+ elif tA and not tB:
+ sout = (sA[1], sB[1])
+ elif not tA and tB:
+ sout = (sA[0], sB[0])
+ elif len(sA) == 3 and len(sB) == 2:
+ if not tA and not tB:
+ sout = (sA[0], sA[1], sB[1])
+ elif tA and tB:
+ sout = (sA[0], sA[2], sB[0])
+ elif tA and not tB:
+ sout = (sA[0], sA[2], sB[1])
+ elif not tA and tB:
+ sout = (sA[0], sA[1], sB[0])
+ elif len(sA) == 3 and len(sB) == 3:
+ if not tA and not tB:
+ sout = (sA[0], sA[1], sB[2])
+ elif tA and tB:
+ sout = (sA[0], sA[2], sB[1])
+ elif tA and not tB:
+ sout = (sA[0], sA[2], sB[2])
+ elif not tA and tB:
+ sout = (sA[0], sA[1], sB[1])
+
+ if not correct:
+ raise ValueError(
+ f"Tensor dimensions incorrect for matrix mulitiplication: A x B: {sA} x {sB} with transpose for A x B: {tA} x {tB}.",
+ )
+
+ return sout
+
+
+def gemv_4bit(
+ A: Tensor,
+ B: Tensor,
+ out: Optional[torch.Tensor] = None,
+ transposed_A=False,
+ transposed_B=False,
+ state=None,
+):
+ if state is None:
+ raise ValueError("state cannot be None. gemv_4bit() requires the state from quantize_4bit()")
+
+ absmax = state.absmax
+ if state.nested:
+ absmax = dequantize_blockwise(absmax, state.state2) + state.offset
+
+ if out is not None:
+ torch.ops.bitsandbytes.gemv_4bit.out(
+ A,
+ B,
+ state.shape,
+ absmax,
+ state.code,
+ state.blocksize,
+ out=out,
+ )
+ return out
+
+ return torch.ops.bitsandbytes.gemv_4bit.default(
+ A,
+ B,
+ state.shape,
+ absmax,
+ state.code,
+ state.blocksize,
+ )
+
+
+@deprecated("This function is deprecated and will be removed in a future release.", category=FutureWarning)
+def igemm(
+ A: Tensor,
+ B: Tensor,
+ out: Optional[torch.Tensor] = None,
+ transposed_A=False,
+ transposed_B=False,
+):
+ sout = check_matmul(A, B, out, transposed_A, transposed_B)
+ if out is None:
+ out = torch.zeros(size=sout, dtype=torch.int32, device=A.device)
+ if len(A.shape) == 3 and len(B.shape) == 3:
+ if A.shape[0] == B.shape[0] and A.shape[2] == B.shape[1]:
+ return batched_igemm(A, B, out)
+
+ sA = A.shape
+ sB = B.shape
+ if transposed_A and len(sA) == 2:
+ sA = (sA[1], sA[0])
+ elif transposed_A and len(sA) == 3:
+ sA = (sA[0], sA[2], sA[0])
+ if transposed_B and len(sB) == 2:
+ sB = (sB[1], sB[0])
+ elif transposed_B and len(sB) == 3:
+ sB = (sB[0], sB[2], sB[0])
+ # this is a mess: cuBLAS expect column major, but PyTorch is row major.
+ # So to perform the matrix multiplication, we have to treat A, B, and C matrices
+ # (transpose of row major is column major)
+ # This means we compute B^T A^T = C^T and we explicitly switch the dimensions of each of these
+
+ # matrices in the input arguments for cuBLAS
+ # column major: A @ B = C: [m, k] @ [k, n] = [m, n]
+ # row major: B^T @ A^T = C^T: [m, k] @ [k, n] = [m, n]
+ # column major with row major layout: B^T @ A^T = C^T: [k, m] @ [n, k] = [n, m]
+ if len(sB) == 2:
+ if B.stride()[0] == B.shape[1]:
+ transposed_B = False
+ elif B.stride()[1] == B.shape[0]:
+ transposed_B = True
+ if len(A.shape) == 2:
+ if A.stride()[0] == A.shape[1]:
+ transposed_A = False
+ elif A.stride()[1] == A.shape[0]:
+ transposed_A = True
+ else:
+ if A.stride()[1] == A.shape[2]:
+ transposed_A = False
+ elif A.stride()[2] == A.shape[1]:
+ transposed_A = True
+
+ if len(sA) == 2:
+ n = sA[0]
+ ldb = A.stride()[1 if transposed_A else 0]
+ elif len(sA) == 3 and len(sB) == 2:
+ n = sA[0] * sA[1]
+ ldb = sA[2]
+
+ m = sB[1]
+ k = sB[0]
+ lda = B.stride()[(1 if transposed_B else 0)]
+ ldc = sB[1]
+ elif len(sB) == 3:
+ # special case
+ assert len(sA) == 3
+ if not (sA[0] == sB[0] and sA[1] == sB[1]):
+ raise ValueError(
+ f"Only bsi,bso->io supported for tensor contractions, but dims for A x B were: {sA} x {sB}",
+ )
+
+ transposed_A = True
+ transposed_B = False
+
+ m = sB[2]
+ n = sA[2]
+ k = sB[0] * sB[1]
+
+ lda = m
+ ldb = sA[2]
+ ldc = m
+
+ ptr = CUBLAS_Context.get_instance().get_context(A.device)
+
+ # B^T @ A^T = C^T
+ # [km, nk -> mn]
+ is_on_gpu([B, A, out])
+ lib.cigemm(
+ ptr,
+ ct.c_bool(transposed_B),
+ ct.c_bool(transposed_A),
+ ct.c_int32(m),
+ ct.c_int32(n),
+ ct.c_int32(k),
+ get_ptr(B),
+ get_ptr(A),
+ get_ptr(out),
+ ct.c_int32(lda),
+ ct.c_int32(ldb),
+ ct.c_int32(ldc),
+ )
+ return out
+
+
+@deprecated("This function is deprecated and will be removed in a future release.", category=FutureWarning)
+def batched_igemm(
+ A: Tensor,
+ B: Tensor,
+ out: Optional[torch.Tensor] = None,
+ transposed_A=False,
+ transposed_B=False,
+):
+ if not len(A.shape) == 3 or not len(B.shape) == 3:
+ raise ValueError(f"Expected 3-dimensional tensors for bmm, but got shapes A and B: {A.shape} and {B.shape}")
+ sout = check_matmul(A, B, out, transposed_A, transposed_B)
+ if out is None:
+ out = torch.zeros(size=sout, dtype=torch.int32, device=A.device)
+
+ if B.is_contiguous():
+ lda = B.stride()[1]
+ transposed_A = False
+ else:
+ s = B.stride()
+ if s[0] != B.shape[0]:
+ B = B.contiguous()
+ lda = B.stride()[1]
+ elif s[2] == B.shape[1]:
+ transposed_A = True
+ lda = B.stride()[2]
+ else:
+ if s[2] == 1:
+ B = B.contiguous()
+ lda = B.stride()[1]
+ elif s[1] == 1:
+ B = B.contiguous()
+ lda = B.stride()[1]
+ else:
+ B = B.contiguous()
+ lda = B.stride()[1]
+
+ if A.is_contiguous():
+ ldb = A.stride()[1]
+ transposed_B = False
+ else:
+ s = A.stride()
+ if s[0] != A.shape[0]:
+ A = A.contiguous()
+ ldb = A.stride()[1]
+ transposed_B = False
+ elif s[2] == A.shape[1]:
+ ldb = A.stride()[2]
+ transposed_B = True
+ else:
+ A = A.contiguous()
+ ldb = A.stride()[1]
+ transposed_B = False
+
+ # this is a mess: cuBLAS expect column major, but PyTorch is row major.
+ # So to perform the matrix multiplication, we have to treat A, B, and C matrices
+ # (transpose of row major is column major)
+ # This means we compute B^T A^T = C^T and we explicitly switch the dimensions of each of these
+ # matrices in the input arguments for cuBLAS
+
+ # column major: A @ B = C: [batch, m, k] @ [batch, k, n] = [batch, m, n]
+ # row major: B^T @ A^T = C^T: [batch, m, k] @ [batch, k, n] = [batch, m, n]
+ # column major with row major layout: B^T @ A^T = C^T: [batch, k, m] @ [batch, n, k] = [batch, n, m]
+ num_batch = A.shape[0]
+ n = A.shape[1]
+ m = B.shape[2]
+ k = B.shape[1]
+
+ ldc = m
+
+ strideA = B.shape[1] * B.shape[2]
+ strideB = A.shape[1] * A.shape[2]
+ strideC = A.shape[1] * B.shape[2]
+
+ ptr = CUBLAS_Context.get_instance().get_context(A.device)
+
+ is_on_gpu([B, A, out])
+ lib.cbatched_igemm(
+ ptr,
+ ct.c_bool(transposed_B),
+ ct.c_bool(transposed_A),
+ ct.c_int32(m),
+ ct.c_int32(n),
+ ct.c_int32(k),
+ get_ptr(B),
+ get_ptr(A),
+ get_ptr(out),
+ ct.c_int32(lda),
+ ct.c_int32(ldb),
+ ct.c_int32(ldc),
+ ct.c_long(strideA),
+ ct.c_long(strideB),
+ ct.c_long(strideC),
+ ct.c_uint32(num_batch),
+ )
+ return out
+
+
+def int8_linear_matmul(A: torch.Tensor, B: torch.Tensor, out: Optional[torch.Tensor] = None, dtype=torch.int32):
+ """Performs an 8-bit integer matrix multiplication.
+
+ A linear transformation is applied such that `out = A @ B.T`. When possible, integer tensor core hardware is
+ utilized to accelerate the operation.
+
+ Args:
+ A (`torch.Tensor`): The first matrix operand with the data type `torch.int8`.
+ B (`torch.Tensor`): The second matrix operand with the data type `torch.int8`.
+ out (`torch.Tensor`, *optional*): A pre-allocated tensor used to store the result.
+ dtype (`torch.dtype`, *optional*): The expected data type of the output. Defaults to `torch.int32`.
+
+ Raises:
+ `NotImplementedError`: The operation is not supported in the current environment.
+ `RuntimeError`: Raised when the cannot be completed for any other reason.
+
+ Returns:
+ `torch.Tensor`: The result of the operation.
+ """
+ if out is not None:
+ torch.ops.bitsandbytes.int8_linear_matmul.out(A, B, out)
+ return out
+
+ return torch.ops.bitsandbytes.int8_linear_matmul.default(A, B)
+
+
+def int8_mm_dequant(
+ A: torch.Tensor,
+ row_stats: torch.Tensor,
+ col_stats: torch.Tensor,
+ out: Optional[torch.Tensor] = None,
+ bias: Optional[torch.Tensor] = None,
+):
+ """Performs dequantization on the result of a quantized int8 matrix multiplication.
+
+ Args:
+ A (`torch.Tensor` with dtype `torch.int32`): The result of a quantized int8 matrix multiplication.
+ row_stats (`torch.Tensor`): The row-wise quantization statistics for the lhs operand of the matrix multiplication.
+ col_stats (`torch.Tensor`): The column-wise quantization statistics for the rhs operand of the matrix multiplication.
+ out (`torch.Tensor`, *optional*): A pre-allocated tensor to store the output of the operation.
+ bias (`torch.Tensor`, *optional*): An optional bias vector to add to the result.
+
+ Returns:
+ `torch.Tensor`: The dequantized result with an optional bias, with dtype `torch.float16`.
+ """
+ result = torch.ops.bitsandbytes.int8_mm_dequant.default(A, row_stats, col_stats, dtype=torch.float16, bias=bias)
+
+ # TODO(matthewdouglas): Deprecate out kwarg
+ if out is not None:
+ return out.copy_(result)
+
+ return result
+
+
+def int8_double_quant(
+ A: torch.Tensor,
+ col_stats: Optional[torch.Tensor] = None,
+ row_stats: Optional[torch.Tensor] = None,
+ out_col: Optional[torch.Tensor] = None,
+ out_row: Optional[torch.Tensor] = None,
+ threshold=0.0,
+):
+ """Determine the quantization statistics for input matrix `A` in accordance to the `LLM.int8()` algorithm.
+
+ The statistics are determined both row-wise and column-wise (transposed).
+
+ For more information, see the [LLM.int8() paper](https://arxiv.org/abs/2208.07339).
+
+
+ This function is useful for training, but for inference it is advised to use [`int8_vectorwise_quant`] instead.
+ This implementation performs additional column-wise transposed calculations which are not optimized.
+
+
+ Args:
+ A (`torch.Tensor` with dtype `torch.float16`): The input matrix.
+ col_stats (`torch.Tensor`, *optional*): A pre-allocated tensor to hold the column-wise quantization scales.
+ row_stats (`torch.Tensor`, *optional*): A pre-allocated tensor to hold the row-wise quantization scales.
+ out_col (`torch.Tensor`, *optional*): A pre-allocated tensor to hold the column-wise quantized data.
+ out_row (`torch.Tensor`, *optional*): A pre-allocated tensor to hold the row-wise quantized data.
+ threshold (`float`, *optional*):
+ An optional threshold for sparse decomposition of outlier features.
+
+ No outliers are held back when 0.0. Defaults to 0.0.
+
+ Returns:
+ `Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]`: A tuple containing the quantized tensor and relevant statistics.
+ - `torch.Tensor` with dtype `torch.int8`: The row-wise quantized data.
+ - `torch.Tensor` with dtype `torch.int8`: The column-wise quantized data.
+ - `torch.Tensor` with dtype `torch.float32`: The row-wise quantization scales.
+ - `torch.Tensor` with dtype `torch.float32`: The column-wise quantization scales.
+ - `torch.Tensor` with dtype `torch.int32`, *optional*: A list of column indices which contain outlier features.
+ """
+
+ if row_stats is not None:
+ raise ValueError("row_stats must be None. int8_double_quant() does not support pre-allocated row_stats.")
+ if col_stats is not None:
+ raise ValueError("col_stats must be None. int8_double_quant() does not support pre-allocated col_stats.")
+ if out_col is not None:
+ raise ValueError("out_col must be None. int8_double_quant() does not support pre-allocated out_col.")
+ if out_row is not None:
+ raise ValueError("out_row must be None. int8_double_quant() does not support pre-allocated out_row.")
+
+ return torch.ops.bitsandbytes.int8_double_quant.default(A, threshold=threshold)
+
+
+def int8_vectorwise_dequant(A: torch.Tensor, stats: torch.Tensor):
+ """Dequantizes a tensor with dtype `torch.int8` to `torch.float32`.
+
+ Args:
+ A (`torch.Tensor` with dtype `torch.int8`): The quantized int8 tensor.
+ stats (`torch.Tensor` with dtype `torch.float32`): The row-wise quantization statistics.
+
+ Returns:
+ `torch.Tensor` with dtype `torch.float32`: The dequantized tensor.
+ """
+ # To dequantize we divide by 127, or multiply by the reciprocal.
+ return torch.ops.bitsandbytes.int8_vectorwise_dequant.default(A, stats)
+
+
+def int8_vectorwise_quant(A: torch.Tensor, threshold=0.0):
+ """Quantizes a tensor with dtype `torch.float16` to `torch.int8` in accordance to the `LLM.int8()` algorithm.
+
+ For more information, see the [LLM.int8() paper](https://arxiv.org/abs/2208.07339).
+
+ Args:
+ A (`torch.Tensor` with dtype `torch.float16`): The input tensor.
+ threshold (`float`, *optional*):
+ An optional threshold for sparse decomposition of outlier features.
+
+ No outliers are held back when 0.0. Defaults to 0.0.
+
+ Returns:
+ `Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]`: A tuple containing the quantized tensor and relevant statistics.
+ - `torch.Tensor` with dtype `torch.int8`: The quantized data.
+ - `torch.Tensor` with dtype `torch.float32`: The quantization scales.
+ - `torch.Tensor` with dtype `torch.int32`, *optional*: A list of column indices which contain outlier features.
+ """
+ return torch.ops.bitsandbytes.int8_vectorwise_quant.default(A, threshold)
+
+
+def _convert_weight_packed_for_cpu(qweight: torch.Tensor, quant_state: QuantState, block_n: int = 32):
+ """
+ qweight: (K * N / 2) uint8
+ return: packed_weight
+ """
+ if qweight.dtype != torch.uint8:
+ quant_state.original_storage_type = qweight.dtype
+ qweight = qweight.view(torch.uint8)
+ quant_state.original_dtype = quant_state.dtype
+ quant_state.original_nested = quant_state.nested
+ quant_state.original_qshape = qweight.shape
+
+ qweight = qweight.reshape(-1)
+ unpacked_w = torch.empty(qweight.shape[0] * 2, dtype=torch.int32, device=qweight.device)
+ unpacked_w[1::2] = qweight & 0xF
+ unpacked_w[::2] = qweight >> 4
+ qweight_final = unpacked_w.reshape(quant_state.shape).to(torch.uint8) # (*, N, K)
+ # pack weight: [*, N, K] -> [*, N, K/2] combine low and high bit
+ assert len(qweight_final.shape) == 2
+ N, K = qweight_final.shape[0], qweight_final.shape[1]
+ assert N % block_n == 0, "N must be divisible by block_n"
+ assert K % 2 == 0, "K must be even"
+ BLOCK_N = block_n
+ BIT_COUNT = 32 # (=32 low +32 high)
+ new_shape = [N // BLOCK_N, BLOCK_N, K // 2, 2]
+ out_shape = [N, K // 2]
+ qw = qweight_final.reshape(new_shape) # (..., N/B, B, K/2, 2)
+ qw = qw.transpose(-3, -2).contiguous() # (..., N/B, K/2, B, 2)
+ qw = qw.reshape(-1, BIT_COUNT * 2) # [-1, 64]
+ high = qw[:, BIT_COUNT:] # high 32
+ low = qw[:, :BIT_COUNT] # low 32
+ packed = ((high << 4) | low).to(torch.uint8) # combine
+ final_qweight = packed.reshape(out_shape)
+ if quant_state.nested:
+ absmax = dequantize_blockwise(quant_state.absmax, quant_state.state2)
+ absmax += quant_state.offset
+ if absmax.dtype != torch.float32:
+ absmax = absmax.float()
+
+ quant_state.absmax = absmax
+ quant_state.nested = False
+ delattr(quant_state, "state2")
+
+ quant_state.absmax = (
+ quant_state.absmax.reshape(quant_state.shape[0], quant_state.shape[1] // quant_state.blocksize)
+ .T.to(torch.bfloat16)
+ .contiguous()
+ )
+
+ quant_state.dtype = torch.bfloat16
+ quant_state.packing_format_for_cpu = True
+ return final_qweight, quant_state
+
+
+def _convert_weight_packed_for_cpu_inverse(
+ packed_weight: torch.Tensor,
+ quant_state: QuantState,
+ block_n: int = 32,
+) -> tuple[torch.Tensor, QuantState]:
+ """
+ packed_weight: [N, K/2] uint8, output of `_convert_weight_packed_for_cpu` (final_qweight)
+ quant_state: QuantState that was modified by `_convert_weight_packed_for_cpu`
+ Returns:
+ qweight: [*, N, K] uint8, original qweight shape (quant_state.shape)
+ recovered_state: QuantState with partially restored fields (best-effort inverse)
+ """
+ assert quant_state.packing_format_for_cpu, "only for packing format"
+ assert packed_weight.dtype == torch.uint8
+ assert len(packed_weight.shape) == 2, "packed_weight should be [N, K/2]"
+ N, K_half = packed_weight.shape
+ K = K_half * 2
+
+ # 1) packed [N, K/2] -> [N//BLOCK_N, BLOCK_N, K/2, 2]
+ BLOCK_N = block_n
+ BIT_COUNT = 32 # (=32 low + 32 high)
+
+ assert N % BLOCK_N == 0, "N must be divisible by block_n"
+ assert K % 2 == 0, "K must be even"
+
+ # [N, K/2] -> [-1, 64] (32 low + 32 high)
+ packed = packed_weight.reshape(-1, BIT_COUNT) # [-1, 64]
+ # split high/low nibbles
+ high = (packed >> 4) & 0xF
+ low = packed & 0xF
+ # concatenate to [..., 64], first 32 are low, last 32 are high
+ qw = torch.cat([low, high], dim=-1).to(torch.uint8) # [..., 64]
+
+ # -> [N/BLOCK_N, K/2, BLOCK_N, 2] -> [N, K]
+ qw = qw.reshape(N // BLOCK_N, K_half, BLOCK_N, 2) # [N/B, K/2, B, 2]
+ qw = qw.transpose(-3, -2).contiguous() # [N/B, B, K/2, 2]
+ qw = qw.reshape(N, K) # [N, K]
+
+ qweight = qw # [N, K]
+
+ unpacked_w = qweight.reshape(-1).to(torch.int32) # [K*N]
+ high4 = (unpacked_w[::2] & 0xF).to(torch.uint8)
+ low4 = (unpacked_w[1::2] & 0xF).to(torch.uint8)
+ qweight = (high4 << 4) | low4 # [K*N/2]
+
+ # 2) Best-effort restore of quant_state fields (absmax / dtype / nested flags, etc.)
+ recovered_state = quant_state
+ qweight = qweight.to(torch.uint8).reshape(recovered_state.original_qshape)
+
+ # quantize absmax
+ if recovered_state.original_nested:
+ absmax = recovered_state.absmax.T.reshape(-1).to(recovered_state.original_dtype)
+ offset = absmax.mean()
+ qabsmax, state2 = quantize_blockwise(absmax - offset, blocksize=256)
+ recovered_state.absmax = qabsmax
+ recovered_state.offset = offset
+ recovered_state.state2 = state2
+ recovered_state.nested = True
+
+ recovered_state.dtype = recovered_state.original_dtype
+ recovered_state.packing_format_for_cpu = False
+
+ if getattr(recovered_state, "original_storage_type", None):
+ qweight = qweight.view(recovered_state.original_storage_type)
+
+ return qweight, recovered_state
+
+
+def has_avx512bf16():
+ """
+ Try calling native lib.has_avx512bf16_cpu().
+ Return False explicitly if symbol missing or call fails.
+ """
+ try:
+ support_avx_bf16 = lib.has_avx512bf16_cpu()
+ except (AttributeError, RuntimeError, OSError):
+ support_avx_bf16 = False
+ return support_avx_bf16
+
+
+C = 127.0
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cpu.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cpu.so
new file mode 100644
index 0000000000000000000000000000000000000000..3cd79d6652ef77793170ffc5dbf79e0588a42212
Binary files /dev/null and b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cpu.so differ
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda118.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda118.so
new file mode 100644
index 0000000000000000000000000000000000000000..8e1c61e0aacf78cbd851992e4d5243045c5cc965
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda118.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:cabbf87a1fbb8f9e70acacd6651f4e19f326f9e3ba555b491a8415ea630b6f51
+size 21879184
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda121.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda121.so
new file mode 100644
index 0000000000000000000000000000000000000000..ed5c8224a24985966dab4e147d1561a3e1467a0c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda121.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:9b1b7457d976b4709eab7f4363b9c9719bf66b535fd3283f98beccbebc06b9ed
+size 21580176
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda124.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda124.so
new file mode 100644
index 0000000000000000000000000000000000000000..39d279025db49950e1307c2f041b0bc1e4420c0c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda124.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:4ecee5956fbff357246c441ac88a15a03c5a918be97e21ec51afdb89725b3861
+size 21294432
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda126.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda126.so
new file mode 100644
index 0000000000000000000000000000000000000000..3704b965a7cfb4536ccea3bc7e664d2d287a0db5
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda126.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:6483fc97012c1025a4b7831081c95b0b00bebac9baf2b26ea5b63cf4a0a4e827
+size 21382936
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda128.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda128.so
new file mode 100644
index 0000000000000000000000000000000000000000..be602325a2ec47da868895860091bed5d671edcd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda128.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:d1a039f3caee90f95a7e46e95a8f4c9714ba57084b17d340d5b6710be70f224a
+size 26535848
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda130.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda130.so
new file mode 100644
index 0000000000000000000000000000000000000000..8175560e539ac61cc58e85e37691cc01cfd4e373
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda130.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:66859ccc66523b15069396a6540777366a437748dab0801dac320cfd6bd5e62a
+size 3901168
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda132.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda132.so
new file mode 100644
index 0000000000000000000000000000000000000000..6f4951f23a5abff947a12b2e0bebffd9743490bd
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_cuda132.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:d2c59fff929bed5695265c397e15eabd9e8be4f2eb158b1e4bf23b78f2408bbe
+size 3984920
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm64.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm64.so
new file mode 100644
index 0000000000000000000000000000000000000000..4b2251261aaf564e764162c1f4a0b15315ccc05c
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm64.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:47a725d5de697f7cd08fca5e0fa93a7dca76ac979ba275148f877a4cfe92ee03
+size 894216
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm70.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm70.so
new file mode 100644
index 0000000000000000000000000000000000000000..b6148114573374cc8c9a8156bb73953fa690ef31
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm70.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:8889f82aa21d013a886c869eebf1c6e79135e218f9fca896e1b8ecaf58c23398
+size 922896
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm71.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm71.so
new file mode 100644
index 0000000000000000000000000000000000000000..1007ff943a10e4ee93a04a754b52ef2102c19b7a
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm71.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:74fda185cd68ee3610abf09b14b8b6bbb66a45a5de9d7216c0ab06dcbe20b71f
+size 922888
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm714.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm714.so
new file mode 100644
index 0000000000000000000000000000000000000000..08d506ed7a88ad8fc8329f626c18a7b74c35ad47
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm714.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:977a27ebfdfe54138de2ab84de8dc3c35b4692afd56590d0de1c7cf4a7ceb72f
+size 1164704
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm72.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm72.so
new file mode 100644
index 0000000000000000000000000000000000000000..cee5e9cb6a505230f39b80aa3d476e5c4cedb068
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_rocm72.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:b00e4f535d136ef2412ef006088b62266b67308e4dee5cc4505a1dc90d0b8362
+size 972048
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2025.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2025.so
new file mode 100644
index 0000000000000000000000000000000000000000..e5b8e0a5d5559203ad0754cd6848c1bb746d0775
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2025.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:8d92324d5bf4cd682ae1dc48d462a54cf1f046866bff162f99ef7493d8f44a4b
+size 320024
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2026.so b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2026.so
new file mode 100644
index 0000000000000000000000000000000000000000..fbf49a08cbf266bbe19a373cd3ec8e25bc7cdc0f
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/libbitsandbytes_xpu2026.so
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:99f9d4adf046d9530b3483a0bbfd07eadd2a198760f839cb9c6f7913099e5838
+size 303496
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/nn/__init__.py b/venv/lib/python3.11/site-packages/bitsandbytes/nn/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..54c2614bdd981e57165cd267015e38081c46ffe6
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/nn/__init__.py
@@ -0,0 +1,19 @@
+# Copyright (c) Facebook, Inc. and its affiliates.
+#
+# This source code is licensed under the MIT license found in the
+# LICENSE file in the root directory of this source tree.
+from .modules import (
+ Embedding,
+ Embedding4bit,
+ Embedding8bit,
+ EmbeddingFP4,
+ EmbeddingNF4,
+ Int8Params,
+ Linear4bit,
+ Linear8bitLt,
+ LinearFP4,
+ LinearNF4,
+ OutlierAwareLinear,
+ Params4bit,
+ StableEmbedding,
+)
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/nn/modules.py b/venv/lib/python3.11/site-packages/bitsandbytes/nn/modules.py
new file mode 100644
index 0000000000000000000000000000000000000000..ebc0b09439c4620ef015b18b242328aa65f14c18
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/nn/modules.py
@@ -0,0 +1,1220 @@
+# Copyright (c) Facebook, Inc. and its affiliates.
+#
+# This source code is licensed under the MIT license found in the
+# LICENSE file in the root directory of this source tree.
+import copy
+import logging
+from typing import Any, Optional, TypeVar, Union, overload
+
+import torch
+from torch import Tensor, device, dtype, nn
+import torch.nn.functional as F
+
+import bitsandbytes as bnb
+from bitsandbytes.functional import (
+ QuantState,
+ _convert_weight_packed_for_cpu,
+ _convert_weight_packed_for_cpu_inverse,
+ has_avx512bf16,
+)
+from bitsandbytes.optim import GlobalOptimManager
+from bitsandbytes.utils import INVERSE_LINEAR_8BIT_WEIGHTS_FORMAT_MAPPING, OutlierTracer
+
+logger = logging.getLogger(__name__)
+
+T = TypeVar("T", bound="torch.nn.Module")
+
+
+class StableEmbedding(torch.nn.Embedding):
+ """
+ Custom embedding layer designed to improve stability during training for NLP tasks by using 32-bit optimizer states. It is designed to reduce gradient variations that can result from quantization. This embedding layer is initialized with Xavier uniform initialization followed by layer normalization.
+
+ Example:
+
+ ```
+ # Initialize StableEmbedding layer with vocabulary size 1000, embedding dimension 300
+ embedding_layer = StableEmbedding(num_embeddings=1000, embedding_dim=300)
+
+ # Reset embedding parameters
+ embedding_layer.reset_parameters()
+
+ # Perform a forward pass with input tensor
+ input_tensor = torch.tensor([1, 2, 3])
+ output_embedding = embedding_layer(input_tensor)
+ ```
+
+ Attributes:
+ norm (`torch.nn.LayerNorm`): Layer normalization applied after the embedding.
+
+ Methods:
+ reset_parameters(): Reset embedding parameters using Xavier uniform initialization.
+ forward(input: Tensor) -> Tensor: Forward pass through the stable embedding layer.
+ """
+
+ def __init__(
+ self,
+ num_embeddings: int,
+ embedding_dim: int,
+ padding_idx: Optional[int] = None,
+ max_norm: Optional[float] = None,
+ norm_type: float = 2.0,
+ scale_grad_by_freq: bool = False,
+ sparse: bool = False,
+ _weight: Optional[Tensor] = None,
+ device=None,
+ dtype=None,
+ ) -> None:
+ """
+ Args:
+ num_embeddings (`int`):
+ The number of unique embeddings (vocabulary size).
+ embedding_dim (`int`):
+ The dimensionality of the embedding.
+ padding_idx (`Optional[int]`):
+ Pads the output with zeros at the given index.
+ max_norm (`Optional[float]`):
+ Renormalizes embeddings to have a maximum L2 norm.
+ norm_type (`float`, defaults to `2.0`):
+ The p-norm to compute for the `max_norm` option.
+ scale_grad_by_freq (`bool`, defaults to `False`):
+ Scale gradient by frequency during backpropagation.
+ sparse (`bool`, defaults to `False`):
+ Computes dense gradients. Set to `True` to compute sparse gradients instead.
+ _weight (`Optional[Tensor]`):
+ Pretrained embeddings.
+ """
+ super().__init__(
+ num_embeddings,
+ embedding_dim,
+ padding_idx,
+ max_norm,
+ norm_type,
+ scale_grad_by_freq,
+ sparse,
+ _weight,
+ device,
+ dtype,
+ )
+ self.norm = torch.nn.LayerNorm(embedding_dim, device=device)
+ GlobalOptimManager.get_instance().register_module_override(self, "weight", {"optim_bits": 32})
+
+ def reset_parameters(self) -> None:
+ torch.nn.init.xavier_uniform_(self.weight)
+ self._fill_padding_idx_with_zero()
+
+ """ !!! This is a redefinition of _fill_padding_idx_with_zero in torch.nn.Embedding
+ to make the Layer compatible with Pytorch < 1.9.
+ This means that if this changes in future PyTorch releases this need to change too
+ which is cumbersome. However, with this we can ensure compatibility with previous
+ PyTorch releases.
+ """
+
+ def _fill_padding_idx_with_zero(self) -> None:
+ if self.padding_idx is not None:
+ with torch.no_grad():
+ self.weight[self.padding_idx].fill_(0)
+
+ def forward(self, input: Tensor) -> Tensor:
+ emb = F.embedding(
+ input,
+ self.weight,
+ self.padding_idx,
+ self.max_norm,
+ self.norm_type,
+ self.scale_grad_by_freq,
+ self.sparse,
+ )
+
+ # always apply layer norm in full precision
+ emb = emb.to(torch.get_default_dtype())
+
+ return self.norm(emb).to(self.weight.dtype)
+
+
+class Embedding(torch.nn.Embedding):
+ """
+ Embedding class to store and retrieve word embeddings from their indices.
+ """
+
+ def __init__(
+ self,
+ num_embeddings: int,
+ embedding_dim: int,
+ padding_idx: Optional[int] = None,
+ max_norm: Optional[float] = None,
+ norm_type: float = 2.0,
+ scale_grad_by_freq: bool = False,
+ sparse: bool = False,
+ _weight: Optional[Tensor] = None,
+ device: Optional[device] = None,
+ ) -> None:
+ """
+ Args:
+ num_embeddings (`int`):
+ The number of unique embeddings (vocabulary size).
+ embedding_dim (`int`):
+ The dimensionality of the embedding.
+ padding_idx (`Optional[int]`):
+ Pads the output with zeros at the given index.
+ max_norm (`Optional[float]`):
+ Renormalizes embeddings to have a maximum L2 norm.
+ norm_type (`float`, defaults to `2.0`):
+ The p-norm to compute for the `max_norm` option.
+ scale_grad_by_freq (`bool`, defaults to `False`):
+ Scale gradient by frequency during backpropagation.
+ sparse (`bool`, defaults to `False`):
+ Computes dense gradients. Set to `True` to compute sparse gradients instead.
+ _weight (`Optional[Tensor]`):
+ Pretrained embeddings.
+ """
+ super().__init__(
+ num_embeddings,
+ embedding_dim,
+ padding_idx,
+ max_norm,
+ norm_type,
+ scale_grad_by_freq,
+ sparse,
+ _weight,
+ device=device,
+ )
+ GlobalOptimManager.get_instance().register_module_override(self, "weight", {"optim_bits": 32})
+
+ def reset_parameters(self) -> None:
+ torch.nn.init.xavier_uniform_(self.weight)
+ self._fill_padding_idx_with_zero()
+
+ """ !!! This is a redefinition of _fill_padding_idx_with_zero in torch.nn.Embedding
+ to make the Layer compatible with Pytorch < 1.9.
+ This means that if this changes in future PyTorch releases this need to change too
+ which is cumbersome. However, with this we can ensure compatibility with previous
+ PyTorch releases.
+ """
+
+ def _fill_padding_idx_with_zero(self) -> None:
+ if self.padding_idx is not None:
+ with torch.no_grad():
+ self.weight[self.padding_idx].fill_(0)
+
+ def forward(self, input: Tensor) -> Tensor:
+ emb = F.embedding(
+ input,
+ self.weight,
+ self.padding_idx,
+ self.max_norm,
+ self.norm_type,
+ self.scale_grad_by_freq,
+ self.sparse,
+ )
+
+ return emb
+
+
+class Params4bit(torch.nn.Parameter):
+ def __new__(
+ cls,
+ data: Optional[torch.Tensor] = None,
+ requires_grad=False, # quantized weights should be frozen by default
+ quant_state: Optional[QuantState] = None,
+ blocksize: Optional[int] = None,
+ compress_statistics: bool = True,
+ quant_type: str = "fp4",
+ quant_storage: torch.dtype = torch.uint8,
+ module: Optional["Linear4bit"] = None,
+ bnb_quantized: bool = False,
+ **kwargs,
+ ) -> "Params4bit":
+ if data is None:
+ data = torch.empty(0)
+
+ if blocksize is None:
+ blocksize = 64
+
+ self = torch.Tensor._make_subclass(cls, data, requires_grad)
+ self.blocksize = blocksize
+ self.compress_statistics = compress_statistics
+ self.quant_type = quant_type
+ self.quant_state = quant_state
+ self.quant_storage = quant_storage
+ self.bnb_quantized = bnb_quantized
+ self.data = data
+ self.module = module
+ return self
+
+ def __getstate__(self):
+ state = self.__dict__.copy()
+ state["data"] = self.data
+ state["requires_grad"] = self.requires_grad
+ return state
+
+ def __setstate__(self, state):
+ self.requires_grad = state["requires_grad"]
+ self.blocksize = state["blocksize"]
+ self.compress_statistics = state["compress_statistics"]
+ self.quant_type = state["quant_type"]
+ self.quant_state = state["quant_state"]
+ self.data = state["data"]
+ self.quant_storage = state["quant_storage"]
+ self.bnb_quantized = state["bnb_quantized"]
+ self.module = state["module"]
+
+ # Properties that proxy QuantState attributes for FSDP state_dict traversal.
+ # FSDP's _get_fqns() resolves dotted FQN keys via getattr, e.g. "weight.absmax"
+ # becomes getattr(weight, "absmax"). Using @property instead of __getattr__
+ # avoids torch.compile graph breaks (see #1904), since Dynamo can trace
+ # descriptor protocol access but not __getattr__ on Tensor subclasses.
+ #
+ # Note: attributes that collide with Params4bit instance attrs (blocksize,
+ # quant_type) or Tensor attrs (dtype, shape) are intentionally omitted —
+ # they are packed into the bitsandbytes__* blob and not traversed by FSDP.
+
+ @property
+ def absmax(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None:
+ return qs.absmax
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'absmax'")
+
+ @property
+ def code(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None:
+ return qs.code
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'code'")
+
+ @property
+ def quant_map(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None:
+ return qs.code
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'quant_map'")
+
+ @property
+ def offset(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None:
+ return qs.offset
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'offset'")
+
+ @property
+ def state2(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None:
+ return qs.state2
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'state2'")
+
+ @property
+ def nested_absmax(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None and qs.state2 is not None:
+ return qs.state2.absmax
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'nested_absmax'")
+
+ @property
+ def nested_blocksize(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None and qs.state2 is not None:
+ return qs.state2.blocksize
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'nested_blocksize'")
+
+ @property
+ def nested_quant_map(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None and qs.state2 is not None:
+ return qs.state2.code
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'nested_quant_map'")
+
+ @property
+ def nested_dtype(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None and qs.state2 is not None:
+ return qs.state2.dtype
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'nested_dtype'")
+
+ @property
+ def nested_offset(self):
+ qs = self.__dict__.get("quant_state")
+ if qs is not None:
+ return qs.offset
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute 'nested_offset'")
+
+ def __deepcopy__(self, memo):
+ new_instance = type(self).__new__(type(self))
+ state = self.__getstate__()
+ new_instance.__setstate__(state)
+ new_instance.quant_state = copy.deepcopy(state["quant_state"])
+ new_instance.data = copy.deepcopy(state["data"])
+ return new_instance
+
+ def __copy__(self):
+ new_instance = type(self).__new__(type(self))
+ state = self.__getstate__()
+ new_instance.__setstate__(state)
+ return new_instance
+
+ @classmethod
+ def from_prequantized(
+ cls,
+ data: torch.Tensor,
+ quantized_stats: dict[str, Any],
+ requires_grad: bool = False,
+ device="cuda",
+ module: Optional["Linear4bit"] = None,
+ **kwargs,
+ ) -> "Params4bit":
+ self = torch.Tensor._make_subclass(cls, data.to(device))
+ self.requires_grad = requires_grad
+ self.quant_state = QuantState.from_dict(qs_dict=quantized_stats, device=device)
+ self.blocksize = self.quant_state.blocksize
+ self.compress_statistics = self.quant_state.nested
+ self.quant_type = self.quant_state.quant_type
+ self.bnb_quantized = True
+
+ self.quant_storage = data.dtype
+ self.module = module
+
+ if self.module is not None:
+ self.module.quant_state = self.quant_state
+
+ return self
+
+ def _quantize(self, device):
+ w = self.data.contiguous().to(device)
+ w_4bit, quant_state = bnb.functional.quantize_4bit(
+ w,
+ blocksize=self.blocksize,
+ compress_statistics=self.compress_statistics,
+ quant_type=self.quant_type,
+ quant_storage=self.quant_storage,
+ )
+ self.data = w_4bit
+ self.quant_state = quant_state
+ if self.module is not None:
+ self.module.quant_state = quant_state
+ self.bnb_quantized = True
+ return self
+
+ def cpu(self):
+ return self.to(device="cpu")
+
+ def cuda(self, device: Optional[int | device | str] = None, non_blocking: bool = False):
+ if getattr(self.quant_state, "packing_format_for_cpu", False):
+ self.data, self.quant_state = _convert_weight_packed_for_cpu_inverse(self.data, self.quant_state)
+ return self.to(device="cuda" if device is None else device, non_blocking=non_blocking)
+
+ def xpu(self, device: Optional[int | device | str] = None, non_blocking: bool = False):
+ if getattr(self.quant_state, "packing_format_for_cpu", False):
+ self.data, self.quant_state = _convert_weight_packed_for_cpu_inverse(self.data, self.quant_state)
+ return self.to(device="xpu" if device is None else device, non_blocking=non_blocking)
+
+ @overload
+ def to(
+ self: T,
+ device: Optional[int | device] = ...,
+ dtype: Optional[dtype | str] = ...,
+ non_blocking: bool = ...,
+ ) -> T: ...
+
+ @overload
+ def to(self: T, dtype: dtype | str, non_blocking: bool = ...) -> T: ...
+
+ @overload
+ def to(self: T, tensor: Tensor, non_blocking: bool = ...) -> T: ...
+
+ def to(self, *args, **kwargs):
+ device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs)
+
+ if device is not None and device.type != "meta" and not self.bnb_quantized:
+ return self._quantize(device)
+ else:
+ if self.quant_state is not None:
+ self.quant_state.to(device)
+
+ new_param = Params4bit(
+ super().to(device=device, dtype=dtype, non_blocking=non_blocking),
+ requires_grad=self.requires_grad,
+ quant_state=self.quant_state,
+ blocksize=self.blocksize,
+ compress_statistics=self.compress_statistics,
+ quant_type=self.quant_type,
+ quant_storage=self.quant_storage,
+ bnb_quantized=self.bnb_quantized,
+ )
+
+ return new_param
+
+ @classmethod
+ def __torch_function__(cls, func, types, args=(), kwargs=None):
+ if kwargs is None:
+ kwargs = {}
+
+ if func in [torch.chunk, torch.split]:
+ tensor = args[0]
+
+ result = super().__torch_function__(func, types, args, kwargs)
+
+ if isinstance(result, tuple):
+ return tuple(
+ cls(
+ data=chunk,
+ requires_grad=tensor.requires_grad,
+ quant_state=tensor.quant_state,
+ blocksize=tensor.blocksize,
+ compress_statistics=tensor.compress_statistics,
+ quant_type=tensor.quant_type,
+ quant_storage=tensor.quant_storage,
+ module=tensor.module,
+ bnb_quantized=tensor.bnb_quantized,
+ )
+ for chunk in result
+ )
+ else:
+ return cls(
+ data=result,
+ requires_grad=tensor.requires_grad,
+ quant_state=tensor.quant_state,
+ blocksize=tensor.blocksize,
+ compress_statistics=tensor.compress_statistics,
+ quant_type=tensor.quant_type,
+ quant_storage=tensor.quant_storage,
+ module=tensor.module,
+ bnb_quantized=tensor.bnb_quantized,
+ )
+
+ return super().__torch_function__(func, types, args, kwargs)
+
+
+def fix_4bit_weight_quant_state_from_module(module: Union["Embedding4bit", "Linear4bit"]):
+ if getattr(module.weight, "quant_state", None) is not None:
+ return
+
+ if getattr(module, "quant_state", None) is None:
+ logger.warning(
+ "FP4 quantization state not initialized. Please call .cuda() or .to(device) on the LinearFP4 layer first.",
+ )
+
+ # the quant state got lost when the parameter got converted. This happens for example for fsdp
+ # since we registered the module, we can recover the state here
+ assert module.weight.shape[1] == 1
+ if not isinstance(module.weight, Params4bit):
+ module.weight = Params4bit(module.weight, quant_storage=module.quant_storage, bnb_quantized=True)
+ module.weight.quant_state = module.quant_state
+
+
+class Linear4bit(nn.Linear):
+ """
+ This class is the base module for the 4-bit quantization algorithm presented in [QLoRA](https://arxiv.org/abs/2305.14314).
+ QLoRA 4-bit linear layers uses blockwise k-bit quantization under the hood, with the possibility of selecting various
+ compute datatypes such as FP4 and NF4.
+
+ In order to quantize a linear layer one should first load the original fp16 / bf16 weights into
+ the Linear4bit module, then call `quantized_module.to("cuda")` to quantize the fp16 / bf16 weights.
+
+ Example:
+
+ ```python
+ import torch
+ import torch.nn as nn
+
+ import bitsandbytes as bnb
+ from bitsandbytes.nn import Linear4bit
+
+ fp16_model = nn.Sequential(
+ nn.Linear(64, 64),
+ nn.Linear(64, 64)
+ )
+
+ quantized_model = nn.Sequential(
+ Linear4bit(64, 64),
+ Linear4bit(64, 64)
+ )
+
+ quantized_model.load_state_dict(fp16_model.state_dict())
+ quantized_model = quantized_model.to(0) # Quantization happens here
+ ```
+ """
+
+ def __init__(
+ self,
+ input_features,
+ output_features,
+ bias=True,
+ compute_dtype=None,
+ compress_statistics=True,
+ quant_type="fp4",
+ quant_storage=torch.uint8,
+ device=None,
+ ):
+ """
+ Initialize Linear4bit class.
+
+ Args:
+ input_features (`str`):
+ Number of input features of the linear layer.
+ output_features (`str`):
+ Number of output features of the linear layer.
+ bias (`bool`, defaults to `True`):
+ Whether the linear class uses the bias term as well.
+ """
+ super().__init__(input_features, output_features, bias, device)
+ self.weight = Params4bit(
+ self.weight.data,
+ requires_grad=False,
+ compress_statistics=compress_statistics,
+ quant_type=quant_type,
+ quant_storage=quant_storage,
+ module=self,
+ )
+ # self.persistent_buffers = [] # TODO consider as way to save quant state
+ self.compute_dtype = compute_dtype
+ self.compute_type_is_set = compute_dtype is not None
+ self.quant_state = None
+ self.quant_storage = quant_storage
+ self.support_avx512bf16_for_cpu = has_avx512bf16()
+
+ def set_compute_type(self, x):
+ if x.dtype in [torch.float32, torch.bfloat16]:
+ # the input is in a dtype that is safe to compute in, we switch
+ # to this type for speed and stability
+ self.compute_dtype = x.dtype
+ elif x.dtype == torch.float16:
+ # we take the compoute dtype passed into the layer
+ if self.compute_dtype in [None, torch.float32] and (x.numel() == x.shape[-1]):
+ # single batch inference with input torch.float16 and compute_dtype float32 -> slow inference when it could be fast
+ # warn the user about this
+ logger.warning(
+ "Input type into Linear4bit is torch.float16, but bnb_4bit_compute_dtype=torch.float32 (default). This will lead to slow inference.",
+ )
+ if self.compute_dtype in [None, torch.float32] and (x.numel() != x.shape[-1]):
+ logger.warning(
+ "Input type into Linear4bit is torch.float16, but bnb_4bit_compute_dtype=torch.float32 (default). This will lead to slow inference or training speed.",
+ )
+
+ def _save_to_state_dict(self, destination, prefix, keep_vars):
+ """
+ save weight and bias,
+ then fill state_dict with components of quant_state
+ """
+ if getattr(self.weight, "quant_state", None) is not None and getattr(
+ self.weight.quant_state, "packing_format_for_cpu", False
+ ):
+ self.weight.data, self.weight.quant_state = _convert_weight_packed_for_cpu_inverse(
+ self.weight.data, self.weight.quant_state
+ )
+ super()._save_to_state_dict(destination, prefix, keep_vars) # saving weight and bias
+ if getattr(self.weight, "quant_state", None) is not None:
+ for k, v in self.weight.quant_state.as_dict(packed=True).items():
+ destination[prefix + "weight." + k] = v if keep_vars else v.detach()
+
+ def forward(self, x: torch.Tensor):
+ fix_4bit_weight_quant_state_from_module(self)
+ quant_state = self.weight.quant_state
+
+ if (
+ x.device.type == "cpu"
+ and self.support_avx512bf16_for_cpu
+ and not self.training
+ and x.requires_grad == False
+ and not getattr(quant_state, "packing_format_for_cpu", False)
+ ):
+ self.weight.data, quant_state = _convert_weight_packed_for_cpu(self.weight.data, quant_state)
+
+ if not self.compute_type_is_set:
+ self.set_compute_type(x)
+ self.compute_type_is_set = True
+
+ inp_dtype = x.dtype
+ if self.compute_dtype is not None:
+ x = x.to(self.compute_dtype)
+
+ bias = self.bias
+ if bias is not None:
+ if bias.dtype != x.dtype:
+ # TODO: do we need to cast bias like this?
+ bias.data = bias.data.to(x.dtype)
+ bias = bias.to(self.compute_dtype)
+
+ return bnb.matmul_4bit(x, self.weight, bias=bias, quant_state=quant_state).to(inp_dtype)
+
+
+class LinearFP4(Linear4bit):
+ """
+ Implements the FP4 data type.
+ """
+
+ def __init__(
+ self,
+ input_features,
+ output_features,
+ bias=True,
+ compute_dtype=None,
+ compress_statistics=True,
+ quant_storage=torch.uint8,
+ device=None,
+ ):
+ """
+ Args:
+ input_features (`str`):
+ Number of input features of the linear layer.
+ output_features (`str`):
+ Number of output features of the linear layer.
+ bias (`bool`, defaults to `True`):
+ Whether the linear class uses the bias term as well.
+ """
+ super().__init__(
+ input_features,
+ output_features,
+ bias,
+ compute_dtype,
+ compress_statistics,
+ "fp4",
+ quant_storage,
+ device,
+ )
+
+
+class LinearNF4(Linear4bit):
+ """Implements the NF4 data type.
+
+ Constructs a quantization data type where each bin has equal area under a standard normal distribution N(0, 1) that
+ is normalized into the range [-1, 1].
+
+ For more information read the paper: QLoRA: Efficient Finetuning of Quantized LLMs (https://arxiv.org/abs/2305.14314)
+
+ Implementation of the NF4 data type in bitsandbytes can be found in the `create_normal_map` function in
+ the `functional.py` file: https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L236.
+ """
+
+ def __init__(
+ self,
+ input_features,
+ output_features,
+ bias=True,
+ compute_dtype=None,
+ compress_statistics=True,
+ quant_storage=torch.uint8,
+ device=None,
+ ):
+ """
+ Args:
+ input_features (`str`):
+ Number of input features of the linear layer.
+ output_features (`str`):
+ Number of output features of the linear layer.
+ bias (`bool`, defaults to `True`):
+ Whether the linear class uses the bias term as well.
+ """
+ super().__init__(
+ input_features,
+ output_features,
+ bias,
+ compute_dtype,
+ compress_statistics,
+ "nf4",
+ quant_storage,
+ device,
+ )
+
+
+class Int8Params(torch.nn.Parameter):
+ def __new__(
+ cls,
+ data: Optional[torch.Tensor] = None,
+ requires_grad=True,
+ has_fp16_weights=False,
+ CB: Optional[torch.Tensor] = None,
+ SCB: Optional[torch.Tensor] = None,
+ **kwargs,
+ ):
+ if data is None:
+ data = torch.empty(0)
+ obj = torch.Tensor._make_subclass(cls, data, requires_grad)
+ obj.CB = CB
+ obj.SCB = SCB
+ obj.has_fp16_weights = has_fp16_weights
+ return obj
+
+ def _quantize(self, device):
+ if self.has_fp16_weights:
+ return super().to(device)
+
+ # We quantize the weight and store in 8bit row-major
+ B = self.data.contiguous().to(device=device, dtype=torch.float16)
+ CB, SCB, _ = bnb.functional.int8_vectorwise_quant(B)
+ self.data = CB
+ self.CB = CB
+ self.SCB = SCB
+
+ return self
+
+ def cpu(self):
+ return self.to(device="cpu")
+
+ def cuda(self, device: Optional[int | device | str] = None, non_blocking: bool = False):
+ return self.to(device="cuda" if device is None else device, non_blocking=non_blocking)
+
+ def xpu(self, device: Optional[int | device | str] = None, non_blocking: bool = False):
+ return self.to(device="xpu" if device is None else device, non_blocking=non_blocking)
+
+ def __deepcopy__(self, memo):
+ # adjust this if new arguments are added to the constructor
+ new_instance = type(self).__new__(
+ type(self),
+ data=copy.deepcopy(self.data, memo),
+ requires_grad=self.requires_grad,
+ has_fp16_weights=self.has_fp16_weights,
+ CB=copy.deepcopy(self.CB, memo),
+ SCB=copy.deepcopy(self.SCB, memo),
+ )
+ return new_instance
+
+ @overload
+ def to(
+ self: T,
+ device: Optional[int | device] = ...,
+ dtype: Optional[dtype | str] = ...,
+ non_blocking: bool = ...,
+ ) -> T: ...
+
+ @overload
+ def to(self: T, dtype: dtype | str, non_blocking: bool = ...) -> T: ...
+
+ @overload
+ def to(self: T, tensor: Tensor, non_blocking: bool = ...) -> T: ...
+
+ def to(self, *args, **kwargs):
+ device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs)
+
+ is_quantized = self.data.dtype == torch.int8
+
+ if not is_quantized and device is not None and device.type != "meta" and self.data.device.type == "cpu":
+ # We're moving from a CPU device to a non-meta device.
+ # In this circumstance, we want to quantize if we haven't already.
+ return self._quantize(device)
+
+ # Create a new parameter on the target device.
+ new_param = Int8Params(
+ super().to(device=device, dtype=dtype, non_blocking=non_blocking),
+ requires_grad=self.requires_grad,
+ has_fp16_weights=self.has_fp16_weights,
+ )
+
+ # If we had already quantized, move the statistics appropriately.
+ if is_quantized:
+ new_param.CB = new_param.data
+
+ if device is not None and self.SCB is not None and self.SCB.device.type != "meta":
+ new_param.SCB = self.SCB.to(device)
+
+ return new_param
+
+
+def maybe_rearrange_weight(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
+ weight = state_dict.get(f"{prefix}weight")
+ if weight is None:
+ # if the state dict has no weights for this layer (e.g., LoRA finetuning), do nothing
+ return
+ weight_format = state_dict.pop(f"{prefix}weight_format", "row")
+
+ if isinstance(weight_format, torch.Tensor):
+ weight_format = weight_format.item()
+
+ # For new weights format storage type, we explicitly check
+ # if weights_format is on the mapping
+ if isinstance(weight_format, int) and weight_format not in INVERSE_LINEAR_8BIT_WEIGHTS_FORMAT_MAPPING:
+ raise ValueError(f"Expected supported weight format - got {weight_format}")
+ elif isinstance(weight_format, int) and weight_format in INVERSE_LINEAR_8BIT_WEIGHTS_FORMAT_MAPPING:
+ weight_format = INVERSE_LINEAR_8BIT_WEIGHTS_FORMAT_MAPPING[weight_format]
+
+ if weight_format != "row":
+ raise ValueError(f"Only 'row' weight format is supported, got {weight_format}")
+
+
+class Embedding8bit(nn.Embedding):
+ """
+ This class implements [LLM.int8()](https://arxiv.org/abs/2208.07339) algorithm for embedding layer
+
+ Quantization API is similar to Linear8bitLt:
+ ```python
+ import torch
+ import torch.nn as nn
+
+ from bitsandbytes.nn import Embedding8bit
+
+ fp16_module = nn.Embedding(128, 64)
+ int8_module = Embedding8bit(128, 64)
+
+ int8_module.load_state_dict(fp16_module.state_dict())
+
+ int8_module = int8_module.to(0) # Quantization happens here
+ ```
+ """
+
+ def __init__(self, num_embeddings, embedding_dim, device=None, dtype=None):
+ super().__init__(num_embeddings, embedding_dim, device=device, dtype=dtype)
+ self.dtype = self.weight.data.dtype
+
+ self.weight = Int8Params(self.weight.data, has_fp16_weights=False, requires_grad=False)
+
+ def _save_to_state_dict(self, destination, prefix, keep_vars):
+ raise NotImplementedError("Saving Embedding8bit module is not implemented")
+
+ def forward(self, input: Tensor) -> Tensor:
+ if not hasattr(self.weight, "SCB"):
+ raise RuntimeError("Embedding layer is not quantized. Please call .cuda() or .to(device) first.")
+
+ rows = self.weight.data
+ row_stats = self.weight.SCB
+
+ assert rows.shape == (self.num_embeddings, self.embedding_dim)
+ assert row_stats.shape == (self.num_embeddings,)
+
+ compressed_output = F.embedding(input, rows)
+ compressed_output_stats = F.embedding(input, row_stats.view(self.num_embeddings, 1))
+
+ output = compressed_output * (compressed_output_stats / 127.0)
+
+ return output.to(self.dtype)
+
+
+class Embedding4bit(nn.Embedding):
+ """
+ This is the base class similar to Linear4bit. It implements the 4-bit quantization algorithm presented in
+ [QLoRA](https://arxiv.org/abs/2305.14314) for embeddings.
+
+ Quantization API is similar to Linear4bit:
+ ```python
+ import torch
+ import torch.nn as nn
+
+ from bitsandbytes.nn import Embedding4bit
+
+ fp16_module = nn.Embedding(128, 64)
+ quantized_module = Embedding4bit(128, 64)
+
+ quantized_module.load_state_dict(fp16_module.state_dict())
+
+ quantized_module = quantized_module.to(0) # Quantization happens here
+ ```
+ """
+
+ def __init__(
+ self,
+ num_embeddings,
+ embedding_dim,
+ dtype=None,
+ quant_type="fp4",
+ quant_storage=torch.uint8,
+ device=None,
+ ):
+ super().__init__(num_embeddings, embedding_dim, device=device, dtype=dtype)
+ self.dtype = self.weight.data.dtype
+
+ self.weight = Params4bit(
+ self.weight.data,
+ requires_grad=False,
+ compress_statistics=None,
+ quant_type=quant_type,
+ quant_storage=quant_storage,
+ module=self,
+ )
+
+ blocksize = self.weight.blocksize
+
+ if embedding_dim % blocksize != 0:
+ logger.warning(
+ f"Embedding size {embedding_dim} is not divisible by block size {blocksize}. "
+ "This will lead to slow inference.",
+ )
+
+ def _forward_with_partial_dequantize(self, input: Tensor):
+ assert self.embedding_dim % self.weight.quant_state.blocksize == 0
+
+ w_4bit_uint8 = self.weight.data.view(torch.uint8).view(self.num_embeddings * self.embedding_dim // 2, 1)
+
+ output_4bit = torch.nn.functional.embedding(
+ weight=w_4bit_uint8.view(self.num_embeddings, self.embedding_dim // 2),
+ input=input,
+ ).view(-1, 1)
+ assert output_4bit.shape == (input.numel() * self.embedding_dim // 2, 1)
+
+ blocks_per_emb = self.embedding_dim // self.weight.blocksize
+
+ absmax = self.weight.quant_state.absmax
+ assert absmax.shape == (self.num_embeddings * blocks_per_emb,)
+
+ output_absmax = torch.nn.functional.embedding(
+ weight=absmax.view(self.num_embeddings, blocks_per_emb),
+ input=input,
+ ).view(
+ -1,
+ )
+ assert output_absmax.shape == (input.numel() * blocks_per_emb,)
+
+ output_quant_state = copy.deepcopy(self.weight.quant_state)
+ output_quant_state.absmax = output_absmax
+ output_quant_state.shape = torch.Size((*input.shape, self.embedding_dim))
+
+ output = bnb.functional.dequantize_4bit(output_4bit, output_quant_state)
+ assert output.shape == (*input.shape, self.embedding_dim)
+
+ return output.to(self.dtype)
+
+ def _save_to_state_dict(self, destination, prefix, keep_vars):
+ raise NotImplementedError("Saving Embedding4bit module is not implemented")
+
+ def forward(self, input: Tensor) -> Tensor:
+ fix_4bit_weight_quant_state_from_module(self)
+
+ if self.embedding_dim % self.weight.quant_state.blocksize == 0:
+ return self._forward_with_partial_dequantize(input)
+
+ dequantized_weight = bnb.functional.dequantize_4bit(self.weight.data, self.weight.quant_state)
+
+ return torch.nn.functional.embedding(
+ weight=dequantized_weight,
+ input=input,
+ ).to(self.dtype)
+
+
+class EmbeddingFP4(Embedding4bit):
+ def __init__(
+ self,
+ num_embeddings,
+ embedding_dim,
+ dtype=None,
+ quant_storage=torch.uint8,
+ device=None,
+ ):
+ super().__init__(
+ num_embeddings,
+ embedding_dim,
+ dtype=dtype,
+ quant_type="fp4",
+ quant_storage=quant_storage,
+ device=device,
+ )
+
+
+class EmbeddingNF4(Embedding4bit):
+ def __init__(
+ self,
+ num_embeddings,
+ embedding_dim,
+ dtype=None,
+ quant_storage=torch.uint8,
+ device=None,
+ ):
+ super().__init__(
+ num_embeddings,
+ embedding_dim,
+ dtype=dtype,
+ quant_type="nf4",
+ quant_storage=quant_storage,
+ device=device,
+ )
+
+
+class Linear8bitLt(nn.Linear):
+ """
+ This class is the base module for the [LLM.int8()](https://arxiv.org/abs/2208.07339) algorithm.
+ To read more about it, have a look at the paper.
+
+ In order to quantize a linear layer one should first load the original fp16 / bf16 weights into
+ the Linear8bitLt module, then call `int8_module.to("cuda")` to quantize the fp16 weights.
+
+ Example:
+
+ ```python
+ import torch
+ import torch.nn as nn
+
+ import bitsandbytes as bnb
+ from bitsandbytes.nn import Linear8bitLt
+
+ fp16_model = nn.Sequential(
+ nn.Linear(64, 64),
+ nn.Linear(64, 64)
+ )
+
+ int8_model = nn.Sequential(
+ Linear8bitLt(64, 64, has_fp16_weights=False),
+ Linear8bitLt(64, 64, has_fp16_weights=False)
+ )
+
+ int8_model.load_state_dict(fp16_model.state_dict())
+ int8_model = int8_model.to(0) # Quantization happens here
+ ```
+ """
+
+ def __init__(
+ self,
+ input_features: int,
+ output_features: int,
+ bias=True,
+ has_fp16_weights=True,
+ threshold=0.0,
+ index=None,
+ device=None,
+ ):
+ """
+ Initialize Linear8bitLt class.
+
+ Args:
+ input_features (`int`):
+ Number of input features of the linear layer.
+ output_features (`int`):
+ Number of output features of the linear layer.
+ bias (`bool`, defaults to `True`):
+ Whether the linear class uses the bias term as well.
+ has_fp16_weights (`bool`, defaults to `True`):
+ If False, weights are quantized to int8 on ``.to(device)``. If True,
+ weights remain in fp16 and are quantized on-the-fly during each forward pass.
+ threshold (`float`, defaults to `0.0`):
+ Outlier threshold for mixed-precision decomposition (LLM.int8()). During the
+ forward pass, activation columns where any value exceeds this threshold are
+ computed in fp16, while the remaining columns use int8. This operates on
+ **activations** (inputs), not on weight values. Set to 0.0 to disable
+ mixed-precision decomposition and quantize all columns to int8.
+ index: Indices for weight reordering (used internally).
+ device: Device to initialize the layer on.
+ """
+ super().__init__(input_features, output_features, bias, device)
+ self.state = bnb.MatmulLtState()
+ self.index = index
+
+ self.state.threshold = threshold
+ self.state.has_fp16_weights = has_fp16_weights
+
+ if threshold > 0.0 and not has_fp16_weights:
+ self.state.use_pool = True
+
+ self.weight = Int8Params(self.weight.data, has_fp16_weights=has_fp16_weights, requires_grad=has_fp16_weights)
+ self._register_load_state_dict_pre_hook(maybe_rearrange_weight)
+
+ def _save_to_state_dict(self, destination, prefix, keep_vars):
+ super()._save_to_state_dict(destination, prefix, keep_vars)
+
+ # we only need to save SCB as extra data, because CB for quantized weights is already stored in weight.data
+ scb_name = "SCB"
+
+ # case 1: .cuda was called, SCB is in self.weight
+ param_from_weight = getattr(self.weight, scb_name, None)
+ # case 2: self.init_8bit_state was called, SCB is in self.state
+ param_from_state = getattr(self.state, scb_name)
+
+ key_name = prefix + f"{scb_name}"
+
+ # We now only save in row-major. This format information is stored for backwards compatibility.
+ format_name = prefix + "weight_format"
+
+ if not self.state.has_fp16_weights:
+ if param_from_weight is not None:
+ destination[key_name] = param_from_weight if keep_vars else param_from_weight.detach()
+ destination[format_name] = torch.tensor(0, dtype=torch.uint8)
+ elif param_from_state is not None:
+ destination[key_name] = param_from_state if keep_vars else param_from_state.detach()
+ destination[format_name] = torch.tensor(0, dtype=torch.uint8)
+
+ def _load_from_state_dict(
+ self,
+ state_dict,
+ prefix,
+ local_metadata,
+ strict,
+ missing_keys,
+ unexpected_keys,
+ error_msgs,
+ ):
+ super()._load_from_state_dict(
+ state_dict,
+ prefix,
+ local_metadata,
+ strict,
+ missing_keys,
+ unexpected_keys,
+ error_msgs,
+ )
+ unexpected_copy = list(unexpected_keys)
+
+ for key in unexpected_copy:
+ input_name = key[len(prefix) :]
+ if input_name == "SCB":
+ weight_scb = getattr(self.weight, "SCB", None)
+ if weight_scb is None:
+ # buffers not yet initialized, can't access them directly without quantizing first
+ raise RuntimeError(
+ "Loading a quantized checkpoint into non-quantized Linear8bitLt is "
+ "not supported. Please call module.cuda() before module.load_state_dict()",
+ )
+
+ input_param = state_dict[key]
+ weight_scb.copy_(input_param)
+
+ if self.state.SCB is not None:
+ self.state.SCB = self.weight.SCB
+
+ unexpected_keys.remove(key)
+
+ def init_8bit_state(self):
+ self.state.CB = self.weight.CB
+ self.state.SCB = self.weight.SCB
+ self.weight.CB = None
+ self.weight.SCB = None
+
+ def to(self, *args, **kwargs):
+ # Call the parent to() method to handle standard parameter/buffer movement
+ result = super().to(*args, **kwargs)
+
+ device, _, _, _ = torch._C._nn._parse_to(*args, **kwargs)
+
+ # Handle state tensors if needed.
+ if device is not None:
+ if result.state.CB is not None:
+ result.state.CB = result.state.CB.to(device)
+ if result.state.SCB is not None:
+ result.state.SCB = result.state.SCB.to(device)
+
+ return result
+
+ def forward(self, x: torch.Tensor):
+ self.state.is_training = self.training
+ if self.weight.CB is not None:
+ self.init_8bit_state()
+
+ # weights are cast automatically as Int8Params, but the bias has to be cast manually
+ if self.bias is not None and self.bias.dtype != x.dtype:
+ self.bias.data = self.bias.data.to(x.dtype)
+
+ out = bnb.matmul(x, self.weight, bias=self.bias, state=self.state)
+
+ if not self.state.has_fp16_weights and self.state.CB is not None:
+ self.weight.data = self.state.CB
+
+ return out
+
+
+class OutlierAwareLinear(nn.Linear):
+ def __init__(self, input_features, output_features, bias=True, device=None):
+ super().__init__(input_features, output_features, bias, device)
+ self.outlier_dim = None
+ self.is_quantized = False
+
+ def forward_with_outliers(self, x, outlier_idx):
+ raise NotImplementedError("Please override the `forward_with_outliers(self, x, outlier_idx)` function")
+
+ def quantize_weight(self, w, outlier_idx):
+ raise NotImplementedError("Please override the `quantize_weights(self, w, outlier_idx)` function")
+
+ def forward(self, x):
+ if self.outlier_dim is None:
+ tracer = OutlierTracer.get_instance()
+ if not tracer.is_initialized():
+ logger.warning("Please use OutlierTracer.initialize(model) before using the OutlierAwareLinear layer")
+ outlier_idx = tracer.get_outliers(self.weight)
+ self.outlier_dim = outlier_idx
+
+ if not self.is_quantized:
+ w = self.quantize_weight(self.weight, self.outlier_dim)
+ self.weight.data.copy_(w)
+ self.is_quantized = True
diff --git a/venv/lib/python3.11/site-packages/bitsandbytes/nn/parametrize.py b/venv/lib/python3.11/site-packages/bitsandbytes/nn/parametrize.py
new file mode 100644
index 0000000000000000000000000000000000000000..55f679cae7e504b6a16bcbdcaea03700255757c8
--- /dev/null
+++ b/venv/lib/python3.11/site-packages/bitsandbytes/nn/parametrize.py
@@ -0,0 +1,206 @@
+from functools import partial
+from typing import Any, Literal, Optional
+
+import torch
+import torch.nn as nn
+import torch.nn.utils.parametrize as P
+
+from .. import functional as F
+
+
+class Bnb4bitParametrization(nn.Module):
+ """
+ A parametrization module that handles dequantization of a 4-bit quantized parameter.
+
+ The parameter data is expected to be already quantized when this parametrization is applied.
+ This module will dequantize the parameter data to its original floating-point representation
+ when the forward method is called (i.e. when the parameter is accessed).
+
+ Args:
+ quant_state (`F.QuantState`):
+ The quantization state containing the necessary information for dequantization.
+ """
+
+ def __init__(self, quant_state: F.QuantState):
+ super().__init__()
+ self.quant_state = quant_state
+
+ @torch.no_grad()
+ def forward(self, quantized_param: torch.Tensor) -> torch.Tensor:
+ """
+ Forward pass to dequantize the parameter.
+
+ Args:
+ quantized_param (`torch.Tensor`): The quantized parameter tensor (from .original)
+
+ Returns:
+ `torch.Tensor`: The dequantized parameter tensor in the original shape and dtype.
+ """
+ return F.dequantize_4bit(quantized_param, self.quant_state)
+
+
+def replace_parameter_4bit_prequantized(
+ module: nn.Module, param_name: str, qs_dict: dict[str, Any], device: torch.device
+):
+ if not hasattr(module, param_name):
+ raise AttributeError(f"Module does not have parameter '{param_name}'")
+
+ original_param = getattr(module, param_name)
+
+ if not isinstance(original_param, nn.Parameter):
+ raise TypeError(f"Parameter '{param_name}' is not an instance of nn.Parameter")
+
+ quant_state = F.QuantState.from_dict(qs_dict, device=device)
+
+ # Apply a parametrization to the module to handle dequantization.
+ P.register_parametrization(module, param_name, Bnb4bitParametrization(quant_state), unsafe=True)
+
+ # Next, register hooks.
+ _register_parametrization_hooks(module, param_name)
+
+
+def replace_parameter_4bit(
+ module: nn.Module,
+ param_name: str,
+ compress_statistics: bool = False,
+ quant_type: Literal["nf4", "fp4"] = "nf4",
+ blocksize: Optional[int] = None,
+):
+ """
+ Replace a module parameter with a 4-bit quantized version using parametrization.
+
+ This function quantizes an existing parameter in a PyTorch module to 4-bit precision
+ and sets up parametrization to handle automatic dequantization during forward passes.
+ The original parameter is replaced with quantized data, and a parametrization layer
+ is registered to manage the quantization state and dequantization process.
+
+ Additional, it registers a state dict post-hook to ensure that the quantization state
+ is saved correctly when the model's state dict is saved.
+
+ It is useful for MoE models or other scenarios where you want to quantize parameters
+ outside of nn.Linear layers without changing the model's architecture.
+
+ This feature is experimental and may change in future releases.
+
+ Args:
+ module (`nn.Module`):
+ The PyTorch module containing the parameter to be quantized.
+ param_name (`str`):
+ The name of the parameter within the module to quantize.
+ compress_statistics (`bool`, *optional*, defaults to `False`):
+ Whether to compress quantization statistics to reduce memory usage.
+ quant_type (`Literal["nf4", "fp4"]`, *optional*, defaults to `"nf4"`):
+ The quantization format to use.
+ blocksize (`int`, *optional*, defaults to `None`):
+ The block size for quantization. If None, uses the default block size.
+
+ Raises:
+ AttributeError: If the module does not have the specified parameter.
+ TypeError: If the specified attribute is not an instance of nn.Parameter.
+ """
+
+ if not hasattr(module, param_name):
+ raise AttributeError(f"Module does not have parameter '{param_name}'")
+
+ original_param = getattr(module, param_name)
+
+ if not isinstance(original_param, nn.Parameter):
+ raise TypeError(f"Parameter '{param_name}' is not an instance of nn.Parameter")
+
+ # Quantize the original parameter.
+ quantized_data, quant_state = F.quantize_4bit(
+ original_param.data,
+ blocksize=blocksize,
+ compress_statistics=compress_statistics,
+ quant_type=quant_type,
+ )
+
+ # Replace the parameter with the quantized data.
+ setattr(module, param_name, nn.Parameter(quantized_data, requires_grad=False))
+ del original_param
+
+ # Apply a parametrization to the module to handle dequantization.
+ P.register_parametrization(module, param_name, Bnb4bitParametrization(quant_state), unsafe=True)
+
+ # Next, register hooks.
+ _register_parametrization_hooks(module, param_name)
+
+
+def _disable_parametrization_cache(module: nn.Module, inputs: tuple[Any, ...], output: Any):
+ # Clamp instead of a bare decrement: with ``always_call=True`` this hook also runs
+ # when the forward raised before the pre-hook incremented (e.g. an earlier pre-hook
+ # failed), and the counter must never go negative — a negative value is truthy, so
+ # ``if not P._cache_enabled`` would stop clearing the cache forever.
+ P._cache_enabled = max(0, P._cache_enabled - 1)
+ if not P._cache_enabled:
+ P._cache = {}
+
+
+def _enable_parametrization_cache(module: nn.Module, inputs: tuple[Any, ...]):
+ P._cache_enabled += 1
+
+
+def _register_parametrization_hooks(module: nn.Module, param_name: str):
+ # Register a state dict hook for saving. Note that this requires torch >= 2.5.0.
+ if torch.__version__ >= (2, 5):
+ module.register_state_dict_post_hook(
+ partial(
+ _parametrized_state_dict_post_hook,
+ param_name=param_name,
+ )
+ )
+
+ # Register hooks to enable caching for the dequantization parametrization.
+ # This helps preserve time and memory when the same quantized parameter
+ # is accessed multiple times in the forward computation.
+ #
+ # ``always_call=True`` is load-bearing: activation checkpointing with
+ # ``use_reentrant=False`` aborts its backward recompute mid-forward by design
+ # (early stop, via an internal exception) once the last needed activation has
+ # been rematerialized. A plain forward hook is skipped in that case, so the
+ # global ``parametrize._cache_enabled`` counter leaks upward once per
+ # checkpointed region per step, after which the cache is enabled (and never
+ # cleared) for the remainder of training — every dequantized parameter this
+ # module produces stays resident, i.e. a memory leak of the full dequantized
+ # model size (4x the packed 4-bit bytes).
+ module.register_forward_pre_hook(_enable_parametrization_cache)
+ module.register_forward_hook(_disable_parametrization_cache, always_call=True)
+
+
+def _parametrized_state_dict_post_hook(
+ module: nn.Module,
+ state_dict: dict[str, Any],
+ prefix: str,
+ local_metadata: Any,
+ *,
+ param_name: str = "weight",
+ **kwargs: dict[str, Any],
+) -> None:
+ """
+ Hook to modify the state dict to include the quantization state.
+ """
+
+ original_key = f"{prefix}parametrizations.{param_name}.original"
+
+ if original_key in state_dict:
+ # Create a clean entry.
+ # The `parametrizations.{param_name}.original` key will have the quantized data,
+ # but we would like it to keep it in the state_dict as `{param_name}`.
+ clean_key = f"{prefix}{param_name}"
+ state_dict[clean_key] = state_dict.pop(original_key)
+
+ assert P.is_parametrized(module, param_name)
+
+ # Find the parametrization, which should have the quantization state.
+ parametrization: Bnb4bitParametrization = next(
+ filter(lambda x: isinstance(x, Bnb4bitParametrization), module.parametrizations[param_name]), None
+ )
+
+ assert parametrization is not None, "Parametrization not found for the parameter."
+
+ quant_state = parametrization.quant_state
+
+ # Next, we need to store the quantization state.
+ if quant_state is not None:
+ for k, v in quant_state.as_dict(packed=True).items():
+ state_dict[f"{prefix}{param_name}.{k}"] = v