diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..33437ccc5715af7f29d2766acda1dc1692262cf8 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,10 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +examples/I2V/00_cat_vac/image.png filter=lfs diff=lfs merge=lfs -text +examples/I2V/01_socrates/image.png filter=lfs diff=lfs merge=lfs -text +examples/I2V/02_chestnut/image.png filter=lfs diff=lfs merge=lfs -text +examples/I2V/06_waterfall/image.png filter=lfs diff=lfs merge=lfs -text +examples/I2V/10_case061/image.png filter=lfs diff=lfs merge=lfs -text +examples/I2V/13_burrow/image.png filter=lfs diff=lfs merge=lfs -text +examples/I2V/15_case104/image.png filter=lfs diff=lfs merge=lfs -text diff --git a/LICENSE.txt b/LICENSE.txt new file mode 100644 index 0000000000000000000000000000000000000000..62bc46170f4808332edba75501edb493b258ed7e --- /dev/null +++ b/LICENSE.txt @@ -0,0 +1,500 @@ +Tencent is pleased to support the community by making WorldCrafter available. + +Copyright (C) 2026 Tencent. All rights reserved. + +The open-source software and/or model(s) included in this distribution may have been modified by Tencent ("Tencent Modifications"). All Tencent Modifications are Copyright (C) Tencent. + +WorldCrafter is licensed under License Term of WorldCrafter, except for the third-party components listed below, which remain licensed under their respective original terms. WorldCrafter does not impose any additional restrictions beyond those specified in the original licenses of these third-party components. Users are required to comply with all applicable terms and conditions of the original licenses and to ensure that the use of these third-party components conforms to all relevant laws and regulations. + +For the avoidance of doubt, WorldCrafter refers solely to code, parameters, and weights made publicly available by Tencent in accordance with License Term of WorldCrafter. + +Terms of License Term of WorldCrafter: +-------------------------------------------------------------------- +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, and /or sublicense copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions: + +- You agree to use the WorldCrafter only for academic purposes, and refrain from using it for any commercial or production purposes under any circumstances. + +- 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. + + +Dependencies and Licenses: + +This open-source project builds upon the following open-source model(s) and/or software, each of which remains licensed under its original license(s). Certain component(s) may include modifications made by Tencent ("Tencent Modifications"), which are Copyright (C) Tencent. + +In case you believe there have been errors in the attribution below, you may submit the concerns to us for review and correction. + + +Open Source Model(s)/Software Licensed under the Apache-2.0: +-------------------------------------------------------------------- +1. Helios-Base +Copyright (c) Helios-Base Original author and authors +Terms of the Apache-2.0: +-------------------------------------------------------------------- +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: + +You must give any other recipients of the Work or Derivative Works a copy of this License; and +You must cause any modified files to carry prominent notices stating that You changed the files; and +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 +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 + + + + + +Open Source Model(s)/Software Licensed under the CC-BY-NC-4.0: +-------------------------------------------------------------------- +1. VGGT +Copyright (c) VGGT Original author and authors +Please note this software has been modified by Tencent in this distribution. +You may find the original code here: https://huggingface.co/facebook/VGGT-1B +Terms of the CC-BY-NC-4.0: +-------------------------------------------------------------------- +Attribution-NonCommercial 4.0 International + +======================================================================= + +Creative Commons Corporation ("Creative Commons") is not a law firm and +does not provide legal services or legal advice. Distribution of +Creative Commons public licenses does not create a lawyer-client or +other relationship. Creative Commons makes its licenses and related +information available on an "as-is" basis. Creative Commons gives no +warranties regarding its licenses, any material licensed under their +terms and conditions, or any related information. Creative Commons +disclaims all liability for damages resulting from their use to the +fullest extent possible. + +Using Creative Commons Public Licenses + +Creative Commons public licenses provide a standard set of terms and +conditions that creators and other rights holders may use to share +original works of authorship and other material subject to copyright +and certain other rights specified in the public license below. The +following considerations are for informational purposes only, are not +exhaustive, and do not form part of our licenses. + + Considerations for licensors: Our public licenses are + intended for use by those authorized to give the public + permission to use material in ways otherwise restricted by + copyright and certain other rights. Our licenses are + irrevocable. Licensors should read and understand the terms + and conditions of the license they choose before applying it. + Licensors should also secure all rights necessary before + applying our licenses so that the public can reuse the + material as expected. Licensors should clearly mark any + material not subject to the license. This includes other CC- + licensed material, or material used under an exception or + limitation to copyright. More considerations for licensors: + wiki.creativecommons.org/Considerations_for_licensors + + Considerations for the public: By using one of our public + licenses, a licensor grants the public permission to use the + licensed material under specified terms and conditions. If + the licensor's permission is not necessary for any reason--for + example, because of any applicable exception or limitation to + copyright--then that use is not regulated by the license. Our + licenses grant only permissions under copyright and certain + other rights that a licensor has authority to grant. Use of + the licensed material may still be restricted for other + reasons, including because others have copyright or other + rights in the material. A licensor may make special requests, + such as asking that all changes be marked or described. + Although not required by our licenses, you are encouraged to + respect those requests where reasonable. More_considerations + for the public: + wiki.creativecommons.org/Considerations_for_licensees + +======================================================================= + +Creative Commons Attribution-NonCommercial 4.0 International Public +License + +By exercising the Licensed Rights (defined below), You accept and agree +to be bound by the terms and conditions of this Creative Commons +Attribution-NonCommercial 4.0 International Public License ("Public +License"). To the extent this Public License may be interpreted as a +contract, You are granted the Licensed Rights in consideration of Your +acceptance of these terms and conditions, and the Licensor grants You +such rights in consideration of benefits the Licensor receives from +making the Licensed Material available under these terms and +conditions. + +Section 1 -- Definitions. + +a. Adapted Material means material subject to Copyright and Similar + Rights that is derived from or based upon the Licensed Material + and in which the Licensed Material is translated, altered, + arranged, transformed, or otherwise modified in a manner requiring + permission under the Copyright and Similar Rights held by the + Licensor. For purposes of this Public License, where the Licensed + Material is a musical work, performance, or sound recording, + Adapted Material is always produced where the Licensed Material is + synched in timed relation with a moving image. + +b. Adapter's License means the license You apply to Your Copyright + and Similar Rights in Your contributions to Adapted Material in + accordance with the terms and conditions of this Public License. + +c. Copyright and Similar Rights means copyright and/or similar rights + closely related to copyright including, without limitation, + performance, broadcast, sound recording, and Sui Generis Database + Rights, without regard to how the rights are labeled or + categorized. For purposes of this Public License, the rights + specified in Section 2(b)(1)-(2) are not Copyright and Similar + Rights. +d. Effective Technological Measures means those measures that, in the + absence of proper authority, may not be circumvented under laws + fulfilling obligations under Article 11 of the WIPO Copyright + Treaty adopted on December 20, 1996, and/or similar international + agreements. + +e. Exceptions and Limitations means fair use, fair dealing, and/or + any other exception or limitation to Copyright and Similar Rights + that applies to Your use of the Licensed Material. + +f. Licensed Material means the artistic or literary work, database, + or other material to which the Licensor applied this Public + License. + +g. Licensed Rights means the rights granted to You subject to the + terms and conditions of this Public License, which are limited to + all Copyright and Similar Rights that apply to Your use of the + Licensed Material and that the Licensor has authority to license. + +h. Licensor means the individual(s) or entity(ies) granting rights + under this Public License. + +i. NonCommercial means not primarily intended for or directed towards + commercial advantage or monetary compensation. For purposes of + this Public License, the exchange of the Licensed Material for + other material subject to Copyright and Similar Rights by digital + file-sharing or similar means is NonCommercial provided there is + no payment of monetary compensation in connection with the + exchange. + +j. Share means to provide material to the public by any means or + process that requires permission under the Licensed Rights, such + as reproduction, public display, public performance, distribution, + dissemination, communication, or importation, and to make material + available to the public including in ways that members of the + public may access the material from a place and at a time + individually chosen by them. + +k. Sui Generis Database Rights means rights other than copyright + resulting from Directive 96/9/EC of the European Parliament and of + the Council of 11 March 1996 on the legal protection of databases, + as amended and/or succeeded, as well as other essentially + equivalent rights anywhere in the world. + +l. You means the individual or entity exercising the Licensed Rights + under this Public License. Your has a corresponding meaning. + +Section 2 -- Scope. + +a. License grant. + + 1. Subject to the terms and conditions of this Public License, + the Licensor hereby grants You a worldwide, royalty-free, + non-sublicensable, non-exclusive, irrevocable license to + exercise the Licensed Rights in the Licensed Material to: + + a. reproduce and Share the Licensed Material, in whole or + in part, for NonCommercial purposes only; and + + b. produce, reproduce, and Share Adapted Material for + NonCommercial purposes only. + + 2. Exceptions and Limitations. For the avoidance of doubt, where + Exceptions and Limitations apply to Your use, this Public + License does not apply, and You do not need to comply with + its terms and conditions. + + 3. Term. The term of this Public License is specified in Section + 6(a). + + 4. Media and formats; technical modifications allowed. The + Licensor authorizes You to exercise the Licensed Rights in + all media and formats whether now known or hereafter created, + and to make technical modifications necessary to do so. The + Licensor waives and/or agrees not to assert any right or + authority to forbid You from making technical modifications + necessary to exercise the Licensed Rights, including + technical modifications necessary to circumvent Effective + Technological Measures. For purposes of this Public License, + simply making modifications authorized by this Section 2(a) + (4) never produces Adapted Material. + + 5. Downstream recipients. + + a. Offer from the Licensor -- Licensed Material. Every + recipient of the Licensed Material automatically + receives an offer from the Licensor to exercise the + Licensed Rights under the terms and conditions of this + Public License. + + b. No downstream restrictions. You may not offer or impose + any additional or different terms or conditions on, or + apply any Effective Technological Measures to, the + Licensed Material if doing so restricts exercise of the + Licensed Rights by any recipient of the Licensed + Material. + + 6. No endorsement. Nothing in this Public License constitutes or + may be construed as permission to assert or imply that You + are, or that Your use of the Licensed Material is, connected + with, or sponsored, endorsed, or granted official status by, + the Licensor or others designated to receive attribution as + provided in Section 3(a)(1)(A)(i). + +b. Other rights. + + 1. Moral rights, such as the right of integrity, are not + licensed under this Public License, nor are publicity, + privacy, and/or other similar personality rights; however, to + the extent possible, the Licensor waives and/or agrees not to + assert any such rights held by the Licensor to the limited + extent necessary to allow You to exercise the Licensed + Rights, but not otherwise. + + 2. Patent and trademark rights are not licensed under this + Public License. + + 3. To the extent possible, the Licensor waives any right to + collect royalties from You for the exercise of the Licensed + Rights, whether directly or through a collecting society + under any voluntary or waivable statutory or compulsory + licensing scheme. In all other cases the Licensor expressly + reserves any right to collect such royalties, including when + the Licensed Material is used other than for NonCommercial + purposes. + +Section 3 -- License Conditions. + +Your exercise of the Licensed Rights is expressly made subject to the +following conditions. + +a. Attribution. + + 1. If You Share the Licensed Material (including in modified + form), You must: + + a. retain the following if it is supplied by the Licensor + with the Licensed Material: + + i. identification of the creator(s) of the Licensed + Material and any others designated to receive + attribution, in any reasonable manner requested by + the Licensor (including by pseudonym if + designated); + + ii. a copyright notice; + + iii. a notice that refers to this Public License; + + iv. a notice that refers to the disclaimer of + warranties; + + v. a URI or hyperlink to the Licensed Material to the + extent reasonably practicable; + + b. indicate if You modified the Licensed Material and + retain an indication of any previous modifications; and + + c. indicate the Licensed Material is licensed under this + Public License, and include the text of, or the URI or + hyperlink to, this Public License. + + 2. You may satisfy the conditions in Section 3(a)(1) in any + reasonable manner based on the medium, means, and context in + which You Share the Licensed Material. For example, it may be + reasonable to satisfy the conditions by providing a URI or + hyperlink to a resource that includes the required + information. + + 3. If requested by the Licensor, You must remove any of the + information required by Section 3(a)(1)(A) to the extent + reasonably practicable. + + 4. If You Share Adapted Material You produce, the Adapter's + License You apply must not prevent recipients of the Adapted + Material from complying with this Public License. + +Section 4 -- Sui Generis Database Rights. + +Where the Licensed Rights include Sui Generis Database Rights that +apply to Your use of the Licensed Material: + +a. for the avoidance of doubt, Section 2(a)(1) grants You the right + to extract, reuse, reproduce, and Share all or a substantial + portion of the contents of the database for NonCommercial purposes + only; + +b. if You include all or a substantial portion of the database + contents in a database in which You have Sui Generis Database + Rights, then the database in which You have Sui Generis Database + Rights (but not its individual contents) is Adapted Material; and + +c. You must comply with the conditions in Section 3(a) if You Share + all or a substantial portion of the contents of the database. + +For the avoidance of doubt, this Section 4 supplements and does not +replace Your obligations under this Public License where the Licensed +Rights include other Copyright and Similar Rights. + +Section 5 -- Disclaimer of Warranties and Limitation of Liability. + +a. UNLESS OTHERWISE SEPARATELY UNDERTAKEN BY THE LICENSOR, TO THE + EXTENT POSSIBLE, THE LICENSOR OFFERS THE LICENSED MATERIAL AS-IS + AND AS-AVAILABLE, AND MAKES NO REPRESENTATIONS OR WARRANTIES OF + ANY KIND CONCERNING THE LICENSED MATERIAL, WHETHER EXPRESS, + IMPLIED, STATUTORY, OR OTHER. THIS INCLUDES, WITHOUT LIMITATION, + WARRANTIES OF TITLE, MERCHANTABILITY, FITNESS FOR A PARTICULAR + PURPOSE, NON-INFRINGEMENT, ABSENCE OF LATENT OR OTHER DEFECTS, + ACCURACY, OR THE PRESENCE OR ABSENCE OF ERRORS, WHETHER OR NOT + KNOWN OR DISCOVERABLE. WHERE DISCLAIMERS OF WARRANTIES ARE NOT + ALLOWED IN FULL OR IN PART, THIS DISCLAIMER MAY NOT APPLY TO YOU. + +b. TO THE EXTENT POSSIBLE, IN NO EVENT WILL THE LICENSOR BE LIABLE + TO YOU ON ANY LEGAL THEORY (INCLUDING, WITHOUT LIMITATION, + NEGLIGENCE) OR OTHERWISE FOR ANY DIRECT, SPECIAL, INDIRECT, + INCIDENTAL, CONSEQUENTIAL, PUNITIVE, EXEMPLARY, OR OTHER LOSSES, + COSTS, EXPENSES, OR DAMAGES ARISING OUT OF THIS PUBLIC LICENSE OR + USE OF THE LICENSED MATERIAL, EVEN IF THE LICENSOR HAS BEEN + ADVISED OF THE POSSIBILITY OF SUCH LOSSES, COSTS, EXPENSES, OR + DAMAGES. WHERE A LIMITATION OF LIABILITY IS NOT ALLOWED IN FULL OR + IN PART, THIS LIMITATION MAY NOT APPLY TO YOU. + +c. The disclaimer of warranties and limitation of liability provided + above shall be interpreted in a manner that, to the extent + possible, most closely approximates an absolute disclaimer and + waiver of all liability. + +Section 6 -- Term and Termination. + +a. This Public License applies for the term of the Copyright and + Similar Rights licensed here. However, if You fail to comply with + this Public License, then Your rights under this Public License + terminate automatically. + +b. Where Your right to use the Licensed Material has terminated under + Section 6(a), it reinstates: + + 1. automatically as of the date the violation is cured, provided + it is cured within 30 days of Your discovery of the + violation; or + + 2. upon express reinstatement by the Licensor. + + For the avoidance of doubt, this Section 6(b) does not affect any + right the Licensor may have to seek remedies for Your violations + of this Public License. + +c. For the avoidance of doubt, the Licensor may also offer the + Licensed Material under separate terms or conditions or stop + distributing the Licensed Material at any time; however, doing so + will not terminate this Public License. + +d. Sections 1, 5, 6, 7, and 8 survive termination of this Public + License. + +Section 7 -- Other Terms and Conditions. + +a. The Licensor shall not be bound by any additional or different + terms or conditions communicated by You unless expressly agreed. + +b. Any arrangements, understandings, or agreements regarding the + Licensed Material not stated herein are separate from and + independent of the terms and conditions of this Public License. + +Section 8 -- Interpretation. + +a. For the avoidance of doubt, this Public License does not, and + shall not be interpreted to, reduce, limit, restrict, or impose + conditions on any use of the Licensed Material that could lawfully + be made without permission under this Public License. + +b. To the extent possible, if any provision of this Public License is + deemed unenforceable, it shall be automatically reformed to the + minimum extent necessary to make it enforceable. If the provision + cannot be reformed, it shall be severed from this Public License + without affecting the enforceability of the remaining terms and + conditions. + +c. No term or condition of this Public License will be waived and no + failure to comply consented to unless expressly agreed to by the + Licensor. + +d. Nothing in this Public License constitutes or may be interpreted + as a limitation upon, or waiver of, any privileges and immunities + that apply to the Licensor or You, including from the legal + processes of any jurisdiction or authority. + +======================================================================= + +Creative Commons is not a party to its public +licenses. Notwithstanding, Creative Commons may elect to apply one of +its public licenses to material it publishes and in those instances +will be considered the “Licensor.” The text of the Creative Commons +public licenses is dedicated to the public domain under the CC0 Public +Domain Dedication. Except for the limited purpose of indicating that +material is shared under a Creative Commons public license or as +otherwise permitted by the Creative Commons policies published at +creativecommons.org/policies, Creative Commons does not authorize the +use of the trademark "Creative Commons" or any other trademark or logo +of Creative Commons without its prior written consent including, +without limitation, in connection with any unauthorized modifications +to any of its public licenses or any other arrangements, +understandings, or agreements concerning use of licensed material. For +the avoidance of doubt, this paragraph does not form part of the +public licenses. + +Creative Commons may be contacted at creativecommons.org. + +================================================== +End of the Attribution Notice of this project. diff --git a/README.md b/README.md index 994b902312665cc02ec2c9e13f92b55e03bb0a0d..8566eb23ce0c5dd106fcf4106cf9f6c40c3a20cb 100644 --- a/README.md +++ b/README.md @@ -1,13 +1,67 @@ --- -title: Worldcrafter Demo -emoji: ⚡ -colorFrom: red -colorTo: green +title: WorldCrafter +emoji: 🌍 +colorFrom: gray +colorTo: red sdk: gradio sdk_version: 6.28.0 -python_version: '3.12' app_file: app.py -pinned: false +python_version: "3.12" +startup_duration_timeout: 1h +short_description: Camera-controlled video world model with 3D-aware memory --- -Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference +# WorldCrafter: Consistent Video World Model with Implicit 3D-aware Memory + +Interactive demo of [`TencentARC/WorldCrafter-Fast`](https://huggingface.co/TencentARC/WorldCrafter-Fast) — +the distilled 6-step variant of WorldCrafter. Give it a start image (or just a prompt) plus a +camera action script, and it explores the scene, keeping geometry consistent across chunks +through a camera-queryable implicit 3D-aware memory. + +- Paper: https://huggingface.co/papers/2609.24984 +- Code: https://github.com/TencentARC/WorldCrafter +- Weights: [`WorldCrafter-Fast`](https://huggingface.co/TencentARC/WorldCrafter-Fast) (distilled, 6 steps, CFG 1.0) + +Output is 384×640 at 16 fps, generated in 33-frame chunks. + +## Camera actions + +One action per chunk, e.g. + +``` +forward1x2 +yaw_left30x3 +backward1 +``` + +| Movement | Actions | Short forms | +| --- | --- | --- | +| Forward / backward | `forward1`, `backward1` | `f1`, `b1` | +| Left / right | `left1`, `right1` | `l1`, `r1` | +| Up / down | `up1`, `down1` | same | +| Turn left / right | `yaw_left30`, `yaw_right30` | `yl30`, `yr30` | +| Look up / down | `pitch_up15`, `pitch_down15` | `pu15`, `pd15` | + +`xN` repeats an action, `&` combines movement and rotation in one chunk +(`forward2&right2&yaw_left45`), `reverseN` retraces the previous N chunks. Keep total +translation per chunk at or below 5. Optional headers `@dtype`, `@sampling`, `@last_frame` +must precede the actions. + +## Credits and license + +Model, example images, prompts and camera scripts are from the official +[WorldCrafter repository](https://github.com/TencentARC/WorldCrafter) and are redistributed +here under the WorldCrafter license (see `LICENSE.txt`), which permits copying and +distribution **for academic purposes only**. The underlying Helios-Base components are +Apache-2.0. Copyright (C) 2026 THL A29 Limited, a Tencent company. + +```bibtex +@article{yu2026worldcrafter, + title = {WorldCrafter: Consistent Video World Model with Implicit 3D-aware Memory}, + author = {Yu, Wangbo and Liu, Kunhao and Hu, Wenbo and Yuan, Shenghai and Feng, Chaoran + and Zhou, Haiyang and Huang, Yukun and Wang, Yiran and Zhao, Wang + and Luo, Yingmin and Shan, Ying}, + journal = {arXiv preprint arXiv:2609.24984}, + year = {2026} +} +``` diff --git a/app.py b/app.py new file mode 100644 index 0000000000000000000000000000000000000000..96799e84965fb5eaf4df2c1c18b6e1c27261e2de --- /dev/null +++ b/app.py @@ -0,0 +1,119 @@ +import os + +os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") + +import spaces # noqa: E402 (must precede torch) +import torch # noqa: E402 +import gradio as gr # noqa: E402 + +import shutil +import subprocess +import sys + + +def _diagnostics() -> str: + lines = [f"python: {sys.version}"] + for name in ( + "torch", + "triton", + "diffusers", + "transformers", + "peft", + "accelerate", + "numpy", + "timm", + "kernels", + "imageio", + "huggingface_hub", + "gradio", + "spaces", + ): + try: + mod = __import__(name) + lines.append(f"{name}: {getattr(mod, '__version__', '?')}") + except Exception as exc: # noqa: BLE001 + lines.append(f"{name}: IMPORT FAILED {exc!r}") + + lines.append("") + lines.append(f"cpu_count: {os.cpu_count()}") + lines.append(f"HF_HOME={os.environ.get('HF_HOME')}") + lines.append(f"HF_HUB_CACHE={os.environ.get('HF_HUB_CACHE')}") + lines.append(f"torch.cuda.is_available(): {torch.cuda.is_available()}") + + lines.append("") + for path in ("/", "/tmp", "/home/user", "/data", os.getcwd()): + try: + total, used, free = shutil.disk_usage(path) + lines.append( + f"disk {path}: total={total / 2**30:.1f}G " + f"used={used / 2**30:.1f}G free={free / 2**30:.1f}G" + ) + except Exception as exc: # noqa: BLE001 + lines.append(f"disk {path}: {exc!r}") + + for cmd in (["df", "-h"], ["free", "-g"]): + try: + out = subprocess.run(cmd, capture_output=True, text=True, timeout=30) + lines.append("") + lines.append(f"$ {' '.join(cmd)}\n{out.stdout}{out.stderr}") + except Exception as exc: # noqa: BLE001 + lines.append(f"{' '.join(cmd)}: {exc!r}") + + lines.append("") + try: + with open("/sys/fs/cgroup/memory.max") as fh: + lines.append(f"cgroup memory.max: {fh.read().strip()}") + except Exception as exc: # noqa: BLE001 + lines.append(f"cgroup memory.max: {exc!r}") + + lines.append("") + try: + from worldcrafter import WorldCrafter # noqa: F401 + + lines.append("import worldcrafter: OK") + except Exception as exc: # noqa: BLE001 + import traceback + + lines.append(f"import worldcrafter: FAILED {exc!r}\n{traceback.format_exc()}") + + try: + from worldcrafter.fast.resident import ResidentBranches # noqa: F401 + + lines.append("import worldcrafter.fast.resident (triton): OK") + except Exception as exc: # noqa: BLE001 + lines.append(f"import worldcrafter.fast.resident: FAILED {exc!r}") + + try: + from worldcrafter.camera import build_trajectory, parse_trajectory + + events, options = parse_trajectory("forward1x2\nyaw_left30") + camera, records = build_trajectory(events, **options) + lines.append( + f"camera smoke: events={events} camera={camera.shape} chunks={len(records)}" + ) + except Exception as exc: # noqa: BLE001 + import traceback + + lines.append(f"camera smoke: FAILED {exc!r}\n{traceback.format_exc()}") + + return "\n".join(lines) + + +REPORT = _diagnostics() +print("==== WORLDCRAFTER SPACE DIAGNOSTICS ====", flush=True) +print(REPORT, flush=True) +print("==== END DIAGNOSTICS ====", flush=True) + + +def report() -> str: + """Return the environment diagnostics collected at startup.""" + return REPORT + + +with gr.Blocks(theme=gr.themes.Citrus(), title="WorldCrafter (provisioning)") as demo: + gr.Markdown("# WorldCrafter — provisioning\nEnvironment diagnostics:") + out = gr.Textbox(value=REPORT, lines=40, label="diagnostics") + gr.Button("Refresh").click(fn=report, inputs=None, outputs=out) + +if __name__ == "__main__": + demo.launch() diff --git a/examples/I2V/00_cat_vac/actions.txt b/examples/I2V/00_cat_vac/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..add0ac37e4a482bed1b88f16aa48677a3486ad0a --- /dev/null +++ b/examples/I2V/00_cat_vac/actions.txt @@ -0,0 +1,9 @@ +@last_frame include + +forward1x2 +backward1x4 +yaw_right45x2 +yaw_left45x4 +right1x2 +yaw_right45x4 +yaw_left45x2 diff --git a/examples/I2V/00_cat_vac/camera.npy b/examples/I2V/00_cat_vac/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..808a9dba873f2cc16d2142d3cb0937f5a38d47b1 --- /dev/null +++ b/examples/I2V/00_cat_vac/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:810e574e8bbbd938abcbdd4561885463ed683a1ad052f1222c5f7e22eb31e998 +size 63488 diff --git a/examples/I2V/00_cat_vac/image.png b/examples/I2V/00_cat_vac/image.png new file mode 100644 index 0000000000000000000000000000000000000000..1b395343c85903d8ee5e54f1cb567c7c2b1448a0 --- /dev/null +++ b/examples/I2V/00_cat_vac/image.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0cd435eeeca8898260105991c84364fc9b941002d598dc40b3ef82aea9586bf1 +size 493694 diff --git a/examples/I2V/00_cat_vac/prompt.txt b/examples/I2V/00_cat_vac/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..971fed5ed997a3bfe8c7b179107bcae36858b568 --- /dev/null +++ b/examples/I2V/00_cat_vac/prompt.txt @@ -0,0 +1 @@ +A third-person gameplay-like camera closely follows a gray robot vacuum moving through a modern interior with reflective hardwood floors and beautiful rays of light. An adult brown tabby sits upright on the circular vacuum with green eyes, striped fur, white paws, and its tail curled beside the shell. The machine has a matte gray body, raised sensor turret, rubber bumper, and a small control panel. It passes between a sofa, low wooden tables, kitchen cabinetry, rugs, potted plants, and scattered household objects. The cat shifts its paws and body to remain balanced while the vacuum turns around furniture and crosses changes in floor material. diff --git a/examples/I2V/01_socrates/actions.txt b/examples/I2V/01_socrates/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..a10e20934c92e5c6e4ca8035b6a15b8c61b45197 --- /dev/null +++ b/examples/I2V/01_socrates/actions.txt @@ -0,0 +1,15 @@ +forward2 +backward2 +left1 +right1 +forward2&right2&yaw_left45 +reverse1 +left1.5 +forward2x3 +reverse4 +up2&forward1.5&pitch_down30 +reverse1 +right1.5 +forward1.5&left2&yaw_right45 +reverse1 +left1.5 diff --git a/examples/I2V/01_socrates/camera.npy b/examples/I2V/01_socrates/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..e9447ab2573f8ce09ee8981a8fa300b64d0b0279 --- /dev/null +++ b/examples/I2V/01_socrates/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9146aad4b53b9f9c458ba21361c7f32469906d475e47c252292a39eb31ee9d8f +size 63488 diff --git a/examples/I2V/01_socrates/image.png b/examples/I2V/01_socrates/image.png new file mode 100644 index 0000000000000000000000000000000000000000..7bad39ec7ac5f78b4c474d77d6c3fb572027d308 --- /dev/null +++ b/examples/I2V/01_socrates/image.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4f16295556a54cee42809952ba527be2c2d38379610cc83344fa9f10f0992671 +size 309740 diff --git a/examples/I2V/01_socrates/prompt.txt b/examples/I2V/01_socrates/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..2a04d1bbdb1712fcc983c795f2707b23ed6db0b8 --- /dev/null +++ b/examples/I2V/01_socrates/prompt.txt @@ -0,0 +1 @@ +A scene of static, painted sculptures depicts a solemn stone prison chamber, with classical figures neatly arranged around a low wooden bed. An elderly philosopher sculpture sits upright in a white robe, one hand extended toward a cup and the other raised in a fixed rhetorical gesture. Companion sculptures wear red, blue, yellow, gray, and ochre garments, with sculpted expressions of grief, disbelief, and contemplation. All figures remain completely motionless, with rigid poses and fixed garment folds. Scrolls, shackles, cups, sandals, stools, and other props are neatly placed in clearly organized groups. Stone walls, orderly steps, and an arched passage frame the scene. The chamber is clean, tidy, and carefully arranged, forming a coherent historical tableau of painted sculptures. diff --git a/examples/I2V/02_chestnut/actions.txt b/examples/I2V/02_chestnut/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..516ad91afd18cbe3b047c2538b0ce0afe188fc9d --- /dev/null +++ b/examples/I2V/02_chestnut/actions.txt @@ -0,0 +1,9 @@ +@dtype float32 +@last_frame include + +backward1.5x4 +yaw_left45x2 +forward1.5x4 +yaw_right45x2 +forward1.5x2 +right1.5x4 diff --git a/examples/I2V/02_chestnut/camera.npy b/examples/I2V/02_chestnut/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..a846c8f276317df5063d0b26942611dcc67313a9 --- /dev/null +++ b/examples/I2V/02_chestnut/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:24573a503f14d7ddb593b4b1bc7902e575160b19649bdd13858727ad551f76fe +size 28640 diff --git a/examples/I2V/02_chestnut/image.png b/examples/I2V/02_chestnut/image.png new file mode 100644 index 0000000000000000000000000000000000000000..6e0b3204042435681346d5dd4457862b1d0ec77d --- /dev/null +++ b/examples/I2V/02_chestnut/image.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cb74295a4b8e83ac934fc1e9e65898ca4a12087a4e3bb4016a06f14bf2b3fecc +size 320275 diff --git a/examples/I2V/02_chestnut/prompt.txt b/examples/I2V/02_chestnut/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..b424821f4dfa9492711228b76404698cf2be8af4 --- /dev/null +++ b/examples/I2V/02_chestnut/prompt.txt @@ -0,0 +1 @@ +A chestnut horse stands in a rural paddock with its ears upright and its attention directed forward. The horse has a broad irregular white blaze running down its face, a dark muzzle, large alert eyes, short whiskers, and a tousled black forelock between its ears. Its reddish-brown coat continues across the neck and shoulders. A rough field, low stable buildings, fencing, distant trees, and wooded hills surround the animal, creating a simple working-farm environment. The animal stands within a complete farm landscape of worn ground, fences, low buildings, open field, mature trees, and wooded hills extending behind the paddock. The scene remains spatially coherent from nearby surfaces and vegetation to the architecture, terrain, and distant boundaries of the environment. diff --git a/examples/I2V/06_waterfall/actions.txt b/examples/I2V/06_waterfall/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..82cced6e512414d90e89eb6b05f96b8ca79edc78 --- /dev/null +++ b/examples/I2V/06_waterfall/actions.txt @@ -0,0 +1,6 @@ +@dtype float32 + +forward1x4 +yaw_right45x4 +forward1x4 +yaw_right45x4 diff --git a/examples/I2V/06_waterfall/camera.npy b/examples/I2V/06_waterfall/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..e84545203f49ae528d58a48979c9d67bd2be19ee --- /dev/null +++ b/examples/I2V/06_waterfall/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:14c7a1dfd01811dbd1625aa3c53c4d76895e4a8ba8dbd7cd732d72987bfcca82 +size 25472 diff --git a/examples/I2V/06_waterfall/image.png b/examples/I2V/06_waterfall/image.png new file mode 100644 index 0000000000000000000000000000000000000000..ee970346c88f69b9c5df302d01037b2d801bfbf1 --- /dev/null +++ b/examples/I2V/06_waterfall/image.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e6a4d699e18b051b03491719d158b9db4cf96d3712a0d968d7506e75a213bbb8 +size 485643 diff --git a/examples/I2V/06_waterfall/prompt.txt b/examples/I2V/06_waterfall/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..3a4cb395531981a1ccaf01091911a838779bbf42 --- /dev/null +++ b/examples/I2V/06_waterfall/prompt.txt @@ -0,0 +1 @@ +A broad garden waterfall pours over layered dark rocks into a shallow pool surrounded by dense subtropical plants. Several parallel curtains of water descend from a ledge beneath mossy boulders, with smaller channels passing between stones and clumps of grass. Pines, broad-leaf shrubs, ferns, and long narrow leaves grow around the banks and across the rock formation. A large flat stone borders the pool on one side, while additional boulders form a natural boundary behind the falling water. The compact arrangement of water, stone, varied foliage, and concealed pond edges resembles a carefully designed botanical garden feature. diff --git a/examples/I2V/10_case061/actions.txt b/examples/I2V/10_case061/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..460daa63313e22abc9400016be01f66105ff5a2b --- /dev/null +++ b/examples/I2V/10_case061/actions.txt @@ -0,0 +1,8 @@ +forward1x4 +yaw_right45x4 +forward1x4 +yaw_right45x4 +forward1.5x6 +pitch_up30 +forward1.5x6 +pitch_down45 diff --git a/examples/I2V/10_case061/camera.npy b/examples/I2V/10_case061/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..5b0ff621bd1a2a46c7bbcdaa25983377a6dcb725 --- /dev/null +++ b/examples/I2V/10_case061/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e61560bb1fb82b50b547e46303426927efe8b5234945f3de563da478a38b2811 +size 95168 diff --git a/examples/I2V/10_case061/image.png b/examples/I2V/10_case061/image.png new file mode 100644 index 0000000000000000000000000000000000000000..c98877556a8da58ff0821f4372d92a363e440bbf --- /dev/null +++ b/examples/I2V/10_case061/image.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b7335efe74c5f244228324014703ac48596a35944639058bc78d40fc9adad5cc +size 395425 diff --git a/examples/I2V/10_case061/prompt.txt b/examples/I2V/10_case061/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..050eda004a482b9cd9a928b6d9a9ff229ba746b3 --- /dev/null +++ b/examples/I2V/10_case061/prompt.txt @@ -0,0 +1 @@ +A third-person trailing gameplay-like view closely follows a majestic blue Amazonian parrot as it flies high through the vibrant Amazon rainforest on a beautiful clear sunny day. The parrot has layered cobalt feathers, a golden throat patch, a curved black beak, and broad articulated wings. Below, a winding river divides dense tiers of palms, ceiba trees, hanging vines, and exposed roots. An ancient stepped stone temple occupies a clearing ahead, its terraces cracked and overgrown with moss, orchids, and tangled foliage. Small birds cross between the treetops, mist gathers over distant ridges, and the riverbank contains fallen trunks, ferns, and scattered stone fragments. diff --git a/examples/I2V/13_burrow/actions.txt b/examples/I2V/13_burrow/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..05c3a7d49a6dcfc3d9902d2739633a0a563192fb --- /dev/null +++ b/examples/I2V/13_burrow/actions.txt @@ -0,0 +1,12 @@ +@dtype float32 +@sampling smooth_turns + +forward1x2 +left1x2 +yaw_left30x3 +left1x3 +yaw_right30x3 +right1x5 +yaw_right30x4 +yaw_left30x3 +reverse_frames25 diff --git a/examples/I2V/13_burrow/camera.npy b/examples/I2V/13_burrow/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..ee1bb1fee243e7ad3d53d5f90175efd094fe3edd --- /dev/null +++ b/examples/I2V/13_burrow/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:78a589f811cfc190ad00ee109387466c853a0c270321115fe46f18e76f59e3d3 +size 79328 diff --git a/examples/I2V/13_burrow/image.png b/examples/I2V/13_burrow/image.png new file mode 100644 index 0000000000000000000000000000000000000000..8b8ccfdb1e6fe1b213e39e46ff1117867807dfd5 --- /dev/null +++ b/examples/I2V/13_burrow/image.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9d51c767a7faaf71a14e4acd7765a0c2ba11473db2d5445db3c969486c8e508f +size 582180 diff --git a/examples/I2V/13_burrow/prompt.txt b/examples/I2V/13_burrow/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..c2b217e81c2d555923eb72fef677abb1990325e5 --- /dev/null +++ b/examples/I2V/13_burrow/prompt.txt @@ -0,0 +1 @@ +A first-person view of a whimsical earthen cottage built directly into a grassy hillside. Uneven stone steps cross the foreground through thick lawn, leafy plants, and clusters of purple and pink flowers. A large round green wooden door sits beneath a broad brick arch in the middle ground, flanked by two circular divided windows and curved timber supports. Ivy and dense shrubs cover much of the plaster facade and turf roof, while mature branches spread overhead. A bright orange pumpkin rests at the left edge, and damp greenery, weathered wood, masonry, and soft daylight give the dwelling a secluded rural character. diff --git a/examples/I2V/15_case104/actions.txt b/examples/I2V/15_case104/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..ce5700c43667c7689c801b35779a056bd73ab52a --- /dev/null +++ b/examples/I2V/15_case104/actions.txt @@ -0,0 +1,4 @@ +forward1.5x10 +yaw_right45x4 +yaw_left45x4 +forward1.5x11 diff --git a/examples/I2V/15_case104/camera.npy b/examples/I2V/15_case104/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..0113014f538f8144361fe63e0ebbc3748b4e969a --- /dev/null +++ b/examples/I2V/15_case104/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b914b0f7552133204adf037b92db1ba99b24abd683e8c3ce0637f811ded64612 +size 92000 diff --git a/examples/I2V/15_case104/image.png b/examples/I2V/15_case104/image.png new file mode 100644 index 0000000000000000000000000000000000000000..a325e5d9b46c65aa13f93e310b42209866c14523 --- /dev/null +++ b/examples/I2V/15_case104/image.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:59cb1cd4b961dab1a84ff5f3d094ac3923971946d772c0a1b205e2a258fbd29e +size 498851 diff --git a/examples/I2V/15_case104/prompt.txt b/examples/I2V/15_case104/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..19d27f71a7585f90c19b415c8f5c96ddf3f7b0ce --- /dev/null +++ b/examples/I2V/15_case104/prompt.txt @@ -0,0 +1 @@ +A third-person fantasy gameplay view follows a small winged fairy exploring a floating sky island. The fairy has long golden hair, a short white dress, delicate translucent wings with iridescent veins, and a light, graceful silhouette. Lush grass and colorful wildflowers cover the island's rocky surface, while waterfalls descend from sheer edges into layers of soft clouds. A branching crystal tree carries luminous jewel-like fruit, and distant floating islands create a broad aerial landscape. Glowing vines, mossy stones, and scattered blossoms add detail to the traversable ground. Vibrant anime colors and soft atmospheric light connect the fairy to the surrounding fantasy world. The view stays closely linked to the fairy while preserving nearby terrain, cliff edges, trees, waterfalls, and distant islands as one continuous explorable environment. diff --git a/examples/README.md b/examples/README.md new file mode 100644 index 0000000000000000000000000000000000000000..01d86e9597c9204268a028630980ded2830e66d8 --- /dev/null +++ b/examples/README.md @@ -0,0 +1,148 @@ +# Camera and prompt guide + +Each example contains `prompt.txt`, `camera.npy`, and `actions.txt`. I2V examples +also include `image.png`. The shared `negative_prompt.txt` is loaded by default. +Tokyo street includes its original `negative_prompt.txt`; select it with +`--negative-prompt-path test/T2V/02_tokyo_street/negative_prompt.txt`. +Run the commands below from the repository root. + +## Choose an example + +| Mode | Example | Description | +| --- | --- | --- | +| I2V | [Cat](I2V/00_cat_vac) | Default; a cat riding a moving robot vacuum | +| I2V | [Socrates](I2V/01_socrates) | Motionless painted sculptures in a stone chamber | +| T2V | [Red balloon](T2V/00_red_balloon) | Default; a balloon floating through an abandoned street | +| T2V | [Tokyo street](T2V/02_tokyo_street) | A woman walking through a neon-lit street | + +Run the Tokyo street example with its original prompt and negative prompt: + +```bash +python inference.py --model-type fast --mode t2v \ + --prompt-path test/T2V/02_tokyo_street/prompt.txt \ + --negative-prompt-path test/T2V/02_tokyo_street/negative_prompt.txt \ + --actions-file test/T2V/02_tokyo_street/actions.txt +``` + +Additional examples: + +| Mode | Cases | +| --- | --- | +| I2V | `02_chestnut`, `06_waterfall`, `10_case061`, `13_burrow`, `15_case104` | +| T2V | `01_t2v-mind131-00` | + +## Camera inputs + +Choose either the saved poses or the action description for the same example: + +```bash +python inference.py --model-type fast \ + --image-path test/I2V/01_socrates/image.png \ + --prompt-path test/I2V/01_socrates/prompt.txt \ + --camera-path test/I2V/01_socrates/camera.npy +``` + +Replace the last argument with `--actions-file test/I2V/01_socrates/actions.txt` +to generate the poses from actions. For T2V, use `--mode t2v`, omit `--image-path`, +and select a T2V example's prompt and trajectory. + +### Write actions + +```text +forward1x2 +yaw_left30x3 +backward1 +``` + +This generates six chunks: two forward moves, three left turns, and one backward +move. Each chunk has 33 frames. Movement values are distances; rotation values +are degrees. Use `--num-chunks` to run only the beginning of a sequence. + +| Movement | Actions | Short forms | +| --- | --- | --- | +| Forward / backward | `forward1`, `backward1` | `f1`, `b1` | +| Left / right | `left1`, `right1` | `l1`, `r1` | +| Up / down | `up1`, `down1` | Same | +| Turn left / right | `yaw_left30`, `yaw_right30` | `yl30`, `yr30` | +| Look up / down | `pitch_up15`, `pitch_down15` | `pu15`, `pd15` | + +The camera starts at the origin, facing +Z, with +X to the right and +Y down. +Forward/backward and left/right follow its heading on the horizontal plane; +pitch does not change movement height. Up/down follows the world vertical axis. +Yaw turns in place. Keep the total translation distance per chunk at most 5; +split longer movements into repeated actions. + +Use spaces, commas, or newlines between actions, and `#` for comments. `xN` +repeats an action. `&` combines movements and rotations in one chunk, such as +`forward2&right2&yaw_left45`; translation follows the heading at the chunk's +start. `reverseN` retraces the preceding N chunks. `reverse_framesN` replays their +sampled poses in reverse frame order. Each reverse command generates N chunks. + +Some examples include headers to preserve their original sampling: + +| Header | Meaning | +| --- | --- | +| `@dtype float32` | Store poses in float32 instead of the default float64 | +| `@sampling smooth_turns` | Ease motion at action changes instead of using linear sampling | +| `@last_frame include` | Include the final endpoint instead of excluding it | + +Keep these headers when reproducing an example. To build poses separately: + +```bash +python tools/build_trajectory.py \ + --actions-file test/I2V/01_socrates/actions.txt \ + --output-dir output/socrates_camera +``` + +### Supply camera poses + +`camera.npy` stores global camera-to-world matrices with shape `[T, 3, 4]` or +`[T, 4, 4]`, using the same right/down/forward convention. Supply one pose per +frame and 33 frames per chunk, with translations in the model's metric scale. +The inference code derives the internal camera representations; do not +pre-normalize the file separately for UCPE or RepEncoder. + +## Prompt styles + +### Dynamic subjects: describe following and motion + +For a moving subject that should stay in view, begin with +**“A third-person ... view closely follows ...”**. This encourages subject +following; it is a prompt cue, not a tracking constraint. Describe the subject's +appearance, its movement, and how it interacts with the surroundings. Keep +nearby obstacles and background landmarks identifiable as the subject moves. + +The [Cat prompt](I2V/00_cat_vac/prompt.txt) starts: + +> A third-person gameplay-like camera closely follows a gray robot vacuum moving through a modern interior with reflective hardwood floors and beautiful rays of light. + +It then describes the cat, the vacuum, the furniture, and how the cat balances +during movement. Adapt the opening to the subject, for example +“A third-person trailing view closely follows a cyclist ...”. +For an environment with moving water or foliage but no followed subject, use +the scene-focused style below and describe that environmental motion directly. + +### Static scenes: describe space and fixed appearance + +Describe the scene as a coherent environment: its layout, foreground and +background, materials, lighting, and relationships between objects. Camera +motion comes from the trajectory. Avoid adding subject movement when the scene +should remain static. + +The [Socrates prompt](I2V/01_socrates/prompt.txt) identifies the people as +**static, painted sculptures** and explicitly says that all figures remain +motionless, with rigid poses and fixed garment folds. This helps distinguish +lifelike sculptures from living people. For an ordinary room or landscape, +describe its actual contents rather than calling everything a sculpture. + +### Length and consistency + +Use one focused English paragraph. Around **80–120 words** is a useful starting +point; dynamic subject-following prompts often need **100–130 words** to cover +both motion and environment. These are writing guidelines, not input limits. + +For I2V, keep the description consistent with the input image. For T2V, describe +the subject and setting explicitly because there is no starting image. Keep +appearance and lighting consistent throughout the paragraph, and avoid cuts, +shot changes, or camera directions that compete with the supplied trajectory. +The bundled prompts preserve the wording used for their original examples. diff --git a/examples/T2V/00_red_balloon/actions.txt b/examples/T2V/00_red_balloon/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..05c3a7d49a6dcfc3d9902d2739633a0a563192fb --- /dev/null +++ b/examples/T2V/00_red_balloon/actions.txt @@ -0,0 +1,12 @@ +@dtype float32 +@sampling smooth_turns + +forward1x2 +left1x2 +yaw_left30x3 +left1x3 +yaw_right30x3 +right1x5 +yaw_right30x4 +yaw_left30x3 +reverse_frames25 diff --git a/examples/T2V/00_red_balloon/camera.npy b/examples/T2V/00_red_balloon/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..ee1bb1fee243e7ad3d53d5f90175efd094fe3edd --- /dev/null +++ b/examples/T2V/00_red_balloon/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:78a589f811cfc190ad00ee109387466c853a0c270321115fe46f18e76f59e3d3 +size 79328 diff --git a/examples/T2V/00_red_balloon/prompt.txt b/examples/T2V/00_red_balloon/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..aa7df065097db9f711532ddab6c36fcaf7fb7bdd --- /dev/null +++ b/examples/T2V/00_red_balloon/prompt.txt @@ -0,0 +1 @@ +A third person view closely follows a red balloon floating above the ground in an abandoned street. The balloon drifts gracefully, its bright red color contrasting sharply against the decaying urban backdrop. The street is littered with debris and graffiti-covered walls, with broken windows and rusted cars scattered about. Shadows dance across the scene as sunlight filters through gaps in the buildings. The camera moves fluidly, capturing the balloon's gentle ascent and descent, emphasizing its playful motion. A close-up of the balloon transitions to a wider shot, showcasing the desolate environment. diff --git a/examples/T2V/01_t2v-mind131-00/actions.txt b/examples/T2V/01_t2v-mind131-00/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..fa50361b65bd1eb853004b8fce6f97246ff435dc --- /dev/null +++ b/examples/T2V/01_t2v-mind131-00/actions.txt @@ -0,0 +1,8 @@ +@dtype float32 + +left2x2 +right2x4 +left2x2 +forward2x4 +backward2x4 +yaw_left45x8 diff --git a/examples/T2V/01_t2v-mind131-00/camera.npy b/examples/T2V/01_t2v-mind131-00/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..7415c38856716ce0ffff91b96fca45e0da1daa57 --- /dev/null +++ b/examples/T2V/01_t2v-mind131-00/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4519540636c34443ccc1d0a4ba60984828de07c065a608dd6b0f1ffae0a71827 +size 38144 diff --git a/examples/T2V/01_t2v-mind131-00/prompt.txt b/examples/T2V/01_t2v-mind131-00/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..b3ec45db40d61985111bd0b55f39076f061da663 --- /dev/null +++ b/examples/T2V/01_t2v-mind131-00/prompt.txt @@ -0,0 +1 @@ +A third person view closely follows a vibrant tropical fish swimming gracefully among colorful coral reefs in a clear, turquoise ocean. The fish has bright blue and yellow scales with a small, distinctive orange spot on its side, its fins moving fluidly. The coral reefs are alive with a variety of marine life, including small schools of colorful fish and sea turtles gliding by. The water is crystal clear, allowing for a view of the sandy ocean floor below. The reef itself is adorned with a mix of hard and soft corals in shades of red, orange, and green. The camera moves smoothly with the fish, keeping it in view as it weaves between the coral formations and explores the reef. diff --git a/examples/T2V/02_tokyo_street/actions.txt b/examples/T2V/02_tokyo_street/actions.txt new file mode 100644 index 0000000000000000000000000000000000000000..82cced6e512414d90e89eb6b05f96b8ca79edc78 --- /dev/null +++ b/examples/T2V/02_tokyo_street/actions.txt @@ -0,0 +1,6 @@ +@dtype float32 + +forward1x4 +yaw_right45x4 +forward1x4 +yaw_right45x4 diff --git a/examples/T2V/02_tokyo_street/camera.npy b/examples/T2V/02_tokyo_street/camera.npy new file mode 100644 index 0000000000000000000000000000000000000000..e84545203f49ae528d58a48979c9d67bd2be19ee --- /dev/null +++ b/examples/T2V/02_tokyo_street/camera.npy @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:14c7a1dfd01811dbd1625aa3c53c4d76895e4a8ba8dbd7cd732d72987bfcca82 +size 25472 diff --git a/examples/T2V/02_tokyo_street/negative_prompt.txt b/examples/T2V/02_tokyo_street/negative_prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..d159af17a18cf57ece66747d39cd178bb035ca45 --- /dev/null +++ b/examples/T2V/02_tokyo_street/negative_prompt.txt @@ -0,0 +1 @@ +overexposed, static, blurred details, subtitles, style, artwork, painting, picture, still, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, malformed limbs, fused fingers, motionless picture, messy background, three legs, many people in the background, walking backwards diff --git a/examples/T2V/02_tokyo_street/prompt.txt b/examples/T2V/02_tokyo_street/prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..a99d6172193d9787ce16470e72de65049ba1ce74 --- /dev/null +++ b/examples/T2V/02_tokyo_street/prompt.txt @@ -0,0 +1 @@ +A stylish woman strolls down a bustling Tokyo street, the warm glow of neon lights and animated city signs casting vibrant reflections. She wears a sleek black leather jacket paired with a flowing red dress and black boots, her black purse slung over her shoulder. Sunglasses perched on her nose and a bold red lipstick add to her confident, casual demeanor. The street is damp and reflective, creating a mirror-like effect that enhances the colorful lights and shadows. Pedestrians move about, adding to the lively atmosphere. The scene is captured in a dynamic medium shot with the woman walking slightly to one side, highlighting her graceful strides. diff --git a/examples/negative_prompt.txt b/examples/negative_prompt.txt new file mode 100644 index 0000000000000000000000000000000000000000..6341e5d503278ece814c38bc98325ce2a410869f --- /dev/null +++ b/examples/negative_prompt.txt @@ -0,0 +1 @@ +Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..8d25d3ebfece4b8c5e3ebfe583da462de19840ea --- /dev/null +++ b/requirements.txt @@ -0,0 +1,18 @@ +torch==2.10.0 +torchvision==0.25.0 +diffusers==0.37.0 +transformers==5.3.0 +accelerate==1.12.0 +peft==0.18.1 +kernels==0.13.0 +timm==1.0.25 +safetensors +numpy<2.0.0 +Pillow +imageio==2.37.3 +imageio-ffmpeg==0.6.0 +ftfy +regex +einops +packaging +sentencepiece diff --git a/worldcrafter/__init__.py b/worldcrafter/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..2f525c7be5771f109ec7cf396f9e902bc1248fd5 --- /dev/null +++ b/worldcrafter/__init__.py @@ -0,0 +1,27 @@ +from importlib import import_module + + +__all__ = [ + "InferenceResult", + "RepEncoder", + "WorldCrafter", + "WorldCrafterPipeline", + "WorldCrafterScheduler", + "WorldCrafterTransformer3DModel", +] + + +def __getattr__(name): + modules = { + "InferenceResult": ".inference", + "WorldCrafter": ".inference", + "RepEncoder": ".repencoder", + "WorldCrafterPipeline": ".diffusers", + "WorldCrafterScheduler": ".diffusers", + "WorldCrafterTransformer3DModel": ".diffusers", + } + if name not in modules: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + value = getattr(import_module(modules[name], __name__), name) + globals()[name] = value + return value diff --git a/worldcrafter/camera.py b/worldcrafter/camera.py new file mode 100644 index 0000000000000000000000000000000000000000..d2de57d3ff2e9ff096ab0d516a234b532297e56f --- /dev/null +++ b/worldcrafter/camera.py @@ -0,0 +1,349 @@ +"""Camera actions in a right/down/forward coordinate system.""" + +from __future__ import annotations + +from dataclasses import asdict, dataclass +import hashlib +import json +import math +from pathlib import Path +import re +from functools import partial + +import numpy as np + + +CHUNK_FRAMES = 33 +FPS = 16 +MAX_TRANSLATION = 5.0 +ACTION_FIELDS = { + "forward": ("forward", 1), + "backward": ("forward", -1), + "left": ("right", -1), + "right": ("right", 1), + "up": ("up", 1), + "down": ("up", -1), + "yaw_left": ("yaw", -1), + "yaw_right": ("yaw", 1), + "pitch_up": ("pitch", 1), + "pitch_down": ("pitch", -1), +} +ALIASES = { + "f": "forward", "b": "backward", "l": "left", "r": "right", + "yl": "yaw_left", "yr": "yaw_right", "pu": "pitch_up", "pd": "pitch_down", +} + + +@dataclass(frozen=True) +class Action: + forward: float = 0.0 + right: float = 0.0 + yaw: float = 0.0 + pitch: float = 0.0 + speed: float = 1.0 + up: float = 0.0 + + def validate(self): + values = (self.forward, self.right, self.up, self.yaw, self.pitch) + if not all(math.isfinite(v) for v in (*values, self.speed)): + raise ValueError("Control values must be finite") + if sum(v != 0 for v in values) > 1: + raise ValueError("Only one movement or rotation may be active per chunk") + + def normalized(self): + """Apply the interactive controls' slider limits.""" + self.validate() + return Action( + forward=max(-1.0, min(1.0, self.forward)), + right=max(-1.0, min(1.0, self.right)), + up=max(-1.0, min(1.0, self.up)), + yaw=max(-30.0, min(30.0, self.yaw)), + pitch=max(-30.0, min(30.0, self.pitch)), + speed=max(0.1, min(MAX_TRANSLATION, self.speed)), + ) + + def json(self): + return asdict(self) + + +def parse_event(event: str) -> tuple[str, float]: + match = re.fullmatch(r"([a-z_]+)([0-9]+(?:\.[0-9]+)?)", event.lower()) + if match is None: + raise ValueError(f"Invalid action {event!r}; use e.g. forward1 or yaw_left30") + name, value = match.groups() + name = ALIASES.get(name, name) + if name not in ACTION_FIELDS: + raise ValueError(f"Unknown action {name!r}; choose from {', '.join(ACTION_FIELDS)}") + amount = float(value) + if not math.isfinite(amount): + raise ValueError(f"Action amount must be finite: {event}") + if ACTION_FIELDS[name][0] in {"forward", "right", "up"} and amount > MAX_TRANSLATION: + raise ValueError(f"{event}: translation must not exceed {MAX_TRANSLATION:g} per chunk") + return name, amount + + +def parse_actions(text: str) -> list[str]: + """Expand space/comma-separated actions, xN repetitions, and # comments.""" + text = re.sub(r"#[^\n]*", "", text) + text = re.sub(r"\s*&\s*", "&", text) + events = [] + for token in re.split(r"[\s,]+", text.strip()): + if not token: + continue + match = re.fullmatch(r"(.+?)(?:x([1-9][0-9]*))?", token.lower()) + event, repeat = match.groups() + if re.fullmatch(r"reverse(?:_frames)?[1-9][0-9]*", event): + canonical = event + else: + parts = [] + axes = set() + for component in event.split("&"): + name, _ = parse_event(component) + field = ACTION_FIELDS[name][0] + if field in axes: + raise ValueError(f"An action may use each axis only once: {event}") + axes.add(field) + value = re.search(r"[0-9].*", component).group() + parts.append(name + value) + canonical = "&".join(parts) + events.extend([canonical] * int(repeat or 1)) + if not events: + raise ValueError("Provide at least one camera action") + return events + + +def parse_trajectory(text: str) -> tuple[list[str], dict[str, str]]: + """Read actions and optional @dtype, @sampling, and @last_frame headers.""" + choices = { + "dtype": {"float32", "float64"}, + "sampling": {"linear", "smooth_turns"}, + "last_frame": {"exclude", "include"}, + } + options, lines = {}, [] + for line in text.splitlines(): + line = line.split("#", 1)[0].strip() + if line.startswith("@"): + fields = line[1:].split() + if len(fields) != 2 or fields[0] not in choices or fields[1] not in choices[fields[0]]: + raise ValueError(f"Invalid trajectory setting: {line}") + if lines: + raise ValueError("Trajectory settings must precede the actions") + options[fields[0]] = fields[1] + elif line: + lines.append(line) + return parse_actions("\n".join(lines)), options + + +def count_chunks(events: list[str]) -> int: + return sum( + int(re.search(r"[0-9]+$", event).group()) if event.startswith("reverse") else 1 + for event in events + ) + + +def action_from_event(event: str) -> Action: + name, amount = parse_event(event) + field, sign = ACTION_FIELDS[name] + if field in {"yaw", "pitch"}: + return Action(**{field: sign * amount}) + return Action(**{field: sign}, speed=amount) + + +def rotation_y(degrees: float) -> np.ndarray: + angle = math.radians(degrees) + cosine, sine = math.cos(angle), math.sin(angle) + return np.asarray( + ((cosine, 0.0, sine), (0.0, 1.0, 0.0), (-sine, 0.0, cosine)), + dtype=np.float64, + ) + + +def rotation_x(degrees: float) -> np.ndarray: + angle = math.radians(degrees) + cosine, sine = math.cos(angle), math.sin(angle) + return np.asarray( + ((1.0, 0.0, 0.0), (0.0, cosine, -sine), (0.0, sine, cosine)), + dtype=np.float64, + ) + + +def horizontal_direction(rotation: np.ndarray, forward: bool) -> np.ndarray: + axis = rotation[:, 2 if forward else 0].copy() + axis[1] = 0.0 + norm = np.linalg.norm(axis) + if norm < 1e-8: + # Keep a horizontal heading when looking straight up or down. + other = rotation[:, 0 if forward else 2] + axis = np.array([-other[2], 0.0, other[0]]) + if not forward: + axis = -axis + norm = np.linalg.norm(axis) + return axis / norm + + +def sample_chunk( + start: np.ndarray, action: Action, *, fractions: np.ndarray | None = None, +) -> tuple[np.ndarray, np.ndarray]: + """Sample one action; the logical endpoint starts the next chunk.""" + action.validate() + alpha = np.arange(CHUNK_FRAMES, dtype=np.float64) / CHUNK_FRAMES if fractions is None else fractions + poses = np.repeat(start[None], len(alpha), axis=0) + end = start.copy() + if action.yaw or action.pitch: + def rotated(fraction): + if action.yaw: + return rotation_y(action.yaw * fraction) @ start[:3, :3] + return start[:3, :3] @ rotation_x(action.pitch * fraction) + + for index, fraction in enumerate(alpha): + poses[index, :3, :3] = rotated(fraction) + end[:3, :3] = rotated(1.0) + else: + if action.up: + direction = np.array([0.0, -action.up, 0.0]) + elif action.forward: + direction = action.forward * horizontal_direction(start[:3, :3], True) + else: + direction = action.right * horizontal_direction(start[:3, :3], False) + delta = action.speed * direction + if np.linalg.norm(delta) > MAX_TRANSLATION + 1e-12: + raise ValueError(f"Translation must not exceed {MAX_TRANSLATION:g} per chunk") + poses[:, :3, 3] = start[:3, 3] + alpha[:, None] * delta + end[:3, 3] = start[:3, 3] + delta + return poses, end + + +def _sample_event(start: np.ndarray, event: str, fractions: np.ndarray) -> np.ndarray: + components = event.split("&") + if len(components) == 1: + return sample_chunk(start, action_from_event(event), fractions=fractions)[0] + poses = np.repeat(start[None], len(fractions), axis=0) + delta = np.zeros(3) + for component in components: + action = action_from_event(component) + if action.yaw: + for index, fraction in enumerate(fractions): + poses[index, :3, :3] = rotation_y(action.yaw * fraction) @ poses[index, :3, :3] + elif action.pitch: + for index, fraction in enumerate(fractions): + poses[index, :3, :3] = poses[index, :3, :3] @ rotation_x(action.pitch * fraction) + else: + _, end = sample_chunk(start, action) + delta += end[:3, 3] - start[:3, 3] + if np.linalg.norm(delta) > MAX_TRANSLATION + 1e-12: + raise ValueError(f"{event}: combined translation must not exceed {MAX_TRANSLATION:g} per chunk") + poses[:, :3, 3] = start[:3, 3] + fractions[:, None] * delta + return poses + + +def _sample_curve(start, event, tangent_start, tangent_end, times): + fractions = times + if tangent_start is not None: + fractions = ( + -2 * times**3 + 3 * times**2 + + (times**3 - 2 * times**2 + times) * tangent_start + + (times**3 - times**2) * tangent_end + ) + return _sample_event(start, event, fractions) + + +def _reverse_curve(curve, times): + return curve(1.0 - times) + + +def _reverse_sampled_curve(curve, last_time, times): + return curve(last_time * (1.0 - times)) + + +def build_trajectory( + events: list[str], *, dtype: str = "float64", sampling: str = "linear", + last_frame: str = "exclude", +) -> tuple[np.ndarray, list[dict[str, object]]]: + if not events: + raise ValueError("Provide at least one camera action") + world = np.eye(4, dtype=np.float64) + chunks, records, curves, sample_times = [], [], [], [] + + def append(event, curve, times, poses=None): + nonlocal world + start, end = curve(np.array([0.0, 1.0])) + chunks.append(curve(times) if poses is None else poses) + records.append({ + "chunk_index": len(records), "event": event, + "logical_start_c2w": start.tolist(), "logical_end_c2w": end.tolist(), + }) + curves.append(curve) + sample_times.append(times) + world = end + + for index, event in enumerate(events): + if event.startswith("reverse"): + count = int(re.search(r"[0-9]+$", event).group()) + if count > len(chunks): + raise ValueError(f"{event} needs {count} preceding chunks; only {len(chunks)} exist") + indices = list(range(len(chunks) - 1, len(chunks) - count - 1, -1)) + for source in indices: + if event.startswith("reverse_frames"): + curve = partial(_reverse_sampled_curve, curves[source], sample_times[source][-1]) + times = np.linspace(0.0, 1.0, CHUNK_FRAMES) + append(event, curve, times, chunks[source][::-1].copy()) + else: + curve = partial(_reverse_curve, curves[source]) + include_end = last_frame == "include" and index == len(events) - 1 and source == indices[-1] + times = ( + np.linspace(0.0, 1.0, CHUNK_FRAMES) if include_end + else np.arange(CHUNK_FRAMES, dtype=np.float64) / CHUNK_FRAMES + ) + poses = None + original = records[source]["event"] + if sampling == "linear" and "&" not in original and not original.startswith("reverse"): + action = action_from_event(original) + if not action.yaw and not action.pitch and sample_times[source][-1] < 1.0 and not include_end: + # Reuse linear translation samples without another interpolation roundoff. + endpoint = np.array(records[source]["logical_end_c2w"]) + poses = np.concatenate([endpoint[None], chunks[source][1:][::-1]]) + append(event, curve, times, poses) + continue + smooth = sampling == "smooth_turns" + entering = index > 0 and events[index - 1] == event + leaving = index + 1 < len(events) and events[index + 1] == event + include_end = (smooth and not leaving) or (last_frame == "include" and index == len(events) - 1) + times = ( + np.linspace(0.0, 1.0, CHUNK_FRAMES) if include_end + else np.arange(CHUNK_FRAMES, dtype=np.float64) / CHUNK_FRAMES + ) + curve = partial( + _sample_curve, world.copy(), event, + float(entering) if smooth else None, float(leaving), + ) + append(event, curve, times) + return np.concatenate(chunks).astype(dtype), records + + +def save_trajectory( + directory: Path, camera: np.ndarray, records: list[dict[str, object]], *, fps: int = FPS, + events: list[str] | None = None, options: dict[str, str] | None = None, +) -> Path: + directory.mkdir(parents=True, exist_ok=True) + camera_path = directory / "camera.npy" + np.save(camera_path, camera[:, :3, :4]) + events = events if events is not None else [record["event"] for record in records] + headers = [f"@{key} {value}" for key, value in (options or {}).items()] + (directory / "actions.txt").write_text("\n".join(headers + events) + "\n", encoding="utf-8") + manifest = { + "format": "worldcrafter_camera_trajectory_v1", + "fps": fps, + "chunk_frames": CHUNK_FRAMES, + "num_chunks": len(records), + "num_frames": len(camera), + "motion_sequence": events, + "options": options or {}, + "camera": camera_path.name, + "camera_semantics": "global metric c2w; x right, y down, z forward", + "sha256": {"camera": hashlib.sha256(camera_path.read_bytes()).hexdigest()}, + "chunks": records, + } + (directory / "trajectory.json").write_text( + json.dumps(manifest, indent=2) + "\n", encoding="utf-8", + ) + return camera_path diff --git a/worldcrafter/cli.py b/worldcrafter/cli.py new file mode 100644 index 0000000000000000000000000000000000000000..576b22a3c33ec07eec8babfa1fb287460c5320c8 --- /dev/null +++ b/worldcrafter/cli.py @@ -0,0 +1,166 @@ +from __future__ import annotations + +import argparse +from datetime import datetime +from uuid import uuid4 +from pathlib import Path +from typing import Sequence + + +ROOT = Path(__file__).resolve().parents[1] +DEFAULT_MODEL = ROOT / "weights" / "WorldCrafter-Base" +DEFAULT_I2V_CASE = ROOT / "test" / "I2V" / "00_cat_vac" +DEFAULT_T2V_CASE = ROOT / "test" / "T2V" / "00_red_balloon" +DEFAULT_IMAGE = DEFAULT_I2V_CASE / "image.png" +DEFAULT_I2V_CAMERA = DEFAULT_I2V_CASE / "camera.npy" +DEFAULT_T2V_CAMERA = DEFAULT_T2V_CASE / "camera.npy" +DEFAULT_I2V_PROMPT = DEFAULT_I2V_CASE / "prompt.txt" +DEFAULT_T2V_PROMPT = DEFAULT_T2V_CASE / "prompt.txt" +DEFAULT_NEGATIVE_PROMPT = ROOT / "test" / "negative_prompt.txt" + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="WorldCrafter camera-controlled image-to-video and text-to-video inference" + ) + parser.add_argument("--mode", choices=("i2v", "t2v"), default="i2v") + parser.add_argument("--model-type", choices=("base", "fast"), default="base") + parser.add_argument("--model-path", type=Path) + parser.add_argument( + "--local-camera-path", + type=Path, + help="Optional precomputed chunk-local UCPE poses for fast; --camera-path always supplies global metric poses", + ) + parser.add_argument("--image-path", type=Path) + camera = parser.add_mutually_exclusive_group() + camera.add_argument("--camera-path", type=Path, help="Global c2w trajectory (.npy)") + camera.add_argument("--actions", help='Camera actions, e.g. "forward1x2 yaw_left30x3 backward1"') + camera.add_argument("--actions-file", type=Path, help="TXT file of camera actions") + parser.add_argument("--prompt") + parser.add_argument("--prompt-path", type=Path) + parser.add_argument("--negative-prompt") + parser.add_argument( + "--negative-prompt-path", type=Path, default=DEFAULT_NEGATIVE_PROMPT + ) + parser.add_argument("--output-path", type=Path) + parser.add_argument("--chunk-output-dir", type=Path) + parser.add_argument("--state-output-dir", type=Path) + parser.add_argument("--resume-from", type=Path) + parser.add_argument("--num-chunks", type=int) + parser.add_argument("--stop-after-chunk", type=int) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--height", type=int, default=384) + parser.add_argument("--width", type=int, default=640) + parser.add_argument("--num-inference-steps", type=int) + parser.add_argument("--guidance-scale", type=float) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--fps", type=int, default=16) + parser.add_argument("--image-noise-sigma-min", type=float, default=0.111) + parser.add_argument("--image-noise-sigma-max", type=float, default=0.135) + parser.add_argument("--camera-x-fov", type=float, default=100.0) + parser.add_argument("--camera-xi", type=float, default=0.0) + parser.add_argument("--memory-fov-h-deg", type=float, default=100.0) + parser.add_argument("--memory-fov-v-deg", type=float, default=71.13349068444832) + parser.add_argument("--memory-fov-samples-per-axis", type=int, default=10) + parser.add_argument( + "--attention-backend", + choices=("native", "auto", "flash_hub", "_flash_3_hub"), + default="native", + ) + parser.add_argument( + "--enable-compile", + action="store_true", + help="Enable torch.compile (off by default); Fast compiles both transformer block stacks", + ) + return parser + + +def _read_text(path: Path) -> str: + if not path.is_file(): + raise FileNotFoundError(path) + value = path.read_text(encoding="utf-8").strip() + if not value: + raise ValueError(f"prompt file is empty: {path}") + return value + + +def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + args = build_parser().parse_args(argv) + fast = args.model_type == "fast" + args.model_path = args.model_path or ROOT / "weights" / ( + "WorldCrafter-Fast" if fast else "WorldCrafter-Base" + ) + if args.num_inference_steps is None: + args.num_inference_steps = 6 if fast else 50 + if args.guidance_scale is None: + args.guidance_scale = 1.0 if fast else 5.0 + if fast and (args.num_inference_steps != 6 or args.guidance_scale != 1.0): + raise ValueError( + "Fast requires CFG=1 and six regular steps; the first T2V chunk uses twelve steps" + ) + if fast and (args.resume_from or args.state_output_dir): + raise ValueError("Fast does not support resume/state export") + if args.num_chunks is not None and args.num_chunks <= 0: + raise ValueError("--num-chunks must be positive") + if args.stop_after_chunk is not None and args.stop_after_chunk < 0: + raise ValueError("--stop-after-chunk must be non-negative") + if args.resume_from is not None and args.chunk_output_dir is None: + raise ValueError("--resume-from requires --chunk-output-dir") + + args.camera_events = None + args.camera_options = {} + if args.actions is not None or args.actions_file is not None: + from .camera import count_chunks, parse_trajectory + + if args.local_camera_path is not None: + raise ValueError("--local-camera-path cannot be combined with camera actions") + text = args.actions if args.actions is not None else args.actions_file.read_text(encoding="utf-8-sig") + args.camera_events, args.camera_options = parse_trajectory(text) + if args.num_chunks is not None: + total = count_chunks(args.camera_events) + if args.num_chunks > total: + raise ValueError(f"Actions provide {total} chunks, but {args.num_chunks} were requested") + elif args.camera_path is None: + args.camera_path = ( + DEFAULT_I2V_CAMERA if args.mode == "i2v" else DEFAULT_T2V_CAMERA + ) + if args.prompt is None: + prompt_path = args.prompt_path + if prompt_path is None: + prompt_path = ( + DEFAULT_I2V_PROMPT if args.mode == "i2v" else DEFAULT_T2V_PROMPT + ) + args.prompt_path = prompt_path + args.prompt = _read_text(prompt_path) + elif args.prompt_path is not None: + raise ValueError("use either --prompt or --prompt-path, not both") + + if args.negative_prompt is None: + args.negative_prompt = _read_text(args.negative_prompt_path) + if args.mode == "i2v": + args.image_path = args.image_path or DEFAULT_IMAGE + elif args.image_path is not None: + raise ValueError("--image-path is only valid with --mode i2v") + + if args.output_path is None: + run_id = f"{datetime.now():%Y%m%d_%H%M%S}_{uuid4().hex[:8]}" + args.output_path = ( + ROOT / "output" / args.model_type / args.mode / run_id / "video.mp4" + ) + return args + + +def prepare_camera(args: argparse.Namespace) -> None: + if args.camera_events is not None: + from .camera import build_trajectory, save_trajectory + + camera, records = build_trajectory(args.camera_events, **args.camera_options) + directory = args.output_path.parent / f"{args.output_path.stem}_trajectory" + args.camera_path = save_trajectory( + directory, camera, records, fps=args.fps, + events=args.camera_events, options=args.camera_options, + ) + print(f"[worldcrafter] saved {len(records)} camera chunks to {args.camera_path}") + + +__all__ = ["build_parser", "parse_args", "prepare_camera"] diff --git a/worldcrafter/diffusers/__init__.py b/worldcrafter/diffusers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..de31b375ea1ede7fa37eae5d5eb100055f079f46 --- /dev/null +++ b/worldcrafter/diffusers/__init__.py @@ -0,0 +1,5 @@ +from .pipeline import WorldCrafterPipeline +from .scheduler import WorldCrafterScheduler +from .transformer import WorldCrafterTransformer3DModel + +__all__ = ["WorldCrafterPipeline", "WorldCrafterScheduler", "WorldCrafterTransformer3DModel"] diff --git a/worldcrafter/diffusers/pipeline.py b/worldcrafter/diffusers/pipeline.py new file mode 100644 index 0000000000000000000000000000000000000000..672b26a29a1a4a6f8129abdf319b51ca40ff9c67 --- /dev/null +++ b/worldcrafter/diffusers/pipeline.py @@ -0,0 +1,1628 @@ +import html +from itertools import accumulate +from typing import Any, Callable + +import numpy as np +import regex as re +import torch +from transformers import AutoTokenizer, UMT5EncoderModel + +from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback +from diffusers.image_processor import PipelineImageInput +from diffusers.loaders import HeliosLoraLoaderMixin as _BaseLoraLoaderMixin +from diffusers.models import AutoencoderKLWan +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.utils import ( + is_ftfy_available, + is_torch_xla_available, + logging, + replace_example_docstring, +) +from diffusers.utils.torch_utils import randn_tensor +from diffusers.video_processor import VideoProcessor + +from ..ucpe.bridge import build_ucpe_attention_kwargs_for_chunk +from .pipeline_output import WorldCrafterPipelineOutput +from .scheduler import WorldCrafterScheduler +from .transformer import WorldCrafterTransformer3DModel + + +if is_torch_xla_available(): + import torch_xla.core.xla_model as xm + + XLA_AVAILABLE = True +else: + XLA_AVAILABLE = False + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +def _render_repencoder_memory_latents( + *, + memory_provider: Any, + generated_latents: torch.Tensor, + camera_trajectory: dict[str, Any], + chunk_index: int, + num_latent_frames_per_chunk: int, + vae_scale_factor_temporal: int, + generator: torch.Generator | list[torch.Generator] | None, +) -> torch.Tensor: + if chunk_index <= 0: + raise ValueError("RepEncoder rendering is only valid after chunk 0") + if generated_latents.ndim != 5 or generated_latents.shape[2] == 0: + raise RuntimeError( + "RepEncoder requires a non-empty [B,C,T,H,W] bank of already generated WorldCrafter latents" + ) + render_memory = getattr(memory_provider, "render_memory", None) + if not callable(render_memory): + raise TypeError( + "memory_provider must expose a callable render_memory(...) method" + ) + + recent_latents = generated_latents[:, :, -1:, :, :] + try: + memory_latents = render_memory( + generated_latents=generated_latents, + recent_latents=recent_latents, + camera_trajectory=camera_trajectory, + chunk_index=chunk_index, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + vae_scale_factor_temporal=vae_scale_factor_temporal, + generator=generator, + ) + except Exception as exc: + raise RuntimeError( + f"RepEncoder memory rendering failed for chunk_index={chunk_index}" + ) from exc + + expected_shape = ( + generated_latents.shape[0], + generated_latents.shape[1], + 4, + generated_latents.shape[3], + generated_latents.shape[4], + ) + if not isinstance(memory_latents, torch.Tensor): + raise TypeError( + "memory_provider.render_memory(...) must return a torch.Tensor, " + f"got {type(memory_latents)!r}" + ) + if tuple(memory_latents.shape) != expected_shape: + raise ValueError( + "RepEncoder memory must have shape [B,C,4,H,W]; " + f"expected {expected_shape}, got {tuple(memory_latents.shape)}" + ) + if memory_latents.device != generated_latents.device: + raise ValueError( + "RepEncoder memory must stay on the WorldCrafter latent device; " + f"expected {generated_latents.device}, got {memory_latents.device}" + ) + if not memory_latents.is_floating_point(): + raise TypeError( + f"RepEncoder memory must be floating point, got {memory_latents.dtype}" + ) + if not torch.isfinite(memory_latents).all(): + raise FloatingPointError( + f"RepEncoder memory contains non-finite values at chunk_index={chunk_index}" + ) + return memory_latents + + +if is_ftfy_available(): + import ftfy + + +EXAMPLE_DOC_STRING = """ + Examples: + Run camera-controlled generation through the repository entry point, + which loads the local weights, camera trajectory, and memory provider: + + ```bash + python inference.py --model-path weights/WorldCrafter-Base + python inference.py --model-type fast --model-path weights/WorldCrafter-Fast + ``` +""" + + +def optimized_scale(positive_flat, negative_flat): + positive_flat = positive_flat.float() + negative_flat = negative_flat.float() + # Compute the dot product. + dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True) + # Squared norm of the unconditional prediction. + squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8 + # st_star = v_cond^T * v_uncond / ||v_uncond||^2 + st_star = dot_product / squared_norm + return st_star + + +def basic_clean(text): + text = ftfy.fix_text(text) + text = html.unescape(html.unescape(text)) + return text.strip() + + +def whitespace_clean(text): + text = re.sub(r"\s+", " ", text) + text = text.strip() + return text + + +def prompt_clean(text): + text = whitespace_clean(basic_clean(text)) + return text + + +# Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift +def calculate_shift( + image_seq_len, + base_seq_len: int = 256, + max_seq_len: int = 4096, + base_shift: float = 0.5, + max_shift: float = 1.15, +): + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return mu + + +class WorldCrafterPipeline(DiffusionPipeline, _BaseLoraLoaderMixin): + r""" + Pipeline for text-to-video / image-to-video / video-to-video generation using WorldCrafter. + + This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods + implemented for all pipelines (downloading, saving, running on a particular device, etc.). + + Args: + tokenizer: + Tokenizer loaded from the local model directory with `AutoTokenizer`. + text_encoder ([`UMT5EncoderModel`]): + Text encoder loaded from the local model directory. + transformer ([`WorldCrafterTransformer3DModel`]): + Conditional Transformer to denoise the input latents. + scheduler ([`WorldCrafterScheduler`]): + A scheduler to be used in combination with `transformer` to denoise the encoded image latents. + vae ([`AutoencoderKLWan`]): + Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations. + """ + + model_cpu_offload_seq = "text_encoder->transformer->vae" + _callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"] + _optional_components = ["transformer"] + + def __init__( + self, + tokenizer: AutoTokenizer, + text_encoder: UMT5EncoderModel, + vae: AutoencoderKLWan, + scheduler: WorldCrafterScheduler, + transformer: WorldCrafterTransformer3DModel, + is_cfg_zero_star: bool = False, + is_distilled: bool = False, + ): + super().__init__() + + self.register_modules( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer, + scheduler=scheduler, + ) + self.register_to_config(is_cfg_zero_star=is_cfg_zero_star) + self.register_to_config(is_distilled=is_distilled) + self.vae_scale_factor_temporal = ( + self.vae.config.scale_factor_temporal if getattr(self, "vae", None) else 4 + ) + self.vae_scale_factor_spatial = ( + self.vae.config.scale_factor_spatial if getattr(self, "vae", None) else 8 + ) + self.video_processor = VideoProcessor( + vae_scale_factor=self.vae_scale_factor_spatial + ) + + def _get_t5_prompt_embeds( + self, + prompt: str | list[str] = None, + num_videos_per_prompt: int = 1, + max_sequence_length: int = 226, + device: torch.device | None = None, + dtype: torch.dtype | None = None, + ): + device = device or self._execution_device + dtype = dtype or self.text_encoder.dtype + + prompt = [prompt] if isinstance(prompt, str) else prompt + prompt = [prompt_clean(u) for u in prompt] + batch_size = len(prompt) + + text_inputs = self.tokenizer( + prompt, + padding="max_length", + max_length=max_sequence_length, + truncation=True, + add_special_tokens=True, + return_attention_mask=True, + return_tensors="pt", + ) + text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask + seq_lens = mask.gt(0).sum(dim=1).long() + + prompt_embeds = self.text_encoder( + text_input_ids.to(device), mask.to(device) + ).last_hidden_state + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)] + prompt_embeds = torch.stack( + [ + torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) + for u in prompt_embeds + ], + dim=0, + ) + + # duplicate text embeddings for each generation per prompt, using mps friendly method + _, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) + prompt_embeds = prompt_embeds.view( + batch_size * num_videos_per_prompt, seq_len, -1 + ) + + return prompt_embeds, text_inputs.attention_mask.bool() + + def encode_prompt( + self, + prompt: str | list[str], + negative_prompt: str | list[str] | None = None, + do_classifier_free_guidance: bool = True, + num_videos_per_prompt: int = 1, + prompt_embeds: torch.Tensor | None = None, + negative_prompt_embeds: torch.Tensor | None = None, + max_sequence_length: int = 226, + device: torch.device | None = None, + dtype: torch.dtype | None = None, + ): + r""" + Encodes the prompt into text encoder hidden states. + + Args: + prompt (`str` or `list[str]`, *optional*): + prompt to be encoded + negative_prompt (`str` or `list[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is + less than or equal to `1`). + do_classifier_free_guidance (`bool`, *optional*, defaults to `True`): + Whether to use classifier free guidance or not. + num_videos_per_prompt (`int`, *optional*, defaults to 1): + Number of videos to generate per prompt. + prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + device: (`torch.device`, *optional*): + torch device + dtype: (`torch.dtype`, *optional*): + torch dtype + """ + device = device or self._execution_device + + prompt = [prompt] if isinstance(prompt, str) else prompt + if prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + prompt_embeds, _ = self._get_t5_prompt_embeds( + prompt=prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + if do_classifier_free_guidance and negative_prompt_embeds is None: + negative_prompt = negative_prompt or "" + negative_prompt = ( + batch_size * [negative_prompt] + if isinstance(negative_prompt, str) + else negative_prompt + ) + + if prompt is not None and type(prompt) is not type(negative_prompt): + raise TypeError( + f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + f" {type(prompt)}." + ) + elif batch_size != len(negative_prompt): + raise ValueError( + f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" + " the batch size of `prompt`." + ) + + negative_prompt_embeds, _ = self._get_t5_prompt_embeds( + prompt=negative_prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + return prompt_embeds, negative_prompt_embeds + + def check_inputs( + self, + prompt, + negative_prompt, + height, + width, + prompt_embeds=None, + negative_prompt_embeds=None, + callback_on_step_end_tensor_inputs=None, + image=None, + video=None, + use_interpolate_prompt=False, + num_videos_per_prompt=None, + interpolate_time_list=None, + interpolation_steps=None, + guidance_scale=None, + ): + if height % 16 != 0 or width % 16 != 0: + raise ValueError( + f"`height` and `width` have to be divisible by 16 but are {height} and {width}." + ) + + if callback_on_step_end_tensor_inputs is not None and not all( + k in self._callback_tensor_inputs + for k in callback_on_step_end_tensor_inputs + ): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" + ) + + if prompt is not None and prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to" + " only forward one of the two." + ) + elif negative_prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to" + " only forward one of the two." + ) + elif prompt is None and prompt_embeds is None: + raise ValueError( + "Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined." + ) + elif prompt is not None and ( + not isinstance(prompt, str) and not isinstance(prompt, list) + ): + raise ValueError( + f"`prompt` has to be of type `str` or `list` but is {type(prompt)}" + ) + elif negative_prompt is not None and ( + not isinstance(negative_prompt, str) + and not isinstance(negative_prompt, list) + ): + raise ValueError( + f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}" + ) + + if image is not None and video is not None: + raise ValueError("image and video cannot be provided simultaneously") + + if use_interpolate_prompt: + assert ( + num_videos_per_prompt == 1 + ), f"num_videos_per_prompt must be 1, got {num_videos_per_prompt}" + assert isinstance(prompt, list), "prompt must be a list" + assert len(prompt) == len( + interpolate_time_list + ), f"Length mismatch: {len(prompt)} vs {len(interpolate_time_list)}" + assert ( + min(interpolate_time_list) > interpolation_steps + ), f"Minimum value {min(interpolate_time_list)} must be greater than {interpolation_steps}" + + if guidance_scale > 1.0 and self.config.is_distilled: + logger.warning( + f"Guidance scale {guidance_scale} is ignored for step-wise distilled models." + ) + + def prepare_latents( + self, + batch_size: int, + num_channels_latents: int = 16, + height: int = 384, + width: int = 640, + num_frames: int = 33, + dtype: torch.dtype | None = None, + device: torch.device | None = None, + generator: torch.Generator | list[torch.Generator] | None = None, + latents: torch.Tensor | None = None, + ) -> torch.Tensor: + if latents is not None: + return latents.to(device=device, dtype=dtype) + + num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1 + shape = ( + batch_size, + num_channels_latents, + num_latent_frames, + int(height) // self.vae_scale_factor_spatial, + int(width) // self.vae_scale_factor_spatial, + ) + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + return latents + + def prepare_image_latents( + self, + image: torch.Tensor, + latents_mean: torch.Tensor, + latents_std: torch.Tensor, + num_latent_frames_per_chunk: int, + dtype: torch.dtype | None = None, + device: torch.device | None = None, + generator: torch.Generator | list[torch.Generator] | None = None, + latents: torch.Tensor | None = None, + fake_latents: torch.Tensor | None = None, + ) -> torch.Tensor: + device = device or self._execution_device + if latents is None: + image = image.unsqueeze(2).to(device=device, dtype=self.vae.dtype) + latents = self.vae.encode(image).latent_dist.sample(generator=generator) + latents = (latents - latents_mean) * latents_std + if fake_latents is None: + min_frames = ( + num_latent_frames_per_chunk - 1 + ) * self.vae_scale_factor_temporal + 1 + fake_video = image.repeat(1, 1, min_frames, 1, 1).to( + device=device, dtype=self.vae.dtype + ) + fake_latents_full = self.vae.encode(fake_video).latent_dist.sample( + generator=generator + ) + fake_latents_full = (fake_latents_full - latents_mean) * latents_std + fake_latents = fake_latents_full[:, :, -1:, :, :] + return latents.to(device=device, dtype=dtype), fake_latents.to( + device=device, dtype=dtype + ) + + def prepare_video_latents( + self, + video: torch.Tensor, + latents_mean: torch.Tensor, + latents_std: torch.Tensor, + num_latent_frames_per_chunk: int, + dtype: torch.dtype | None = None, + device: torch.device | None = None, + generator: torch.Generator | list[torch.Generator] | None = None, + latents: torch.Tensor | None = None, + ) -> torch.Tensor: + device = device or self._execution_device + video = video.to(device=device, dtype=self.vae.dtype) + if latents is None: + num_frames = video.shape[2] + min_frames = ( + num_latent_frames_per_chunk - 1 + ) * self.vae_scale_factor_temporal + 1 + num_chunks = num_frames // min_frames + if num_chunks == 0: + raise ValueError( + f"Video must have at least {min_frames} frames " + f"(got {num_frames} frames). " + f"Required: (num_latent_frames_per_chunk - 1) * {self.vae_scale_factor_temporal} + 1 = ({num_latent_frames_per_chunk} - 1) * {self.vae_scale_factor_temporal} + 1 = {min_frames}" + ) + total_valid_frames = num_chunks * min_frames + start_frame = num_frames - total_valid_frames + + first_frame = video[:, :, 0:1, :, :] + first_frame_latent = self.vae.encode(first_frame).latent_dist.sample( + generator=generator + ) + first_frame_latent = (first_frame_latent - latents_mean) * latents_std + + latents_chunks = [] + for i in range(num_chunks): + chunk_start = start_frame + i * min_frames + chunk_end = chunk_start + min_frames + video_chunk = video[:, :, chunk_start:chunk_end, :, :] + chunk_latents = self.vae.encode(video_chunk).latent_dist.sample( + generator=generator + ) + chunk_latents = (chunk_latents - latents_mean) * latents_std + latents_chunks.append(chunk_latents) + latents = torch.cat(latents_chunks, dim=2) + return first_frame_latent.to(device=device, dtype=dtype), latents.to( + device=device, dtype=dtype + ) + + def interpolate_prompt_embeds( + self, + prompt_embeds_1: torch.Tensor, + prompt_embeds_2: torch.Tensor, + interpolation_steps: int = 3, + ): + x = torch.lerp( + prompt_embeds_1, + prompt_embeds_2, + torch.linspace(0, 1, steps=interpolation_steps) + .unsqueeze(1) + .unsqueeze(2) + .to(prompt_embeds_1), + ) + interpolated_prompt_embeds = list(x.chunk(interpolation_steps, dim=0)) + return interpolated_prompt_embeds + + def sample_block_noise( + self, + batch_size, + channel, + num_frames, + height, + width, + patch_size: tuple[int, ...] = (1, 2, 2), + device: torch.device | None = None, + generator: torch.Generator | None = None, + ): + # The default generator is independent of the trajectory RNG. + if generator is None: + generator = torch.Generator(device=device) + elif isinstance(generator, list): + generator = generator[0] + + gamma = self.scheduler.config.gamma + _, ph, pw = patch_size + block_size = ph * pw + + cov = ( + torch.eye(block_size, device=device) * (1 + gamma) + - torch.ones(block_size, block_size, device=device) * gamma + ) + cov += torch.eye(block_size, device=device) * 1e-8 + cov = ( + cov.float() + ) # Upcast to fp32 for numerical stability — cholesky is unreliable in fp16/bf16. + + L = torch.linalg.cholesky(cov) + block_number = ( + batch_size * channel * num_frames * (height // ph) * (width // pw) + ) + z = torch.randn( + block_number, block_size, generator=generator, device=generator.device + ).to(device=device) + noise = z @ L.T + + noise = noise.view( + batch_size, channel, num_frames, height // ph, width // pw, ph, pw + ) + noise = noise.permute(0, 1, 2, 3, 5, 4, 6).reshape( + batch_size, channel, num_frames, height, width + ) + + return noise + + def stage1_sample( + self, + latents: torch.Tensor = None, + prompt_embeds: torch.Tensor = None, + negative_prompt_embeds: torch.Tensor = None, + timesteps: torch.Tensor = None, + guidance_scale: float | None = 5.0, + indices_hidden_states: torch.Tensor = None, + indices_latents_history_short: torch.Tensor = None, + indices_latents_history_mid: torch.Tensor = None, + indices_latents_history_long: torch.Tensor = None, + latents_history_short: torch.Tensor = None, + latents_history_mid: torch.Tensor = None, + latents_history_long: torch.Tensor = None, + attention_kwargs: dict | None = None, + device: torch.device | None = None, + transformer_dtype: torch.dtype = None, + generator: torch.Generator | None = None, + num_warmup_steps: int | None = None, + # ------------ CFG Zero ------------ + use_zero_init: bool | None = True, + zero_steps: int | None = 1, + # ------------ Callback ------------ + callback_on_step_end: ( + Callable[[int, int], None] + | PipelineCallback + | MultiPipelineCallbacks + | None + ) = None, + callback_on_step_end_tensor_inputs: list[str] = ["latents"], + progress_bar=None, + ): + batch_size = latents.shape[0] + + for i, t in enumerate(timesteps): + if self.interrupt: + continue + + self._current_timestep = t + timestep = t.expand(latents.shape[0]) + + latent_model_input = latents.to(transformer_dtype) + with self.transformer.cache_context("cond"): + noise_pred = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=prompt_embeds, + indices_hidden_states=indices_hidden_states, + indices_latents_history_short=indices_latents_history_short, + indices_latents_history_mid=indices_latents_history_mid, + indices_latents_history_long=indices_latents_history_long, + latents_history_short=latents_history_short.to(transformer_dtype), + latents_history_mid=latents_history_mid.to(transformer_dtype), + latents_history_long=latents_history_long.to(transformer_dtype), + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] + + if self.do_classifier_free_guidance: + with self.transformer.cache_context("uncond"): + noise_uncond = self.transformer( + hidden_states=latent_model_input, + timestep=timestep, + encoder_hidden_states=negative_prompt_embeds, + indices_hidden_states=indices_hidden_states, + indices_latents_history_short=indices_latents_history_short, + indices_latents_history_mid=indices_latents_history_mid, + indices_latents_history_long=indices_latents_history_long, + latents_history_short=latents_history_short.to( + transformer_dtype + ), + latents_history_mid=latents_history_mid.to(transformer_dtype), + latents_history_long=latents_history_long.to(transformer_dtype), + attention_kwargs=attention_kwargs, + return_dict=False, + )[0] + + if self.config.is_cfg_zero_star: + noise_pred_text = noise_pred + positive_flat = noise_pred_text.view(batch_size, -1) + negative_flat = noise_uncond.view(batch_size, -1) + + alpha = optimized_scale(positive_flat, negative_flat) + alpha = alpha.view( + batch_size, *([1] * (len(noise_pred_text.shape) - 1)) + ) + alpha = alpha.to(noise_pred_text.dtype) + + if (i <= zero_steps) and use_zero_init: + noise_pred = noise_pred_text * 0.0 + else: + noise_pred = noise_uncond * alpha + guidance_scale * ( + noise_pred_text - noise_uncond * alpha + ) + else: + noise_pred = noise_uncond + guidance_scale * ( + noise_pred - noise_uncond + ) + + latents = self.scheduler.step( + noise_pred, + t, + latents, + return_dict=False, + )[0] + + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + negative_prompt_embeds = callback_outputs.pop( + "negative_prompt_embeds", negative_prompt_embeds + ) + + if i == len(timesteps) - 1 or ( + (i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0 + ): + progress_bar.update() + + if XLA_AVAILABLE: + xm.mark_step() + + return latents + + @property + def guidance_scale(self): + return self._guidance_scale + + @property + def do_classifier_free_guidance(self): + return self._guidance_scale > 1.0 + + @property + def num_timesteps(self): + return self._num_timesteps + + @property + def current_timestep(self): + return self._current_timestep + + @property + def interrupt(self): + return self._interrupt + + @property + def attention_kwargs(self): + return self._attention_kwargs + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: str | list[str] = None, + negative_prompt: str | list[str] = None, + height: int = 384, + width: int = 640, + num_frames: int = 132, + num_inference_steps: int = 50, + sigmas: list[float] = None, + guidance_scale: float = 5.0, + num_videos_per_prompt: int | None = 1, + generator: torch.Generator | list[torch.Generator] | None = None, + latents: torch.Tensor | None = None, + prompt_embeds: torch.Tensor | None = None, + negative_prompt_embeds: torch.Tensor | None = None, + output_type: str | None = "np", + return_dict: bool = True, + attention_kwargs: dict[str, Any] | None = None, + callback_on_step_end: ( + Callable[[int, int], None] + | PipelineCallback + | MultiPipelineCallbacks + | None + ) = None, + callback_on_step_end_tensor_inputs: list[str] = ["latents"], + callback_on_chunk_end: Callable[[int, torch.Tensor], None] | None = None, + callback_on_chunk_state: Callable[[int, dict[str, Any]], None] | None = None, + resume_state: dict[str, Any] | None = None, + stop_after_chunk: int | None = None, + max_sequence_length: int = 512, + # ------------ I2V ------------ + image: PipelineImageInput | None = None, + image_latents: torch.Tensor | None = None, + fake_image_latents: torch.Tensor | None = None, + add_noise_to_image_latents: bool = True, + image_noise_sigma_min: float = 0.111, + image_noise_sigma_max: float = 0.135, + # ------------ V2V ------------ + video: PipelineImageInput | None = None, + video_latents: torch.Tensor | None = None, + add_noise_to_video_latents: bool = True, + video_noise_sigma_min: float = 0.111, + video_noise_sigma_max: float = 0.135, + # ------------ Interactive ------------ + use_interpolate_prompt: bool = False, + interpolate_time_list: list = [7, 7, 7], + interpolation_steps: int = 3, + # ------------ Stage 1 ------------ + memory_size: int = 4, + history_sizes: list = [2, 1], + num_latent_frames_per_chunk: int = 9, + keep_first_frame: bool = True, + is_skip_first_chunk: bool = False, + # ------------ Camera control ------------ + camera_trajectory: dict[str, Any] | None = None, + # ------------ RepEncoder 3D memory ------------ + memory_provider: Any | None = None, + # ------------ Stage 2 ------------ + is_enable_stage2: bool = False, + pyramid_num_stages: int = 3, + pyramid_num_inference_steps_list: list = [10, 10, 10], + # ------------ CFG Zero ------------ + use_zero_init: bool | None = True, + zero_steps: int | None = 1, + # ------------ DMD ------------ + is_amplify_first_chunk: bool = False, + ): + r""" + The call function to the pipeline for generation. + + Args: + prompt (`str` or `list[str]`, *optional*): + The prompt or prompts to guide the image generation. If not defined, pass `prompt_embeds` instead. + negative_prompt (`str` or `list[str]`, *optional*): + The prompt or prompts to avoid during image generation. If not defined, pass `negative_prompt_embeds` + instead. Ignored when not using guidance (`guidance_scale` <= `1`). + height (`int`, defaults to `384`): + The height in pixels of the generated image. + width (`int`, defaults to `640`): + The width in pixels of the generated image. + num_frames (`int`, defaults to `132`): + The number of frames in the generated video. + num_inference_steps (`int`, defaults to `50`): + The number of denoising steps. More denoising steps usually lead to a higher quality image at the + expense of slower inference. + guidance_scale (`float`, defaults to `5.0`): + Guidance scale as defined in [Classifier-Free Diffusion + Guidance](https://huggingface.co/papers/2207.12598). `guidance_scale` is defined as `w` of equation 2. + of [Imagen Paper](https://huggingface.co/papers/2205.11487). Guidance scale is enabled by setting + `guidance_scale > 1`. Higher guidance scale encourages to generate images that are closely linked to + the text `prompt`, usually at the expense of lower image quality. + num_videos_per_prompt (`int`, *optional*, defaults to 1): + The number of videos to generate per prompt. + generator (`torch.Generator` or `list[torch.Generator]`, *optional*): + A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make + generation deterministic. + latents (`torch.Tensor`, *optional*): + Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image + generation. Can be used to tweak the same generation with different prompts. If not provided, a latents + tensor is generated by sampling using the supplied random `generator`. + prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not + provided, text embeddings are generated from the `prompt` input argument. + output_type (`str`, *optional*, defaults to `"np"`): + Video output format: `"np"`, `"pt"`, or `"pil"`; `"latent"` returns latent tensors. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`WorldCrafterPipelineOutput`] instead of a plain tuple. + attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*): + A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of + each denoising step during the inference. with the following arguments: `callback_on_step_end(self: + DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a + list of all tensors as specified by `callback_on_step_end_tensor_inputs`. + callback_on_step_end_tensor_inputs (`list`, *optional*): + The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list + will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the + `._callback_tensor_inputs` attribute of your pipeline class. + max_sequence_length (`int`, defaults to `512`): + The maximum sequence length of the text encoder. If the prompt is longer than this, it will be + truncated. If the prompt is shorter, it will be padded to this length. + + Examples: + + Returns: + [`~WorldCrafterPipelineOutput`] or `tuple`: + If `return_dict` is `True`, [`WorldCrafterPipelineOutput`] is returned, otherwise a `tuple` is returned where + the only element contains the generated video batch. No safety-classification flags are returned. + """ + + if image is not None and video is not None: + raise ValueError("image and video cannot be provided simultaneously") + use_fast = bool(self.config.is_distilled) + if use_fast: + if not is_enable_stage2 or guidance_scale != 1.0: + raise ValueError("Fast inference requires pyramid sampling and CFG=1") + if pyramid_num_inference_steps_list is not None: + raise ValueError("Fast steps are owned by the checkpoint DMD contract") + if resume_state is not None: + raise ValueError("Fast resume is not yet validated") + elif camera_trajectory is not None and is_enable_stage2: + raise ValueError("Base camera inference requires stage1 sampling") + if memory_size != 4: + raise ValueError( + f"RepEncoder memory contract requires memory_size=4, got {memory_size}" + ) + if num_latent_frames_per_chunk != 9: + raise ValueError( + "RepEncoder target slots [2,4,6,8] require num_latent_frames_per_chunk=9, " + f"got {num_latent_frames_per_chunk}" + ) + + requested_window_num_frames = ( + num_latent_frames_per_chunk - 1 + ) * self.vae_scale_factor_temporal + 1 + requested_num_chunks = max( + 1, + (max(num_frames, 1) + requested_window_num_frames - 1) + // requested_window_num_frames, + ) + if use_interpolate_prompt: + requested_num_chunks = max(requested_num_chunks, sum(interpolate_time_list)) + if requested_num_chunks > 1: + if camera_trajectory is None: + raise ValueError( + "Multi-chunk RepEncoder inference requires a global metric camera trajectory" + ) + if memory_provider is None: + raise ValueError("Multi-chunk inference requires memory_provider") + + history_sizes = sorted(history_sizes, reverse=True) # From big to small + assert ( + memory_size <= num_latent_frames_per_chunk + ), f"memory_size={memory_size} must be <= num_latent_frames_per_chunk={num_latent_frames_per_chunk}" + + if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): + callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + negative_prompt, + height, + width, + prompt_embeds, + negative_prompt_embeds, + callback_on_step_end_tensor_inputs, + image, + video, + use_interpolate_prompt, + num_videos_per_prompt, + interpolate_time_list, + interpolation_steps, + guidance_scale, + ) + + num_frames = max(num_frames, 1) + + self._guidance_scale = guidance_scale + self._attention_kwargs = attention_kwargs + self._current_timestep = None + self._interrupt = False + + device = self._execution_device + vae_dtype = self.vae.dtype + + latents_mean = ( + torch.tensor(self.vae.config.latents_mean) + .view(1, self.vae.config.z_dim, 1, 1, 1) + .to(device, self.vae.dtype) + ) + latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view( + 1, self.vae.config.z_dim, 1, 1, 1 + ).to(device, self.vae.dtype) + + # 2. Define call parameters + if use_interpolate_prompt or (prompt is not None and isinstance(prompt, str)): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + # 3. Encode input prompt + if use_interpolate_prompt: + interpolate_interval_idx = None + interpolate_embeds = None + interpolate_cumulative_list = list(accumulate(interpolate_time_list)) + + all_prompt_embeds, negative_prompt_embeds = self.encode_prompt( + prompt=prompt, + negative_prompt=negative_prompt, + do_classifier_free_guidance=self.do_classifier_free_guidance, + num_videos_per_prompt=num_videos_per_prompt, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + max_sequence_length=max_sequence_length, + device=device, + ) + + transformer_dtype = self.transformer.dtype + all_prompt_embeds = all_prompt_embeds.to(transformer_dtype) + if negative_prompt_embeds is not None: + if use_interpolate_prompt: + negative_prompt_embeds = negative_prompt_embeds[0].unsqueeze(0) + negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype) + + # 4. Prepare image or video + if image is not None: + image = self.video_processor.preprocess(image, height=height, width=width) + image_latents, fake_image_latents = self.prepare_image_latents( + image, + latents_mean=latents_mean, + latents_std=latents_std, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + dtype=torch.float32, + device=device, + generator=generator, + latents=image_latents, + fake_latents=fake_image_latents, + ) + + if image_latents is not None and add_noise_to_image_latents: + image_noise_sigma = ( + torch.rand(1, device=device, generator=generator) + * (image_noise_sigma_max - image_noise_sigma_min) + + image_noise_sigma_min + ) + image_latents = ( + image_noise_sigma + * randn_tensor(image_latents.shape, generator=generator, device=device) + + (1 - image_noise_sigma) * image_latents + ) + fake_image_noise_sigma = ( + torch.rand(1, device=device, generator=generator) + * (video_noise_sigma_max - video_noise_sigma_min) + + video_noise_sigma_min + ) + fake_image_latents = ( + fake_image_noise_sigma + * randn_tensor( + fake_image_latents.shape, generator=generator, device=device + ) + + (1 - fake_image_noise_sigma) * fake_image_latents + ) + + if video is not None: + video = self.video_processor.preprocess_video( + video, height=height, width=width + ) + image_latents, video_latents = self.prepare_video_latents( + video, + latents_mean=latents_mean, + latents_std=latents_std, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + dtype=torch.float32, + device=device, + generator=generator, + latents=video_latents, + ) + + if video_latents is not None and add_noise_to_video_latents: + image_noise_sigma = ( + torch.rand(1, device=device, generator=generator) + * (image_noise_sigma_max - image_noise_sigma_min) + + image_noise_sigma_min + ) + image_latents = ( + image_noise_sigma + * randn_tensor(image_latents.shape, generator=generator, device=device) + + (1 - image_noise_sigma) * image_latents + ) + + noisy_latents_chunks = [] + num_latent_chunks = video_latents.shape[2] // num_latent_frames_per_chunk + for i in range(num_latent_chunks): + chunk_start = i * num_latent_frames_per_chunk + chunk_end = chunk_start + num_latent_frames_per_chunk + latent_chunk = video_latents[:, :, chunk_start:chunk_end, :, :] + + chunk_frames = latent_chunk.shape[2] + frame_sigmas = ( + torch.rand(chunk_frames, device=device, generator=generator) + * (video_noise_sigma_max - video_noise_sigma_min) + + video_noise_sigma_min + ) + frame_sigmas = frame_sigmas.view(1, 1, chunk_frames, 1, 1) + + noisy_chunk = ( + frame_sigmas + * randn_tensor( + latent_chunk.shape, generator=generator, device=device + ) + + (1 - frame_sigmas) * latent_chunk + ) + noisy_latents_chunks.append(noisy_chunk) + video_latents = torch.cat(noisy_latents_chunks, dim=2) + + # 5. Prepare latent variables + num_channels_latents = self.transformer.config.in_channels + window_num_frames = ( + num_latent_frames_per_chunk - 1 + ) * self.vae_scale_factor_temporal + 1 + num_latent_chunk = max( + 1, (num_frames + window_num_frames - 1) // window_num_frames + ) + history_video = None + total_generated_latent_frames = 0 + + if not keep_first_frame: + history_sizes[-1] = history_sizes[-1] + 1 + history_latents = torch.zeros( + batch_size, + num_channels_latents, + sum(history_sizes), + height // self.vae_scale_factor_spatial, + width // self.vae_scale_factor_spatial, + device=device, + dtype=torch.float32, + ) + if fake_image_latents is not None: + history_latents = torch.cat([history_latents, fake_image_latents], dim=2) + total_generated_latent_frames += 1 + if video_latents is not None: + history_frames = history_latents.shape[2] + video_frames = video_latents.shape[2] + if video_frames < history_frames: + keep_frames = history_frames - video_frames + history_latents = torch.cat( + [history_latents[:, :, :keep_frames, :, :], video_latents], dim=2 + ) + else: + history_latents = video_latents + total_generated_latent_frames += video_latents.shape[2] + + generated_memory_latents = history_latents[:, :, :0, :, :] + + start_chunk = 0 + if resume_state is not None: + if resume_state.get("format") != "worldcrafter_chunk_state_v1": + raise ValueError("unsupported WorldCrafter resume-state format") + start_chunk = int(resume_state["next_chunk_index"]) + if start_chunk <= 0 or start_chunk >= num_latent_chunk: + raise ValueError( + f"resume next_chunk_index must be in [1, {num_latent_chunk - 1}], got {start_chunk}" + ) + generated_memory_latents = resume_state["generated_memory_latents"].to( + device=device, dtype=torch.float32 + ) + expected_generated = start_chunk * num_latent_frames_per_chunk + if tuple(generated_memory_latents.shape) != ( + batch_size, + num_channels_latents, + expected_generated, + height // self.vae_scale_factor_spatial, + width // self.vae_scale_factor_spatial, + ): + raise ValueError( + "resume generated_memory_latents shape does not match next_chunk_index: " + f"{tuple(generated_memory_latents.shape)}" + ) + history_latents = resume_state["history_latents"].to( + device=device, dtype=torch.float32 + ) + expected_history = ( + batch_size, + num_channels_latents, + sum(history_sizes), + height // self.vae_scale_factor_spatial, + width // self.vae_scale_factor_spatial, + ) + if tuple(history_latents.shape) != expected_history: + raise ValueError( + f"resume history_latents must have shape {expected_history}, " + f"got {tuple(history_latents.shape)}" + ) + saved_image_latents = resume_state.get("image_latents") + if saved_image_latents is None: + if keep_first_frame: + raise ValueError("resume state is missing fixed image_latents") + image_latents = None + else: + image_latents = saved_image_latents.to( + device=device, dtype=torch.float32 + ) + if not isinstance(generator, torch.Generator): + raise TypeError( + "resumable WorldCrafter inference requires one torch.Generator" + ) + generator.set_state(resume_state["generator_state"].cpu()) + total_generated_latent_frames = expected_generated + + final_chunk_index = num_latent_chunk - 1 + if stop_after_chunk is not None: + final_chunk_index = int(stop_after_chunk) + if final_chunk_index < start_chunk or final_chunk_index >= num_latent_chunk: + raise ValueError("stop_after_chunk is outside this inference interval") + + # 6. Denoising loop + if use_interpolate_prompt: + if num_latent_chunk < max(interpolate_cumulative_list): + num_latent_chunk = sum(interpolate_cumulative_list) + print(f"Update num_latent_chunk to: {num_latent_chunk}") + + if not is_enable_stage2: + patch_size = self.transformer.config.patch_size + image_seq_len = ( + num_latent_frames_per_chunk + * (height // self.vae_scale_factor_spatial) + * (width // self.vae_scale_factor_spatial) + // (patch_size[0] * patch_size[1] * patch_size[2]) + ) + sigmas = ( + np.linspace(0.999, 0.0, num_inference_steps + 1)[:-1] + if sigmas is None + else sigmas + ) + mu = calculate_shift( + image_seq_len, + self.scheduler.config.get("base_image_seq_len", 256), + self.scheduler.config.get("max_image_seq_len", 4096), + self.scheduler.config.get("base_shift", 0.5), + self.scheduler.config.get("max_shift", 1.15), + ) + + for k in range(start_chunk, num_latent_chunk): + if use_interpolate_prompt: + assert num_latent_chunk >= max(interpolate_cumulative_list) + + current_interval_idx = 0 + for idx, cumulative_val in enumerate(interpolate_cumulative_list): + if k < cumulative_val: + current_interval_idx = idx + break + + if current_interval_idx == 0: + prompt_embeds = all_prompt_embeds[0].unsqueeze(0) + else: + interval_start = interpolate_cumulative_list[ + current_interval_idx - 1 + ] + position_in_interval = k - interval_start + + if position_in_interval < interpolation_steps: + if ( + interpolate_embeds is None + or interpolate_interval_idx != current_interval_idx + ): + interpolate_embeds = self.interpolate_prompt_embeds( + prompt_embeds_1=all_prompt_embeds[ + current_interval_idx - 1 + ].unsqueeze(0), + prompt_embeds_2=all_prompt_embeds[ + current_interval_idx + ].unsqueeze(0), + interpolation_steps=interpolation_steps, + ) + interpolate_interval_idx = current_interval_idx + + prompt_embeds = interpolate_embeds[position_in_interval] + else: + prompt_embeds = all_prompt_embeds[ + current_interval_idx + ].unsqueeze(0) + else: + prompt_embeds = all_prompt_embeds + + is_first_chunk = k == 0 + is_second_chunk = k == 1 + if is_first_chunk: + first_memory_latents = generated_memory_latents.new_zeros( + batch_size, + num_channels_latents, + memory_size, + height // self.vae_scale_factor_spatial, + width // self.vae_scale_factor_spatial, + ) + else: + first_memory_latents = _render_repencoder_memory_latents( + memory_provider=memory_provider, + generated_latents=generated_memory_latents, + camera_trajectory=camera_trajectory, + chunk_index=k, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + vae_scale_factor_temporal=self.vae_scale_factor_temporal, + generator=generator, + ) + if keep_first_frame: + if is_first_chunk: + history_sizes_first_chunk = [1] + history_sizes.copy() + history_latents_first_chunk = torch.zeros( + batch_size, + num_channels_latents, + sum(history_sizes_first_chunk), + height // self.vae_scale_factor_spatial, + width // self.vae_scale_factor_spatial, + device=device, + dtype=torch.float32, + ) + if fake_image_latents is not None: + history_latents_first_chunk = torch.cat( + [history_latents_first_chunk, fake_image_latents], dim=2 + ) + if video_latents is not None: + history_frames = history_latents_first_chunk.shape[2] + video_frames = video_latents.shape[2] + if video_frames < history_frames: + keep_frames = history_frames - video_frames + history_latents_first_chunk = torch.cat( + [ + history_latents_first_chunk[ + :, :, :keep_frames, :, : + ], + video_latents, + ], + dim=2, + ) + else: + history_latents_first_chunk = video_latents + + indices = torch.arange( + 0, + sum( + [ + 1, + memory_size, + *history_sizes, + num_latent_frames_per_chunk, + ] + ), + ) + ( + indices_prefix, + indices_latents_memory, + indices_latents_history_mid, + indices_latents_history_1x, + indices_hidden_states, + ) = indices.split( + [1, memory_size, *history_sizes, num_latent_frames_per_chunk], + dim=0, + ) + indices_latents_history_short = torch.cat( + [indices_prefix, indices_latents_history_1x], dim=0 + ) + + latents_memory = first_memory_latents + latents_prefix, latents_history_mid, latents_history_1x = ( + history_latents_first_chunk[ + :, :, -sum(history_sizes_first_chunk) : + ].split(history_sizes_first_chunk, dim=2) + ) + if image_latents is not None: + latents_prefix = image_latents + latents_history_short = torch.cat( + [latents_prefix, latents_history_1x], dim=2 + ) + else: + indices = torch.arange( + 0, + sum( + [ + 1, + memory_size, + *history_sizes, + num_latent_frames_per_chunk, + ] + ), + ) + ( + indices_prefix, + indices_latents_memory, + indices_latents_history_mid, + indices_latents_history_1x, + indices_hidden_states, + ) = indices.split( + [1, memory_size, *history_sizes, num_latent_frames_per_chunk], + dim=0, + ) + indices_latents_history_short = torch.cat( + [indices_prefix, indices_latents_history_1x], dim=0 + ) + + latents_prefix = image_latents + latents_memory = first_memory_latents + latents_history_mid, latents_history_1x = history_latents[ + :, :, -sum(history_sizes) : + ].split(history_sizes, dim=2) + latents_history_short = torch.cat( + [latents_prefix, latents_history_1x], dim=2 + ) + else: + indices = torch.arange( + 0, sum([memory_size, *history_sizes, num_latent_frames_per_chunk]) + ) + ( + indices_latents_memory, + indices_latents_history_mid, + indices_latents_history_short, + indices_hidden_states, + ) = indices.split( + [memory_size, *history_sizes, num_latent_frames_per_chunk], dim=0 + ) + latents_memory = first_memory_latents + latents_history_mid, latents_history_short = history_latents[ + :, :, -sum(history_sizes) : + ].split(history_sizes, dim=2) + + indices_hidden_states = indices_hidden_states.unsqueeze(0) + indices_latents_history_short = indices_latents_history_short.unsqueeze(0) + indices_latents_history_mid = indices_latents_history_mid.unsqueeze(0) + indices_latents_memory = indices_latents_memory.unsqueeze(0) + + latents = self.prepare_latents( + batch_size, + num_channels_latents, + height, + width, + window_num_frames, + dtype=torch.float32, + device=device, + generator=generator, + latents=None, + ) + + if not is_enable_stage2: + self.scheduler.set_timesteps( + num_inference_steps, device=device, sigmas=sigmas, mu=mu + ) + timesteps = self.scheduler.timesteps + num_warmup_steps = ( + len(timesteps) - num_inference_steps * self.scheduler.order + ) + self._num_timesteps = len(timesteps) + else: + if use_fast: + from ..fast.contract import resolve_dmd_inference_trace + + num_inference_steps = resolve_dmd_inference_trace( + self.dmd_timestep_contract, + latent_shape=latents.shape[1:], + history_tensors=( + latents_history_short, + latents_history_mid, + latents_memory, + ), + num_stages=pyramid_num_stages, + ).num_steps + else: + num_inference_steps = sum(pyramid_num_inference_steps_list) + + with self.progress_bar(total=num_inference_steps) as progress_bar: + current_attention_kwargs = attention_kwargs + if camera_trajectory is not None and not use_fast: + current_attention_kwargs = dict(attention_kwargs or {}) + ucpe_attention_kwargs = build_ucpe_attention_kwargs_for_chunk( + transformer=self.transformer, + camera_trajectory=camera_trajectory, + height=height, + width=width, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + chunk_index=k, + vae_scale_factor_temporal=self.vae_scale_factor_temporal, + ) + if ucpe_attention_kwargs is None: + raise ValueError( + f"UCPE camera control could not be built for latent chunk {k}; " + "check pose length and camera adapter patching" + ) + current_attention_kwargs.update(ucpe_attention_kwargs) + if is_enable_stage2: + from ..fast.sampling import sample_fast + + # Upsample block noise uses a separate default generator. + # Forwarding the trajectory generator here changes its RNG + # consumption and the generated video. + latents = sample_fast( + self, + latents=latents, + pyramid_num_stages=pyramid_num_stages, + pyramid_num_inference_steps_list=pyramid_num_inference_steps_list, + prompt_embeds=prompt_embeds, + guidance_scale=guidance_scale, + indices_hidden_states=indices_hidden_states, + indices_latents_history_short=indices_latents_history_short, + indices_latents_history_mid=indices_latents_history_mid, + indices_latents_history_long=indices_latents_memory, + latents_history_short=latents_history_short, + latents_history_mid=latents_history_mid, + latents_history_long=latents_memory, + attention_kwargs=current_attention_kwargs, + device=device, + transformer_dtype=transformer_dtype, + camera_trajectory=camera_trajectory, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + chunk_index=k, + camera_restart_each_chunk=False, + ucpe_pixel_center=True, + callback_on_step_end=callback_on_step_end, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + progress_bar=progress_bar, + ) + else: + latents = self.stage1_sample( + latents=latents, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + timesteps=timesteps, + guidance_scale=guidance_scale, + indices_hidden_states=indices_hidden_states, + indices_latents_history_short=indices_latents_history_short, + indices_latents_history_mid=indices_latents_history_mid, + indices_latents_history_long=indices_latents_memory, + latents_history_short=latents_history_short, + latents_history_mid=latents_history_mid, + latents_history_long=latents_memory, + attention_kwargs=current_attention_kwargs, + device=device, + transformer_dtype=transformer_dtype, + generator=generator, + num_warmup_steps=num_warmup_steps, + # ------------ CFG Zero ------------ + use_zero_init=use_zero_init, + zero_steps=zero_steps, + # ------------ Callback ------------ + callback_on_step_end=callback_on_step_end, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + progress_bar=progress_bar, + ) + + if keep_first_frame and ( + (is_first_chunk and image_latents is None) + or (is_skip_first_chunk and is_second_chunk) + ): + image_latents = latents[:, :, 0:1, :, :] + + generated_memory_latents = torch.cat( + [generated_memory_latents, latents], dim=2 + ) + + total_generated_latent_frames += latents.shape[2] + history_latents = torch.cat([history_latents, latents], dim=2) + real_history_latents = history_latents[ + :, :, -total_generated_latent_frames: + ] + current_latents = ( + real_history_latents[:, :, -num_latent_frames_per_chunk:].to( + vae_dtype + ) + / latents_std + + latents_mean + ) + current_video = self.vae.decode(current_latents, return_dict=False)[0] + + if callback_on_chunk_end is not None: + callback_on_chunk_end(k, current_video) + + if callback_on_chunk_state is not None: + if not isinstance(generator, torch.Generator): + raise TypeError( + "resumable WorldCrafter inference requires one torch.Generator" + ) + callback_on_chunk_state( + k, + { + "format": "worldcrafter_chunk_state_v1", + "completed_chunk_index": int(k), + "next_chunk_index": int(k + 1), + "generated_memory_latents": generated_memory_latents.detach().cpu(), + "history_latents": history_latents[ + :, :, -sum(history_sizes) : + ] + .detach() + .cpu(), + "image_latents": ( + image_latents.detach().cpu() + if image_latents is not None + else None + ), + "generator_state": generator.get_state().cpu(), + }, + ) + + if history_video is None: + history_video = current_video + else: + history_video = torch.cat([history_video, current_video], dim=2) + if k == final_chunk_index: + break + + self._current_timestep = None + + if output_type != "latent": + if not use_fast: + # Preserve the existing base output contract. Fast decodes each + # complete 33-frame chunk independently; applying the latent + # length rule again would drop valid RGB frames (330 -> 329). + generated_frames = history_video.size(2) + generated_frames = ( + (generated_frames - 1) + // self.vae_scale_factor_temporal + * self.vae_scale_factor_temporal + + 1 + ) + history_video = history_video[:, :, :generated_frames] + video = self.video_processor.postprocess_video( + history_video, output_type=output_type + ) + else: + video = real_history_latents + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + return (video,) + + return WorldCrafterPipelineOutput(frames=video) diff --git a/worldcrafter/diffusers/pipeline_output.py b/worldcrafter/diffusers/pipeline_output.py new file mode 100644 index 0000000000000000000000000000000000000000..1c37c3f149b649da4cfe31dccc3fb95471a351a3 --- /dev/null +++ b/worldcrafter/diffusers/pipeline_output.py @@ -0,0 +1,10 @@ +from dataclasses import dataclass + +import torch + +from diffusers.utils import BaseOutput + + +@dataclass +class WorldCrafterPipelineOutput(BaseOutput): + frames: torch.Tensor diff --git a/worldcrafter/diffusers/scheduler.py b/worldcrafter/diffusers/scheduler.py new file mode 100644 index 0000000000000000000000000000000000000000..3a4ad6bc0fea20ab497e2b10f68c9fb65dcc0a95 --- /dev/null +++ b/worldcrafter/diffusers/scheduler.py @@ -0,0 +1,931 @@ +import math +from dataclasses import dataclass +from typing import Literal + +import numpy as np +import torch + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.schedulers.scheduling_utils import SchedulerMixin +from diffusers.utils import BaseOutput, deprecate + + +@dataclass +class WorldCrafterSchedulerOutput(BaseOutput): + prev_sample: torch.FloatTensor + model_outputs: torch.FloatTensor | None = None + last_sample: torch.FloatTensor | None = None + this_order: int | None = None + + +class WorldCrafterScheduler(SchedulerMixin, ConfigMixin): + _compatibles = [] + order = 1 + + @register_to_config + def __init__( + self, + num_train_timesteps: int = 1000, + shift: float = 1.0, # Following Stable diffusion 3, + stages: int = 3, + stage_range: list = [0, 1 / 3, 2 / 3, 1], + gamma: float = 1 / 3, + # For UniPC + thresholding: bool = False, + prediction_type: str = "flow_prediction", + solver_order: int = 2, + predict_x0: bool = True, + solver_type: str = "bh2", + lower_order_final: bool = True, + disable_corrector: list[int] = [], + solver_p: SchedulerMixin = None, + use_flow_sigmas: bool = True, + scheduler_type: str = "unipc", # ["euler", "unipc", "dmd"] + use_dynamic_shifting: bool = False, + time_shift_type: Literal["exponential", "linear"] = "linear", + ): + self.timestep_ratios = {} # The timestep ratio for each stage + self.timesteps_per_stage = {} # The detailed timesteps per stage (fix max and min per stage) + self.sigmas_per_stage = {} # always uniform [1000, 0] + self.start_sigmas = {} # for start point / upsample renoise + self.end_sigmas = {} # for end point + self.ori_start_sigmas = {} + + self.init_sigmas_for_each_stage() + self.sigma_min = self.sigmas[-1].item() + self.sigma_max = self.sigmas[0].item() + self.gamma = gamma + + if solver_type not in ["bh1", "bh2"]: + if solver_type in ["midpoint", "heun", "logrho"]: + self.register_to_config(solver_type="bh2") + else: + raise NotImplementedError(f"{solver_type} is not implemented for {self.__class__}") + + self.predict_x0 = predict_x0 + self.model_outputs = [None] * solver_order + self.timestep_list = [None] * solver_order + self.lower_order_nums = 0 + self.disable_corrector = disable_corrector + self.solver_p = solver_p + self.last_sample = None + self._step_index = None + self._begin_index = None + + def init_sigmas(self): + """ + initialize the global timesteps and sigmas + """ + num_train_timesteps = self.config.num_train_timesteps + shift = self.config.shift + + alphas = np.linspace(1, 1 / num_train_timesteps, num_train_timesteps + 1) + sigmas = 1.0 - alphas + sigmas = np.flip(shift * sigmas / (1 + (shift - 1) * sigmas))[:-1].copy() + sigmas = torch.from_numpy(sigmas) + timesteps = (sigmas * num_train_timesteps).clone() + + self._step_index = None + self._begin_index = None + self.timesteps = timesteps + self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication + + def init_sigmas_for_each_stage(self): + """ + Init the timesteps for each stage + """ + self.init_sigmas() + + stage_distance = [] + stages = self.config.stages + training_steps = self.config.num_train_timesteps + stage_range = self.config.stage_range + + # Init the start and end point of each stage + for i_s in range(stages): + # To decide the start and ends point + start_indice = int(stage_range[i_s] * training_steps) + start_indice = max(start_indice, 0) + end_indice = int(stage_range[i_s + 1] * training_steps) + end_indice = min(end_indice, training_steps) + start_sigma = self.sigmas[start_indice].item() + end_sigma = self.sigmas[end_indice].item() if end_indice < training_steps else 0.0 + self.ori_start_sigmas[i_s] = start_sigma + + if i_s != 0: + ori_sigma = 1 - start_sigma + gamma = self.config.gamma + corrected_sigma = (1 / (math.sqrt(1 + (1 / gamma)) * (1 - ori_sigma) + ori_sigma)) * ori_sigma + # corrected_sigma = 1 / (2 - ori_sigma) * ori_sigma + start_sigma = 1 - corrected_sigma + + stage_distance.append(start_sigma - end_sigma) + self.start_sigmas[i_s] = start_sigma + self.end_sigmas[i_s] = end_sigma + + # Determine the ratio of each stage according to flow length + tot_distance = sum(stage_distance) + for i_s in range(stages): + if i_s == 0: + start_ratio = 0.0 + else: + start_ratio = sum(stage_distance[:i_s]) / tot_distance + if i_s == stages - 1: + end_ratio = 0.9999999999999999 + else: + end_ratio = sum(stage_distance[: i_s + 1]) / tot_distance + + self.timestep_ratios[i_s] = (start_ratio, end_ratio) + + # Determine the timesteps and sigmas for each stage + for i_s in range(stages): + timestep_ratio = self.timestep_ratios[i_s] + timestep_max = min(self.timesteps[int(timestep_ratio[0] * training_steps)], 999) + timestep_min = self.timesteps[min(int(timestep_ratio[1] * training_steps), training_steps - 1)] + timesteps = np.linspace(timestep_max, timestep_min, training_steps + 1) + self.timesteps_per_stage[i_s] = ( + timesteps[:-1] if isinstance(timesteps, torch.Tensor) else torch.from_numpy(timesteps[:-1]) + ) + stage_sigmas = np.linspace(0.999, 0, training_steps + 1) + self.sigmas_per_stage[i_s] = torch.from_numpy(stage_sigmas[:-1]) + + @property + def step_index(self): + """ + The index counter for current timestep. It will increase 1 after each scheduler step. + """ + return self._step_index + + @property + def begin_index(self): + """ + The index for the first timestep. It should be set from pipeline with `set_begin_index` method. + """ + return self._begin_index + + def set_begin_index(self, begin_index: int = 0): + """ + Sets the begin index for the scheduler. This function should be run from pipeline before the inference. + + Args: + begin_index (`int`): + The begin index for the scheduler. + """ + self._begin_index = begin_index + + def _sigma_to_t(self, sigma): + return sigma * self.config.num_train_timesteps + + def set_timesteps( + self, + num_inference_steps: int, + stage_index: int | None = None, + device: str | torch.device = None, + sigmas: bool | None = None, + mu: bool | None = None, + is_amplify_first_chunk: bool = False, + ): + """ + Setting the timesteps and sigmas for each stage + """ + if self.config.scheduler_type == "dmd": + if is_amplify_first_chunk: + num_inference_steps = num_inference_steps * 2 + 1 + else: + num_inference_steps = num_inference_steps + 1 + + self.num_inference_steps = num_inference_steps + self.init_sigmas() + + if self.config.stages == 1: + if sigmas is None: + sigmas = np.linspace(1, 1 / self.config.num_train_timesteps, num_inference_steps + 1)[:-1].astype( + np.float32 + ) + if self.config.shift != 1.0: + assert not self.config.use_dynamic_shifting + sigmas = self.time_shift(self.config.shift, 1.0, sigmas) + timesteps = (sigmas * self.config.num_train_timesteps).copy() + sigmas = torch.from_numpy(sigmas) + else: + stage_timesteps = self.timesteps_per_stage[stage_index] + timesteps = np.linspace( + stage_timesteps[0].item(), + stage_timesteps[-1].item(), + num_inference_steps, + ) + + stage_sigmas = self.sigmas_per_stage[stage_index] + ratios = np.linspace(stage_sigmas[0].item(), stage_sigmas[-1].item(), num_inference_steps) + sigmas = torch.from_numpy(ratios) + + self.timesteps = torch.from_numpy(timesteps).to(device=device) + self.sigmas = torch.cat([sigmas, torch.zeros(1)]).to(device=device) + + self._step_index = None + self.reset_scheduler_history() + + if self.config.scheduler_type == "dmd": + self.timesteps = self.timesteps[:-1] + self.sigmas = torch.cat([self.sigmas[:-2], self.sigmas[-1:]]) + + if self.config.use_dynamic_shifting: + assert self.config.shift == 1.0 + self.sigmas = self.time_shift(mu, 1.0, self.sigmas) + if self.config.stages == 1: + self.timesteps = self.sigmas[:-1] * self.config.num_train_timesteps + else: + self.timesteps = self.timesteps_per_stage[stage_index].min() + self.sigmas[:-1] * ( + self.timesteps_per_stage[stage_index].max() - self.timesteps_per_stage[stage_index].min() + ) + + # Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler.time_shift + def time_shift(self, mu: float, sigma: float, t: torch.Tensor): + """ + Apply time shifting to the sigmas. + + Args: + mu (`float`): + The mu parameter for the time shift. + sigma (`float`): + The sigma parameter for the time shift. + t (`torch.Tensor`): + The input timesteps. + + Returns: + `torch.Tensor`: + The time-shifted timesteps. + """ + if self.config.time_shift_type == "exponential": + return self._time_shift_exponential(mu, sigma, t) + elif self.config.time_shift_type == "linear": + return self._time_shift_linear(mu, sigma, t) + + # Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler._time_shift_exponential + def _time_shift_exponential(self, mu, sigma, t): + return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma) + + # Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler._time_shift_linear + def _time_shift_linear(self, mu, sigma, t): + return mu / (mu + (1 / t - 1) ** sigma) + + # ---------------------------------- Euler ---------------------------------- + def index_for_timestep(self, timestep, schedule_timesteps=None): + if schedule_timesteps is None: + schedule_timesteps = self.timesteps + + indices = (schedule_timesteps == timestep).nonzero() + + # The sigma index that is taken for the **very** first `step` + # is always the second index (or the last index if there is only 1) + # This way we can ensure we don't accidentally skip a sigma in + # case we start in the middle of the denoising schedule (e.g. for image-to-image) + pos = 1 if len(indices) > 1 else 0 + + return indices[pos].item() + + def _init_step_index(self, timestep): + if self.begin_index is None: + if isinstance(timestep, torch.Tensor): + timestep = timestep.to(self.timesteps.device) + self._step_index = self.index_for_timestep(timestep) + else: + self._step_index = self._begin_index + + def step_euler( + self, + model_output: torch.FloatTensor, + timestep: float | torch.FloatTensor = None, + sample: torch.FloatTensor = None, + generator: torch.Generator | None = None, + sigma: torch.FloatTensor | None = None, + sigma_next: torch.FloatTensor | None = None, + return_dict: bool = True, + ) -> WorldCrafterSchedulerOutput | tuple: + assert (sigma is None) == (sigma_next is None), "sigma and sigma_next must both be None or both be not None" + + if sigma is None and sigma_next is None: + if ( + isinstance(timestep, int) + or isinstance(timestep, torch.IntTensor) + or isinstance(timestep, torch.LongTensor) + ): + raise ValueError( + ( + "Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to" + " `EulerDiscreteScheduler.step()` is not supported. Make sure to pass" + " one of the `scheduler.timesteps` as a timestep." + ), + ) + + if self.step_index is None: + self._step_index = 0 + + # Upcast to avoid precision issues when computing prev_sample + sample = sample.to(torch.float32) + + if sigma is None and sigma_next is None: + sigma = self.sigmas[self.step_index] + sigma_next = self.sigmas[self.step_index + 1] + + prev_sample = sample + (sigma_next - sigma) * model_output + + # Cast sample back to model compatible dtype + prev_sample = prev_sample.to(model_output.dtype) + + # upon completion increase step index by one + self._step_index += 1 + + if not return_dict: + return (prev_sample,) + + return WorldCrafterSchedulerOutput(prev_sample=prev_sample) + + # ---------------------------------- UniPC ---------------------------------- + def _sigma_to_alpha_sigma_t(self, sigma): + if self.config.use_flow_sigmas: + alpha_t = 1 - sigma + sigma_t = torch.clamp(sigma, min=1e-8) + else: + alpha_t = 1 / ((sigma**2 + 1) ** 0.5) + sigma_t = sigma * alpha_t + + return alpha_t, sigma_t + + def convert_model_output( + self, + model_output: torch.Tensor, + *args, + sample: torch.Tensor = None, + sigma: torch.Tensor = None, + **kwargs, + ) -> torch.Tensor: + r""" + Convert the model output to the corresponding type the UniPC algorithm needs. + + Args: + model_output (`torch.Tensor`): + The direct output from the learned diffusion model. + timestep (`int`): + The current discrete timestep in the diffusion chain. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + + Returns: + `torch.Tensor`: + The converted model output. + """ + timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None) + if sample is None: + if len(args) > 1: + sample = args[1] + else: + raise ValueError("missing `sample` as a required keyword argument") + if timestep is not None: + deprecate( + "timesteps", + "1.0.0", + "Passing `timesteps` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + flag = False + if sigma is None: + flag = True + sigma = self.sigmas[self.step_index] + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma) + + if self.predict_x0: + if self.config.prediction_type == "epsilon": + x0_pred = (sample - sigma_t * model_output) / alpha_t + elif self.config.prediction_type == "sample": + x0_pred = model_output + elif self.config.prediction_type == "v_prediction": + x0_pred = alpha_t * sample - sigma_t * model_output + elif self.config.prediction_type == "flow_prediction": + if flag: + sigma_t = self.sigmas[self.step_index] + else: + sigma_t = sigma + x0_pred = sample - sigma_t * model_output + else: + raise ValueError( + f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`, " + "`v_prediction`, or `flow_prediction` for the UniPCMultistepScheduler." + ) + + if self.config.thresholding: + x0_pred = self._threshold_sample(x0_pred) + + return x0_pred + else: + if self.config.prediction_type == "epsilon": + return model_output + elif self.config.prediction_type == "sample": + epsilon = (sample - alpha_t * model_output) / sigma_t + return epsilon + elif self.config.prediction_type == "v_prediction": + epsilon = alpha_t * model_output + sigma_t * sample + return epsilon + else: + raise ValueError( + f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`, or" + " `v_prediction` for the UniPCMultistepScheduler." + ) + + def multistep_uni_p_bh_update( + self, + model_output: torch.Tensor, + *args, + sample: torch.Tensor = None, + order: int = None, + sigma: torch.Tensor = None, + sigma_next: torch.Tensor = None, + **kwargs, + ) -> torch.Tensor: + """ + One step for the UniP (B(h) version). Alternatively, `self.solver_p` is used if is specified. + + Args: + model_output (`torch.Tensor`): + The direct output from the learned diffusion model at the current timestep. + prev_timestep (`int`): + The previous discrete timestep in the diffusion chain. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + order (`int`): + The order of UniP at this timestep (corresponds to the *p* in UniPC-p). + + Returns: + `torch.Tensor`: + The sample tensor at the previous timestep. + """ + prev_timestep = args[0] if len(args) > 0 else kwargs.pop("prev_timestep", None) + if sample is None: + if len(args) > 1: + sample = args[1] + else: + raise ValueError("missing `sample` as a required keyword argument") + if order is None: + if len(args) > 2: + order = args[2] + else: + raise ValueError("missing `order` as a required keyword argument") + if prev_timestep is not None: + deprecate( + "prev_timestep", + "1.0.0", + "Passing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + model_output_list = self.model_outputs + + s0 = self.timestep_list[-1] + m0 = model_output_list[-1] + x = sample + + if self.solver_p: + x_t = self.solver_p.step(model_output, s0, x).prev_sample + return x_t + + if sigma_next is None and sigma is None: + sigma_t, sigma_s0 = self.sigmas[self.step_index + 1], self.sigmas[self.step_index] + else: + sigma_t, sigma_s0 = sigma_next, sigma + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t) + alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0) + + lambda_t = torch.log(alpha_t) - torch.log(sigma_t) + lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0) + + h = lambda_t - lambda_s0 + device = sample.device + + rks = [] + D1s = [] + for i in range(1, order): + si = self.step_index - i + mi = model_output_list[-(i + 1)] + alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si]) + lambda_si = torch.log(alpha_si) - torch.log(sigma_si) + rk = (lambda_si - lambda_s0) / h + rks.append(rk) + D1s.append((mi - m0) / rk) + + rks.append(1.0) + rks = torch.tensor(rks, device=device) + + R = [] + b = [] + + hh = -h if self.predict_x0 else h + h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1 + h_phi_k = h_phi_1 / hh - 1 + + factorial_i = 1 + + if self.config.solver_type == "bh1": + B_h = hh + elif self.config.solver_type == "bh2": + B_h = torch.expm1(hh) + else: + raise NotImplementedError() + + for i in range(1, order + 1): + R.append(torch.pow(rks, i - 1)) + b.append(h_phi_k * factorial_i / B_h) + factorial_i *= i + 1 + h_phi_k = h_phi_k / hh - 1 / factorial_i + + R = torch.stack(R) + b = torch.tensor(b, device=device) + + if len(D1s) > 0: + D1s = torch.stack(D1s, dim=1) # (B, K) + # for order 2, we use a simplified version + if order == 2: + rhos_p = torch.tensor([0.5], dtype=x.dtype, device=device) + else: + rhos_p = torch.linalg.solve(R[:-1, :-1], b[:-1]).to(device).to(x.dtype) + else: + D1s = None + + if self.predict_x0: + x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0 + if D1s is not None: + pred_res = torch.einsum("k,bkc...->bc...", rhos_p, D1s) + else: + pred_res = 0 + x_t = x_t_ - alpha_t * B_h * pred_res + else: + x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0 + if D1s is not None: + pred_res = torch.einsum("k,bkc...->bc...", rhos_p, D1s) + else: + pred_res = 0 + x_t = x_t_ - sigma_t * B_h * pred_res + + x_t = x_t.to(x.dtype) + return x_t + + def multistep_uni_c_bh_update( + self, + this_model_output: torch.Tensor, + *args, + last_sample: torch.Tensor = None, + this_sample: torch.Tensor = None, + order: int = None, + sigma_before: torch.Tensor = None, + sigma: torch.Tensor = None, + **kwargs, + ) -> torch.Tensor: + """ + One step for the UniC (B(h) version). + + Args: + this_model_output (`torch.Tensor`): + The model outputs at `x_t`. + this_timestep (`int`): + The current timestep `t`. + last_sample (`torch.Tensor`): + The generated sample before the last predictor `x_{t-1}`. + this_sample (`torch.Tensor`): + The generated sample after the last predictor `x_{t}`. + order (`int`): + The `p` of UniC-p at this step. The effective order of accuracy should be `order + 1`. + + Returns: + `torch.Tensor`: + The corrected sample tensor at the current timestep. + """ + this_timestep = args[0] if len(args) > 0 else kwargs.pop("this_timestep", None) + if last_sample is None: + if len(args) > 1: + last_sample = args[1] + else: + raise ValueError("missing `last_sample` as a required keyword argument") + if this_sample is None: + if len(args) > 2: + this_sample = args[2] + else: + raise ValueError("missing `this_sample` as a required keyword argument") + if order is None: + if len(args) > 3: + order = args[3] + else: + raise ValueError("missing `order` as a required keyword argument") + if this_timestep is not None: + deprecate( + "this_timestep", + "1.0.0", + "Passing `this_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + model_output_list = self.model_outputs + + m0 = model_output_list[-1] + x = last_sample + x_t = this_sample + model_t = this_model_output + + if sigma_before is None and sigma is None: + sigma_t, sigma_s0 = self.sigmas[self.step_index], self.sigmas[self.step_index - 1] + else: + sigma_t, sigma_s0 = sigma, sigma_before + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t) + alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0) + + lambda_t = torch.log(alpha_t) - torch.log(sigma_t) + lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0) + + h = lambda_t - lambda_s0 + device = this_sample.device + + rks = [] + D1s = [] + for i in range(1, order): + si = self.step_index - (i + 1) + mi = model_output_list[-(i + 1)] + alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si]) + lambda_si = torch.log(alpha_si) - torch.log(sigma_si) + rk = (lambda_si - lambda_s0) / h + rks.append(rk) + D1s.append((mi - m0) / rk) + + rks.append(1.0) + rks = torch.tensor(rks, device=device) + + R = [] + b = [] + + hh = -h if self.predict_x0 else h + h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1 + h_phi_k = h_phi_1 / hh - 1 + + factorial_i = 1 + + if self.config.solver_type == "bh1": + B_h = hh + elif self.config.solver_type == "bh2": + B_h = torch.expm1(hh) + else: + raise NotImplementedError() + + for i in range(1, order + 1): + R.append(torch.pow(rks, i - 1)) + b.append(h_phi_k * factorial_i / B_h) + factorial_i *= i + 1 + h_phi_k = h_phi_k / hh - 1 / factorial_i + + R = torch.stack(R) + b = torch.tensor(b, device=device) + + if len(D1s) > 0: + D1s = torch.stack(D1s, dim=1) + else: + D1s = None + + # for order 1, we use a simplified version + if order == 1: + rhos_c = torch.tensor([0.5], dtype=x.dtype, device=device) + else: + rhos_c = torch.linalg.solve(R, b).to(device).to(x.dtype) + + if self.predict_x0: + x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0 + if D1s is not None: + corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s) + else: + corr_res = 0 + D1_t = model_t - m0 + x_t = x_t_ - alpha_t * B_h * (corr_res + rhos_c[-1] * D1_t) + else: + x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0 + if D1s is not None: + corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s) + else: + corr_res = 0 + D1_t = model_t - m0 + x_t = x_t_ - sigma_t * B_h * (corr_res + rhos_c[-1] * D1_t) + x_t = x_t.to(x.dtype) + return x_t + + def step_unipc( + self, + model_output: torch.Tensor, + timestep: int | torch.Tensor = None, + sample: torch.Tensor = None, + return_dict: bool = True, + model_outputs: list = None, + timestep_list: list = None, + sigma_before: torch.Tensor = None, + sigma: torch.Tensor = None, + sigma_next: torch.Tensor = None, + cus_step_index: int = None, + cus_lower_order_num: int = None, + cus_this_order: int = None, + cus_last_sample: torch.Tensor = None, + ) -> WorldCrafterSchedulerOutput | tuple: + if self.num_inference_steps is None: + raise ValueError( + "Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler" + ) + + if cus_step_index is None: + if self.step_index is None: + self._step_index = 0 + else: + self._step_index = cus_step_index + + if cus_lower_order_num is not None: + self.lower_order_nums = cus_lower_order_num + + if cus_this_order is not None: + self.this_order = cus_this_order + + if cus_last_sample is not None: + self.last_sample = cus_last_sample + + use_corrector = ( + self.step_index > 0 and self.step_index - 1 not in self.disable_corrector and self.last_sample is not None + ) + + # Convert model output using the proper conversion method + model_output_convert = self.convert_model_output(model_output, sample=sample, sigma=sigma) + + if model_outputs is not None and timestep_list is not None: + self.model_outputs = model_outputs[:-1] + self.timestep_list = timestep_list[:-1] + + if use_corrector: + sample = self.multistep_uni_c_bh_update( + this_model_output=model_output_convert, + last_sample=self.last_sample, + this_sample=sample, + order=self.this_order, + sigma_before=sigma_before, + sigma=sigma, + ) + + if model_outputs is not None and timestep_list is not None: + model_outputs[-1] = model_output_convert + self.model_outputs = model_outputs[1:] + self.timestep_list = timestep_list[1:] + else: + for i in range(self.config.solver_order - 1): + self.model_outputs[i] = self.model_outputs[i + 1] + self.timestep_list[i] = self.timestep_list[i + 1] + self.model_outputs[-1] = model_output_convert + self.timestep_list[-1] = timestep + + if self.config.lower_order_final: + this_order = min(self.config.solver_order, len(self.timesteps) - self.step_index) + else: + this_order = self.config.solver_order + self.this_order = min(this_order, self.lower_order_nums + 1) # warmup for multistep + assert self.this_order > 0 + + self.last_sample = sample + prev_sample = self.multistep_uni_p_bh_update( + model_output=model_output, # pass the original non-converted model output, in case solver-p is used + sample=sample, + order=self.this_order, + sigma=sigma, + sigma_next=sigma_next, + ) + + if cus_lower_order_num is None: + if self.lower_order_nums < self.config.solver_order: + self.lower_order_nums += 1 + + # upon completion increase step index by one + if cus_step_index is None: + self._step_index += 1 + + if not return_dict: + return (prev_sample, model_outputs, self.last_sample, self.this_order) + + return WorldCrafterSchedulerOutput( + prev_sample=prev_sample, + model_outputs=model_outputs, + last_sample=self.last_sample, + this_order=self.this_order, + ) + + # ---------------------------------- For DMD ---------------------------------- + def add_noise(self, original_samples, noise, timestep, sigmas, timesteps): + sigmas = sigmas.to(noise.device) + timesteps = timesteps.to(noise.device) + timestep_id = torch.argmin((timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) + sigma = sigmas[timestep_id].reshape(-1, 1, 1, 1, 1) + sample = (1 - sigma) * original_samples + sigma * noise + return sample.type_as(noise) + + def convert_flow_pred_to_x0(self, flow_pred, xt, timestep, sigmas, timesteps): + # use higher precision for calculations + original_dtype = flow_pred.dtype + device = flow_pred.device + flow_pred, xt, sigmas, timesteps = (x.double().to(device) for x in (flow_pred, xt, sigmas, timesteps)) + + timestep_id = torch.argmin((timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) + sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1, 1) + x0_pred = xt - sigma_t * flow_pred + return x0_pred.to(original_dtype) + + def step_dmd( + self, + model_output: torch.FloatTensor, + timestep: float | torch.FloatTensor = None, + sample: torch.FloatTensor = None, + generator: torch.Generator | None = None, + return_dict: bool = True, + cur_sampling_step: int = 0, + dmd_noisy_tensor: torch.FloatTensor | None = None, + dmd_sigmas: torch.FloatTensor | None = None, + dmd_timesteps: torch.FloatTensor | None = None, + all_timesteps: torch.FloatTensor | None = None, + ): + pred_image_or_video = self.convert_flow_pred_to_x0( + flow_pred=model_output, + xt=sample, + timestep=torch.full((model_output.shape[0],), timestep, dtype=torch.long, device=model_output.device), + sigmas=dmd_sigmas, + timesteps=dmd_timesteps, + ) + if cur_sampling_step < len(all_timesteps) - 1: + prev_sample = self.add_noise( + pred_image_or_video, + dmd_noisy_tensor, + torch.full( + (model_output.shape[0],), + all_timesteps[cur_sampling_step + 1], + dtype=torch.long, + device=model_output.device, + ), + sigmas=dmd_sigmas, + timesteps=dmd_timesteps, + ) + else: + prev_sample = pred_image_or_video + + if not return_dict: + return (prev_sample,) + + return WorldCrafterSchedulerOutput(prev_sample=prev_sample) + + # ---------------------------------- Merge ---------------------------------- + def step( + self, + model_output: torch.FloatTensor, + timestep: float | torch.FloatTensor = None, + sample: torch.FloatTensor = None, + generator: torch.Generator | None = None, + return_dict: bool = True, + # For DMD + cur_sampling_step: int = 0, + dmd_noisy_tensor: torch.FloatTensor | None = None, + dmd_sigmas: torch.FloatTensor | None = None, + dmd_timesteps: torch.FloatTensor | None = None, + all_timesteps: torch.FloatTensor | None = None, + ) -> WorldCrafterSchedulerOutput | tuple: + if self.config.scheduler_type == "euler": + return self.step_euler( + model_output=model_output, + timestep=timestep, + sample=sample, + generator=generator, + return_dict=return_dict, + ) + elif self.config.scheduler_type == "unipc": + return self.step_unipc( + model_output=model_output, + timestep=timestep, + sample=sample, + return_dict=return_dict, + ) + elif self.config.scheduler_type == "dmd": + return self.step_dmd( + model_output=model_output, + timestep=timestep, + sample=sample, + generator=generator, + return_dict=return_dict, + cur_sampling_step=cur_sampling_step, + dmd_noisy_tensor=dmd_noisy_tensor, + dmd_sigmas=dmd_sigmas, + dmd_timesteps=dmd_timesteps, + all_timesteps=all_timesteps, + ) + else: + raise NotImplementedError + + def reset_scheduler_history(self): + self.model_outputs = [None] * self.config.solver_order + self.timestep_list = [None] * self.config.solver_order + self.lower_order_nums = 0 + self.disable_corrector = self.config.disable_corrector + self.solver_p = self.config.solver_p + self.last_sample = None + self._step_index = None + self._begin_index = None + + def __len__(self): + return self.config.num_train_timesteps diff --git a/worldcrafter/diffusers/transformer.py b/worldcrafter/diffusers/transformer.py new file mode 100644 index 0000000000000000000000000000000000000000..59c60b03cbfd54fdad35e25aa8fd18f816b9bc3d --- /dev/null +++ b/worldcrafter/diffusers/transformer.py @@ -0,0 +1,1097 @@ +import json +import math +import os +from functools import lru_cache +from typing import Any + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin +from diffusers.models._modeling_parallel import ( + ContextParallelInput, + ContextParallelOutput, +) +from diffusers.models.attention import AttentionMixin, AttentionModuleMixin, FeedForward +from diffusers.models.attention_dispatch import dispatch_attention_fn +from diffusers.models.cache_utils import CacheMixin +from diffusers.models.embeddings import ( + PixArtAlphaTextProjection, + TimestepEmbedding, + Timesteps, +) +from diffusers.models.modeling_outputs import Transformer2DModelOutput +from diffusers.models.modeling_utils import ModelMixin +from diffusers.models.normalization import FP32LayerNorm +from diffusers.utils import apply_lora_scale, logging +from diffusers.utils.torch_utils import maybe_allow_in_graph + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +def pad_for_3d_conv(x, kernel_size): + b, c, t, h, w = x.shape + pt, ph, pw = kernel_size + pad_t = (pt - (t % pt)) % pt + pad_h = (ph - (h % ph)) % ph + pad_w = (pw - (w % pw)) % pw + return torch.nn.functional.pad(x, (0, pad_w, 0, pad_h, 0, pad_t), mode="replicate") + + +def center_down_sample_3d(x, kernel_size): + return torch.nn.functional.avg_pool3d(x, kernel_size, stride=kernel_size) + + +def apply_rotary_emb_transposed( + hidden_states: torch.Tensor, + freqs_cis: torch.Tensor, +): + x_1, x_2 = hidden_states.unflatten(-1, (-1, 2)).unbind(-1) + cos, sin = freqs_cis.unsqueeze(-2).chunk(2, dim=-1) + out = torch.empty_like(hidden_states) + out[..., 0::2] = x_1 * cos[..., 0::2] - x_2 * sin[..., 1::2] + out[..., 1::2] = x_1 * sin[..., 1::2] + x_2 * cos[..., 0::2] + return out.type_as(hidden_states) + + +def _get_qkv_projections( + attn: "WorldCrafterAttention", + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, +): + # encoder_hidden_states is only passed for cross-attention + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + + if attn.fused_projections: + if not attn.is_cross_attention: + # In self-attention layers, we can fuse the entire QKV projection into a single linear + query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1) + else: + # In cross-attention layers, we can only fuse the KV projections into a single linear + query = attn.to_q(hidden_states) + key, value = attn.to_kv(encoder_hidden_states).chunk(2, dim=-1) + else: + query = attn.to_q(hidden_states) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + return query, key, value + + +class WorldCrafterOutputNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = False): + super().__init__() + self.scale_shift_table = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5) + self.norm = FP32LayerNorm(dim, eps, elementwise_affine=False) + + def forward( + self, + hidden_states: torch.Tensor, + temb: torch.Tensor, + original_context_length: int, + ): + temb = temb[:, -original_context_length:, :] + shift, scale = ( + self.scale_shift_table.unsqueeze(0).to(temb.device) + temb.unsqueeze(2) + ).chunk(2, dim=2) + shift, scale = shift.squeeze(2).to(hidden_states.device), scale.squeeze(2).to( + hidden_states.device + ) + hidden_states = hidden_states[:, -original_context_length:, :] + hidden_states = ( + self.norm(hidden_states.float()) * (1 + scale) + shift + ).type_as(hidden_states) + return hidden_states + + +class WorldCrafterAttnProcessor: + _attention_backend = None + _parallel_config = None + + def __init__(self): + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError( + "WorldCrafterAttnProcessor requires PyTorch 2.0. To use it, please upgrade PyTorch to version 2.0 or higher." + ) + + def __call__( + self, + attn: "WorldCrafterAttention", + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor | None = None, + attention_mask: torch.Tensor | None = None, + rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, + original_context_length: int = None, + ) -> torch.Tensor: + query, key, value = _get_qkv_projections( + attn, hidden_states, encoder_hidden_states + ) + + query = attn.norm_q(query) + key = attn.norm_k(key) + + query = query.unflatten(2, (attn.heads, -1)) + key = key.unflatten(2, (attn.heads, -1)) + value = value.unflatten(2, (attn.heads, -1)) + + if rotary_emb is not None: + query = apply_rotary_emb_transposed(query, rotary_emb) + key = apply_rotary_emb_transposed(key, rotary_emb) + + if not attn.is_cross_attention and attn.is_amplify_history: + history_seq_len = hidden_states.shape[1] - original_context_length + + if history_seq_len > 0: + scale_key = 1.0 + torch.sigmoid(attn.history_key_scale) * ( + attn.max_scale - 1.0 + ) + if attn.history_scale_mode == "per_head": + scale_key = scale_key.view(1, 1, -1, 1) + key = torch.cat( + [key[:, :history_seq_len] * scale_key, key[:, history_seq_len:]], + dim=1, + ) + + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + backend=self._attention_backend, + # Reference: https://github.com/huggingface/diffusers/pull/12909 + parallel_config=( + self._parallel_config if encoder_hidden_states is None else None + ), + ) + hidden_states = hidden_states.flatten(2, 3) + hidden_states = hidden_states.type_as(query) + + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + return hidden_states + + +class WorldCrafterAttention(torch.nn.Module, AttentionModuleMixin): + _default_processor_cls = WorldCrafterAttnProcessor + _available_processors = [WorldCrafterAttnProcessor] + + def __init__( + self, + dim: int, + heads: int = 8, + dim_head: int = 64, + eps: float = 1e-5, + dropout: float = 0.0, + added_kv_proj_dim: int | None = None, + cross_attention_dim_head: int | None = None, + processor=None, + is_cross_attention=None, + is_amplify_history=False, + history_scale_mode="per_head", # [scalar, per_head] + ): + super().__init__() + + self.inner_dim = dim_head * heads + self.heads = heads + self.added_kv_proj_dim = added_kv_proj_dim + self.cross_attention_dim_head = cross_attention_dim_head + self.kv_inner_dim = ( + self.inner_dim + if cross_attention_dim_head is None + else cross_attention_dim_head * heads + ) + + self.to_q = torch.nn.Linear(dim, self.inner_dim, bias=True) + self.to_k = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) + self.to_v = torch.nn.Linear(dim, self.kv_inner_dim, bias=True) + self.to_out = torch.nn.ModuleList( + [ + torch.nn.Linear(self.inner_dim, dim, bias=True), + torch.nn.Dropout(dropout), + ] + ) + self.norm_q = torch.nn.RMSNorm( + dim_head * heads, eps=eps, elementwise_affine=True + ) + self.norm_k = torch.nn.RMSNorm( + dim_head * heads, eps=eps, elementwise_affine=True + ) + + self.add_k_proj = self.add_v_proj = None + if added_kv_proj_dim is not None: + self.add_k_proj = torch.nn.Linear( + added_kv_proj_dim, self.inner_dim, bias=True + ) + self.add_v_proj = torch.nn.Linear( + added_kv_proj_dim, self.inner_dim, bias=True + ) + self.norm_added_k = torch.nn.RMSNorm(dim_head * heads, eps=eps) + + if is_cross_attention is not None: + self.is_cross_attention = is_cross_attention + else: + self.is_cross_attention = cross_attention_dim_head is not None + + self.set_processor(processor) + + self.is_amplify_history = is_amplify_history + if is_amplify_history: + if history_scale_mode == "scalar": + self.history_key_scale = nn.Parameter(torch.ones(1)) + elif history_scale_mode == "per_head": + self.history_key_scale = nn.Parameter(torch.ones(heads)) + else: + raise ValueError(f"Unknown history_scale_mode: {history_scale_mode}") + self.history_scale_mode = history_scale_mode + self.max_scale = 10.0 + + def fuse_projections(self): + if getattr(self, "fused_projections", False): + return + + if not self.is_cross_attention: + concatenated_weights = torch.cat( + [self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data] + ) + concatenated_bias = torch.cat( + [self.to_q.bias.data, self.to_k.bias.data, self.to_v.bias.data] + ) + out_features, in_features = concatenated_weights.shape + with torch.device("meta"): + self.to_qkv = nn.Linear(in_features, out_features, bias=True) + self.to_qkv.load_state_dict( + {"weight": concatenated_weights, "bias": concatenated_bias}, + strict=True, + assign=True, + ) + else: + concatenated_weights = torch.cat( + [self.to_k.weight.data, self.to_v.weight.data] + ) + concatenated_bias = torch.cat([self.to_k.bias.data, self.to_v.bias.data]) + out_features, in_features = concatenated_weights.shape + with torch.device("meta"): + self.to_kv = nn.Linear(in_features, out_features, bias=True) + self.to_kv.load_state_dict( + {"weight": concatenated_weights, "bias": concatenated_bias}, + strict=True, + assign=True, + ) + + if self.added_kv_proj_dim is not None: + concatenated_weights = torch.cat( + [self.add_k_proj.weight.data, self.add_v_proj.weight.data] + ) + concatenated_bias = torch.cat( + [self.add_k_proj.bias.data, self.add_v_proj.bias.data] + ) + out_features, in_features = concatenated_weights.shape + with torch.device("meta"): + self.to_added_kv = nn.Linear(in_features, out_features, bias=True) + self.to_added_kv.load_state_dict( + {"weight": concatenated_weights, "bias": concatenated_bias}, + strict=True, + assign=True, + ) + + self.fused_projections = True + + @torch.no_grad() + def unfuse_projections(self): + if not getattr(self, "fused_projections", False): + return + + if hasattr(self, "to_qkv"): + delattr(self, "to_qkv") + if hasattr(self, "to_kv"): + delattr(self, "to_kv") + if hasattr(self, "to_added_kv"): + delattr(self, "to_added_kv") + + self.fused_projections = False + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor | None = None, + attention_mask: torch.Tensor | None = None, + rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, + original_context_length: int = None, + **kwargs, + ) -> torch.Tensor: + return self.processor( + self, + hidden_states, + encoder_hidden_states, + attention_mask, + rotary_emb, + original_context_length, + **kwargs, + ) + + +class WorldCrafterTimeTextEmbedding(nn.Module): + def __init__( + self, + dim: int, + time_freq_dim: int, + time_proj_dim: int, + text_embed_dim: int, + ): + super().__init__() + + self.timesteps_proj = Timesteps( + num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0 + ) + self.time_embedder = TimestepEmbedding( + in_channels=time_freq_dim, time_embed_dim=dim + ) + self.act_fn = nn.SiLU() + self.time_proj = nn.Linear(dim, time_proj_dim) + self.text_embedder = PixArtAlphaTextProjection( + text_embed_dim, dim, act_fn="gelu_tanh" + ) + + def forward( + self, + timestep: torch.Tensor, + encoder_hidden_states: torch.Tensor | None = None, + is_return_encoder_hidden_states: bool = True, + ): + timestep = self.timesteps_proj(timestep) + + time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype + if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8: + timestep = timestep.to(time_embedder_dtype) + temb = self.time_embedder(timestep).type_as(encoder_hidden_states) + timestep_proj = self.time_proj(self.act_fn(temb)) + + if encoder_hidden_states is not None and is_return_encoder_hidden_states: + encoder_hidden_states = self.text_embedder(encoder_hidden_states) + + return temb, timestep_proj, encoder_hidden_states + + +class WorldCrafterRotaryPosEmbed(nn.Module): + def __init__(self, rope_dim, theta): + super().__init__() + self.DT, self.DY, self.DX = rope_dim + self.theta = theta + self.register_buffer( + "freqs_base_t", self._get_freqs_base(self.DT), persistent=False + ) + self.register_buffer( + "freqs_base_y", self._get_freqs_base(self.DY), persistent=False + ) + self.register_buffer( + "freqs_base_x", self._get_freqs_base(self.DX), persistent=False + ) + + def _get_freqs_base(self, dim): + return 1.0 / ( + self.theta + ** (torch.arange(0, dim, 2, dtype=torch.float32)[: (dim // 2)] / dim) + ) + + @torch.no_grad() + def get_frequency_batched(self, freqs_base, pos): + freqs = torch.einsum("d,bthw->dbthw", freqs_base, pos) + freqs = freqs.repeat_interleave(2, dim=0) + return freqs.cos(), freqs.sin() + + @torch.no_grad() + @lru_cache(maxsize=32) + def _get_spatial_meshgrid(self, height, width, device_str): + device = torch.device(device_str) + grid_y_coords = torch.arange(height, device=device, dtype=torch.float32) + grid_x_coords = torch.arange(width, device=device, dtype=torch.float32) + grid_y, grid_x = torch.meshgrid(grid_y_coords, grid_x_coords, indexing="ij") + return grid_y, grid_x + + @torch.no_grad() + def forward(self, frame_indices, height, width, device): + batch_size = frame_indices.shape[0] + num_frames = frame_indices.shape[1] + + frame_indices = frame_indices.to(device=device, dtype=torch.float32) + grid_y, grid_x = self._get_spatial_meshgrid(height, width, str(device)) + + grid_t = frame_indices[:, :, None, None].expand( + batch_size, num_frames, height, width + ) + grid_y_batch = grid_y[None, None, :, :].expand(batch_size, num_frames, -1, -1) + grid_x_batch = grid_x[None, None, :, :].expand(batch_size, num_frames, -1, -1) + + freqs_cos_t, freqs_sin_t = self.get_frequency_batched(self.freqs_base_t, grid_t) + freqs_cos_y, freqs_sin_y = self.get_frequency_batched( + self.freqs_base_y, grid_y_batch + ) + freqs_cos_x, freqs_sin_x = self.get_frequency_batched( + self.freqs_base_x, grid_x_batch + ) + + result = torch.cat( + [ + freqs_cos_t, + freqs_cos_y, + freqs_cos_x, + freqs_sin_t, + freqs_sin_y, + freqs_sin_x, + ], + dim=0, + ) + + return result.permute(1, 0, 2, 3, 4) + + +@maybe_allow_in_graph +class WorldCrafterTransformerBlock(nn.Module): + def __init__( + self, + dim: int, + ffn_dim: int, + num_heads: int, + qk_norm: str = "rms_norm_across_heads", + cross_attn_norm: bool = False, + eps: float = 1e-6, + added_kv_proj_dim: int | None = None, + guidance_cross_attn: bool = False, + is_amplify_history: bool = False, + history_scale_mode: str = "per_head", # [scalar, per_head] + ): + super().__init__() + + # 1. Self-attention + self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False) + self.attn1 = WorldCrafterAttention( + dim=dim, + heads=num_heads, + dim_head=dim // num_heads, + eps=eps, + cross_attention_dim_head=None, + processor=WorldCrafterAttnProcessor(), + is_amplify_history=is_amplify_history, + history_scale_mode=history_scale_mode, + ) + + # 2. Cross-attention + self.attn2 = WorldCrafterAttention( + dim=dim, + heads=num_heads, + dim_head=dim // num_heads, + eps=eps, + added_kv_proj_dim=added_kv_proj_dim, + cross_attention_dim_head=dim // num_heads, + processor=WorldCrafterAttnProcessor(), + ) + self.norm2 = ( + FP32LayerNorm(dim, eps, elementwise_affine=True) + if cross_attn_norm + else nn.Identity() + ) + + # 3. Feed-forward + self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate") + self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False) + + self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) + + # 4. Guidance cross-attention + self.guidance_cross_attn = guidance_cross_attn + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + temb: torch.Tensor, + rotary_emb: torch.Tensor, + original_context_length: int = None, + camera_control_ucpe_input: dict | None = None, + ) -> torch.Tensor: + if temb.ndim == 4: + shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( + self.scale_shift_table.unsqueeze(0) + temb.float() + ).chunk(6, dim=2) + # batch_size, seq_len, 1, inner_dim + shift_msa = shift_msa.squeeze(2) + scale_msa = scale_msa.squeeze(2) + gate_msa = gate_msa.squeeze(2) + c_shift_msa = c_shift_msa.squeeze(2) + c_scale_msa = c_scale_msa.squeeze(2) + c_gate_msa = c_gate_msa.squeeze(2) + else: + shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( + self.scale_shift_table + temb.float() + ).chunk(6, dim=1) + + # 1. Self-attention + norm_hidden_states = ( + self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa + ).type_as(hidden_states) + history_seq_len = ( + hidden_states.shape[1] - original_context_length + if original_context_length is not None + else 0 + ) + current_norm_hidden_states = norm_hidden_states[:, history_seq_len:, :] + + if hasattr(self, "cam_self_attn") and camera_control_ucpe_input is not None: + if self.cam_self_attn.adaptation_method == "before": + cam_current = self.cam_self_attn( + current_norm_hidden_states, camera_control_ucpe_input + ) + cam_full = torch.zeros_like(norm_hidden_states) + cam_full[:, history_seq_len:, :] = cam_current + norm_hidden_states = norm_hidden_states + cam_full + + attn_output = self.attn1( + norm_hidden_states, None, None, rotary_emb, original_context_length + ) + + if hasattr(self, "cam_self_attn") and camera_control_ucpe_input is not None: + if self.cam_self_attn.adaptation_method == "parallel": + cam_current = self.cam_self_attn( + current_norm_hidden_states, camera_control_ucpe_input + ) + cam_full = torch.zeros_like(attn_output) + cam_full[:, history_seq_len:, :] = cam_current + attn_output = attn_output + cam_full + + hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as( + hidden_states + ) + + if hasattr(self, "cam_self_attn") and camera_control_ucpe_input is not None: + if self.cam_self_attn.adaptation_method == "after": + cam_current = self.cam_self_attn( + hidden_states[:, history_seq_len:, :], camera_control_ucpe_input + ) + hidden_states = hidden_states.clone() + hidden_states[:, history_seq_len:, :] = ( + hidden_states[:, history_seq_len:, :] + cam_current + ) + + # 2. Cross-attention + if self.guidance_cross_attn: + history_seq_len = hidden_states.shape[1] - original_context_length + + history_hidden_states, hidden_states = torch.split( + hidden_states, [history_seq_len, original_context_length], dim=1 + ) + norm_hidden_states = self.norm2(hidden_states.float()).type_as( + hidden_states + ) + attn_output = self.attn2( + norm_hidden_states, + encoder_hidden_states, + None, + None, + original_context_length, + ) + hidden_states = hidden_states + attn_output + hidden_states = torch.cat([history_hidden_states, hidden_states], dim=1) + else: + norm_hidden_states = self.norm2(hidden_states.float()).type_as( + hidden_states + ) + attn_output = self.attn2( + norm_hidden_states, + encoder_hidden_states, + None, + None, + original_context_length, + ) + hidden_states = hidden_states + attn_output + + # 3. Feed-forward + norm_hidden_states = ( + self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa + ).type_as(hidden_states) + ff_output = self.ffn(norm_hidden_states) + hidden_states = ( + hidden_states.float() + ff_output.float() * c_gate_msa + ).type_as(hidden_states) + + return hidden_states + + +class WorldCrafterTransformer3DModel( + ModelMixin, + ConfigMixin, + PeftAdapterMixin, + FromOriginalModelMixin, + CacheMixin, + AttentionMixin, +): + r""" + A Transformer model for video-like data used in the WorldCrafter model. + + Args: + patch_size (`tuple[int]`, defaults to `(1, 2, 2)`): + 3D patch dimensions for video embedding (t_patch, h_patch, w_patch). + num_attention_heads (`int`, defaults to `40`): + Fixed length for text embeddings. + attention_head_dim (`int`, defaults to `128`): + The number of channels in each head. + in_channels (`int`, defaults to `16`): + The number of channels in the input. + out_channels (`int`, defaults to `16`): + The number of channels in the output. + text_dim (`int`, defaults to `512`): + Input dimension for text embeddings. + freq_dim (`int`, defaults to `256`): + Dimension for sinusoidal time embeddings. + ffn_dim (`int`, defaults to `13824`): + Intermediate dimension in feed-forward network. + num_layers (`int`, defaults to `40`): + The number of layers of transformer blocks to use. + window_size (`tuple[int]`, defaults to `(-1, -1)`): + Window size for local attention (-1 indicates global attention). + cross_attn_norm (`bool`, defaults to `True`): + Enable cross-attention normalization. + qk_norm (`bool`, defaults to `True`): + Enable query/key normalization. + eps (`float`, defaults to `1e-6`): + Epsilon value for normalization layers. + add_img_emb (`bool`, defaults to `False`): + Whether to use img_emb. + added_kv_proj_dim (`int`, *optional*, defaults to `None`): + The number of channels to use for the added key and value projections. If `None`, no projection is used. + """ + + _supports_gradient_checkpointing = True + _skip_layerwise_casting_patterns = [ + "patch_embedding", + "patch_short", + "patch_mid", + "patch_memory", + "condition_embedder", + "norm", + ] + _no_split_modules = ["WorldCrafterTransformerBlock", "WorldCrafterOutputNorm"] + _keep_in_fp32_modules = [ + "time_embedder", + "scale_shift_table", + "norm1", + "norm2", + "norm3", + "history_key_scale", + ] + _keys_to_ignore_on_load_unexpected = ["norm_added_q", r"patch_long\..*"] + _repeated_blocks = ["WorldCrafterTransformerBlock"] + _cp_plan = { + # Input split at attn level and ffn level. + "blocks.*.attn1": { + "hidden_states": ContextParallelInput( + split_dim=1, expected_dims=3, split_output=False + ), + "rotary_emb": ContextParallelInput( + split_dim=1, expected_dims=3, split_output=False + ), + }, + "blocks.*.attn2": { + "hidden_states": ContextParallelInput( + split_dim=1, expected_dims=3, split_output=False + ), + }, + "blocks.*.ffn": { + "hidden_states": ContextParallelInput( + split_dim=1, expected_dims=3, split_output=False + ), + }, + # Output gather at attn level and ffn level. + **{ + f"blocks.{i}.attn1": ContextParallelOutput(gather_dim=1, expected_dims=3) + for i in range(40) + }, + **{ + f"blocks.{i}.attn2": ContextParallelOutput(gather_dim=1, expected_dims=3) + for i in range(40) + }, + **{ + f"blocks.{i}.ffn": ContextParallelOutput(gather_dim=1, expected_dims=3) + for i in range(40) + }, + } + + @staticmethod + def _local_checkpoint_keys(checkpoint_dir: str) -> set[str] | None: + index_names = ( + "diffusion_pytorch_model.safetensors.index.json", + "model.safetensors.index.json", + "diffusion_pytorch_model.bin.index.json", + "pytorch_model.bin.index.json", + ) + for index_name in index_names: + index_path = os.path.join(checkpoint_dir, index_name) + if os.path.exists(index_path): + with open(index_path, "r", encoding="utf-8") as f: + index = json.load(f) + weight_map = index.get("weight_map", {}) + return set(weight_map.keys()) + + for weights_name in ( + "diffusion_pytorch_model.safetensors", + "model.safetensors", + ): + weights_path = os.path.join(checkpoint_dir, weights_name) + if os.path.exists(weights_path): + from safetensors import safe_open + + with safe_open(weights_path, framework="pt", device="cpu") as f: + return set(f.keys()) + + return None + + @staticmethod + def _resolve_local_checkpoint_dir( + pretrained_model_name_or_path: str, subfolder: str | None + ) -> str | None: + if not isinstance(pretrained_model_name_or_path, str): + return None + if not os.path.isdir(pretrained_model_name_or_path): + return None + checkpoint_dir = pretrained_model_name_or_path + if subfolder is not None: + checkpoint_dir = os.path.join(checkpoint_dir, subfolder) + if not os.path.isdir(checkpoint_dir): + return None + return checkpoint_dir + + @classmethod + def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): + subfolder = kwargs.get("subfolder") + checkpoint_dir = cls._resolve_local_checkpoint_dir( + pretrained_model_name_or_path, subfolder + ) + checkpoint_keys = ( + cls._local_checkpoint_keys(checkpoint_dir) + if checkpoint_dir is not None + else None + ) + init_memory_from_short = ( + checkpoint_keys is not None + and "patch_memory.weight" not in checkpoint_keys + and "patch_short.weight" in checkpoint_keys + ) + + loaded = super().from_pretrained( + pretrained_model_name_or_path, *model_args, **kwargs + ) + model = loaded[0] if isinstance(loaded, tuple) else loaded + + if ( + init_memory_from_short + and hasattr(model, "patch_memory") + and hasattr(model, "patch_short") + ): + with torch.no_grad(): + model.patch_memory.weight.copy_(model.patch_short.weight) + if ( + model.patch_memory.bias is not None + and model.patch_short.bias is not None + ): + model.patch_memory.bias.copy_(model.patch_short.bias) + logger.info("Initialized patch_memory weights from patch_short.") + + return loaded + + @register_to_config + def __init__( + self, + patch_size: tuple[int, ...] = (1, 2, 2), + num_attention_heads: int = 40, + attention_head_dim: int = 128, + in_channels: int = 16, + out_channels: int = 16, + text_dim: int = 4096, + freq_dim: int = 256, + ffn_dim: int = 13824, + num_layers: int = 40, + cross_attn_norm: bool = True, + qk_norm: str | None = "rms_norm_across_heads", + eps: float = 1e-6, + added_kv_proj_dim: int | None = None, + rope_dim: tuple[int, ...] = (44, 42, 42), + rope_theta: float = 10000.0, + guidance_cross_attn: bool = True, + zero_history_timestep: bool = True, + has_multi_term_memory_patch: bool = True, + is_amplify_history: bool = False, + history_scale_mode: str = "per_head", # [scalar, per_head] + ) -> None: + super().__init__() + + inner_dim = num_attention_heads * attention_head_dim + out_channels = out_channels or in_channels + + # 1. Patch & position embedding + self.rope = WorldCrafterRotaryPosEmbed(rope_dim=rope_dim, theta=rope_theta) + self.patch_embedding = nn.Conv3d( + in_channels, inner_dim, kernel_size=patch_size, stride=patch_size + ) + + # 2. Initial Multi Term Memory Patch + self.zero_history_timestep = zero_history_timestep + self.inner_dim = inner_dim + if has_multi_term_memory_patch: + self.patch_short = nn.Conv3d( + in_channels, self.inner_dim, kernel_size=patch_size, stride=patch_size + ) + self.patch_mid = nn.Conv3d( + in_channels, + self.inner_dim, + kernel_size=tuple(2 * p for p in patch_size), + stride=tuple(2 * p for p in patch_size), + ) + self.patch_memory = nn.Conv3d( + in_channels, self.inner_dim, kernel_size=patch_size, stride=patch_size + ) + + # 3. Condition embeddings + self.condition_embedder = WorldCrafterTimeTextEmbedding( + dim=inner_dim, + time_freq_dim=freq_dim, + time_proj_dim=inner_dim * 6, + text_embed_dim=text_dim, + ) + + # 4. Transformer blocks + self.blocks = nn.ModuleList( + [ + WorldCrafterTransformerBlock( + inner_dim, + ffn_dim, + num_attention_heads, + qk_norm, + cross_attn_norm, + eps, + added_kv_proj_dim, + guidance_cross_attn=guidance_cross_attn, + is_amplify_history=is_amplify_history, + history_scale_mode=history_scale_mode, + ) + for _ in range(num_layers) + ] + ) + + # 5. Output norm & projection + self.norm_out = WorldCrafterOutputNorm(inner_dim, eps, elementwise_affine=False) + self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size)) + + self.gradient_checkpointing = False + + @apply_lora_scale("attention_kwargs") + def forward( + self, + hidden_states: torch.Tensor, + timestep: torch.LongTensor, + encoder_hidden_states: torch.Tensor, + # ------------ Stage 1 ------------ + indices_hidden_states=None, + indices_latents_history_short=None, + indices_latents_history_mid=None, + indices_latents_memory=None, + latents_history_short=None, + latents_history_mid=None, + latents_memory=None, + # Backward-compatible aliases for existing stage1 callers. + indices_latents_history_long=None, + latents_history_long=None, + return_dict: bool = True, + attention_kwargs: dict[str, Any] | None = None, + ) -> torch.Tensor | dict[str, torch.Tensor]: + if indices_latents_memory is None and indices_latents_history_long is not None: + indices_latents_memory = indices_latents_history_long + if latents_memory is None and latents_history_long is not None: + latents_memory = latents_history_long + + camera_control_ucpe_input = None + if attention_kwargs is not None: + camera_control_ucpe_input = attention_kwargs.get( + "camera_control_ucpe_input", None + ) + + # 1. Input + batch_size = hidden_states.shape[0] + p_t, p_h, p_w = self.config.patch_size + + # 2. Process noisy latents + hidden_states = self.patch_embedding(hidden_states) + _, _, post_patch_num_frames, post_patch_height, post_patch_width = ( + hidden_states.shape + ) + + if indices_hidden_states is None: + indices_hidden_states = ( + torch.arange(0, post_patch_num_frames) + .unsqueeze(0) + .expand(batch_size, -1) + ) + + hidden_states = hidden_states.flatten(2).transpose(1, 2) + rotary_emb = self.rope( + frame_indices=indices_hidden_states, + height=post_patch_height, + width=post_patch_width, + device=hidden_states.device, + ) + rotary_emb = rotary_emb.flatten(2).transpose(1, 2) + original_context_length = hidden_states.shape[1] + + # 3. Process short history latents + if ( + latents_history_short is not None + and indices_latents_history_short is not None + ): + latents_history_short = latents_history_short.to(hidden_states) + latents_history_short = self.patch_short(latents_history_short) + _, _, _, H1, W1 = latents_history_short.shape + latents_history_short = latents_history_short.flatten(2).transpose(1, 2) + + rotary_emb_history_short = self.rope( + frame_indices=indices_latents_history_short, + height=H1, + width=W1, + device=latents_history_short.device, + ) + rotary_emb_history_short = rotary_emb_history_short.flatten(2).transpose( + 1, 2 + ) + + hidden_states = torch.cat([latents_history_short, hidden_states], dim=1) + rotary_emb = torch.cat([rotary_emb_history_short, rotary_emb], dim=1) + + # 4. Process mid history latents + if latents_history_mid is not None and indices_latents_history_mid is not None: + latents_history_mid = latents_history_mid.to(hidden_states) + latents_history_mid = pad_for_3d_conv(latents_history_mid, (2, 4, 4)) + latents_history_mid = self.patch_mid(latents_history_mid) + latents_history_mid = latents_history_mid.flatten(2).transpose(1, 2) + + rotary_emb_history_mid = self.rope( + frame_indices=indices_latents_history_mid, + height=H1, + width=W1, + device=latents_history_mid.device, + ) + rotary_emb_history_mid = pad_for_3d_conv(rotary_emb_history_mid, (2, 2, 2)) + rotary_emb_history_mid = center_down_sample_3d( + rotary_emb_history_mid, (2, 2, 2) + ) + rotary_emb_history_mid = rotary_emb_history_mid.flatten(2).transpose(1, 2) + + hidden_states = torch.cat([latents_history_mid, hidden_states], dim=1) + rotary_emb = torch.cat([rotary_emb_history_mid, rotary_emb], dim=1) + + # 5. Process memory latents. These frames are not temporally/spatially downsampled. + if latents_memory is not None and indices_latents_memory is not None: + latents_memory = latents_memory.to(hidden_states) + latents_memory = self.patch_memory(latents_memory) + _, _, _, HM, WM = latents_memory.shape + latents_memory = latents_memory.flatten(2).transpose(1, 2) + + rotary_emb_memory = self.rope( + frame_indices=indices_latents_memory, + height=HM, + width=WM, + device=latents_memory.device, + ) + rotary_emb_memory = rotary_emb_memory.flatten(2).transpose(1, 2) + + hidden_states = torch.cat([latents_memory, hidden_states], dim=1) + rotary_emb = torch.cat([rotary_emb_memory, rotary_emb], dim=1) + + history_context_length = hidden_states.shape[1] - original_context_length + + if indices_hidden_states is not None and self.zero_history_timestep: + timestep_t0 = torch.zeros((1), dtype=timestep.dtype, device=timestep.device) + temb_t0, timestep_proj_t0, _ = self.condition_embedder( + timestep_t0, + encoder_hidden_states, + is_return_encoder_hidden_states=False, + ) + temb_t0 = temb_t0.unsqueeze(1).expand( + batch_size, history_context_length, -1 + ) + timestep_proj_t0 = ( + timestep_proj_t0.unflatten(-1, (6, -1)) + .view(1, 6, 1, -1) + .expand(batch_size, -1, history_context_length, -1) + ) + + temb, timestep_proj, encoder_hidden_states = self.condition_embedder( + timestep, encoder_hidden_states + ) + timestep_proj = timestep_proj.unflatten(-1, (6, -1)) + + if indices_hidden_states is not None and not self.zero_history_timestep: + main_repeat_size = hidden_states.shape[1] + else: + main_repeat_size = original_context_length + temb = temb.view(batch_size, 1, -1).expand(batch_size, main_repeat_size, -1) + timestep_proj = timestep_proj.view(batch_size, 6, 1, -1).expand( + batch_size, 6, main_repeat_size, -1 + ) + + if indices_hidden_states is not None and self.zero_history_timestep: + temb = torch.cat([temb_t0, temb], dim=1) + timestep_proj = torch.cat([timestep_proj_t0, timestep_proj], dim=2) + + if timestep_proj.ndim == 4: + timestep_proj = timestep_proj.permute(0, 2, 1, 3) + + # 6. Transformer blocks + hidden_states = hidden_states.contiguous() + encoder_hidden_states = encoder_hidden_states.contiguous() + rotary_emb = rotary_emb.contiguous() + if torch.is_grad_enabled() and self.gradient_checkpointing: + for block in self.blocks: + hidden_states = self._gradient_checkpointing_func( + block, + hidden_states, + encoder_hidden_states, + timestep_proj, + rotary_emb, + original_context_length, + camera_control_ucpe_input, + ) + else: + for block in self.blocks: + hidden_states = block( + hidden_states, + encoder_hidden_states, + timestep_proj, + rotary_emb, + original_context_length, + camera_control_ucpe_input, + ) + + # 7. Normalization + hidden_states = self.norm_out(hidden_states, temb, original_context_length) + hidden_states = self.proj_out(hidden_states) + + # 8. Unpatchify + hidden_states = hidden_states.reshape( + batch_size, + post_patch_num_frames, + post_patch_height, + post_patch_width, + p_t, + p_h, + p_w, + -1, + ) + hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) + output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) + + if not return_dict: + return (output,) + + return Transformer2DModelOutput(sample=output) diff --git a/worldcrafter/fast/__init__.py b/worldcrafter/fast/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a2e1429496109612a44e324217300ba65fecd548 --- /dev/null +++ b/worldcrafter/fast/__init__.py @@ -0,0 +1 @@ +"""Exact six-step inference using two independently adapted transformers.""" diff --git a/worldcrafter/fast/attention.py b/worldcrafter/fast/attention.py new file mode 100644 index 0000000000000000000000000000000000000000..d1eb991f9826b625b0200f56b9d969f8c2300501 --- /dev/null +++ b/worldcrafter/fast/attention.py @@ -0,0 +1,57 @@ +"""Pyramid UCPE with the reference FP32 attention and residual semantics. + +Storage dtype is independent from compute dtype. In particular, compact UCPE +must not downcast camera embeddings or Q/K/V to its BF16 storage dtype. +""" + +import torch +from einops import rearrange, repeat +from ..ucpe.camera import UcpeSelfAttention +from ..ucpe.bridge import _flash_attention_sdpa as flash_attention + + +class FastUcpeSelfAttention(UcpeSelfAttention): + def forward(self, x: torch.Tensor, control_camera_dit_input: dict): + """``x`` is ``[B, T, D]`` over the current-chunk tokens only.""" + B, T, D = x.shape + num_cameras = control_camera_dit_input["viewmats"].shape[1] + grid_h = control_camera_dit_input.get("patches_y", self.patches_y) + grid_w = control_camera_dit_input.get("patches_x", self.patches_x) + expected = num_cameras * grid_h * grid_w + assert T == expected or T == num_cameras, ( + f"Expected token count {expected} ({num_cameras}x{grid_h}x{grid_w}) or {num_cameras}, got {T}" + ) + + if hasattr(self, "cam_encoder") and "cam_emb" in control_camera_dit_input: + cam_emb = control_camera_dit_input["cam_emb"] + y = self.cam_encoder(cam_emb) + if y.shape[1] != T: + hw = T // cam_emb.shape[1] + y = repeat(y, "b f d -> b (f hw) d", hw=hw) + x = x + y + + q = self.q_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) + k = self.k_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) + v = self.v_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2) + + self.prope_attn._precompute_and_cache_apply_fns( + viewmats=control_camera_dit_input["viewmats"], + Ks=control_camera_dit_input.get("K", None), + coeffs_x=control_camera_dit_input.get("coeffs_x", None), + coeffs_y=control_camera_dit_input.get("coeffs_y", None), + ) + + q = self.prope_attn._apply_to_q(q) + k = self.prope_attn._apply_to_kv(k) + v = self.prope_attn._apply_to_kv(v) + + q = rearrange(q, "b h t d -> b t (h d)") + k = rearrange(k, "b h t d -> b t (h d)") + v = rearrange(v, "b h t d -> b t (h d)") + + out = flash_attention(q, k, v, num_heads=self.num_heads) + + out = rearrange(out, "b t (h d) -> b h t d", h=self.num_heads) + out = self.prope_attn._apply_to_o(out) + out = out.transpose(1, 2).reshape(B, T, -1) + return self.out_proj(out) diff --git a/worldcrafter/fast/camera.py b/worldcrafter/fast/camera.py new file mode 100644 index 0000000000000000000000000000000000000000..21beca42f9d94ae0b72fdb08785621906dafb6b5 --- /dev/null +++ b/worldcrafter/fast/camera.py @@ -0,0 +1,368 @@ +"""Multiscale UCPE geometry; local control poses and global memory poses stay separate.""" + +from __future__ import annotations +from typing import Any +import torch +from einops import rearrange, repeat +from . import geometry as ucpe_cc +from ..ucpe import prope as prope_torch + + +def rope_axis_positions( + num_patches: int, + reference_num_patches: int | None = None, + scale_invariant: bool = True, + device: torch.device | str = "cpu", +) -> torch.Tensor: + """PRoPE positions along one axis of a token grid. + + With ``scale_invariant`` set, positions are expressed on the reference grid + so a physical location keeps the same angle at every pyramid level: + + pos(m) = (m + 0.5) * (W_ref / W_s) - 0.5 + + Position ``m`` then sits at the centre of the reference columns it covers. + At ``W_s == W_ref``, the expression reduces to ``arange(W)``. + """ + base = torch.arange(num_patches, device=device, dtype=torch.float32) + if not scale_invariant: + return base + scale = (reference_num_patches or num_patches) / num_patches + return (base + 0.5) * scale - 0.5 + + +def _build_stage_rope_coeffs( + cameras: int, + patches_y: int, + patches_x: int, + head_dim: int, + freq_base: float, + freq_scale: float, + device: torch.device, + dtype: torch.dtype, + reference_patches_y: int | None = None, + reference_patches_x: int | None = None, + scale_invariant_positions: bool = True, +): + """PRoPE cos/sin coefficients for one token grid.""" + base_x = rope_axis_positions( + patches_x, reference_patches_x, scale_invariant_positions, device + ) + base_y = rope_axis_positions( + patches_y, reference_patches_y, scale_invariant_positions, device + ) + + x_positions = torch.tile(base_x, (patches_y * cameras,)) + y_positions = torch.tile(torch.repeat_interleave(base_y, patches_x), (cameras,)) + + coeffs_x = prope_torch._rope_precompute_coeffs( + x_positions, + freq_base=freq_base, + freq_scale=freq_scale, + feat_dim=head_dim // 4, + dtype=dtype, + ) + coeffs_y = prope_torch._rope_precompute_coeffs( + y_positions, + freq_base=freq_base, + freq_scale=freq_scale, + feat_dim=head_dim // 4, + dtype=dtype, + ) + return coeffs_x, coeffs_y + + +def _slice_pose_chunk( + pose: torch.Tensor, + chunk_index: int, + window_num_frames: int, + restart_each_chunk: bool = True, +): + if restart_each_chunk: + start, end = 0, window_num_frames + else: + start = chunk_index * window_num_frames + end = start + window_num_frames + if pose.shape[1] < end: + return None + return pose[:, start:end] + + +def resolve_pose_chunk( + camera_trajectory: dict[str, Any], + num_latent_frames_per_chunk: int, + chunk_index: int, + vae_scale_factor_temporal: int = 4, + restart_each_chunk: bool = True, + translation_scale: float = 1.0, + pose_is_chunk_aligned: bool = False, +): + """Extract the ``[B, T, 3, 4]`` camera-to-world poses for one chunk. + + Returns ``(pose_chunk, x_fov, xi)`` or ``None`` when the trajectory is too + short for the requested chunk. + """ + pose = camera_trajectory["pose"] + x_fov = camera_trajectory["x_fov"] + xi = camera_trajectory["xi"] + + if pose.ndim == 3: + pose = pose.unsqueeze(0) + if torch.is_tensor(x_fov) and x_fov.ndim == 0: + x_fov = x_fov.unsqueeze(0) + if torch.is_tensor(xi) and xi.ndim == 0: + xi = xi.unsqueeze(0) + + if pose_is_chunk_aligned: + if pose.shape[1] < num_latent_frames_per_chunk: + return None + pose_chunk = pose[:, :num_latent_frames_per_chunk] + else: + window_num_frames = ( + num_latent_frames_per_chunk - 1 + ) * vae_scale_factor_temporal + 1 + pose_chunk = _slice_pose_chunk( + pose, chunk_index, window_num_frames, restart_each_chunk=restart_each_chunk + ) + if pose_chunk is None: + return None + pose_chunk = pose_chunk[:, ::vae_scale_factor_temporal] + + if pose_chunk.shape[-2:] == (4, 4): + pose_chunk = pose_chunk[..., :3, :4] + elif pose_chunk.shape[-2:] != (3, 4): + raise ValueError( + f"pose_chunk is expected to be [B, T, 3, 4] or [B, T, 4, 4], got shape={tuple(pose_chunk.shape)}" + ) + + if translation_scale != 1.0: + pose_chunk = pose_chunk.clone() + pose_chunk[..., 3] *= pose_chunk.new_tensor(float(translation_scale)) + + return pose_chunk, x_fov, xi + + +def build_ucpe_camera_input( + transformer, + pose_chunk: torch.Tensor, + x_fov: torch.Tensor, + xi: torch.Tensor, + grid_h: int, + grid_w: int, + reference_grid_h: int | None = None, + reference_grid_w: int | None = None, + pixel_center: bool = False, + scale_invariant_positions: bool = True, + force_explicit_coeffs: bool = False, +): + """Build one ``control_camera_dit_input`` dict for a single token grid.""" + method = transformer.camera_condition + if "gta" in method or "prope" in method: + raise NotImplementedError( + "The UCPE bridge implements relray_absmap only." + ) + if "relray" not in method: + raise ValueError(f"Unsupported camera condition: {method}") + + attn = transformer.blocks[0].cam_self_attn + reference_grid_h = reference_grid_h or attn.patches_y + reference_grid_w = reference_grid_w or attn.patches_x + + if grid_h * reference_grid_w != grid_w * reference_grid_h: + raise ValueError( + "UCPE pyramid levels must share an aspect ratio, otherwise fy=fx with cy=H/2 makes the " + f"vertical FOV differ per level. Got level grid {grid_h}x{grid_w} against reference " + f"{reference_grid_h}x{reference_grid_w}." + ) + + c2w = torch.eye(4, device=pose_chunk.device, dtype=pose_chunk.dtype) + c2w = repeat( + c2w, "... -> B T ...", B=pose_chunk.shape[0], T=pose_chunk.shape[1] + ).clone() + c2w[..., :3, :4] = pose_chunk + + d_cam = ucpe_cc.ucm_unproject_grid_fov( + x_fov=x_fov, + xi=xi, + height=grid_h, + width=grid_w, + device=pose_chunk.device, + dtype=pose_chunk.dtype, + pixel_center=pixel_center, + ) + raymats = ucpe_cc.world_to_ray_mats(d_cam, c2w) + viewmats = rearrange(raymats, "B T H W ... -> B (T H W) ...") + + camera_input: dict[str, Any] = { + "viewmats": viewmats, + "patches_y": grid_h, + "patches_x": grid_w, + } + + needs_coeffs = ( + force_explicit_coeffs or grid_h != attn.patches_y or grid_w != attn.patches_x + ) + if needs_coeffs: + coeffs_x, coeffs_y = _build_stage_rope_coeffs( + cameras=pose_chunk.shape[1], + patches_y=grid_h, + patches_x=grid_w, + head_dim=attn.head_dim, + freq_base=attn.freq_base, + freq_scale=attn.freq_scale, + device=pose_chunk.device, + dtype=pose_chunk.dtype, + reference_patches_y=reference_grid_h, + reference_patches_x=reference_grid_w, + scale_invariant_positions=scale_invariant_positions, + ) + camera_input["coeffs_x"] = coeffs_x + camera_input["coeffs_y"] = coeffs_y + + if "absmap" in method: + up_map, lat_map = ucpe_cc.compute_up_lat_map( + R=pose_chunk[..., :3, :3], + x_fov=x_fov, + xi=xi, + width=grid_w, + height=grid_h, + device=pose_chunk.device, + pixel_center=pixel_center, + ) + cam_emb = torch.cat([up_map, lat_map], dim=-1) + camera_input["cam_emb"] = rearrange(cam_emb, "B T H W C -> B (T H W) C") + + expected_token_count = pose_chunk.shape[1] * grid_h * grid_w + if camera_input["viewmats"].shape[1] != expected_token_count: + raise ValueError( + f"camera viewmats token count mismatch: expected {expected_token_count}, " + f"got {camera_input['viewmats'].shape[1]}" + ) + if ( + "cam_emb" in camera_input + and camera_input["cam_emb"].shape[1] != expected_token_count + ): + raise ValueError( + f"camera cam_emb token count mismatch: expected {expected_token_count}, " + f"got {camera_input['cam_emb'].shape[1]}" + ) + + return camera_input + + +def pyramid_token_grids( + reference_grid_h: int, + reference_grid_w: int, + num_stages: int, + low_to_high: bool = False, +) -> list[tuple[int, int]]: + """Token grids for each pyramid level, finest first by default. + + Stage-2 halves the spatial latent size per level, which halves the post-patch + token grid too. Levels that would not halve evenly are rejected, since an + inconsistent aspect ratio changes the vertical FOV. + + Set ``low_to_high`` for the coarsest-first order used by fast sampling. + """ + grids = [(reference_grid_h, reference_grid_w)] + for level in range(1, num_stages): + prev_h, prev_w = grids[-1] + if prev_h % 2 or prev_w % 2: + raise ValueError( + f"Cannot build {num_stages} pyramid levels from a {reference_grid_h}x{reference_grid_w} token " + f"grid: level {level} would halve {prev_h}x{prev_w} unevenly, which changes the aspect ratio " + "and therefore the vertical FOV." + ) + grids.append((prev_h // 2, prev_w // 2)) + return list(reversed(grids)) if low_to_high else grids + + +def build_ucpe_attention_kwargs_pyramid( + transformer, + camera_trajectory: dict[str, Any] | None, + num_latent_frames_per_chunk: int, + chunk_index: int, + token_grids: list[tuple[int, int]], + vae_scale_factor_temporal: int = 4, + restart_each_chunk: bool = True, + translation_scale: float = 1.0, + pose_is_chunk_aligned: bool = False, + pixel_center: bool = True, + scale_invariant_positions: bool = True, +): + """Build per-level UCPE inputs in coarsest-first order. + + The returned list follows ``token_grids``. Every level uses the finest grid + as its common geometric reference. Sampling at pixel centers keeps the + frustum symmetric, including at the coarsest resolution. + """ + if ( + camera_trajectory is None + or getattr(transformer, "camera_condition", "none") == "none" + ): + return None + + resolved = resolve_pose_chunk( + camera_trajectory, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + chunk_index=chunk_index, + vae_scale_factor_temporal=vae_scale_factor_temporal, + restart_each_chunk=restart_each_chunk, + translation_scale=translation_scale, + pose_is_chunk_aligned=pose_is_chunk_aligned, + ) + if resolved is None: + return None + pose_chunk, x_fov, xi = resolved + + reference_grid_h, reference_grid_w = max( + token_grids, key=lambda grid: grid[0] * grid[1] + ) + camera_inputs = [ + build_ucpe_camera_input( + transformer, + pose_chunk=pose_chunk, + x_fov=x_fov, + xi=xi, + grid_h=grid_h, + grid_w=grid_w, + reference_grid_h=reference_grid_h, + reference_grid_w=reference_grid_w, + pixel_center=pixel_center, + scale_invariant_positions=scale_invariant_positions, + force_explicit_coeffs=True, + ) + for grid_h, grid_w in token_grids + ] + + return {"camera_control_ucpe_input_list": camera_inputs} + + +def build_ucpe_attention_kwargs_sequential_pyramid( + transformer, + camera_trajectory: dict[str, Any] | None, + num_latent_frames_per_chunk: int, + chunk_index: int, + token_grids: list[tuple[int, int]], + **kwargs, +): + """One single-camera kwargs dict per pyramid level, coarsest first. + + Fast sampling denoises levels sequentially. Each forward receives the + camera input for its resolution, with the finest grid as the shared + geometric reference. + """ + pyramid_kwargs = build_ucpe_attention_kwargs_pyramid( + transformer, + camera_trajectory, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + chunk_index=chunk_index, + token_grids=token_grids, + **kwargs, + ) + if pyramid_kwargs is None: + return None + return [ + {"camera_control_ucpe_input": camera_input} + for camera_input in pyramid_kwargs["camera_control_ucpe_input_list"] + ] diff --git a/worldcrafter/fast/compact_ucpe.py b/worldcrafter/fast/compact_ucpe.py new file mode 100644 index 0000000000000000000000000000000000000000..d00f7e94ab463a8b199da68cf42c5d7eba2a220e --- /dev/null +++ b/worldcrafter/fast/compact_ucpe.py @@ -0,0 +1,50 @@ +"""Losslessly store BF16-origin UCPE weights; preserve FP32 linear operations.""" + +import types +import torch +from torch import nn +from torch.nn import functional as F + + +def _forward(self, x): + # Keep the input and operator unchanged. No BF16 GEMM substitution. + return F.linear( + x, self.weight.float(), None if self.bias is None else self.bias.float() + ) + + +def compact_ucpe(model): + modules = [] + parameters = [] + for block in model.blocks: + camera = block.cam_self_attn + supported = set() + for layer in camera.modules(): + if isinstance(layer, nn.Linear): + modules.append(layer) + for p in layer.parameters(recurse=False): + if p.dtype != torch.float32: + raise ValueError( + f"Expected FP32 UCPE parameters, got {p.dtype}" + ) + if not torch.equal(p, p.bfloat16().float()): + raise ValueError( + "UCPE storage conversion would lose information" + ) + parameters.append(p) + supported.add(id(p)) + if {id(p) for p in camera.parameters()} != supported: + raise ValueError("Unsupported UCPE parameter type") + saved = sum(p.numel() * 2 for p in parameters) + for p in parameters: + p.data = p.data.bfloat16() + for layer in modules: + layer.forward = types.MethodType(_forward, layer) + return dict( + modules=len(modules), + tensors=len(parameters), + bytes_saved=saved, + storage_dtype="bfloat16", + linear_weight_compute_dtype="float32", + all_weights_exact=True, + ) diff --git a/worldcrafter/fast/contract.py b/worldcrafter/fast/contract.py new file mode 100644 index 0000000000000000000000000000000000000000..a97b93ea0852143c149db2a2932cd001489a82f0 --- /dev/null +++ b/worldcrafter/fast/contract.py @@ -0,0 +1,144 @@ +"""Inference-side enforcement for the checkpoint DMD timestep contract. + +The denoising tables themselves live in :mod:`worldcrafter.fast.timestep_grid`. This +module deliberately contains no schedule math: it only locates a checkpoint's +sidecar, validates the latent shape, and selects the normal or empty-history +table from the conditioning tensors that will actually enter the transformer. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Iterable, Sequence + +import torch + +from .timestep_grid import ( + DmdTimestepContract, + DmdTimestepStep, + resolve_dmd_contract_path, +) + + +@dataclass(frozen=True) +class DmdInferenceTrace: + """The immutable contract trace selected for one denoised chunk.""" + + fingerprint: str + empty_history: bool + stages: tuple[tuple[DmdTimestepStep, ...], ...] + + @property + def num_steps(self) -> int: + return sum(len(stage) for stage in self.stages) + + +def load_dmd_inference_contract( + checkpoint_path: str | Path, + *, + expected_latent_shape: Sequence[int] | None = None, + expected_fingerprint: str | None = None, +) -> DmdTimestepContract: + """Load and authenticate the sidecar next to ``checkpoint_path``. + + A DMD checkpoint without its sidecar is intentionally unusable. Rebuilding + schedule from inference flags could silently change the checkpoint + timesteps and generated output. + """ + + sidecar = resolve_dmd_contract_path(checkpoint_path) + return DmdTimestepContract.load_json( + sidecar, + expected_latent_shape=expected_latent_shape, + expected_fingerprint=expected_fingerprint, + ) + + +def _history_nonempty_mask( + history_tensors: Iterable[torch.Tensor | None], +) -> torch.Tensor: + mask = None + batch_size = None + for history in history_tensors: + if history is None: + continue + if history.ndim < 1: + raise ValueError( + f"DMD history tensor must have a batch dimension, got shape={tuple(history.shape)}" + ) + if batch_size is None: + batch_size = int(history.shape[0]) + mask = torch.zeros(batch_size, dtype=torch.bool, device=history.device) + elif int(history.shape[0]) != batch_size: + raise ValueError( + "DMD history tensors disagree on batch size: " + f"expected {batch_size}, got {int(history.shape[0])}" + ) + current = history.detach().reshape(batch_size, -1).ne(0).any(dim=1) + mask = mask | current.to(device=mask.device) + + if mask is None: + # A missing conditioning bank is the same semantic condition as the + # all-zero placeholders used by the pipelines for the first T2V chunk. + return torch.zeros(1, dtype=torch.bool) + return mask + + +def history_is_empty(history_tensors: Iterable[torch.Tensor | None]) -> bool: + """Classify the real conditioning bank, rejecting a mixed batch. + + Empty/non-empty examples require different step counts, so a batch cannot + share one denoising loop when only some rows have history. + """ + + nonempty = _history_nonempty_mask(history_tensors) + has_nonempty = bool(nonempty.any().item()) + has_empty = bool((~nonempty).any().item()) + if has_nonempty and has_empty: + raise ValueError( + "DMD inference cannot mix empty-history and non-empty-history samples in one batch; " + "split the batch so every sample follows one checkpoint schedule." + ) + return not has_nonempty + + +def resolve_dmd_inference_trace( + contract: DmdTimestepContract | None, + *, + latent_shape: Sequence[int], + history_tensors: Iterable[torch.Tensor | None], + num_stages: int, +) -> DmdInferenceTrace: + """Validate and select the exact full rollout trace for a chunk.""" + + if contract is None: + raise RuntimeError( + "DMD inference requires dmd_timestep_contract.json from the checkpoint; " + "no contract was loaded." + ) + + actual_shape = tuple(int(value) for value in latent_shape) + expected_shape = tuple(int(value) for value in contract.latent_shape) + if actual_shape != expected_shape: + raise ValueError( + "DMD latent shape does not match the checkpoint contract: " + f"checkpoint={expected_shape}, inference={actual_shape}" + ) + + if len(contract.normal_stages) != int(num_stages): + raise ValueError( + "DMD pyramid stage count does not match the checkpoint contract: " + f"checkpoint={len(contract.normal_stages)}, inference={int(num_stages)}" + ) + + empty_history = history_is_empty(history_tensors) + stages = tuple( + tuple(contract.stage(stage_index, empty_history=empty_history)) + for stage_index in range(int(num_stages)) + ) + return DmdInferenceTrace( + fingerprint=contract.fingerprint, + empty_history=empty_history, + stages=stages, + ) diff --git a/worldcrafter/fast/geometry.py b/worldcrafter/fast/geometry.py new file mode 100644 index 0000000000000000000000000000000000000000..bcc4428492be550dfec550cfdc1b5cdf751c033b --- /dev/null +++ b/worldcrafter/fast/geometry.py @@ -0,0 +1,359 @@ +from __future__ import annotations +from typing import Tuple +import torch +import torch.nn.functional as F +from einops import einsum, rearrange, repeat + + +def _to_tensor_1d(value, device, dtype) -> torch.Tensor: + if torch.is_tensor(value): + return value.to(device=device, dtype=dtype).reshape(-1) + return torch.tensor([value], dtype=dtype, device=device) + + +def pixel_coordinates( + height: int, + width: int, + pixel_center: bool = False, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.float32, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Sampling coordinates for a ``height x width`` grid. + + With ``pixel_center=False`` the samples are the integers ``0 .. W-1``, which + is what ``equilib.create_grid`` produces. Combined with ``cx = W / 2`` the + covered extent is ``[-0.5, 0.5 - 1/W]`` in normalized units: asymmetric by + one pixel, and the asymmetry is ``2/W`` relative to the half width. That is + invisible at high resolution but reaches 20% of the half width on a 10-column + grid, so a pyramid built this way samples a progressively offset frustum. + + With ``pixel_center=True`` the samples are ``0.5 .. W-0.5``, giving the + symmetric extent ``[-0.5 + 0.5/W, 0.5 - 0.5/W]``. Each coarse sample then + lands exactly on the centroid of the corresponding block of finer samples, + which makes the geometry consistent across pyramid levels. + """ + offset = 0.5 if pixel_center else 0.0 + xs = torch.arange(width, device=device, dtype=dtype) + offset + ys = torch.arange(height, device=device, dtype=dtype) + offset + grid_y, grid_x = torch.meshgrid(ys, xs, indexing="ij") + return grid_x, grid_y + + +def compute_fx_from_fov_xi( + x_fov: torch.Tensor | float, + xi: torch.Tensor | float, + width: int, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """Focal length in pixels from horizontal FOV (degrees) and the UCM mirror parameter. + + ``fx`` is proportional to ``width``, which is what keeps the covered field of + view constant when the same camera is sampled on a coarser grid. + """ + x_fov = _to_tensor_1d(x_fov, device, dtype) + xi = _to_tensor_1d(xi, device, dtype) + + batch = max(x_fov.shape[0], xi.shape[0]) + x_fov = x_fov.view(-1).expand(batch) + xi = xi.view(-1).expand(batch) + + theta = torch.deg2rad(0.5 * x_fov) + eps = torch.finfo(dtype).eps + denom = torch.sin(theta).clamp_min(eps) + return (width * 0.5) * (torch.cos(theta) + xi) / denom + + +def compute_fov_from_fx_xi( + fx: torch.Tensor | float, + xi: torch.Tensor | float, + width: int, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """Inverse of :func:`compute_fx_from_fov_xi`, returning degrees.""" + fx = _to_tensor_1d(fx, device, dtype) + xi = _to_tensor_1d(xi, device, dtype) + batch = max(fx.shape[0], xi.shape[0]) + fx = fx.expand(batch) + xi = xi.expand(batch) + + a = 2.0 * fx / width + phi = torch.atan(1.0 / a) + denom = torch.sqrt(a * a + 1.0) + ratio = (xi / denom).clamp(-1.0, 1.0) + theta = torch.asin(ratio) + phi + return torch.rad2deg(2.0 * theta) + + +def ucm_unproject_grid( + height: int, + width: int, + fx: torch.Tensor, + fy: torch.Tensor, + cx: float | torch.Tensor, + cy: float | torch.Tensor, + xi: torch.Tensor, + pixel_center: bool = False, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """Unproject a pixel grid to unit-sphere ray directions under the UCM model. + + Returns ``[B, H, W, 3]``. + """ + fx = _to_tensor_1d(fx, device, dtype) + fy = _to_tensor_1d(fy, device, dtype) + cx = _to_tensor_1d(cx, device, dtype) + cy = _to_tensor_1d(cy, device, dtype) + xi = _to_tensor_1d(xi, device, dtype) + batch = max(fx.shape[0], fy.shape[0], cx.shape[0], cy.shape[0], xi.shape[0]) + + grid_x, grid_y = pixel_coordinates( + height, width, pixel_center, device=device, dtype=dtype + ) + grid_x = grid_x.unsqueeze(0).expand(batch, -1, -1) + grid_y = grid_y.unsqueeze(0).expand(batch, -1, -1) + + fx = fx.expand(batch)[:, None, None] + fy = fy.expand(batch)[:, None, None] + cx = cx.expand(batch)[:, None, None] + cy = cy.expand(batch)[:, None, None] + xi = xi.expand(batch)[:, None, None] + + x = (grid_x - cx) / fx + y = (grid_y - cy) / fy + + r2 = x * x + y * y + alpha = xi + torch.sqrt(1 + (1 - xi * xi) * r2) + gamma = alpha / (1 + r2) + + return torch.stack([gamma * x, gamma * y, gamma - xi], dim=-1) + + +def ucm_unproject_grid_fov( + x_fov: float | torch.Tensor, + xi: float | torch.Tensor, + height: int, + width: int, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.float32, + pixel_center: bool = False, +) -> torch.Tensor: + """Ray directions for a ``height x width`` grid covering the given horizontal FOV. + + Intrinsics are always derived from the grid that is passed in, so requesting a + coarser grid resamples the same frustum rather than cropping it. Reusing a + high-resolution ``fx`` with a low-resolution grid would instead shrink the + covered FOV, which is the failure mode this function exists to avoid. + + Returns ``[H, W, 3]`` for scalar parameters, ``[B, H, W, 3]`` otherwise. + """ + is_batched = any( + torch.is_tensor(p) and p.reshape(-1).numel() > 1 for p in (x_fov, xi) + ) + + fx = compute_fx_from_fov_xi(x_fov, xi, width, device, dtype) + d_cam = ucm_unproject_grid( + height=height, + width=width, + fx=fx, + fy=fx, + cx=width / 2, + cy=height / 2, + xi=xi, + pixel_center=pixel_center, + device=device, + dtype=dtype, + ) + return d_cam if is_batched else d_cam[0] + + +def project_ucm_points(X, Y, Z, fx, fy, cx, cy, xi): + """Project camera-frame points onto the UCM image plane.""" + + def broadcast_param(param): + if not torch.is_tensor(param): + return torch.tensor(param, device=X.device, dtype=X.dtype) + param = param.to(device=X.device, dtype=X.dtype) + if param.ndim == 0: + return param + flat = param.reshape(-1) + if flat.numel() == 1: + return flat.view(1) + if X.ndim >= 1 and flat.numel() == X.shape[0]: + return flat.view(flat.shape[0], *([1] * (X.ndim - 1))) + return param + + fx = broadcast_param(fx) + fy = broadcast_param(fy) + cx = broadcast_param(cx) + cy = broadcast_param(cy) + xi = broadcast_param(xi) + + r = torch.sqrt(X * X + Y * Y + Z * Z) + alpha = Z + xi * r + du = fx * (X / alpha) + cx + dv = fy * (Y / alpha) + cy + return du, dv + + +def project_ucm_points_fov(X, Y, Z, x_fov, xi, height, width): + fx = compute_fx_from_fov_xi(x_fov, xi, width, X.device, X.dtype) + return project_ucm_points(X, Y, Z, fx, fx, width / 2, height / 2, xi) + + +def d_cam_to_angles(d_cam: torch.Tensor) -> torch.Tensor: + """Direction vectors to ``[azimuth, elevation]`` in radians.""" + d_unit = F.normalize(d_cam, dim=-1) + x, y, z = d_unit[..., 0], d_unit[..., 1], d_unit[..., 2] + azimuth = torch.atan2(x, z) + elevation = -torch.asin(y.clamp(-1.0, 1.0)) + return torch.stack([azimuth, elevation], dim=-1) + + +def world_to_ray_mats( + d_cam: torch.Tensor, # [B, H, W, 3] + c2w: torch.Tensor, # [B, T, 4, 4] +) -> torch.Tensor: + """Per-ray world-to-ray-local transforms, ``[B, T, H, W, 4, 4]``. + + The ray-local frame is z along the ray, x = cam_y x z, y = z x x. + """ + if d_cam.ndim == 3: + d_cam = d_cam.unsqueeze(0) + if c2w.ndim == 3: + c2w = c2w.unsqueeze(0) + if d_cam.ndim != 4 or d_cam.shape[-1] != 3: + raise ValueError( + f"d_cam must have shape [H,W,3] or [B,H,W,3], got {tuple(d_cam.shape)}" + ) + if c2w.ndim != 4 or c2w.shape[-2:] != (4, 4): + raise ValueError( + f"c2w must have shape [T,4,4] or [B,T,4,4], got {tuple(c2w.shape)}" + ) + if d_cam.shape[0] == 1 and c2w.shape[0] != 1: + d_cam = d_cam.expand(c2w.shape[0], -1, -1, -1) + elif c2w.shape[0] == 1 and d_cam.shape[0] != 1: + c2w = c2w.expand(d_cam.shape[0], -1, -1, -1) + elif d_cam.shape[0] != c2w.shape[0]: + raise ValueError( + f"d_cam and c2w batch mismatch: {d_cam.shape[0]} vs {c2w.shape[0]}" + ) + + B, H, W, _ = d_cam.shape + T = c2w.shape[1] + device = d_cam.device + dtype = d_cam.dtype + + d_cam = repeat(d_cam, "b h w c -> b t h w c", t=T) + R_cam = c2w[..., :3, :3] + t_cam = c2w[..., :3, 3] + + d_world = einsum(R_cam, d_cam, "b t i j, b t h w j -> b t h w i") + + cam_y = R_cam[..., :, 1] + cam_y = repeat(cam_y, "b t c -> b t h w c", h=H, w=W) + + z_ray = F.normalize(d_world, dim=-1, eps=1e-6) + x_ray = F.normalize(torch.cross(cam_y, z_ray, dim=-1), dim=-1, eps=1e-6) + y_ray = F.normalize(torch.cross(z_ray, x_ray, dim=-1), dim=-1, eps=1e-6) + + R_l2w = torch.stack([x_ray, y_ray, z_ray], dim=-1) + R_w2l = rearrange(R_l2w, "b t h w i j -> b t h w j i") + + t_world = repeat(t_cam, "b t c -> b t h w c", h=H, w=W) + t_w2l = -einsum(R_w2l, t_world, "b t h w i j, b t h w j -> b t h w i") + + raymats = torch.zeros(B, T, H, W, 4, 4, device=device, dtype=dtype) + raymats[..., :3, :3] = R_w2l + raymats[..., :3, 3] = t_w2l + raymats[..., 3, 3] = 1.0 + + mask = torch.isnan(d_world).any(-1) + raymats[mask] = torch.eye(4, device=device, dtype=dtype) + + return raymats + + +def compute_up_lat_map( + R: torch.Tensor, # [B, T, 3, 3] + x_fov: torch.Tensor, + xi: torch.Tensor, + height: int, + width: int, + device: torch.device = torch.device("cpu"), + delta: float = 0.1, + pixel_center: bool = False, +): + """World-up direction and latitude maps used by the ``absmap`` conditioning. + + Returns ``(up_map [B,T,H,W,2], lat_map [B,T,H,W,1])``. + """ + B, T, _, _ = R.shape + dtype = R.dtype + R = R.float() + + d_cam = ucm_unproject_grid_fov( + x_fov=x_fov, + xi=xi, + height=height, + width=width, + device=device, + dtype=torch.float32, + pixel_center=pixel_center, + ) + if d_cam.ndim == 3: + d_cam = d_cam.unsqueeze(0) + mask = d_cam.isnan().any(dim=-1, keepdim=True) + + d_cam_exp = repeat(d_cam, "B H W C -> B T H W C", T=T) + d_world = torch.einsum("btij,bthwj->bthwi", R, d_cam_exp) + d_world = d_world / torch.clamp_min(d_world.norm(dim=-1, keepdim=True), 1e-8) + + Xw, Yw, Zw = d_world[..., 0], d_world[..., 1], d_world[..., 2] + lat_map = torch.atan2(-Yw, torch.sqrt(Xw**2 + Zw**2)).unsqueeze(-1) + + v = d_world + up_world = torch.tensor([0, -1, 0], device=device, dtype=torch.float32) + k = torch.cross( + v, up_world.unsqueeze(0).unsqueeze(0).unsqueeze(0).expand_as(v), dim=-1 + ) + k = k / torch.clamp_min(k.norm(dim=-1, keepdim=True), 1e-8) + + delta = torch.tensor(delta, device=device, dtype=torch.float32) + cos_eps = torch.cos(delta) + sin_eps = torch.sin(delta) + v_rot = ( + v * cos_eps + + torch.cross(k, v, dim=-1) * sin_eps + + k * (k * (v * 1).sum(dim=-1, keepdim=True)) * (1 - cos_eps) + ) + + dirs_cam = torch.einsum("btij,bthwj->bthwi", R.transpose(-1, -2), v_rot) + Xs, Ys, Zs = dirs_cam[..., 0], dirs_cam[..., 1], dirs_cam[..., 2] + + du, dv = project_ucm_points_fov( + Xs, + Ys, + Zs, + x_fov=x_fov.float() if torch.is_tensor(x_fov) else x_fov, + xi=xi.float() if torch.is_tensor(xi) else xi, + height=height, + width=width, + ) + + grid_x, grid_y = pixel_coordinates( + height, width, pixel_center, device=device, dtype=torch.float32 + ) + grid_x = grid_x.view(1, 1, height, width) + grid_y = grid_y.view(1, 1, height, width) + + up_map = torch.stack((du - grid_x, dv - grid_y), dim=-1) + up_map = up_map / torch.clamp_min(up_map.norm(dim=-1, keepdim=True), 1e-8) + + up_map = up_map.to(dtype=dtype) + lat_map = lat_map.to(dtype=dtype) + + mask_exp = mask.unsqueeze(1).expand(B, T, height, width, 1) + return up_map.masked_fill(mask_exp, 0.0), lat_map.masked_fill(mask_exp, 0.0) diff --git a/worldcrafter/fast/resident.py b/worldcrafter/fast/resident.py new file mode 100644 index 0000000000000000000000000000000000000000..cca5c86ae533967230f8b85fca4ecfd7aca080e3 --- /dev/null +++ b/worldcrafter/fast/resident.py @@ -0,0 +1,286 @@ +"""Lossless GPU-resident BF16 branch storage. + +A materialized BF16 model + reversible packed integer bit-pattern differences. +LoRA and UCPE tensors remain independent. No CPU copies, decompression rounding, +or new floating-point additions enter a denoiser forward. A stage switch updates +shared storage on the current CUDA stream; concurrent branch forwards are not +supported. Both complete checkpoints remain represented on GPU at all times. +""" + +from dataclasses import dataclass +import gc +import logging +import time +import torch +import torch.nn.functional as F +import triton +import triton.language as tl +from triton.language.extra.cuda import libdevice + +logger = logging.getLogger(__name__) + + +@triton.jit +def _switch_dense( + W, + P, + MASK, + PREFIX, + VALUES, + N: tl.constexpr, + B: tl.constexpr, + BITMAP: tl.constexpr, + SIGN: tl.constexpr, + BLOCK: tl.constexpr, +): + i = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + valid = i < N + if B == 16: + d = tl.load(P + i, valid, other=0).to(tl.int32) + else: + per = 8 // B + byte = tl.load(P + i // per, valid, other=0).to(tl.int32) + code = (byte >> ((i % per) * B)) & ((1 << B) - 1) + d = (code ^ (1 << (B - 1))) - (1 << (B - 1)) + if BITMAP: + flags = tl.load(MASK + i // 32, valid, other=0).to(tl.uint32) + lower = (tl.full((BLOCK,), 1, tl.uint32) << (i % 32)) - 1 + exceptional = ((flags >> (i % 32)) & 1) != 0 + prefix = tl.load(PREFIX + i // 32, valid, other=0).to(tl.int32) + rank = libdevice.popc((flags & lower).to(tl.int32)).to(tl.int32) + exc = tl.load(VALUES + prefix + rank, valid & exceptional, other=0).to( + tl.int32 + ) + d = tl.where(exceptional, exc, d) + before = tl.load(W + i, valid, other=0).to(tl.int32) + tl.store(W + i, (before + SIGN * d).to(tl.int16), valid) + + +@triton.jit +def _switch_sparse( + W, INDEX, VALUES, N: tl.constexpr, SIGN: tl.constexpr, BLOCK: tl.constexpr +): + i = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + idx = tl.load(INDEX + i, i < N, other=0).to(tl.int32) + d = tl.load(VALUES + i, i < N, other=0).to(tl.int32) + before = tl.load(W + idx, i < N, other=0).to(tl.int32) + tl.store(W + idx, (before + SIGN * d).to(tl.int16), i < N) + + +def pack_codes(codes, bits): + per = 8 // bits + n = codes.numel() + pad = (-n) % per + if pad: + codes = F.pad(codes, (0, pad)) + shifts = torch.arange(per, device=codes.device, dtype=torch.int32) * bits + return (codes.reshape(-1, per).int() << shifts).sum(1).to(torch.uint8).contiguous() + + +def choose_format(d): + n = d.numel() + choices = [(n * 2, 16, "dense")] + if bool((d == 0).all()): + return (0, 0, "shared") + for b in [1, 2, 4, 8]: + count = int(((d < -(1 << (b - 1))) | (d >= (1 << (b - 1)))).sum()) + dense = (n * b + 7) // 8 + choices += [ + (dense + count * 6, b, "sparse"), + (dense + ((n + 31) // 32) * 8 + count * 2, b, "bitmap"), + ] + return min(choices) + + +@dataclass +class PackedDelta: + weight: torch.Tensor + bits: int + mode: str + packed: torch.Tensor + indices: torch.Tensor + values: torch.Tensor + bitmap: torch.Tensor + prefix: torch.Tensor + name: str = "" + + @property + def encoded_bytes(self): + return sum( + t.numel() * t.element_size() + for t in [self.packed, self.indices, self.values, self.bitmap, self.prefix] + ) + + def apply(self, sign): + if self.bits == 0: + return + n = self.weight.numel() + _switch_dense[(triton.cdiv(n, 2048),)]( + self.weight, + self.packed, + self.bitmap, + self.prefix, + self.values, + n, + self.bits, + self.mode == "bitmap", + sign, + 2048, + ) + if self.mode == "sparse" and self.indices.numel(): + _switch_sparse[(triton.cdiv(self.indices.numel(), 256),)]( + self.weight, self.indices, self.values, self.indices.numel(), sign, 256 + ) + + +def encode(left, right, name=""): + if not (left.dtype == right.dtype == torch.bfloat16 and left.shape == right.shape): + raise ValueError("Invalid branch weight layout") + if not (left.is_cuda and left.is_contiguous() and right.is_contiguous()): + raise ValueError("Invalid branch weight layout") + right = right.to(left.device) + d = ( + ( + right.view(torch.int16).flatten().int() + - left.view(torch.int16).flatten().int() + ) + .to(torch.int16) + .int() + ) + expected_bytes, bits, mode = choose_format(d) + empty = torch.empty(0, device=left.device, dtype=torch.int32) + p = PackedDelta( + left.view(torch.int16).flatten(), + bits, + mode, + empty, + empty, + empty, + empty, + empty, + name, + ) + if bits == 0: + if not (bool((d == 0).all())): + raise ValueError("Invalid branch weight layout") + return p + if bits == 16: + p.packed = d.to(torch.int16) + return p + exception = (d < -(1 << (bits - 1))) | (d >= (1 << (bits - 1))) + indices = torch.nonzero(exception).flatten() + p.values = d[indices].to(torch.int16) + p.packed = pack_codes(torch.where(exception, 0, d & ((1 << bits) - 1)), bits) + if mode == "sparse": + p.indices = indices.to(torch.int32) + elif mode == "bitmap": + flags = exception.int() + pad = (-flags.numel()) % 32 + if pad: + flags = F.pad(flags, (0, pad)) + flags = flags.reshape(-1, 32) + shifts = torch.arange(32, device=left.device, dtype=torch.int64) + p.bitmap = (flags.to(torch.int64) << shifts).sum(1).to(torch.int32) + counts = flags.sum(1, dtype=torch.int32) + p.prefix = (counts.cumsum(0) - counts).to(torch.int32) + else: + raise ValueError(mode) + if not (p.encoded_bytes == expected_bytes): + raise ValueError((name, p.encoded_bytes, expected_bytes)) + return p + + +class ResidentBranches: + def __init__(self, early, late, use_graph=True): + start = time.perf_counter() + self.plans = [] + self.active = "equal" + self.switches = 0 + self.graphs = {} + self.device = next(early.parameters()).device + early_params = dict(early.named_parameters()) + late_params = dict(late.named_parameters()) + if not (set(early_params) == set(late_params)): + raise ValueError("Invalid branch weight layout") + before = torch.cuda.memory_allocated() + shared_bytes = 0 + independent = [] + for name, left in early_params.items(): + right = late_params[name] + if ( + ".lora_" in name + or ".cam_self_attn." in name + or left.dtype != torch.bfloat16 + ): + independent.append(name) + continue + if not (left.dtype == right.dtype and left.shape == right.shape): + raise ValueError(name) + plan = encode(left.detach(), right.detach(), name) + # Verify reconstruction before sharing the dense parameter storage. + plan.apply(1) + if not ( + torch.equal( + left.detach().view(torch.int16), + right.detach().to(left.device).view(torch.int16), + ) + ): + raise ValueError(name) + plan.apply(-1) + right.data = ( + left.data + ) # Both module objects reference one materialized BF16 tensor. + self.plans.append(plan) + shared_bytes += left.numel() * left.element_size() + if len(self.plans) % 100 == 0: + logger.debug("Packed %d parameter tensors", len(self.plans)) + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize() + # Compile/warm both directions; exact round trip restores the initial branch. + self._apply(1) + self._apply(-1) + torch.cuda.synchronize() + if use_graph: + for sign in [1, -1]: + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + self._apply(sign) + self.graphs[sign] = graph + torch.cuda.synchronize() + self.report = dict( + shared_dense_bytes=shared_bytes, + delta_bytes=sum(p.encoded_bytes for p in self.plans), + base_parameter_tensors=len(self.plans), + independent_parameter_tensors=len(independent), + format_counts={ + mode: sum(p.mode == mode for p in self.plans) + for mode in ["shared", "dense", "sparse", "bitmap"] + }, + gpu_allocated_before=before, + gpu_allocated_after=torch.cuda.memory_allocated(), + gpu_reserved_after=torch.cuda.memory_reserved(), + setup_seconds=time.perf_counter() - start, + gpu_only_switch=True, + lossless_bf16_bit_patterns=True, + cuda_graph=use_graph, + all_reconstructed_base_tensors_bitwise_verified=True, + ) + logger.info("Prepared resident branches in %.1fs", self.report["setup_seconds"]) + + def _apply(self, sign): + for plan in self.plans: + plan.apply(sign) + + def switch(self, branch): + if branch == self.active: + return + if branch not in ["equal", "old"]: + raise ValueError(branch) + sign = 1 if branch == "old" else -1 + if self.graphs: + self.graphs[sign].replay() + else: + self._apply(sign) + self.active = branch + self.switches += 1 diff --git a/worldcrafter/fast/sampling.py b/worldcrafter/fast/sampling.py new file mode 100644 index 0000000000000000000000000000000000000000..f0c7b0209f197a7512f2b3809832a56b777c7d9f --- /dev/null +++ b/worldcrafter/fast/sampling.py @@ -0,0 +1,194 @@ +"""Contract-driven sampler with separate image- and text-to-video routing.""" + +from typing import Any, Callable +import math +import torch +import torch.nn.functional as F +from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback +from .contract import resolve_dmd_inference_trace +from .camera import build_ucpe_attention_kwargs_sequential_pyramid, pyramid_token_grids +from ..diffusers.pipeline import XLA_AVAILABLE + + +def sample_fast( + self, + latents: torch.Tensor = None, + pyramid_num_stages: int = None, + pyramid_num_inference_steps_list: list[int] = None, + prompt_embeds: torch.Tensor = None, + guidance_scale: float | None = 1.0, + indices_hidden_states: torch.Tensor = None, + indices_latents_history_short: torch.Tensor = None, + indices_latents_history_mid: torch.Tensor = None, + indices_latents_history_long: torch.Tensor = None, + latents_history_short: torch.Tensor = None, + latents_history_mid: torch.Tensor = None, + latents_history_long: torch.Tensor = None, + attention_kwargs: dict | None = None, + camera_trajectory: dict[str, Any] | None = None, + num_latent_frames_per_chunk: int | None = None, + chunk_index: int | None = None, + camera_restart_each_chunk: bool = True, + camera_translation_scale: float = 1.0, + ucpe_pixel_center: bool = False, + device: torch.device | None = None, + transformer_dtype: torch.dtype = None, + generator: torch.Generator | None = None, + callback_on_step_end: ( + Callable[[int, int], None] | PipelineCallback | MultiPipelineCallbacks | None + ) = None, + callback_on_step_end_tensor_inputs: list[str] = ["latents"], + progress_bar=None, +): + """Run the checkpoint's three-stage DMD schedule with CFG=1.""" + if not self.config.is_distilled or guidance_scale != 1.0: + raise ValueError("Fast sampling requires a distilled model and CFG=1") + if pyramid_num_inference_steps_list is not None: + raise ValueError("Fast step counts are defined by the checkpoint") + controller = self.resident_branches + stage_models = self.stage_transformers + if len(stage_models) != 3: + raise ValueError("Fast sampling requires three stage models") + dmd_trace = resolve_dmd_inference_trace( + self.dmd_timestep_contract, + latent_shape=latents.shape[1:], + history_tensors=( + latents_history_short, + latents_history_mid, + latents_history_long, + ), + num_stages=pyramid_num_stages, + ) + batch_size, num_channel, num_frames, height, width = latents.shape + patch_size = self.transformer.config.patch_size + ucpe_attention_kwargs_per_stage = None + if ( + camera_trajectory is not None + and num_latent_frames_per_chunk is not None + and (chunk_index is not None) + ): + ucpe_attention_kwargs_per_stage = ( + build_ucpe_attention_kwargs_sequential_pyramid( + self.transformer, + camera_trajectory, + num_latent_frames_per_chunk=num_latent_frames_per_chunk, + chunk_index=chunk_index, + token_grids=pyramid_token_grids( + height // patch_size[1], + width // patch_size[2], + pyramid_num_stages, + low_to_high=True, + ), + vae_scale_factor_temporal=self.vae_scale_factor_temporal, + restart_each_chunk=camera_restart_each_chunk, + translation_scale=camera_translation_scale, + pixel_center=ucpe_pixel_center, + ) + ) + latents = latents.permute(0, 2, 1, 3, 4).reshape( + batch_size * num_frames, num_channel, height, width + ) + for _ in range(pyramid_num_stages - 1): + height //= 2 + width //= 2 + latents = F.interpolate(latents, size=(height, width), mode="bilinear") * 2 + latents = latents.reshape( + batch_size, num_frames, num_channel, height, width + ).permute(0, 2, 1, 3, 4) + batch_size = latents.shape[0] + start_point_list = [latents] + i = 0 + for i_s in range(pyramid_num_stages): + stage_contract = self.dmd_timestep_contract.stage_tensors( + i_s, empty_history=dmd_trace.empty_history, device=device + ) + self.scheduler.timesteps = stage_contract["model_timestep"] + self.scheduler.sigmas = stage_contract["current_sigma"] + timesteps = self.scheduler.timesteps + if i_s > 0: + height *= 2 + width *= 2 + num_frames = latents.shape[2] + latents = latents.permute(0, 2, 1, 3, 4).reshape( + batch_size * num_frames, num_channel, height // 2, width // 2 + ) + latents = F.interpolate(latents, size=(height, width), mode="nearest") + latents = latents.reshape( + batch_size, num_frames, num_channel, height, width + ).permute(0, 2, 1, 3, 4) + ori_sigma = 1 - self.scheduler.ori_start_sigmas[i_s] + gamma = self.scheduler.config.gamma + alpha = 1 / (math.sqrt(1 + 1 / gamma) * (1 - ori_sigma) + ori_sigma) + beta = alpha * (1 - ori_sigma) / math.sqrt(gamma) + batch_size, channel, num_frames, height, width = latents.shape + noise = self.sample_block_noise( + batch_size, + channel, + num_frames, + height, + width, + patch_size, + device, + generator, + ) + noise = noise.to(device=device, dtype=transformer_dtype) + latents = alpha * latents + beta * noise + start_point_list.append(latents) + stage_attention_kwargs = dict(attention_kwargs or {}) + if ucpe_attention_kwargs_per_stage is not None: + stage_attention_kwargs.update(ucpe_attention_kwargs_per_stage[i_s]) + for idx, t in enumerate(timesteps): + use_low_noise = ( + i_s >= 1 + if getattr(self, "fast_inference_mode", "i2v") == "t2v" + else i_s == 2 and idx >= len(timesteps) // 2 + ) + branch = "old" if use_low_noise else "equal" + controller.switch(branch) + stage_transformer = stage_models[2 if branch == "old" else 0] + self.stage_forward_context = ( + int(chunk_index), + int(i_s), + int(idx), + int(len(timesteps)), + ) + timestep = self.dmd_timestep_contract.student_condition( + t, latents.shape[0], device=latents.device + ) + with stage_transformer.cache_context("cond"): + noise_pred = stage_transformer( + hidden_states=latents.to(transformer_dtype), + timestep=timestep, + encoder_hidden_states=prompt_embeds, + attention_kwargs=stage_attention_kwargs, + return_dict=False, + indices_hidden_states=indices_hidden_states, + indices_latents_history_short=indices_latents_history_short, + indices_latents_history_mid=indices_latents_history_mid, + indices_latents_history_long=indices_latents_history_long, + latents_history_short=latents_history_short.to(transformer_dtype), + latents_history_mid=latents_history_mid.to(transformer_dtype), + latents_history_long=latents_history_long.to(transformer_dtype), + )[0] + current_sigma = stage_contract["current_sigma"][idx].float() + next_sigma = stage_contract["next_sigma"][idx].float() + pred_image_or_video = latents.float() - current_sigma * noise_pred.float() + latents = self.dmd_timestep_contract.renoise_x0( + pred_image_or_video, start_point_list[i_s].float(), next_sigma + ) + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + progress_bar.update() + if XLA_AVAILABLE: + from ..diffusers.pipeline import xm + + xm.mark_step() + i += 1 + if not torch.isfinite(latents).all(): + raise FloatingPointError("Fast denoising produced NaN or Inf") + return latents diff --git a/worldcrafter/fast/timestep_grid.py b/worldcrafter/fast/timestep_grid.py new file mode 100644 index 0000000000000000000000000000000000000000..1fea5d839da48d7bd16cbab3c399f54c15001480 --- /dev/null +++ b/worldcrafter/fast/timestep_grid.py @@ -0,0 +1,440 @@ +"""Checkpoint timestep tables for Fast inference.""" + +from __future__ import annotations + +import hashlib +import json +import math +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Mapping, Sequence + +import torch + +DMD_TIMESTEP_CONTRACT_FILENAME = "dmd_timestep_contract.json" +DMD_TIMESTEP_CONTRACT_SCHEMA_VERSION = 3 +_LOCKED_LATENT_SHAPE = (16, 9, 48, 80) + +_LOCKED_ROLLOUT_STEPS = (2, 2, 2) + +_LOCKED_NORMAL_TIMESTEPS = ( + (998.5342, 833.9636), + (742.8216, 547.1926), + (385.4137, 253.9905), +) + +_LOCKED_EMPTY_HISTORY_TIMESTEPS = ( + (998.5342, 902.2183, 833.9636, 783.0660), + (742.8216, 640.0038, 547.1926, 462.9951), + (385.4137, 328.6249, 253.9905, 151.5308), +) + + +def resolve_dmd_contract_path(checkpoint_path: str | Path) -> Path: + """Return the canonical sidecar location for a checkpoint/run directory.""" + + path = Path(checkpoint_path) + if path.name == DMD_TIMESTEP_CONTRACT_FILENAME: + return path + if path.is_dir(): + return path / DMD_TIMESTEP_CONTRACT_FILENAME + if path.suffix: + return path.parent / DMD_TIMESTEP_CONTRACT_FILENAME + return path / DMD_TIMESTEP_CONTRACT_FILENAME + + +def _canonical_json(payload: Mapping[str, Any]) -> bytes: + return json.dumps( + payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True + ).encode("utf-8") + + +def _require_exact_keys( + payload: Mapping[str, Any], expected: set[str], name: str +) -> None: + actual = set(payload) + if actual != expected: + missing = sorted(expected - actual) + unexpected = sorted(actual - expected) + raise ValueError( + f"Invalid {name} keys: missing={missing}, unexpected={unexpected}" + ) + + +@dataclass(frozen=True) +class DmdTimestepStep: + """One student query and the flow interpolation coefficients around it.""" + + model_timestep: float + current_sigma: float + next_sigma: float + + def to_dict(self) -> dict[str, float]: + return { + "model_timestep": self.model_timestep, + "current_sigma": self.current_sigma, + "next_sigma": self.next_sigma, + } + + @classmethod + def from_dict(cls, payload: Mapping[str, Any]) -> "DmdTimestepStep": + _require_exact_keys( + payload, {"model_timestep", "current_sigma", "next_sigma"}, "DMD step" + ) + return cls( + model_timestep=float(payload["model_timestep"]), + current_sigma=float(payload["current_sigma"]), + next_sigma=float(payload["next_sigma"]), + ) + + +@dataclass(frozen=True) +class DmdSchedulerConfig: + """The fixed stage/band math from which the serialized student table came.""" + + num_train_timesteps: int + stages: int + stage_range: tuple[float, ...] + gamma: float + shift: float + version: str + use_dynamic_shifting: bool = True + time_shift_type: str = "linear" + base_seq_len: int = 256 + max_seq_len: int = 4096 + base_shift: float = 0.5 + max_shift: float = 1.15 + + def validate(self) -> None: + if self.num_train_timesteps != 1000: + raise ValueError( + f"DMD requires 1000 scheduler timesteps, got {self.num_train_timesteps}" + ) + if self.stages != 3: + raise ValueError( + f"Direct DMD requires three pyramid stages, got {self.stages}" + ) + expected_range = (0.0, 1.0 / 3.0, 2.0 / 3.0, 1.0) + if len(self.stage_range) != len(expected_range) or any( + not math.isclose(actual, expected, rel_tol=0.0, abs_tol=1e-12) + for actual, expected in zip(self.stage_range, expected_range) + ): + raise ValueError(f"Unsupported DMD stage range: {self.stage_range}") + if not math.isclose(self.gamma, 1.0 / 3.0, rel_tol=0.0, abs_tol=1e-12): + raise ValueError(f"Unsupported DMD stage gamma: {self.gamma}") + if not math.isclose(self.shift, 1.0, rel_tol=0.0, abs_tol=1e-12): + raise ValueError(f"Unsupported DMD stage shift: {self.shift}") + if self.version != "v1": + raise ValueError(f"Unsupported DMD stage scheduler version: {self.version}") + if not self.use_dynamic_shifting or self.time_shift_type != "linear": + raise ValueError( + "Direct DMD student schedules require linear resolution-dependent shifting" + ) + if (self.base_seq_len, self.max_seq_len) != (256, 4096): + raise ValueError("Direct DMD student shift sequence-length anchors changed") + if not math.isclose(self.base_shift, 0.5) or not math.isclose( + self.max_shift, 1.15 + ): + raise ValueError("Direct DMD student shift endpoints changed") + + def to_dict(self) -> dict[str, Any]: + return { + "num_train_timesteps": self.num_train_timesteps, + "stages": self.stages, + "stage_range": list(self.stage_range), + "gamma": self.gamma, + "shift": self.shift, + "version": self.version, + "use_dynamic_shifting": self.use_dynamic_shifting, + "time_shift_type": self.time_shift_type, + "base_seq_len": self.base_seq_len, + "max_seq_len": self.max_seq_len, + "base_shift": self.base_shift, + "max_shift": self.max_shift, + } + + @classmethod + def from_dict(cls, payload: Mapping[str, Any]) -> "DmdSchedulerConfig": + expected = { + "num_train_timesteps", + "stages", + "stage_range", + "gamma", + "shift", + "version", + "use_dynamic_shifting", + "time_shift_type", + "base_seq_len", + "max_seq_len", + "base_shift", + "max_shift", + } + _require_exact_keys(payload, expected, "student scheduler provenance") + provenance = cls( + num_train_timesteps=int(payload["num_train_timesteps"]), + stages=int(payload["stages"]), + stage_range=tuple(float(value) for value in payload["stage_range"]), + gamma=float(payload["gamma"]), + shift=float(payload["shift"]), + version=str(payload["version"]), + use_dynamic_shifting=bool(payload["use_dynamic_shifting"]), + time_shift_type=str(payload["time_shift_type"]), + base_seq_len=int(payload["base_seq_len"]), + max_seq_len=int(payload["max_seq_len"]), + base_shift=float(payload["base_shift"]), + max_shift=float(payload["max_shift"]), + ) + provenance.validate() + return provenance + + +@dataclass(frozen=True) +class DmdTimestepContract: + schema_version: int + latent_shape: tuple[int, int, int, int] + rollout_steps_per_stage: tuple[int, int, int] + amplify_empty_history: bool + normal_stages: tuple[tuple[DmdTimestepStep, ...], ...] + empty_history_stages: tuple[tuple[DmdTimestepStep, ...], ...] + student_scheduler: DmdSchedulerConfig + fingerprint: str + + def __post_init__(self) -> None: + self._validate() + + def _validate(self) -> None: + if self.schema_version not in (2, DMD_TIMESTEP_CONTRACT_SCHEMA_VERSION): + raise ValueError( + f"Unsupported DMD timestep contract schema {self.schema_version}; " + f"expected 2 or {DMD_TIMESTEP_CONTRACT_SCHEMA_VERSION}" + ) + if len(self.latent_shape) != 4 or any( + value <= 0 for value in self.latent_shape + ): + raise ValueError(f"Invalid DMD latent shape: {self.latent_shape}") + if self.rollout_steps_per_stage != _LOCKED_ROLLOUT_STEPS: + raise ValueError( + f"Unsupported DMD rollout step counts {self.rollout_steps_per_stage}; " + f"expected {_LOCKED_ROLLOUT_STEPS}" + ) + self.student_scheduler.validate() + if len(self.normal_stages) != 3 or len(self.empty_history_stages) != 3: + raise ValueError( + "DMD student schedule must contain exactly three pyramid stages" + ) + expected_empty_multiplier = 2 if self.amplify_empty_history else 1 + for stage_index, (normal, empty) in enumerate( + zip(self.normal_stages, self.empty_history_stages) + ): + expected_normal = self.rollout_steps_per_stage[stage_index] + if len(normal) != expected_normal: + raise ValueError( + f"DMD normal stage {stage_index} contains {len(normal)} steps, expected {expected_normal}" + ) + if len(empty) != expected_normal * expected_empty_multiplier: + raise ValueError( + f"DMD empty-history stage {stage_index} contains {len(empty)} steps, expected " + f"{expected_normal * expected_empty_multiplier}" + ) + self._validate_stage(stage_index, normal, "normal") + self._validate_stage(stage_index, empty, "empty_history") + if ( + not self.amplify_empty_history + and self.empty_history_stages != self.normal_stages + ): + raise ValueError( + "Unamplified DMD contracts must use the normal table for empty history" + ) + if self.latent_shape == _LOCKED_LATENT_SHAPE: + self._validate_locked_timesteps( + self.normal_stages, _LOCKED_NORMAL_TIMESTEPS, "normal" + ) + if self.amplify_empty_history: + self._validate_locked_timesteps( + self.empty_history_stages, + _LOCKED_EMPTY_HISTORY_TIMESTEPS, + "empty_history", + ) + + @staticmethod + def _validate_stage( + stage_index: int, steps: Sequence[DmdTimestepStep], name: str + ) -> None: + previous_timestep = float("inf") + for index, step in enumerate(steps): + values = (step.model_timestep, step.current_sigma, step.next_sigma) + if not all(math.isfinite(value) for value in values): + raise ValueError( + f"Non-finite value in {name} stage {stage_index}, step {index}: {values}" + ) + if not 0.0 <= step.next_sigma < step.current_sigma <= 1.0: + raise ValueError( + f"Invalid sigma transition in {name} stage {stage_index}, step {index}: {values}" + ) + if step.model_timestep >= previous_timestep: + raise ValueError( + f"Timesteps are not strictly descending in {name} stage {stage_index}" + ) + previous_timestep = step.model_timestep + if index + 1 < len(steps) and not math.isclose( + step.next_sigma, + steps[index + 1].current_sigma, + rel_tol=0.0, + abs_tol=1e-7, + ): + raise ValueError( + f"Broken x0 re-noise transition in {name} stage {stage_index}, step {index}" + ) + if steps[-1].next_sigma != 0.0: + raise ValueError( + f"The final {name} transition in stage {stage_index} must end at clean x0" + ) + + @staticmethod + def _validate_locked_timesteps( + stages: Sequence[Sequence[DmdTimestepStep]], + expected: Sequence[Sequence[float]], + name: str, + ) -> None: + for stage_index, (actual_stage, expected_stage) in enumerate( + zip(stages, expected) + ): + actual = [step.model_timestep for step in actual_stage] + if len(actual) != len(expected_stage) or any( + not math.isclose(value, target, rel_tol=0.0, abs_tol=1e-3) + for value, target in zip(actual, expected_stage) + ): + raise ValueError( + f"Locked 384x640x9 DMD {name} schedule changed at stage {stage_index}: " + f"expected {list(expected_stage)}, got {actual}" + ) + + def stage( + self, stage_index: int, *, empty_history: bool = False + ) -> tuple[DmdTimestepStep, ...]: + if stage_index not in range(3): + raise IndexError(f"DMD pyramid stage must be 0, 1 or 2, got {stage_index}") + return (self.empty_history_stages if empty_history else self.normal_stages)[ + stage_index + ] + + def stage_tensors( + self, + stage_index: int, + *, + empty_history: bool = False, + device: str | torch.device | None = None, + ) -> dict[str, torch.Tensor]: + steps = self.stage(stage_index, empty_history=empty_history) + return { + "model_timestep": torch.tensor( + [step.model_timestep for step in steps], + device=device, + dtype=torch.float32, + ), + "current_sigma": torch.tensor( + [step.current_sigma for step in steps], + device=device, + dtype=torch.float32, + ), + "next_sigma": torch.tensor( + [step.next_sigma for step in steps], device=device, dtype=torch.float32 + ), + } + + @staticmethod + def student_condition( + model_timestep: float | DmdTimestepStep | torch.Tensor, + batch_size: int, + device: str | torch.device | None = None, + ) -> torch.Tensor: + if isinstance(model_timestep, DmdTimestepStep): + model_timestep = model_timestep.model_timestep + condition = torch.as_tensor(model_timestep, dtype=torch.float32, device=device) + if condition.ndim == 0 or condition.numel() == 1: + return condition.reshape(1).expand(batch_size) + if condition.shape != (batch_size,): + raise ValueError( + f"Student timestep condition must be scalar or shape ({batch_size},), got {tuple(condition.shape)}" + ) + return condition + + @staticmethod + def renoise_x0( + x0: torch.Tensor, + noise: torch.Tensor, + next_sigma: float | DmdTimestepStep | torch.Tensor, + ) -> torch.Tensor: + """Apply the shared flow transition ``(1-sigma)*x0 + sigma*noise``.""" + + if x0.shape != noise.shape: + raise ValueError( + f"x0/noise shapes differ: {tuple(x0.shape)} vs {tuple(noise.shape)}" + ) + if isinstance(next_sigma, DmdTimestepStep): + next_sigma = next_sigma.next_sigma + sigma = torch.as_tensor(next_sigma, device=x0.device, dtype=torch.float32) + if sigma.ndim == 0: + pass + elif sigma.ndim == 1 and sigma.shape[0] in (1, x0.shape[0]): + sigma = sigma.reshape(sigma.shape[0], *([1] * (x0.ndim - 1))) + else: + raise ValueError( + f"next_sigma must be scalar or one value per batch item; got shape {tuple(sigma.shape)}" + ) + # The contract owns the numerical transition as well as its coefficients: + # every caller keeps the rollout state in FP32 between model queries and + # casts only the transformer's input. Returning ``x0.dtype`` here would + # silently make a BF16 caller follow a different trajectory. + return (1.0 - sigma) * x0.to(torch.float32) + sigma * noise.to(torch.float32) + + @classmethod + def load_json( + cls, + path: str | Path, + *, + expected_latent_shape=None, + expected_fingerprint: str | None = None, + ) -> "DmdTimestepContract": + path = Path(path) + source = ( + path if path.suffix.lower() == ".json" else resolve_dmd_contract_path(path) + ) + document = json.loads(source.read_text(encoding="utf-8")) + if not isinstance(document, dict) or "fingerprint" not in document: + raise ValueError(f"Missing DMD contract fingerprint: {source}") + claimed = document.pop("fingerprint") + actual = hashlib.sha256(_canonical_json(document)).hexdigest() + if claimed != actual or ( + expected_fingerprint is not None and actual != expected_fingerprint + ): + raise ValueError(f"DMD contract fingerprint mismatch: {source}") + # Hash the complete original document, including metadata unused by inference. + student = document["student"] + if student["condition_dtype"] != "float32": + raise ValueError("DMD timestep conditions must use float32") + contract = cls( + schema_version=int(document["schema_version"]), + latent_shape=tuple(int(x) for x in document["latent_shape"]), + rollout_steps_per_stage=tuple( + int(x) for x in student["rollout_steps_per_stage"] + ), + amplify_empty_history=bool(student["amplify_empty_history"]), + normal_stages=tuple( + tuple(DmdTimestepStep.from_dict(x) for x in stage) + for stage in student["normal_stages"] + ), + empty_history_stages=tuple( + tuple(DmdTimestepStep.from_dict(x) for x in stage) + for stage in student["empty_history_stages"] + ), + student_scheduler=DmdSchedulerConfig.from_dict( + student["scheduler_provenance"] + ), + fingerprint=actual, + ) + if expected_latent_shape is not None and contract.latent_shape != tuple( + expected_latent_shape + ): + raise ValueError(f"DMD latent shape mismatch: {contract.latent_shape}") + return contract diff --git a/worldcrafter/inference.py b/worldcrafter/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..1cbf2bae3e623e9262c59a61639450eba97a45e5 --- /dev/null +++ b/worldcrafter/inference.py @@ -0,0 +1,573 @@ +from __future__ import annotations + +import hashlib +import json +from dataclasses import dataclass +from collections.abc import Callable +from pathlib import Path + +import numpy as np +import torch +from diffusers.models import AutoencoderKLWan +from diffusers.utils import export_to_video, load_image +from transformers import AutoTokenizer, UMT5EncoderModel + +from .output import sha256, save_chunk_state, assemble_resumed_video +from .diffusers import ( + WorldCrafterPipeline, + WorldCrafterScheduler, + WorldCrafterTransformer3DModel, +) +from .kernels import ( + replace_all_norms_with_flash_norms, + replace_rmsnorm_with_fp32, + replace_rope_with_flash_rope, +) +from .repencoder import ( + RepEncoder, + RepEncoderInferenceMemoryProvider, + RepEncoderInferenceProviderConfig, +) +from .ucpe.bridge import ( + enable_ucpe_inference_sdpa_attention, + load_ucpe_camera_adapter_weights, + patch_worldcrafter_transformer_ucpe, +) + + +CAMERA_CHUNK_FRAMES = 33 +MODEL_HEIGHT = 384 +MODEL_WIDTH = 640 + + +@dataclass(frozen=True) +class InferenceResult: + video_path: Path + summary_path: Path + summary: dict[str, object] + + +def load_camera(path: Path, num_chunks: int | None = None) -> np.ndarray: + pose = np.load(path, allow_pickle=False) + if pose.ndim == 3: + pose = pose[None] + if pose.ndim != 4 or pose.shape[-2:] not in ((3, 4), (4, 4)): + raise ValueError(f"camera must be [B,T,3,4] or [B,T,4,4], got {pose.shape}") + if pose.shape[0] != 1 or pose.shape[1] % CAMERA_CHUNK_FRAMES: + raise ValueError( + "WorldCrafter requires one camera trajectory containing complete 33-frame chunks" + ) + if not np.issubdtype(pose.dtype, np.floating) or not np.isfinite(pose).all(): + raise ValueError("camera must contain finite floating-point c2w matrices") + rotation = np.asarray(pose[..., :3, :3], dtype=np.float64) + gram = np.swapaxes(rotation, -1, -2) @ rotation + if not np.allclose(gram, np.eye(3), atol=5e-3, rtol=0.0): + raise ValueError("camera rotations are not orthonormal") + if not np.allclose(np.linalg.det(rotation), 1.0, atol=5e-3, rtol=0.0): + raise ValueError("camera rotations must have determinant +1") + if num_chunks is not None: + frames = int(num_chunks) * CAMERA_CHUNK_FRAMES + if pose.shape[1] < frames: + raise ValueError( + f"camera has {pose.shape[1] // CAMERA_CHUNK_FRAMES} chunks, " + f"but {num_chunks} were requested" + ) + pose = pose[:, :frames] + return np.ascontiguousarray(pose) + + +def validate_weights(model_path: Path) -> dict[str, Path]: + root = model_path.expanduser().resolve() + config_path = root / "inference_config.json" + config = json.loads(config_path.read_text()) if config_path.is_file() else {} + shared = (root / config.get("shared_components", ".")).resolve() + required = { + "root": root, + "transformer": root / "transformer", + "adapter": root / "adapter", + "repencoder": shared / "repencoder", + "vae": shared / "vae", + "scheduler": shared / "scheduler", + "text_encoder": shared / "text_encoder", + "tokenizer": shared / "tokenizer", + } + missing = [str(path) for path in required.values() if not path.exists()] + for filename in ( + required["adapter"] / "camera_adapter.pth", + required["adapter"] / "pytorch_lora_weights.safetensors", + required["repencoder"] / "model.safetensors", + required["repencoder"] / "config.json", + required["repencoder"] / "manifest.json", + ): + if not filename.is_file(): + missing.append(str(filename)) + if missing: + raise FileNotFoundError( + "WorldCrafter-Base is incomplete: " + ", ".join(missing) + ) + return required + + +def configure_attention( + transformer: WorldCrafterTransformer3DModel, backend: str +) -> str: + if backend != "auto": + transformer.set_attention_backend(backend) + return backend + for candidate in ("native", "_flash_3_hub", "flash_hub"): + try: + transformer.set_attention_backend(candidate) + return candidate + except (ImportError, RuntimeError, ValueError): + continue + raise RuntimeError("no supported attention backend is available") + + +def load_model_adapter( + pipe: WorldCrafterPipeline, adapter_path: Path +) -> dict[str, object]: + from diffusers.loaders.peft import _SET_ADAPTER_SCALE_FN_MAPPING + + _SET_ADAPTER_SCALE_FN_MAPPING.setdefault( + WorldCrafterTransformer3DModel.__name__, lambda _model_class, weights: weights + ) + state = WorldCrafterPipeline.lora_state_dict(str(adapter_path)) + transformer_keys = [key for key in state if key.startswith("transformer.")] + if not transformer_keys: + raise RuntimeError("adapter does not contain transformer low-rank weights") + name = "worldcrafter" + pipe.load_lora_weights(str(adapter_path), adapter_name=name) + pipe.set_adapters([name], adapter_weights=[1.0]) + return {"name": name, "tensor_keys": len(transformer_keys)} + + +class WorldCrafter: + def __init__( + self, + *, + pipeline: WorldCrafterPipeline, + memory_provider: RepEncoderInferenceMemoryProvider, + model_path: Path, + device: torch.device, + attention_backend: str, + adapter_load: dict[str, object], + height: int, + width: int, + ) -> None: + self.model_type = "base" + self.pipeline = pipeline + self.memory_provider = memory_provider + self.model_path = model_path + self.device = device + self.attention_backend = attention_backend + self.adapter_load = adapter_load + self.height = height + self.width = width + + @classmethod + def from_pretrained( + cls, + model_path: Path, + *, + model_type: str = "base", + device: str = "cuda:0", + height: int = MODEL_HEIGHT, + width: int = MODEL_WIDTH, + seed: int = 42, + memory_fov_h_deg: float = 100.0, + memory_fov_v_deg: float = 71.13349068444832, + memory_fov_samples_per_axis: int = 10, + attention_backend: str = "native", + enable_compile: bool = False, + ) -> "WorldCrafter": + if model_type == "fast": + from .model_loading import load_fast + + return load_fast( + cls, + model_path, + device=device, + height=height, + width=width, + seed=seed, + memory_fov_h_deg=memory_fov_h_deg, + memory_fov_v_deg=memory_fov_v_deg, + memory_fov_samples_per_axis=memory_fov_samples_per_axis, + attention_backend=attention_backend, + enable_compile=enable_compile, + ) + if model_type != "base": + raise ValueError(f"Unknown model type: {model_type}") + torch_device = torch.device(device) + if torch_device.type != "cuda" or not torch.cuda.is_available(): + raise RuntimeError("WorldCrafter inference requires CUDA") + if (height, width) != (MODEL_HEIGHT, MODEL_WIDTH): + raise ValueError("WorldCrafter-Base is fixed to 384x640 inference") + torch.cuda.set_device(torch_device) + paths = validate_weights(model_path) + + enable_ucpe_inference_sdpa_attention() + repencoder = RepEncoder.from_pretrained( + paths["repencoder"], + device=torch_device, + compute_dtype="bf16", + target_microbatch=4, + ) + memory_provider = RepEncoderInferenceMemoryProvider( + repencoder, + RepEncoderInferenceProviderConfig( + seed=seed, + trajectory_fov_horizontal_fov_degrees=memory_fov_h_deg, + trajectory_fov_vertical_fov_degrees=memory_fov_v_deg, + trajectory_fov_samples_per_axis=memory_fov_samples_per_axis, + ), + ) + + transformer = WorldCrafterTransformer3DModel.from_pretrained( + paths["transformer"], torch_dtype=torch.bfloat16 + ) + patch_worldcrafter_transformer_ucpe( + transformer=transformer, + method="relray_absmap", + height=height, + width=width, + attn_compress=8, + adaptation_method="parallel", + ) + camera_adapter = load_ucpe_camera_adapter_weights(transformer, paths["adapter"]) + if ( + camera_adapter["loaded_tensor_keys"] + != camera_adapter["expected_tensor_keys"] + ): + raise RuntimeError(f"camera adapter load is incomplete: {camera_adapter}") + adapter_dtypes = { + parameter.dtype + for block in transformer.blocks + for parameter in block.cam_self_attn.parameters() + } + if adapter_dtypes != {torch.float32}: + raise RuntimeError(f"camera adapter dtypes are invalid: {adapter_dtypes}") + + if not enable_compile: + transformer = replace_rmsnorm_with_fp32(transformer) + transformer = replace_all_norms_with_flash_norms(transformer) + replace_rope_with_flash_rope() + resolved_backend = configure_attention(transformer, attention_backend) + pipeline = WorldCrafterPipeline( + tokenizer=AutoTokenizer.from_pretrained(paths["tokenizer"]), + text_encoder=UMT5EncoderModel.from_pretrained( + paths["text_encoder"], torch_dtype=torch.bfloat16 + ), + transformer=transformer, + vae=AutoencoderKLWan.from_pretrained( + paths["vae"], torch_dtype=torch.float32 + ), + scheduler=WorldCrafterScheduler.from_pretrained(paths["scheduler"]), + ) + adapter_load = load_model_adapter(pipeline, paths["adapter"]) + pipeline = pipeline.to(torch_device) + if enable_compile: + torch.backends.cudnn.benchmark = True + pipeline.text_encoder.compile( + mode="max-autotune-no-cudagraphs", dynamic=False + ) + pipeline.vae.compile(mode="max-autotune-no-cudagraphs", dynamic=False) + pipeline.transformer.compile( + mode="max-autotune-no-cudagraphs", dynamic=False + ) + return cls( + pipeline=pipeline, + memory_provider=memory_provider, + model_path=paths["root"], + device=torch_device, + attention_backend=resolved_backend, + adapter_load=adapter_load, + height=height, + width=width, + ) + + def generate( + self, + *, + mode: str, + camera_path: Path, + output_path: Path, + prompt: str, + negative_prompt: str, + image_path: Path | None = None, + num_chunks: int | None = None, + chunk_output_dir: Path | None = None, + state_output_dir: Path | None = None, + resume_from: Path | None = None, + on_chunk_saved: Callable[[int, Path], None] | None = None, + stop_after_chunk: int | None = None, + num_inference_steps: int | None = None, + guidance_scale: float | None = None, + seed: int = 42, + fps: int = 16, + image_noise_sigma_min: float = 0.111, + image_noise_sigma_max: float = 0.135, + camera_x_fov: float = 100.0, + camera_xi: float = 0.0, + local_camera_path: Path | None = None, + ) -> InferenceResult: + if on_chunk_saved is not None and chunk_output_dir is None: + raise ValueError("on_chunk_saved requires chunk_output_dir") + is_fast = self.model_type == "fast" + if num_inference_steps is None: + num_inference_steps = 6 if is_fast else 50 + if guidance_scale is None: + guidance_scale = 1.0 if is_fast else 5.0 + if is_fast: + if guidance_scale != 1.0 or num_inference_steps != 6: + raise ValueError( + "Fast requires CFG=1 and six regular steps; the first T2V chunk uses twelve steps" + ) + if resume_from is not None or state_output_dir is not None: + raise ValueError("Fast resume/state export is not yet validated") + self.pipeline.resident_branches.switch("equal") + self.pipeline.resident_branches.switches = 0 + self.pipeline.stage_model_trace.clear() + self.pipeline.fast_inference_mode = mode + elif local_camera_path is not None: + raise ValueError("--local-camera-path is only used by fast inference") + if mode not in {"i2v", "t2v"}: + raise ValueError(f"unsupported mode: {mode}") + if not camera_path.is_file(): + raise FileNotFoundError(camera_path) + if mode == "i2v": + if image_path is None or not image_path.is_file(): + raise FileNotFoundError(image_path) + image = load_image(str(image_path)).resize((self.width, self.height)) + else: + if image_path is not None: + raise ValueError( + "text-to-video inference does not accept an input image" + ) + image = None + + camera_c2w = load_camera(camera_path, num_chunks=num_chunks) + num_frames = int(camera_c2w.shape[1]) + total_chunks = num_frames // CAMERA_CHUNK_FRAMES + camera = { + "c2w": camera_c2w, + "x_fov": torch.full( + (1,), camera_x_fov, device=self.device, dtype=torch.float32 + ), + "xi": torch.full((1,), camera_xi, device=self.device, dtype=torch.float32), + } + if is_fast: + if local_camera_path is not None: + local_c2w = load_camera(local_camera_path, num_chunks=total_chunks) + else: + from .ucpe.bridge import _relative_pose_chunk + + local_c2w = torch.cat( + [ + _relative_pose_chunk( + camera_c2w, + chunk_index=k, + window_num_frames=33, + device=self.device, + ) + for k in range(total_chunks) + ], + dim=1, + ) + camera["pose"] = torch.as_tensor( + local_c2w, device=self.device, dtype=torch.float32 + ) + camera["c2w"] = torch.as_tensor( + camera_c2w, device=self.device, dtype=torch.float32 + ) + if chunk_output_dir is not None: + chunk_output_dir.mkdir(parents=True, exist_ok=True) + if resume_from is not None and chunk_output_dir is None: + raise ValueError( + "resuming requires --chunk-output-dir with completed prefix chunks" + ) + if state_output_dir is not None: + state_output_dir.mkdir(parents=True, exist_ok=True) + final_chunk_index = total_chunks - 1 + if stop_after_chunk is not None: + if stop_after_chunk < 0 or stop_after_chunk >= total_chunks: + raise ValueError("stop_after_chunk must identify a generated chunk") + final_chunk_index = int(stop_after_chunk) + + run_contract = { + "mode": mode, + "camera_sha256": sha256(camera_path), + "image_sha256": sha256(image_path) if image_path is not None else None, + "prompt_sha256": hashlib.sha256(prompt.encode("utf-8")).hexdigest(), + "negative_prompt_sha256": hashlib.sha256( + negative_prompt.encode("utf-8") + ).hexdigest(), + "num_chunks": total_chunks, + "num_inference_steps": num_inference_steps, + "guidance_scale": guidance_scale, + "seed": seed, + "image_noise_sigma_min": image_noise_sigma_min, + "image_noise_sigma_max": image_noise_sigma_max, + "camera_x_fov": camera_x_fov, + "camera_xi": camera_xi, + "repencoder_model_sha256": self.memory_provider.runtime.report[ + "model_sha256" + ], + } + resume_state: dict[str, object] | None = None + prior_history_selection: list[dict[str, object]] = [] + if resume_from is not None: + if not resume_from.is_file(): + raise FileNotFoundError(resume_from) + resume_state = torch.load( + resume_from, map_location="cpu", weights_only=False + ) + if resume_state.get("run_contract") != run_contract: + raise ValueError( + "resume checkpoint does not match this inference run contract" + ) + prior_history_selection = list(resume_state.get("history_selection", [])) + if final_chunk_index < int(resume_state["next_chunk_index"]): + raise ValueError("stop_after_chunk precedes the resume point") + + def save_chunk(chunk_index: int, current_video: torch.Tensor) -> None: + if chunk_output_dir is None: + return + frames = self.pipeline.video_processor.postprocess_video( + current_video, output_type="np" + )[0] + path = chunk_output_dir / f"chunk_{chunk_index:03d}_33f.mp4" + export_to_video(frames, str(path), fps=fps) + if on_chunk_saved is not None: + on_chunk_saved(chunk_index, path) + print(f"[worldcrafter] completed {path}", flush=True) + + def save_state(chunk_index: int, state: dict[str, object]) -> None: + if state_output_dir is not None: + save_chunk_state( + chunk_index, + state, + state_output_dir=state_output_dir, + run_contract=run_contract, + history_selection=[ + *prior_history_selection, + *[ + record.to_jsonable() + for record in self.memory_provider.render_records + ], + ], + ) + + self.memory_provider.reset_sequence() + with torch.inference_mode(): + frames = self.pipeline( + prompt=prompt, + negative_prompt=negative_prompt, + height=self.height, + width=self.width, + num_frames=num_frames, + num_inference_steps=num_inference_steps, + guidance_scale=guidance_scale, + generator=torch.Generator(device=self.device).manual_seed(seed), + memory_size=4, + history_sizes=[2, 1], + num_latent_frames_per_chunk=9, + keep_first_frame=True, + is_enable_stage2=is_fast, + pyramid_num_inference_steps_list=None if is_fast else [2, 2, 2], + is_skip_first_chunk=False, + is_amplify_first_chunk=False, + use_zero_init=False, + zero_steps=1, + image=image, + image_noise_sigma_min=image_noise_sigma_min, + image_noise_sigma_max=image_noise_sigma_max, + video=None, + video_noise_sigma_min=0.111, + video_noise_sigma_max=0.135, + camera_trajectory=camera, + memory_provider=self.memory_provider, + callback_on_chunk_end=( + save_chunk if chunk_output_dir is not None else None + ), + callback_on_chunk_state=( + save_state if state_output_dir is not None else None + ), + resume_state=resume_state, + stop_after_chunk=final_chunk_index, + ).frames[0] + + start_chunk = ( + int(resume_state["next_chunk_index"]) if resume_state is not None else 0 + ) + expected_records = max(0, final_chunk_index - max(1, start_chunk) + 1) + if len(self.memory_provider.render_records) != expected_records: + raise RuntimeError( + f"expected {expected_records} RepEncoder calls, " + f"got {len(self.memory_provider.render_records)}" + ) + output_path.parent.mkdir(parents=True, exist_ok=True) + if is_fast and len(frames) != (final_chunk_index + 1) * CAMERA_CHUNK_FRAMES: + raise RuntimeError("Fast output must preserve every decoded RGB frame") + if resume_state is None: + export_to_video(frames, str(output_path), fps=fps) + else: + assemble_resumed_video(output_path, chunk_output_dir, final_chunk_index) + history_selection = [ + *prior_history_selection, + *[record.to_jsonable() for record in self.memory_provider.render_records], + ] + summary = { + "format": "worldcrafter_inference_v2", + "mode": mode, + "model_path": str(self.model_path), + "image_path": str(image_path) if image_path is not None else None, + "image_sha256": sha256(image_path) if image_path is not None else None, + "camera_path": str(camera_path), + "camera_sha256": sha256(camera_path), + "camera_semantics": "global metric c2w; UCPE chunk-relative poses are derived internally", + "output_path": str(output_path), + "output_sha256": sha256(output_path), + "prompt": prompt, + "negative_prompt": negative_prompt, + "num_frames": num_frames, + "completed_through_chunk": final_chunk_index, + "fps": fps, + "seed": seed, + "num_inference_steps": num_inference_steps, + "guidance_scale": guidance_scale, + "attention_backend": self.attention_backend, + "adapter": self.adapter_load, + "repencoder_model_sha256": self.memory_provider.runtime.report[ + "model_sha256" + ], + "resumed_from": str(resume_from) if resume_from is not None else None, + "history_selection": history_selection, + } + if is_fast: + first_chunk_steps = 12 if mode == "t2v" else 6 + expected_calls = first_chunk_steps + final_chunk_index * 6 + if len(self.pipeline.stage_model_trace) != expected_calls: + raise RuntimeError( + "Fast forward count differs from the mode-specific DMD contract" + ) + summary.update( + model_type="fast", + local_camera_path=str(local_camera_path), + local_camera_sha256=( + sha256(local_camera_path) if local_camera_path else None + ), + first_chunk_inference_steps=first_chunk_steps, + first_chunk_routing="4+8" if mode == "t2v" else "5+1", + subsequent_chunk_routing="2+4" if mode == "t2v" else "5+1", + fast=self.fast_report, + stage_model_trace=self.pipeline.stage_model_trace, + resident_branch_switches=self.pipeline.resident_branches.switches, + ) + summary_path = output_path.with_suffix(".json") + summary_path.write_text(json.dumps(summary, indent=2, sort_keys=True) + "\n") + print(f"[worldcrafter] saved {output_path}", flush=True) + return InferenceResult(output_path, summary_path, summary) + + +__all__ = ["InferenceResult", "WorldCrafter", "load_camera", "sha256"] diff --git a/worldcrafter/kernels/__init__.py b/worldcrafter/kernels/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e9de170f82e649b18a53b1f67d0177ad1a6599dc --- /dev/null +++ b/worldcrafter/kernels/__init__.py @@ -0,0 +1,9 @@ +from .fp32_rmsnorm import replace_rmsnorm_with_fp32 +from .triton_norm import replace_all_norms_with_flash_norms +from .triton_rope import replace_rope_with_flash_rope + +__all__ = [ + "replace_rmsnorm_with_fp32", + "replace_all_norms_with_flash_norms", + "replace_rope_with_flash_rope", +] diff --git a/worldcrafter/kernels/attention_dispatch.py b/worldcrafter/kernels/attention_dispatch.py new file mode 100644 index 0000000000000000000000000000000000000000..2b8f24d2980325ac6dde907eb8fadacd0e04f76b --- /dev/null +++ b/worldcrafter/kernels/attention_dispatch.py @@ -0,0 +1,195 @@ +import torch +from kernels import get_kernel + + +try: + # FA3 Only support Hopper (SM90, H100/H800) + major, _ = torch.cuda.get_device_capability() + if major < 9: + raise RuntimeError("FA3 requires Hopper (SM90+), current GPU not supported") + flash_attn3 = get_kernel("kernels-community/flash-attn3") + flash_attn_func = flash_attn3.flash_attn_func + flash_attn_varlen_func = flash_attn3.flash_attn_varlen_func +except (ImportError, RuntimeError): + try: + flash_attn2 = get_kernel("kernels-community/flash-attn2") + flash_attn_func = flash_attn2.flash_attn_func + flash_attn_varlen_func = flash_attn2.flash_attn_varlen_func + except ImportError: + flash_attn_varlen_func = None + flash_attn_func = None + + +try: + from sageattention import sageattn, sageattn_varlen + +except ImportError: + sageattn_varlen = None + sageattn = None + +try: + from xformers.ops import memory_efficient_attention as xformers_attn_func + +except ImportError: + xformers_attn_func = None + + +def create_navit_attention_masks( + batch_size: int, + original_context_length_list: list, + history_context_length: int, + encoder_hidden_states_seq_len: int, + device: torch.device, + restrict_self_attn: bool = False, + guidance_cross_attn: bool = False, +): + # For navit_hidden_attention_mask + if restrict_self_attn: + cu_seqlens_q = [0] + for _ in range(batch_size): + for length in original_context_length_list: + cu_seqlens_q.append(cu_seqlens_q[-1] + length) + cu_seqlens_q = torch.tensor(cu_seqlens_q, device=device, dtype=torch.int32) + max_seqlen_q = max(original_context_length_list) + + cu_seqlens_kv = [0] + for _ in range(batch_size): + for length in original_context_length_list: + cu_seqlens_kv.append( + cu_seqlens_kv[-1] + length + history_context_length + ) + cu_seqlens_kv = torch.tensor(cu_seqlens_kv, device=device, dtype=torch.int32) + max_seqlen_kv = max(original_context_length_list) + history_context_length + else: + cu_seqlens_kv = [0] + for _ in range(batch_size): + for length in original_context_length_list: + cu_seqlens_kv.append( + cu_seqlens_kv[-1] + length + history_context_length + ) + cu_seqlens_kv = torch.tensor(cu_seqlens_kv, device=device, dtype=torch.int32) + max_seqlen_kv = max(original_context_length_list) + history_context_length + cu_seqlens_q = cu_seqlens_kv + max_seqlen_q = max_seqlen_kv + navit_hidden_attention_mask = ( + cu_seqlens_q, + cu_seqlens_kv, + max_seqlen_q, + max_seqlen_kv, + ) + + # For navit_history_hidden_attention_mask + navit_history_hidden_attention_mask = None + if restrict_self_attn: + cu_seqlens_kv = [0] + for _ in range(batch_size): + for length in original_context_length_list: + cu_seqlens_kv.append(cu_seqlens_kv[-1] + history_context_length) + cu_seqlens_kv = torch.tensor(cu_seqlens_kv, device=device, dtype=torch.int32) + max_seqlen_kv = history_context_length + cu_seqlens_q = cu_seqlens_kv + max_seqlen_q = max_seqlen_kv + navit_history_hidden_attention_mask = ( + cu_seqlens_q, + cu_seqlens_kv, + max_seqlen_q, + max_seqlen_kv, + ) + + # For navit_encoder_attention_mask + if guidance_cross_attn: + cross_cu_seqlens_q = [0] + for _ in range(batch_size): + for length in original_context_length_list: + cross_cu_seqlens_q.append(cross_cu_seqlens_q[-1] + length) + cross_cu_seqlens_q = torch.tensor( + cross_cu_seqlens_q, device=device, dtype=torch.int32 + ) + cross_max_seqlen_q = max(original_context_length_list) + else: + cross_cu_seqlens_q = [0] + for _ in range(batch_size): + for length in original_context_length_list: + cross_cu_seqlens_q.append( + cross_cu_seqlens_q[-1] + length + history_context_length + ) + cross_cu_seqlens_q = torch.tensor( + cross_cu_seqlens_q, device=device, dtype=torch.int32 + ) + cross_cu_seqlens_q[0] = 0 + cross_max_seqlen_q = max(original_context_length_list) + history_context_length + + cu_seqlens_kv = [0] + for _ in range(batch_size): + for length in original_context_length_list: + cu_seqlens_kv.append(cu_seqlens_kv[-1] + encoder_hidden_states_seq_len) + cu_seqlens_kv = torch.tensor(cu_seqlens_kv, device=device, dtype=torch.int32) + max_seqlen_kv = encoder_hidden_states_seq_len + navit_encoder_attention_mask = ( + cross_cu_seqlens_q, + cu_seqlens_kv, + cross_max_seqlen_q, + max_seqlen_kv, + ) + + return ( + navit_hidden_attention_mask, + navit_encoder_attention_mask, + navit_history_hidden_attention_mask, + ) + + +@torch.compiler.disable +def _flash_attn_wrapper(q, k, v): + return flash_attn_func(q, k, v) + + +@torch.compiler.disable +def _flash_attn_varlen_wrapper( + q, k, v, cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv +): + return flash_attn_varlen_func( + q, k, v, cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv + ) + + +def attn_varlen_func(q, k, v, attention_mask=None): + if attention_mask is None: + if flash_attn_func is not None: + x = _flash_attn_wrapper(q, k, v) + return x + + if sageattn is not None: + x = sageattn(q, k, v, tensor_layout="NHD") + return x + + if xformers_attn_func is not None: + x = xformers_attn_func(q, k, v) + return x + + x = torch.nn.functional.scaled_dot_product_attention( + q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) + ).transpose(1, 2) + return x + + B, L, H, C = q.shape + + q = q.flatten(0, 1) + k = k.flatten(0, 1) + v = v.flatten(0, 1) + + cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv = attention_mask + if flash_attn_varlen_func is not None: + x = _flash_attn_varlen_wrapper( + q, k, v, cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv + ) + elif sageattn_varlen is not None: + x = sageattn_varlen( + q, k, v, cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv + ) + else: + raise NotImplementedError("No Attn Installed!") + + x = x.unflatten(0, (B, L)) + + return x diff --git a/worldcrafter/kernels/fp32_rmsnorm.py b/worldcrafter/kernels/fp32_rmsnorm.py new file mode 100644 index 0000000000000000000000000000000000000000..51d565ffb96534d190714e17922652499f682bb0 --- /dev/null +++ b/worldcrafter/kernels/fp32_rmsnorm.py @@ -0,0 +1,48 @@ +import torch +import torch.nn as nn + +from diffusers.models.normalization import RMSNorm +from diffusers.utils import is_torch_npu_available, is_torch_version + + +# ------------------------------- replace funtion ------------------------------- + + +def replace_rmsnorm_with_fp32(model): + patched_count = 0 + for name, module in model.named_modules(): + if isinstance(module, (torch.nn.RMSNorm, RMSNorm)): + + def new_forward(self, x): + return FP32RMSNorm.forward(self, x) + + module.forward = new_forward.__get__(module, module.__class__) + patched_count += 1 + print(f"Patched {patched_count} FP32_RMSNorm modules\n") + return model + + +# ------------------------------- Tiled MLP ------------------------------- + + +class FP32RMSNorm(RMSNorm): + def forward(self, hidden_states): + if is_torch_npu_available(): + raise ValueError("FP32RMSNorm is not available on NPU") + + if not is_torch_version(">=", "2.4"): + raise ValueError("FP32RMSNorm is only available in PyTorch 2.4 or higher") + + original_dtype = hidden_states.dtype + hidden_states = nn.functional.rms_norm( + hidden_states.float(), + normalized_shape=(hidden_states.shape[-1],), + weight=self.weight.float(), + eps=self.eps, + ) + + bias = getattr(self, "bias", None) + if bias is not None: + hidden_states = hidden_states + bias.float() + + return hidden_states.to(original_dtype) diff --git a/worldcrafter/kernels/triton_norm.py b/worldcrafter/kernels/triton_norm.py new file mode 100644 index 0000000000000000000000000000000000000000..e0c584069089bdaea929d498abca04f0b9df27ca --- /dev/null +++ b/worldcrafter/kernels/triton_norm.py @@ -0,0 +1,410 @@ +import torch +import triton +import triton.language as tl + +from diffusers.models.normalization import FP32LayerNorm, LayerNorm, RMSNorm + +from .fp32_rmsnorm import FP32RMSNorm +from .utils import calculate_settings, torch_gpu_device + + +# ------------------------------- replace funtion ------------------------------- + + +def replace_all_norms_with_flash_norms(model): + patched_count = {"LayerNorm": 0, "RMSNorm": 0} + + for name, module in model.named_modules(): + if isinstance(module, (LayerNorm, FP32LayerNorm)): + if hasattr(module, "elementwise_affine") and module.elementwise_affine: + module.forward = (lambda self, x: flash_layernorm(self, x)).__get__(module, module.__class__) + patched_count["LayerNorm"] += 1 + + if isinstance(module, (torch.nn.RMSNorm, RMSNorm, FP32RMSNorm)): + module.forward = (lambda self, x: flash_rms_layernorm(self, x)).__get__(module, module.__class__) + patched_count["RMSNorm"] += 1 + + print(f"Patched {patched_count['LayerNorm']} Flash_LayerNorm modules\n") + print(f"Patched {patched_count['RMSNorm']} Flash_RMSNorm modules\n") + + return model + + +# ------------------------------- layer norm ------------------------------- + + +@triton.jit +def layernorm_forward( + Y, + Y_row_stride, + X, + X_row_stride, + W, + b, + r, + mu, + n_cols: tl.constexpr, + eps: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + row_idx = tl.program_id(0) + col_offsets = tl.arange(0, BLOCK_SIZE) + mask = col_offsets < n_cols + + Y += row_idx * Y_row_stride + X += row_idx * X_row_stride + r += row_idx + mu += row_idx + + # According to https://pytorch.org/torchtune/stable/_modules/torchtune/modules/layer_norm.html#Fp32LayerNorm, all modules + # are in float32! + X_row = tl.load(X + col_offsets, mask=mask, other=0).to(tl.float32) + W_row = tl.load(W + col_offsets, mask=mask, other=0).to(tl.float32) + b_row = tl.load(b + col_offsets, mask=mask, other=0).to(tl.float32) + + mean_X = tl.sum(X_row, axis=0) / n_cols + # (X[0] - mean) == -mean so we need to mask it out + XX = tl.where(mask, X_row - mean_X, 0) + row_var = tl.sum(XX * XX, axis=0) / n_cols + inv_var = tl.math.rsqrt(row_var + eps) + tl.store(r, inv_var) + tl.store(mu, mean_X) + output = (XX * inv_var) * W_row + b_row + tl.store(Y + col_offsets, output, mask=mask) + + +@triton.jit +def layernorm_backward( + dY, + dY_row_stride, + X, + X_row_stride, + W, + b, + r, + mu, + n_cols: tl.constexpr, + eps: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + # Approximately follows https://github.com/karpathy/llm.c/blob/master/doc/layernorm/layernorm.md + row_idx = tl.program_id(0) + col_offsets = tl.arange(0, BLOCK_SIZE) + mask = col_offsets < n_cols + + dY += row_idx * dY_row_stride + X += row_idx * X_row_stride + r += row_idx + mu += row_idx + + # According to https://pytorch.org/torchtune/stable/_modules/torchtune/modules/layer_norm.html#Fp32LayerNorm, all modules + # are in float32! + dY_row = tl.load(dY + col_offsets, mask=mask, other=0).to(tl.float32) + X_row = tl.load(X + col_offsets, mask=mask, other=0).to(tl.float32) + W_row = tl.load(W + col_offsets, mask=mask, other=0).to(tl.float32) + # b_row = tl.load(b + col_offsets, mask = mask, other = 0).to(tl.float32) + + inv_var = tl.load(r).to(tl.float32) + mean = tl.load(mu).to(tl.float32) + normed = (X_row - mean) * inv_var + dY_W = dY_row * W_row + dX_row = dY_W - tl.sum(dY_W, axis=0) / n_cols - normed * tl.sum(dY_W * normed, axis=0) / n_cols + dX_row = dX_row * inv_var + tl.store(dY + col_offsets, dX_row, mask=mask) + + +class Flash_Layernorm(torch.autograd.Function): + @staticmethod + def forward(ctx, X, W, b, eps): + shape = X.shape + dim = shape[-1] + X = X.view(-1, dim) + n_rows, n_cols = X.shape + BLOCK_SIZE, num_warps = calculate_settings(n_cols) + device = X.device + Y = torch.empty((n_rows, n_cols), dtype=X.dtype, device=device) + r = torch.empty(n_rows, dtype=torch.float32, device=device) + mu = torch.empty(n_rows, dtype=torch.float32, device=device) + + with torch_gpu_device(device): + layernorm_forward[(n_rows,)]( + Y, + Y.stride(0), + X, + X.stride(0), + W, + b, + r, + mu, + n_cols, + eps, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=num_warps, + ) + ctx.eps = eps + ctx.BLOCK_SIZE = BLOCK_SIZE + ctx.num_warps = num_warps + ctx.save_for_backward(X, W, b, r, mu) + return Y.view(*shape) + + @staticmethod + def backward(ctx, dY): + shape = dY.shape + dim = shape[-1] + dY = dY.view(-1, dim) + X, W, b, r, mu = ctx.saved_tensors + n_rows, n_cols = dY.shape + + with torch_gpu_device(dY.device): + layernorm_backward[(n_rows,)]( + dY, + dY.stride(0), + X, + X.stride(0), + W, + b, + r, + mu, + n_cols, + ctx.eps, + BLOCK_SIZE=ctx.BLOCK_SIZE, + num_warps=ctx.num_warps, + ) + dX = dY.view(*shape) + return dX, None, None, None, None + + +def flash_layernorm(layernorm, X): + assert layernorm.elementwise_affine is True + W = layernorm.weight + bias = layernorm.bias + eps = layernorm.variance_epsilon if hasattr(layernorm, "variance_epsilon") else layernorm.eps + out = Flash_Layernorm.apply(X, W, bias, eps) + return out + + +# ------------------------------- layer norm ------------------------------- + + +# ------------------------------- rms norm ------------------------------- + + +@triton.jit +def _rms_layernorm_forward( + Y, + Y_row_stride: tl.constexpr, + X, + X_row_stride: tl.constexpr, + W, + W_row_stride: tl.constexpr, + r, + r_row_stride: tl.constexpr, + n_cols: tl.constexpr, + eps: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + """ + Flash RMS Layernorm kernel + Inspiration from a Triton tutorial: + https://triton-lang.org/main/getting-started/tutorials/05-layer-norm.html + """ + row_idx = tl.program_id(0) + col_offsets = tl.arange(0, BLOCK_SIZE) + mask = col_offsets < n_cols + + Y += row_idx * Y_row_stride + X += row_idx * X_row_stride + r += row_idx * r_row_stride + + X_row = tl.load(X + col_offsets, mask=mask, other=0).to(tl.float32) + W_row = tl.load(W + col_offsets, mask=mask, other=0) # .to(tl.float32) + + row_var = tl.sum(X_row * X_row, axis=0) / n_cols + inv_var = tl.math.rsqrt(row_var + eps) + tl.store(r, inv_var) + normed = X_row * inv_var + normed = normed.to(W_row.dtype) # Exact copy from HF + output = normed * W_row + tl.store(Y + col_offsets, output, mask=mask) + + +def _rms_layernorm_backward( + dY, + dY_row_stride: tl.constexpr, + dX, + dX_row_stride: tl.constexpr, + X, + X_row_stride: tl.constexpr, + W, + W_row_stride: tl.constexpr, + r, + r_row_stride: tl.constexpr, + n_cols: tl.constexpr, + eps: tl.constexpr, + GEMMA: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + """ + Flash RMS Layernorm kernel for the backward pass + Inspiration from a Triton tutorial: + https://triton-lang.org/main/getting-started/tutorials/05-layer-norm.html + """ + row_idx = tl.program_id(0) + col_offsets = tl.arange(0, BLOCK_SIZE) + mask = col_offsets < n_cols + + dY += row_idx * dY_row_stride + X += row_idx * X_row_stride + r += row_idx * r_row_stride + + if GEMMA: + dX += row_idx * dY_row_stride + else: + dX = dY + + dY_row = tl.load(dY + col_offsets, mask=mask, other=0).to(tl.float32) + X_row = tl.load(X + col_offsets, mask=mask, other=0).to(tl.float32) + W_row = tl.load(W + col_offsets, mask=mask, other=0).to(tl.float32) + + # Get saved row variance + inv_var = tl.load(r).to(tl.float32) + normed = X_row * inv_var + + if GEMMA: + dY_W = dY_row * (W_row + 1.0) + else: + dY_W = dY_row * W_row + + rowsum_dY_normed = tl.sum(dY_W * normed, axis=0) + output = inv_var / n_cols * (n_cols * dY_W - normed * rowsum_dY_normed) + tl.store(dX + col_offsets, output, mask=mask) + + +_rms_layernorm_backward = triton.jit(_rms_layernorm_backward) +_rms_layernorm_backward = triton.heuristics( + { + "GEMMA": lambda args: bool(args["GEMMA"]), + } +)(_rms_layernorm_backward) + + +@triton.jit +def _gemma_rms_layernorm_forward( + Y, + Y_row_stride: tl.constexpr, + X, + X_row_stride: tl.constexpr, + W, + W_row_stride: tl.constexpr, + r, + r_row_stride: tl.constexpr, + n_cols: tl.constexpr, + eps: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + # Copies https://github.com/google-deepmind/gemma/blob/main/gemma/layers.py#L31 + # and https://github.com/keras-team/keras-nlp/blob/v0.8.2/keras_nlp/models/gemma/rms_normalization.py#L33 + # exactly. Essentially all in float32! + row_idx = tl.program_id(0) + col_offsets = tl.arange(0, BLOCK_SIZE) + mask = col_offsets < n_cols + + Y += row_idx * Y_row_stride + X += row_idx * X_row_stride + r += row_idx * r_row_stride + + X_row = tl.load(X + col_offsets, mask=mask, other=0).to(tl.float32) + W_row = tl.load(W + col_offsets, mask=mask, other=0).to(tl.float32) + + row_var = tl.sum(X_row * X_row, axis=0) / n_cols + inv_var = tl.math.rsqrt(row_var + eps) + tl.store(r, inv_var) + normed = X_row * inv_var + output = normed * (W_row + 1.0) + + tl.store(Y + col_offsets, output, mask=mask) + + +class Flash_RMS_Layernorm(torch.autograd.Function): + @staticmethod + def forward(ctx, X: torch.Tensor, W: torch.Tensor, eps: float, gemma: bool = False): + shape = X.shape + dim: int = shape[-1] + X = X.reshape(-1, dim) + n_rows: int + n_cols: int + n_rows, n_cols = X.shape + BLOCK_SIZE: int + num_warps: int + BLOCK_SIZE, num_warps = calculate_settings(n_cols) + device = X.device + + Y = torch.empty((n_rows, n_cols), dtype=X.dtype, device=device) + r = torch.empty(n_rows, dtype=torch.float32, device=device) + + fx = _gemma_rms_layernorm_forward if gemma else _rms_layernorm_forward + with torch_gpu_device(device): + fx[(n_rows,)]( + Y, + Y.stride(0), + X, + X.stride(0), + W, + W.stride(0), + r, + r.stride(0), + n_cols, + eps, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=num_warps, + ) + ctx.eps = eps + ctx.BLOCK_SIZE = BLOCK_SIZE + ctx.num_warps = num_warps + ctx.GEMMA = gemma + ctx.save_for_backward(X, W, r) + return Y.view(*shape) + + @staticmethod + def backward(ctx, dY: torch.Tensor): + shape = dY.shape + dim: int = shape[-1] + dY = dY.reshape(-1, dim) + X, W, r = ctx.saved_tensors + n_rows: int + n_cols: int + n_rows, n_cols = dY.shape + dX = torch.empty_like(dY) if ctx.GEMMA else dY + + with torch_gpu_device(dY.device): + _rms_layernorm_backward[(n_rows,)]( + dY, + dY.stride(0), + dX, + dX.stride(0), + X, + X.stride(0), + W, + W.stride(0), + r, + r.stride(0), + n_cols, + ctx.eps, + GEMMA=ctx.GEMMA, + BLOCK_SIZE=ctx.BLOCK_SIZE, + num_warps=ctx.num_warps, + ) + dX = dX.view(*shape) + return dX, None, None, None + + +# Keep the custom autograd RMSNorm operation outside compiled graphs. +@torch.compiler.disable +def flash_rms_layernorm(layernorm, X: torch.Tensor, gemma: bool = False): + W: torch.Tensor = layernorm.weight + eps: float = layernorm.variance_epsilon if hasattr(layernorm, "variance_epsilon") else layernorm.eps + out = Flash_RMS_Layernorm.apply(X, W, eps, gemma) + return out + + +# ------------------------------- rms norm ------------------------------- diff --git a/worldcrafter/kernels/triton_rope.py b/worldcrafter/kernels/triton_rope.py new file mode 100644 index 0000000000000000000000000000000000000000..8e8e2cfb3dfacacebfbe68aafb1ed6bf52dc5427 --- /dev/null +++ b/worldcrafter/kernels/triton_rope.py @@ -0,0 +1,136 @@ +import torch +import triton +import triton.language as tl + +from .utils import calculate_settings, torch_gpu_device + + +# ------------------------------- replace funtion ------------------------------- + + +def apply_rotary_emb_transposed_flash(x, freqs_cis): + return Flash_RoPE_Transposed.apply(x, freqs_cis) + + +def replace_rope_with_flash_rope(): + from ..diffusers import transformer + + transformer.apply_rotary_emb_transposed = apply_rotary_emb_transposed_flash + print("Patched Flash_RoPE globally\n") + + +# ------------------------------- layer norm ------------------------------- + + +@triton.jit +def _apply_rope_transposed_kernel( + X, + Out, + cos, + sin, + n_heads: tl.constexpr, + stride_x: tl.constexpr, + stride_out: tl.constexpr, + stride_freq: tl.constexpr, + head_dim: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + row_idx = tl.program_id(0) + freq_row_idx = row_idx // n_heads + + half_head_dim = head_dim // 2 + col_offsets = tl.arange(0, BLOCK_SIZE) + mask = col_offsets < half_head_dim + + x_ptr = X + row_idx * stride_x + out_ptr = Out + row_idx * stride_out + cos_ptr = cos + freq_row_idx * stride_freq + sin_ptr = sin + freq_row_idx * stride_freq + + x_real = tl.load(x_ptr + col_offsets * 2, mask=mask, other=0.0) + x_imag = tl.load(x_ptr + col_offsets * 2 + 1, mask=mask, other=0.0) + cos_even = tl.load(cos_ptr + col_offsets * 2, mask=mask, other=0.0) + sin_odd = tl.load(sin_ptr + col_offsets * 2 + 1, mask=mask, other=0.0) + + out_even = x_real * cos_even - x_imag * sin_odd + out_odd = x_real * sin_odd + x_imag * cos_even + + tl.store(out_ptr + col_offsets * 2, out_even, mask=mask) + tl.store(out_ptr + col_offsets * 2 + 1, out_odd, mask=mask) + + +class Flash_RoPE_Transposed(torch.autograd.Function): + @staticmethod + def forward(ctx, x, freqs_cis): + # x: [B, seq_len, n_heads, head_dim] + # freqs_cis: [B, seq_len, head_dim*2] + + B, seq_len, n_heads, head_dim = x.shape + + x_flat = x.reshape(-1, head_dim).contiguous() + device = x_flat.device + out = torch.empty_like(x_flat) + + freqs_flat = freqs_cis.reshape(B * seq_len, -1).contiguous() + half_dim = freqs_flat.shape[-1] // 2 + cos = freqs_flat[:, :half_dim].contiguous() # [B*seq_len, head_dim] + sin = freqs_flat[:, half_dim:].contiguous() # [B*seq_len, head_dim] + + n_rows = x_flat.shape[0] # B*seq_len*n_heads + BLOCK_SIZE, num_warps = calculate_settings(head_dim // 2) + + with torch_gpu_device(device): + _apply_rope_transposed_kernel[(n_rows,)]( + x_flat, + out, + cos, + sin, + n_heads, + x_flat.stride(0), + out.stride(0), + cos.stride(0), + head_dim, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=num_warps, + ) + + out = out.reshape(B, seq_len, n_heads, head_dim) + + ctx.save_for_backward(cos, sin) + ctx.n_heads = n_heads + ctx.BLOCK_SIZE = BLOCK_SIZE + ctx.num_warps = num_warps + ctx.head_dim = head_dim + + return out + + @staticmethod + def backward(ctx, grad_output): + cos, sin = ctx.saved_tensors + + B, seq_len, n_heads, head_dim = grad_output.shape + grad_flat = grad_output.reshape(-1, head_dim).contiguous() + device = grad_flat.device + grad_x = torch.empty_like(grad_flat) + + sin_neg = -sin + + n_rows = grad_flat.shape[0] + + with torch_gpu_device(device): + _apply_rope_transposed_kernel[(n_rows,)]( + grad_flat, + grad_x, + cos, + sin_neg, + ctx.n_heads, + grad_flat.stride(0), + grad_x.stride(0), + cos.stride(0), + ctx.head_dim, + BLOCK_SIZE=ctx.BLOCK_SIZE, + num_warps=ctx.num_warps, + ) + + grad_x = grad_x.reshape(B, seq_len, n_heads, head_dim) + return grad_x, None diff --git a/worldcrafter/kernels/utils.py b/worldcrafter/kernels/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..aa5d395cbf611d6b189a5725a8a40b55c6e183ef --- /dev/null +++ b/worldcrafter/kernels/utils.py @@ -0,0 +1,70 @@ +from contextlib import nullcontext + +import torch +import triton + + +def get_device_type(): + if torch.cuda.is_available(): + try: + if torch.version.hip is not None: + return "hip" + except AttributeError: + pass + return "cuda" + + try: + if hasattr(torch, "xpu") and torch.xpu.is_available(): + return "xpu" + except (AttributeError, RuntimeError): + pass + + return "cpu" + + +def get_device_count(device_type): + if device_type == "cuda" or device_type == "hip": + return torch.cuda.device_count() + elif device_type == "xpu": + try: + return torch.xpu.device_count() + except (AttributeError, RuntimeError): + return 0 + return 0 + + +MAX_FUSED_SIZE: int = 65536 +next_power_of_2 = triton.next_power_of_2 +DEVICE_TYPE = get_device_type() +DEVICE_COUNT = get_device_count(DEVICE_TYPE) + +if DEVICE_COUNT > 1: + if DEVICE_TYPE in ("cuda", "hip"): + torch_gpu_device = torch.cuda.device + elif DEVICE_TYPE == "xpu": + torch_gpu_device = torch.xpu.device +else: + + def torch_gpu_device(device): + return nullcontext() + + +def calculate_settings( + n: int, +) -> ( + int, + int, +): + BLOCK_SIZE: int = next_power_of_2(n) + if BLOCK_SIZE > MAX_FUSED_SIZE: + raise RuntimeError( + f"Cannot launch Triton kernel since n = {n} exceeds the maximum CUDA blocksize = {MAX_FUSED_SIZE}." + ) + num_warps: int = 4 + if BLOCK_SIZE >= 32768: + num_warps = 32 + elif BLOCK_SIZE >= 8192: + num_warps = 16 + elif BLOCK_SIZE >= 2048: + num_warps = 8 + return BLOCK_SIZE, num_warps diff --git a/worldcrafter/model_loading.py b/worldcrafter/model_loading.py new file mode 100644 index 0000000000000000000000000000000000000000..d5329a0df0035589fadebfc3b5de4e5efd734ff7 --- /dev/null +++ b/worldcrafter/model_loading.py @@ -0,0 +1,238 @@ +"""Assembly of the independently adapted high- and low-noise fast branches.""" + +from __future__ import annotations + +import copy +import gc +import json +from pathlib import Path + +import torch +from diffusers.models import AutoencoderKLWan +from transformers import AutoTokenizer, UMT5EncoderModel + +from .diffusers import ( + WorldCrafterPipeline, + WorldCrafterScheduler, + WorldCrafterTransformer3DModel, +) +from .fast.compact_ucpe import compact_ucpe +from .fast.attention import FastUcpeSelfAttention +from .fast.contract import load_dmd_inference_contract +from .fast.resident import ResidentBranches +from .kernels import ( + replace_rmsnorm_with_fp32, + replace_all_norms_with_flash_norms, + replace_rope_with_flash_rope, +) +from .repencoder import ( + RepEncoder, + RepEncoderInferenceMemoryProvider, + RepEncoderInferenceProviderConfig, +) +from .ucpe.bridge import ( + enable_ucpe_inference_sdpa_attention, + patch_worldcrafter_transformer_ucpe, + load_ucpe_camera_adapter_weights, +) + + +def load_fast( + cls, + model_path, + *, + device, + height, + width, + seed, + memory_fov_h_deg, + memory_fov_v_deg, + memory_fov_samples_per_axis, + attention_backend, + enable_compile, +): + from .inference import configure_attention, load_model_adapter, sha256 + + if (height, width) != (384, 640): + raise ValueError("Fast weights require height=384 and width=640") + device = torch.device(device) + if device.type != "cuda" or not torch.cuda.is_available(): + raise RuntimeError("Fast inference requires CUDA") + torch.cuda.set_device(device) + root = Path(model_path).expanduser().resolve() + config = json.loads((root / "inference_config.json").read_text()) + manifest = json.loads((root / "manifest.json").read_text()) + expected_config = dict( + steps_per_stage=[2, 2, 2], + guidance_scale=1.0, + ucpe_pixel_center=True, + repencoder_target_microbatch=1, + representation="resident_byte_compact_ucpe", + compile=False, + ) + if any(config.get(key) != value for key, value in expected_config.items()): + raise ValueError( + "Fast inference configuration differs from the validated release contract" + ) + if config["routing"] != [["equal", "equal"], ["equal", "equal"], ["equal", "old"]]: + raise ValueError("Fast I2V requires 5+1 routing") + shared = (root / config["shared_components"]).resolve() + for row in manifest["files"]: + path = root / row["path"] + if not path.is_file() or path.stat().st_size != row["bytes"]: + raise ValueError(f"Incomplete fast checkpoint: {path}") + if path.stat().st_size < 1024 * 1024 and sha256(path) != row["sha256"]: + raise ValueError(f"Fast checkpoint metadata mismatch: {path}") + adapters = [root / "adapter_high_noise", root / "adapter_low_noise"] + contracts = [ + load_dmd_inference_contract(p, expected_latent_shape=(16, 9, 48, 80)) + for p in adapters + ] + contract = contracts[0] + if contract.fingerprint != contracts[ + 1 + ].fingerprint or contract.rollout_steps_per_stage != (2, 2, 2): + raise ValueError( + "Fast branches must have identical native 2/2/2 timestep contracts" + ) + for adapter in adapters: + frozen = json.loads((adapter / "repencoder_frozen.json").read_text()) + if ( + frozen["repencoder"]["model_sha256"] + != manifest["repencoder"]["reference_file_sha256"] + ): + raise ValueError("Fast adapter references an unexpected RepEncoder") + enable_ucpe_inference_sdpa_attention() + repencoder = RepEncoder.from_pretrained( + shared / "repencoder", device=device, compute_dtype="bf16", target_microbatch=1 + ) + if ( + repencoder.report["model_sha256"] + != manifest["repencoder"]["shared_file_sha256"] + ): + raise ValueError( + "Shared RepEncoder differs from the verified renamed checkpoint" + ) + provider = RepEncoderInferenceMemoryProvider( + repencoder, + RepEncoderInferenceProviderConfig( + seed=seed, + trajectory_fov_horizontal_fov_degrees=memory_fov_h_deg, + trajectory_fov_vertical_fov_degrees=memory_fov_v_deg, + trajectory_fov_samples_per_axis=memory_fov_samples_per_axis, + ), + ) + + def transformer(branch): + model = WorldCrafterTransformer3DModel.from_pretrained( + root / f"transformer_{branch}_noise", torch_dtype=torch.bfloat16 + ) + patch_worldcrafter_transformer_ucpe( + model, + method="relray_absmap", + height=height, + width=width, + attn_compress=8, + adaptation_method="parallel", + attention_cls=FastUcpeSelfAttention, + ) + loaded = load_ucpe_camera_adapter_weights( + model, root / f"adapter_{branch}_noise" / "transformer_partial.pth" + ) + if loaded["loaded_tensor_keys"] != loaded["expected_tensor_keys"]: + raise ValueError(f"Incomplete {branch} UCPE state") + model = replace_rmsnorm_with_fp32(model) + model = replace_all_norms_with_flash_norms(model) + configure_attention(model, attention_backend) + return model + + early = transformer("high") + replace_rope_with_flash_rope() + provenance = contract.student_scheduler + scheduler = WorldCrafterScheduler.from_config( + WorldCrafterScheduler.from_pretrained(shared / "scheduler").config, + num_train_timesteps=provenance.num_train_timesteps, + shift=provenance.shift, + stages=provenance.stages, + stage_range=list(provenance.stage_range), + gamma=provenance.gamma, + scheduler_type="dmd", + use_dynamic_shifting=provenance.use_dynamic_shifting, + time_shift_type=provenance.time_shift_type, + ) + pipe = WorldCrafterPipeline( + tokenizer=AutoTokenizer.from_pretrained(shared / "tokenizer"), + text_encoder=UMT5EncoderModel.from_pretrained( + shared / "text_encoder", torch_dtype=torch.bfloat16 + ), + transformer=early, + vae=AutoencoderKLWan.from_pretrained(shared / "vae", torch_dtype=torch.float32), + scheduler=scheduler, + is_distilled=True, + ) + early_lora = load_model_adapter(pipe, adapters[0]) + pipe.dmd_timestep_contract = contract + pipe.to(device) + late = transformer("low") + loader = copy.copy(pipe) + loader.register_modules(transformer=late) + late_lora = load_model_adapter(loader, adapters[1]) + compact = [compact_ucpe(m) for m in (early, late)] + pipe.resident_branches = ResidentBranches(early, late) + late.to(device) + pipe.stage_transformers = (early, early, late) + pipe.stage_model_trace = [] + + def record(branch): + def hook(module, args, kwargs, output): + chunk, stage, step, stage_steps = pipe.stage_forward_context + low_noise = ( + stage >= 1 + if getattr(pipe, "fast_inference_mode", "i2v") == "t2v" + else stage == 2 and step >= stage_steps // 2 + ) + expected = "old" if low_noise else "equal" + if branch != expected or pipe.resident_branches.active != branch: + raise RuntimeError("Fast transformer/adapter routing mismatch") + pipe.stage_model_trace.append( + dict( + chunk=chunk, + stage=stage, + step=step, + stage_steps=stage_steps, + branch=branch, + ) + ) + + return hook + + early.register_forward_hook(record("equal"), with_kwargs=True) + late.register_forward_hook(record("old"), with_kwargs=True) + if enable_compile: + # Keep routing and shared-weight switches outside compiled graphs. + # Compile each branch's blocks without changing parameter storage. + for branch in (early, late): + for block in branch.blocks: + block.compile(mode="default", dynamic=False) + gc.collect() + torch.cuda.empty_cache() + model = cls( + pipeline=pipe, + memory_provider=provider, + model_path=root, + device=device, + attention_backend=attention_backend, + adapter_load={"high": early_lora, "low": late_lora}, + height=height, + width=width, + ) + model.model_type = "fast" + model.fast_config = config + model.fast_report = dict( + compile_enabled=bool(enable_compile), + compile_scope="transformer_blocks" if enable_compile else None, + contract_fingerprint=contract.fingerprint, + compact_ucpe=compact, + resident=pipe.resident_branches.report, + ) + return model diff --git a/worldcrafter/output.py b/worldcrafter/output.py new file mode 100644 index 0000000000000000000000000000000000000000..6657f3a0e2adc34a940fe8d665c0fbdef75c3755 --- /dev/null +++ b/worldcrafter/output.py @@ -0,0 +1,94 @@ +"""Video output and resumable chunk state.""" + +from __future__ import annotations + +import hashlib +import json +import os +import subprocess +from pathlib import Path + +import torch + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def save_chunk_state( + chunk_index: int, + state: dict[str, object], + *, + state_output_dir: Path, + run_contract: dict, + history_selection: list, +) -> None: + state = dict(state) + state["run_contract"] = run_contract + state["history_selection"] = history_selection + checkpoint_path = state_output_dir / f"chunk_{chunk_index:03d}_complete.pt" + temporary_path = checkpoint_path.with_suffix(".pt.tmp") + torch.save(state, temporary_path) + os.replace(temporary_path, checkpoint_path) + metadata = { + "format": state["format"], + "completed_chunk_index": state["completed_chunk_index"], + "next_chunk_index": state["next_chunk_index"], + "checkpoint": checkpoint_path.name, + "checkpoint_sha256": sha256(checkpoint_path), + "run_contract": run_contract, + } + metadata_path = state_output_dir / f"chunk_{chunk_index:03d}_complete.json" + temporary_metadata_path = metadata_path.with_suffix(".json.tmp") + temporary_metadata_path.write_text( + json.dumps(metadata, indent=2, sort_keys=True) + "\n" + ) + os.replace(temporary_metadata_path, metadata_path) + latest_path = state_output_dir / "latest.json" + temporary_latest_path = latest_path.with_suffix(".json.tmp") + temporary_latest_path.write_text( + json.dumps(metadata, indent=2, sort_keys=True) + "\n" + ) + os.replace(temporary_latest_path, latest_path) + print(f"[worldcrafter] saved resumable state {checkpoint_path}", flush=True) + + +def assemble_resumed_video( + output_path: Path, chunk_output_dir: Path, final_chunk_index: int +) -> None: + chunk_paths = [ + chunk_output_dir / f"chunk_{index:03d}_33f.mp4" + for index in range(final_chunk_index + 1) + ] + missing_chunks = [str(path) for path in chunk_paths if not path.is_file()] + if missing_chunks: + raise FileNotFoundError( + "cannot assemble resumed output; missing chunks: " + + ", ".join(missing_chunks) + ) + concat_path = output_path.with_suffix(".concat.txt") + concat_path.write_text( + "".join(f"file '{path.resolve()}'\n" for path in chunk_paths), + encoding="utf-8", + ) + subprocess.run( + [ + "ffmpeg", + "-y", + "-f", + "concat", + "-safe", + "0", + "-i", + str(concat_path), + "-c", + "copy", + str(output_path), + ], + check=True, + ) + concat_path.unlink() diff --git a/worldcrafter/repencoder/__init__.py b/worldcrafter/repencoder/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..cff1fb8885ba8fa84fba5906146040489a288f9e --- /dev/null +++ b/worldcrafter/repencoder/__init__.py @@ -0,0 +1,16 @@ +from .model import RepEncoder +from .trajectory_fov import TrajectoryFovSelection, select_trajectory_fov_history +from .trajectory_memory_provider import ( + RepEncoderInferenceMemoryProvider, + RepEncoderInferenceProviderConfig, + RepEncoderInferenceRenderRecord, +) + +__all__ = [ + "RepEncoder", + "RepEncoderInferenceMemoryProvider", + "RepEncoderInferenceProviderConfig", + "RepEncoderInferenceRenderRecord", + "TrajectoryFovSelection", + "select_trajectory_fov_history", +] diff --git a/worldcrafter/repencoder/camera.py b/worldcrafter/repencoder/camera.py new file mode 100644 index 0000000000000000000000000000000000000000..b5f898b02597e7eaa3080d612400040b433982f2 --- /dev/null +++ b/worldcrafter/repencoder/camera.py @@ -0,0 +1,299 @@ +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch +import torch.nn.functional as F + + +NUM_SOURCE_VIEWS = 9 +NUM_TARGET_VIEWS = 4 +RAY_HEIGHT = 384 +RAY_WIDTH = 640 +RAY_HW = (RAY_HEIGHT, RAY_WIDTH) +CAMERA_SCALE_MULTIPLIER = 1.35 +DEFAULT_NEAR_ZERO_BASELINE_M = 1e-6 + + +@dataclass(frozen=True) +class RepEncoderCameraConditioning: + source_camera_tokens: torch.Tensor + target_rays: torch.Tensor + source_c2w_normalized: torch.Tensor + target_c2w_normalized: torch.Tensor + scale_tokens: torch.Tensor + scene_scale_m: torch.Tensor + + +def _validate_metric_c2w( + source_c2w_metric: torch.Tensor, + target_c2w_metric: torch.Tensor, +) -> None: + if not isinstance(source_c2w_metric, torch.Tensor) or not isinstance( + target_c2w_metric, torch.Tensor + ): + raise TypeError("source and target c2w values must be torch tensors") + if source_c2w_metric.ndim != 4 or source_c2w_metric.shape[1:] != ( + NUM_SOURCE_VIEWS, + 4, + 4, + ): + raise ValueError( + "source_c2w_metric must be [B,9,4,4], got " + f"{tuple(source_c2w_metric.shape)}" + ) + if target_c2w_metric.ndim != 4 or target_c2w_metric.shape[1:] != ( + NUM_TARGET_VIEWS, + 4, + 4, + ): + raise ValueError( + "target_c2w_metric must be [B,4,4,4], got " + f"{tuple(target_c2w_metric.shape)}" + ) + if source_c2w_metric.shape[0] != target_c2w_metric.shape[0]: + raise ValueError("source and target camera batches differ") + if source_c2w_metric.device != target_c2w_metric.device: + raise ValueError("source and target cameras must share one device") + if not torch.is_floating_point(source_c2w_metric) or not torch.is_floating_point( + target_c2w_metric + ): + raise TypeError("camera matrices must use a floating-point dtype") + if not torch.isfinite(source_c2w_metric).all() or not torch.isfinite( + target_c2w_metric + ).all(): + raise ValueError("camera matrices contain NaN or Inf") + + +def normalize_source_target_cameras( + source_c2w_metric: torch.Tensor, + target_c2w_metric: torch.Tensor, + *, + near_zero_baseline_m: float = DEFAULT_NEAR_ZERO_BASELINE_M, +) -> tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, +]: + _validate_metric_c2w(source_c2w_metric, target_c2w_metric) + epsilon = float(near_zero_baseline_m) + if not math.isfinite(epsilon) or epsilon < 0: + raise ValueError("near_zero_baseline_m must be finite and non-negative") + + with torch.autocast(device_type=source_c2w_metric.device.type, enabled=False): + reference_inverse = torch.linalg.inv(source_c2w_metric[:, 0]) + source = reference_inverse[:, None] @ source_c2w_metric + target = reference_inverse[:, None] @ target_c2w_metric + max_source = torch.linalg.vector_norm(source[..., :3, 3], dim=-1).amax(dim=1) + is_zero = max_source <= epsilon + scene_scale = torch.where( + is_zero, + torch.ones_like(max_source), + CAMERA_SCALE_MULTIPLIER * max_source, + ) + zero = torch.zeros_like(max_source) + one = torch.ones_like(max_source) + scale_tokens = torch.stack( + [ + torch.where(is_zero, zero, max_source / scene_scale), + torch.where(is_zero, one, zero), + ], + dim=-1, + ) + source = source.clone() + target = target.clone() + source[..., :3, 3] /= scene_scale[:, None, None] + target[..., :3, 3] /= scene_scale[:, None, None] + + identity = torch.eye(4, dtype=source.dtype, device=source.device) + if not torch.allclose( + source[:, 0], identity.expand(source.shape[0], -1, -1), atol=2e-4, rtol=0 + ): + raise RuntimeError("source[0] is not identity after RepEncoder normalization") + if not torch.isfinite(source).all() or not torch.isfinite(target).all(): + raise FloatingPointError("normalized RepEncoder cameras contain NaN or Inf") + return ( + source.contiguous(), + target.contiguous(), + scale_tokens.contiguous(), + scene_scale.contiguous(), + ) + + +def _sqrt_positive_part(value: torch.Tensor) -> torch.Tensor: + result = torch.zeros_like(value) + positive = value > 0 + if torch.is_grad_enabled(): + result[positive] = torch.sqrt(value[positive]) + return result + return torch.where(positive, torch.sqrt(value), result) + + +def _mat_to_quat_scalar_last(matrix: torch.Tensor) -> torch.Tensor: + if matrix.shape[-2:] != (3, 3): + raise ValueError(f"invalid rotation matrix shape {tuple(matrix.shape)}") + batch_shape = matrix.shape[:-2] + m00, m01, m02, m10, m11, m12, m20, m21, m22 = torch.unbind( + matrix.reshape(*batch_shape, 9), dim=-1 + ) + q_abs = _sqrt_positive_part( + torch.stack( + [ + 1.0 + m00 + m11 + m22, + 1.0 + m00 - m11 - m22, + 1.0 - m00 + m11 - m22, + 1.0 - m00 - m11 + m22, + ], + dim=-1, + ) + ) + quat_by_rijk = torch.stack( + [ + torch.stack( + [q_abs[..., 0] ** 2, m21 - m12, m02 - m20, m10 - m01], + dim=-1, + ), + torch.stack( + [m21 - m12, q_abs[..., 1] ** 2, m10 + m01, m02 + m20], + dim=-1, + ), + torch.stack( + [m02 - m20, m10 + m01, q_abs[..., 2] ** 2, m12 + m21], + dim=-1, + ), + torch.stack( + [m10 - m01, m20 + m02, m21 + m12, q_abs[..., 3] ** 2], + dim=-1, + ), + ], + dim=-2, + ) + floor = torch.tensor(0.1, dtype=q_abs.dtype, device=q_abs.device) + candidates = quat_by_rijk / (2.0 * q_abs[..., None].max(floor)) + selected = candidates[ + F.one_hot(q_abs.argmax(dim=-1), num_classes=4) > 0.5, : + ].reshape(*batch_shape, 4) + selected = selected[..., [1, 2, 3, 0]] + return torch.where(selected[..., 3:4] < 0, -selected, selected) + + +def _nominal_intrinsics( + batch: int, + views: int, + *, + device: torch.device, +) -> torch.Tensor: + matrices = torch.zeros(batch, views, 3, 3, device=device, dtype=torch.float32) + matrices[..., 0, 0] = float(RAY_WIDTH) + matrices[..., 1, 1] = float(RAY_WIDTH) + matrices[..., 0, 2] = RAY_WIDTH / 2.0 + matrices[..., 1, 2] = RAY_HEIGHT / 2.0 + matrices[..., 2, 2] = 1.0 + return matrices + + +def _compute_plucker_rays( + target_c2w: torch.Tensor, + target_intrinsics: torch.Tensor, +) -> torch.Tensor: + batch, views = target_c2w.shape[:2] + device = target_c2w.device + pixel_x = torch.linspace( + 0.5, RAY_WIDTH - 0.5, RAY_WIDTH, device=device, dtype=torch.float32 + ) + pixel_y = torch.linspace( + 0.5, RAY_HEIGHT - 0.5, RAY_HEIGHT, device=device, dtype=torch.float32 + ) + grid_y, grid_x = torch.meshgrid(pixel_y, pixel_x, indexing="ij") + uv = torch.stack([grid_x, grid_y, torch.ones_like(grid_x)], dim=-1) + inverse_intrinsics = torch.linalg.inv(target_intrinsics).float() + directions_local = torch.einsum("bvij,hwj->bvhwi", inverse_intrinsics, uv) + directions_local = directions_local / torch.linalg.vector_norm( + directions_local, dim=-1, keepdim=True + ) + directions_global = torch.einsum( + "bvij,bvhwj->bvhwi", target_c2w[..., :3, :3].float(), directions_local + ) + ray_origin = target_c2w[..., :3, 3].float()[:, :, None, None, :].expand( + batch, views, RAY_HEIGHT, RAY_WIDTH, 3 + ) + moment = torch.cross(ray_origin, directions_global, dim=-1) + return ( + torch.cat([moment, directions_global], dim=-1) + .permute(0, 1, 4, 2, 3) + .contiguous() + ) + + +def prepare_repencoder_camera_conditioning( + source_c2w_metric: torch.Tensor, + target_c2w_metric: torch.Tensor, + *, + near_zero_baseline_m: float = DEFAULT_NEAR_ZERO_BASELINE_M, + device: torch.device | str | None = None, +) -> RepEncoderCameraConditioning: + target_device = ( + source_c2w_metric.device if device is None else torch.device(device) + ) + with torch.autocast(device_type=target_device.type, enabled=False): + source_metric = source_c2w_metric.to(device=target_device, dtype=torch.float32) + target_metric = target_c2w_metric.to(device=target_device, dtype=torch.float32) + source, target, scale_tokens, scene_scale = normalize_source_target_cameras( + source_metric, + target_metric, + near_zero_baseline_m=near_zero_baseline_m, + ) + batch = source.shape[0] + intrinsics = _nominal_intrinsics( + batch, + NUM_SOURCE_VIEWS + NUM_TARGET_VIEWS, + device=target_device, + ) + quaternion = _mat_to_quat_scalar_last(source[..., :3, :3]) + fy = intrinsics[:, :NUM_SOURCE_VIEWS, 1, 1] + fx = intrinsics[:, :NUM_SOURCE_VIEWS, 0, 0] + fov_h = 2.0 * torch.atan((RAY_HEIGHT / 2.0) / fy) + fov_w = 2.0 * torch.atan((RAY_WIDTH / 2.0) / fx) + pose_encoding = torch.cat( + [source[..., :3, 3], quaternion, fov_h[..., None], fov_w[..., None]], + dim=-1, + ).float() + camera_tokens = torch.cat( + [ + pose_encoding, + scale_tokens[:, None, :].expand(batch, NUM_SOURCE_VIEWS, 2), + ], + dim=-1, + ).contiguous() + target_rays = _compute_plucker_rays( + target, + intrinsics[:, NUM_SOURCE_VIEWS:], + ) + if camera_tokens.shape != (batch, NUM_SOURCE_VIEWS, 11): + raise RuntimeError(f"camera token shape is invalid: {camera_tokens.shape}") + if target_rays.shape != (batch, NUM_TARGET_VIEWS, 6, RAY_HEIGHT, RAY_WIDTH): + raise RuntimeError(f"target ray shape is invalid: {target_rays.shape}") + if not torch.isfinite(camera_tokens).all() or not torch.isfinite(target_rays).all(): + raise FloatingPointError("RepEncoder camera conditioning contains NaN or Inf") + return RepEncoderCameraConditioning( + source_camera_tokens=camera_tokens, + target_rays=target_rays, + source_c2w_normalized=source, + target_c2w_normalized=target, + scale_tokens=scale_tokens, + scene_scale_m=scene_scale, + ) + + +__all__ = [ + "CAMERA_SCALE_MULTIPLIER", + "DEFAULT_NEAR_ZERO_BASELINE_M", + "RepEncoderCameraConditioning", + "NUM_SOURCE_VIEWS", + "NUM_TARGET_VIEWS", + "RAY_HW", + "normalize_source_target_cameras", + "prepare_repencoder_camera_conditioning", +] diff --git a/worldcrafter/repencoder/checkpoint.py b/worldcrafter/repencoder/checkpoint.py new file mode 100644 index 0000000000000000000000000000000000000000..ba0516210a0bbfac587f9a24a3f37baef73b7196 --- /dev/null +++ b/worldcrafter/repencoder/checkpoint.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import hashlib +import json +from pathlib import Path +from typing import Any, Mapping + +import torch +import torch.nn as nn +from safetensors.torch import load_file + +from .config import RepEncoderConfig + + +CHECKPOINT_FILENAME = "model.safetensors" +CONFIG_FILENAME = "config.json" +MANIFEST_FILENAME = "manifest.json" + + +def sha256_file(path: str | Path) -> str: + source = Path(path) + digest = hashlib.sha256() + with source.open("rb") as handle: + for chunk in iter(lambda: handle.read(16 * 1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def resolve_checkpoint(path: str | Path) -> tuple[Path, Path, Path]: + source = Path(path).expanduser().resolve() + root = source if source.is_dir() else source.parent + model_path = root / CHECKPOINT_FILENAME if source.is_dir() else source + config_path = root / CONFIG_FILENAME + manifest_path = root / MANIFEST_FILENAME + for required in (model_path, config_path, manifest_path): + if not required.is_file(): + raise FileNotFoundError(f"RepEncoder checkpoint artifact is missing: {required}") + return model_path, config_path, manifest_path + + +def read_checkpoint_metadata( + path: str | Path, + *, + expected_sha256: str | None = None, +) -> tuple[Path, RepEncoderConfig, dict[str, Any], str]: + model_path, config_path, manifest_path = resolve_checkpoint(path) + config_payload = json.loads(config_path.read_text()) + manifest = json.loads(manifest_path.read_text()) + if not isinstance(manifest, Mapping) or manifest.get("format") != "worldcrafter_repencoder_v1": + raise ValueError("unsupported RepEncoder manifest") + actual_sha256 = sha256_file(model_path) + declared_sha256 = str(manifest.get("model_sha256", "")) + if declared_sha256 != actual_sha256: + raise RuntimeError( + f"RepEncoder manifest SHA256 mismatch: {declared_sha256} != {actual_sha256}" + ) + if expected_sha256 is not None and str(expected_sha256) != actual_sha256: + raise RuntimeError( + f"RepEncoder checkpoint SHA256 mismatch: {expected_sha256} != {actual_sha256}" + ) + return model_path, RepEncoderConfig.from_dict(config_payload), dict(manifest), actual_sha256 + + +def load_checkpoint_state( + model: nn.Module, + model_path: str | Path, + *, + device: torch.device | str = "cpu", +) -> None: + state = load_file(str(model_path), device=str(torch.device(device))) + expected = model.state_dict() + if set(state) != set(expected): + raise RuntimeError( + "RepEncoder checkpoint key mismatch: " + f"missing={sorted(set(expected) - set(state))[:8]}, " + f"unexpected={sorted(set(state) - set(expected))[:8]}" + ) + dtype_mismatch = { + name: (value.dtype, expected[name].dtype) + for name, value in state.items() + if value.dtype != expected[name].dtype + } + if dtype_mismatch: + raise TypeError( + "RepEncoder checkpoint violates mixed dtype ownership: " + f"{list(dtype_mismatch.items())[:8]}" + ) + nonfinite = [name for name, value in state.items() if not torch.isfinite(value).all()] + if nonfinite: + raise FloatingPointError(f"RepEncoder checkpoint contains non-finite tensors: {nonfinite[:8]}") + incompatible = model.load_state_dict(state, strict=True) + if incompatible.missing_keys or incompatible.unexpected_keys: + raise RuntimeError(f"strict RepEncoder checkpoint load failed: {incompatible}") + + +__all__ = [ + "CHECKPOINT_FILENAME", + "CONFIG_FILENAME", + "MANIFEST_FILENAME", + "load_checkpoint_state", + "read_checkpoint_metadata", + "resolve_checkpoint", + "sha256_file", +] diff --git a/worldcrafter/repencoder/config.py b/worldcrafter/repencoder/config.py new file mode 100644 index 0000000000000000000000000000000000000000..59cbd89df9566eba397555142d9c0b7e9e9146fd --- /dev/null +++ b/worldcrafter/repencoder/config.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +from dataclasses import asdict, dataclass +from typing import Any, Mapping + + +@dataclass(frozen=True) +class RepEncoderConfig: + format: str = "worldcrafter_repencoder_config_v1" + source_views: int = 9 + target_views: int = 4 + latent_channels: int = 16 + latent_height: int = 48 + latent_width: int = 80 + dino_dim: int = 1024 + dino_input_block: int = 2 + dino_depth: int = 24 + dino_heads: int = 16 + latent_special_tokens: int = 5 + patch_height: int = 22 + patch_width: int = 37 + vggt_depth: int = 24 + vggt_heads: int = 16 + vggt_register_tokens: int = 4 + scene_dim: int = 768 + repfeature_depth: int = 10 + repfeature_heads: int = 12 + target_patch_size: int = 8 + ray_height: int = 384 + ray_width: int = 640 + output_kernel: tuple[int, int, int] = (3, 3, 3) + target_slots: tuple[int, int, int, int] = (2, 4, 6, 8) + + def __post_init__(self) -> None: + if ( + self.format != "worldcrafter_repencoder_config_v1" + or self.source_views != 9 + or self.target_views != 4 + or self.latent_channels != 16 + or (self.latent_height, self.latent_width) != (48, 80) + or self.dino_dim != 1024 + or (self.dino_input_block, self.dino_depth, self.dino_heads) != (2, 24, 16) + or self.latent_special_tokens != 5 + or (self.patch_height, self.patch_width) != (22, 37) + or (self.vggt_depth, self.vggt_heads, self.vggt_register_tokens) + != (24, 16, 4) + or self.scene_dim != 768 + or (self.repfeature_depth, self.repfeature_heads) != (10, 12) + or self.target_patch_size != 8 + or (self.ray_height, self.ray_width) != (384, 640) + or tuple(self.output_kernel) != (3, 3, 3) + or tuple(self.target_slots) != (2, 4, 6, 8) + ): + raise ValueError("RepEncoder v1 has a frozen architecture") + + @property + def patch_tokens(self) -> int: + return int(self.patch_height * self.patch_width) + + def to_dict(self) -> dict[str, Any]: + value = asdict(self) + value["output_kernel"] = list(self.output_kernel) + value["target_slots"] = list(self.target_slots) + return value + + @classmethod + def from_dict(cls, value: Mapping[str, Any]) -> "RepEncoderConfig": + payload = dict(value) + if "output_kernel" in payload: + payload["output_kernel"] = tuple(int(x) for x in payload["output_kernel"]) + if "target_slots" in payload: + payload["target_slots"] = tuple(int(x) for x in payload["target_slots"]) + config = cls(**payload) + if config.to_dict() != cls().to_dict(): + raise ValueError("RepEncoder config does not match the frozen v1 contract") + return config + + +__all__ = ["RepEncoderConfig"] diff --git a/worldcrafter/repencoder/input_layer.py b/worldcrafter/repencoder/input_layer.py new file mode 100644 index 0000000000000000000000000000000000000000..b4101514572552d26aa67d2ba8e85c2b58fb25ea --- /dev/null +++ b/worldcrafter/repencoder/input_layer.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +from collections.abc import Sequence + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .config import RepEncoderConfig + + +class InputLayer(nn.Module): + def __init__( + self, + *, + latents_mean: Sequence[float] | None = None, + latents_std: Sequence[float] | None = None, + config: RepEncoderConfig | None = None, + ) -> None: + super().__init__() + self.config = RepEncoderConfig() if config is None else config + mean = torch.zeros(self.config.latent_channels, dtype=torch.float32) + std = torch.ones(self.config.latent_channels, dtype=torch.float32) + if latents_mean is not None: + mean = torch.as_tensor(latents_mean, dtype=torch.float32).reshape(-1) + if latents_std is not None: + std = torch.as_tensor(latents_std, dtype=torch.float32).reshape(-1) + expected = (self.config.latent_channels,) + if tuple(mean.shape) != expected or tuple(std.shape) != expected: + raise ValueError("latents_mean and latents_std must each contain 16 values") + if not torch.isfinite(mean).all() or not torch.isfinite(std).all() or (std <= 0).any(): + raise ValueError("latent statistics must be finite and std must be positive") + self.register_buffer("latents_mean", mean, persistent=True) + self.register_buffer("latents_std", std, persistent=True) + self.proj = nn.Conv2d( + self.config.latent_channels, + self.config.dino_dim, + kernel_size=3, + stride=2, + padding=0, + bias=True, + ) + + def forward(self, source_latents: torch.Tensor) -> torch.Tensor: + expected_tail = ( + self.config.source_views, + self.config.latent_channels, + self.config.latent_height, + self.config.latent_width, + ) + if source_latents.ndim != 5 or tuple(source_latents.shape[1:]) != expected_tail: + raise ValueError( + f"source_latents must be [B,{','.join(map(str, expected_tail))}], " + f"got {tuple(source_latents.shape)}" + ) + if not torch.is_floating_point(source_latents): + raise TypeError("source_latents must be floating point") + if not torch.isfinite(source_latents).all(): + raise FloatingPointError("source_latents contain NaN or Inf") + batch = source_latents.shape[0] + flat = source_latents.reshape( + batch * self.config.source_views, + self.config.latent_channels, + self.config.latent_height, + self.config.latent_width, + ) + mean = self.latents_mean.to(device=flat.device, dtype=flat.dtype) + std = self.latents_std.to(device=flat.device, dtype=flat.dtype) + raw = flat * std[None, :, None, None] + mean[None, :, None, None] + raw = F.interpolate(raw, size=(44, 74), mode="bilinear", align_corners=True) + prediction = self.proj(F.pad(raw, (1, 1, 1, 1), mode="replicate")) + expected = ( + batch * self.config.source_views, + self.config.dino_dim, + self.config.patch_height, + self.config.patch_width, + ) + if tuple(prediction.shape) != expected: + raise RuntimeError(f"input_layer output {tuple(prediction.shape)} != {expected}") + return prediction.permute(0, 2, 3, 1).reshape( + batch, + self.config.source_views, + self.config.patch_tokens, + self.config.dino_dim, + ) + + +__all__ = ["InputLayer"] diff --git a/worldcrafter/repencoder/low_rank.py b/worldcrafter/repencoder/low_rank.py new file mode 100644 index 0000000000000000000000000000000000000000..75591bea92e57ef86ab1dbc036dd1b2e54709cf7 --- /dev/null +++ b/worldcrafter/repencoder/low_rank.py @@ -0,0 +1,139 @@ +from __future__ import annotations + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +LOW_RANK_DIM = 64 +LOW_RANK_SCALE = 32.0 / LOW_RANK_DIM +EXPECTED_LOW_RANK_MODULES = 438 + + +class LowRankLinear(nn.Linear): + def __init__(self, source: nn.Linear) -> None: + super().__init__( + source.in_features, + source.out_features, + bias=source.bias is not None, + device=source.weight.device, + dtype=source.weight.dtype, + ) + self.low_rank_a = nn.Parameter( + torch.empty(LOW_RANK_DIM, source.in_features, dtype=torch.float32), + requires_grad=False, + ) + self.low_rank_b = nn.Parameter( + torch.empty(source.out_features, LOW_RANK_DIM, dtype=torch.float32), + requires_grad=False, + ) + with torch.no_grad(): + self.weight.copy_(source.weight) + if self.bias is not None: + self.bias.copy_(source.bias) + self.weight.requires_grad_(False) + if self.bias is not None: + self.bias.requires_grad_(False) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + base = F.linear(value, self.weight, self.bias) + update = F.linear(F.linear(value, self.low_rank_a), self.low_rank_b) + return base + update * LOW_RANK_SCALE + + +class LowRankConv2d(nn.Conv2d): + def __init__(self, source: nn.Conv2d) -> None: + if source.groups != 1 or source.kernel_size[0] != source.kernel_size[1]: + raise ValueError("Only square, ungrouped Conv2d layers are supported") + super().__init__( + source.in_channels, + source.out_channels, + source.kernel_size, + stride=source.stride, + padding=source.padding, + dilation=source.dilation, + groups=source.groups, + bias=source.bias is not None, + padding_mode=source.padding_mode, + device=source.weight.device, + dtype=source.weight.dtype, + ) + kernel = source.kernel_size[0] + expanded_rank = LOW_RANK_DIM * kernel + self.low_rank_a = nn.Parameter( + torch.empty(expanded_rank, source.in_channels * kernel, dtype=torch.float32), + requires_grad=False, + ) + self.low_rank_b = nn.Parameter( + torch.empty(source.out_channels * kernel, expanded_rank, dtype=torch.float32), + requires_grad=False, + ) + with torch.no_grad(): + self.weight.copy_(source.weight) + if self.bias is not None: + self.bias.copy_(source.bias) + self.weight.requires_grad_(False) + if self.bias is not None: + self.bias.requires_grad_(False) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + delta = (self.low_rank_b @ self.low_rank_a).reshape_as(self.weight) * LOW_RANK_SCALE + base = self._conv_forward(value, self.weight, self.bias) + update = self._conv_forward(value, delta, None) + return base + update + + +def _uses_low_rank_branch(name: str) -> bool: + if name.startswith("dino_tail.blocks."): + return True + if name.startswith(("vggt.frame_blocks.", "vggt.global_blocks.")): + return True + if name.startswith("vggt.camera_mlp.") or name == "vggt.geo_feature_connector": + return True + if name == "repfeature.target_embedding": + return True + if name.startswith("repfeature.blocks."): + return True + if name.startswith("repfeature.final_block."): + return True + return False + + +def _parent_and_child(model: nn.Module, dotted_name: str) -> tuple[nn.Module, str]: + parts = dotted_name.split(".") + parent = model + for part in parts[:-1]: + parent = getattr(parent, part) + return parent, parts[-1] + + +def inject_low_rank_branches(model: nn.Module) -> int: + count = 0 + for module_name, module in list(model.named_modules()): + if not module_name or not isinstance(module, (nn.Linear, nn.Conv2d)): + continue + if not _uses_low_rank_branch(module_name): + continue + parent, child = _parent_and_child(model, module_name) + replacement: nn.Module + if isinstance(module, nn.Linear): + replacement = LowRankLinear(module) + else: + replacement = LowRankConv2d(module) + setattr(parent, child, replacement) + count += 1 + if count != EXPECTED_LOW_RANK_MODULES: + raise RuntimeError( + f"RepEncoder low-rank graph changed: {count} != {EXPECTED_LOW_RANK_MODULES}" + ) + return count + + +__all__ = [ + "EXPECTED_LOW_RANK_MODULES", + "LOW_RANK_DIM", + "LOW_RANK_SCALE", + "LowRankConv2d", + "LowRankLinear", + "inject_low_rank_branches", +] diff --git a/worldcrafter/repencoder/model.py b/worldcrafter/repencoder/model.py new file mode 100644 index 0000000000000000000000000000000000000000..dc9abee8924bc3bb71c8cd165a38b1c4bcbfadaf --- /dev/null +++ b/worldcrafter/repencoder/model.py @@ -0,0 +1,209 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import torch +import torch.nn as nn + +from .camera import ( + DEFAULT_NEAR_ZERO_BASELINE_M, + RepEncoderCameraConditioning, + prepare_repencoder_camera_conditioning, +) +from .checkpoint import load_checkpoint_state, read_checkpoint_metadata +from .config import RepEncoderConfig +from .input_layer import InputLayer +from .low_rank import LowRankConv2d, LowRankLinear, inject_low_rank_branches +from .output_layer import OutputLayer +from .repfeature import RepFeature +from .vggt import DinoTail, RepresentationBackbone + + +class RepEncoder(nn.Module): + def __init__( + self, + config: RepEncoderConfig | None = None, + *, + compute_dtype: str | torch.dtype = torch.bfloat16, + target_microbatch: int = 4, + near_zero_baseline_m: float = DEFAULT_NEAR_ZERO_BASELINE_M, + ) -> None: + super().__init__() + self.config = RepEncoderConfig() if config is None else config + self.input_layer = InputLayer(config=self.config) + self.dino_tail = DinoTail(config=self.config) + self.vggt = RepresentationBackbone(config=self.config) + self.repfeature = RepFeature(config=self.config) + self.output_layer = OutputLayer(config=self.config) + self.low_rank_branch_count = inject_low_rank_branches(self) + self._apply_checkpoint_dtype_ownership() + self.target_microbatch = int(target_microbatch) + if self.target_microbatch <= 0: + raise ValueError("target_microbatch must be positive") + self.near_zero_baseline_m = float(near_zero_baseline_m) + if self.near_zero_baseline_m < 0: + raise ValueError("near_zero_baseline_m must be non-negative") + self.compute_dtype = self._resolve_compute_dtype(compute_dtype) + self.report: dict[str, Any] = {} + + def _apply_checkpoint_dtype_ownership(self) -> None: + self.dino_tail.to(dtype=torch.bfloat16) + self.vggt.to(dtype=torch.bfloat16) + self.repfeature.to(dtype=torch.bfloat16) + self.dino_tail.latent_cls_token.data = ( + self.dino_tail.latent_cls_token.data.float() + ) + self.dino_tail.latent_register_tokens.data = ( + self.dino_tail.latent_register_tokens.data.float() + ) + for module in self.modules(): + if isinstance(module, (LowRankLinear, LowRankConv2d)): + module.low_rank_a.data = module.low_rank_a.data.float() + module.low_rank_b.data = module.low_rank_b.data.float() + self.input_layer.float() + self.output_layer.float() + + @staticmethod + def _resolve_compute_dtype(value: str | torch.dtype) -> torch.dtype: + dtype = ( + value + if isinstance(value, torch.dtype) + else { + "bfloat16": torch.bfloat16, + "bf16": torch.bfloat16, + }.get(str(value).lower()) + ) + if dtype != torch.bfloat16: + raise ValueError( + "RepEncoder v2 has one exact inference contract: bfloat16 autocast " + "with checkpoint-owned mixed parameter dtypes" + ) + return torch.bfloat16 + + @property + def device(self) -> torch.device: + return next(self.parameters()).device + + @classmethod + def from_pretrained( + cls, + path: str | Path, + *, + expected_sha256: str | None = None, + device: torch.device | str = "cpu", + compute_dtype: str | torch.dtype = torch.bfloat16, + target_microbatch: int = 4, + near_zero_baseline_m: float = DEFAULT_NEAR_ZERO_BASELINE_M, + ) -> "RepEncoder": + model_path, config, manifest, actual_sha256 = read_checkpoint_metadata( + path, expected_sha256=expected_sha256 + ) + model = cls( + config, + compute_dtype=compute_dtype, + target_microbatch=target_microbatch, + near_zero_baseline_m=near_zero_baseline_m, + ) + load_checkpoint_state(model, model_path, device="cpu") + model.requires_grad_(False) + model.eval() + model.to(device=torch.device(device)) + model.report = { + "format": "worldcrafter_repencoder_runtime_v1", + "model_path": str(model_path), + "model_sha256": actual_sha256, + "manifest": manifest, + "compute_dtype": str(model.compute_dtype).removeprefix("torch."), + } + return model + + def prepare_camera_conditioning( + self, + source_c2w_metric: torch.Tensor, + target_c2w_metric: torch.Tensor, + *, + near_zero_baseline_m: float | None = None, + ) -> RepEncoderCameraConditioning: + epsilon = ( + self.near_zero_baseline_m + if near_zero_baseline_m is None + else float(near_zero_baseline_m) + ) + return prepare_repencoder_camera_conditioning( + source_c2w_metric, + target_c2w_metric, + near_zero_baseline_m=epsilon, + device=self.device, + ) + + def forward_conditioned( + self, + source_latents: torch.Tensor, + source_camera_tokens: torch.Tensor, + target_rays: torch.Tensor, + ) -> torch.Tensor: + expected_source = ( + source_latents.shape[0], + self.config.source_views, + self.config.latent_channels, + self.config.latent_height, + self.config.latent_width, + ) + if tuple(source_latents.shape) != expected_source: + raise ValueError(f"source_latents must be {expected_source}") + batch = source_latents.shape[0] + if tuple(source_camera_tokens.shape) != (batch, self.config.source_views, 11): + raise ValueError("source_camera_tokens must be [B,9,11]") + if tuple(target_rays.shape) != ( + batch, + self.config.target_views, + 6, + self.config.ray_height, + self.config.ray_width, + ): + raise ValueError("target_rays must be [B,4,6,384,640]") + if {source_latents.device, source_camera_tokens.device, target_rays.device} != { + self.device + }: + raise ValueError("all RepEncoder inputs must be on the model device") + + if self.device.type != "cuda": + raise RuntimeError( + "RepEncoder exact inference requires CUDA BF16 autocast; " + "CPU forward is not a supported runtime" + ) + autocast_context = torch.autocast(device_type="cuda", dtype=torch.bfloat16) + + with autocast_context: + input_features = self.input_layer(source_latents) + dino_features = self.dino_tail(input_features) + scene = self.vggt(dino_features, source_camera_tokens) + target_features = self.repfeature( + scene, + target_rays, + target_microbatch=self.target_microbatch, + ) + memory4 = self.output_layer(target_features) + return memory4 + + def forward( + self, + source_latents: torch.Tensor, + source_c2w_metric: torch.Tensor, + target_c2w_metric: torch.Tensor, + near_zero_baseline_m: float | None = None, + ) -> torch.Tensor: + conditioning = self.prepare_camera_conditioning( + source_c2w_metric, + target_c2w_metric, + near_zero_baseline_m=near_zero_baseline_m, + ) + return self.forward_conditioned( + source_latents, + conditioning.source_camera_tokens, + conditioning.target_rays, + ) + + +__all__ = ["RepEncoder"] diff --git a/worldcrafter/repencoder/output_layer.py b/worldcrafter/repencoder/output_layer.py new file mode 100644 index 0000000000000000000000000000000000000000..eaa6dcf49b0e9418bf2fa0371dce34c715febaf0 --- /dev/null +++ b/worldcrafter/repencoder/output_layer.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .config import RepEncoderConfig + + +class OutputLayer(nn.Module): + def __init__(self, config: RepEncoderConfig | None = None) -> None: + super().__init__() + self.config = RepEncoderConfig() if config is None else config + self.proj = nn.Conv3d( + self.config.scene_dim, + self.config.latent_channels, + kernel_size=self.config.output_kernel, + stride=1, + padding=0, + bias=True, + ) + + def forward(self, features: torch.Tensor) -> torch.Tensor: + expected = ( + self.config.scene_dim, + self.config.target_views, + self.config.latent_height, + self.config.latent_width, + ) + if features.ndim != 5 or tuple(features.shape[1:]) != expected: + raise ValueError(f"features must be [B,{expected}], got {tuple(features.shape)}") + temporal, height, width = self.config.output_kernel + padded = F.pad( + features, + (width // 2, width // 2, height // 2, height // 2, temporal // 2, temporal // 2), + mode="replicate", + ) + latent = self.proj(padded) + expected_output = ( + features.shape[0], + self.config.latent_channels, + self.config.target_views, + self.config.latent_height, + self.config.latent_width, + ) + if tuple(latent.shape) != expected_output: + raise RuntimeError(f"memory4 shape {tuple(latent.shape)} != {expected_output}") + return latent + + +__all__ = ["OutputLayer"] diff --git a/worldcrafter/repencoder/repfeature/__init__.py b/worldcrafter/repencoder/repfeature/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..c9f67891d296c963de1301b33912ca8c14e03cf0 --- /dev/null +++ b/worldcrafter/repencoder/repfeature/__init__.py @@ -0,0 +1,3 @@ +from .model import RepFeature + +__all__ = ["RepFeature"] diff --git a/worldcrafter/repencoder/repfeature/attention.py b/worldcrafter/repencoder/repfeature/attention.py new file mode 100644 index 0000000000000000000000000000000000000000..13e902aefa987ac302d7321daeabc2e52929d0ad --- /dev/null +++ b/worldcrafter/repencoder/repfeature/attention.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import torch +import torch.nn as nn +import torch.nn.functional as F + +try: + import xformers.ops as xops +except ImportError: # pragma: no cover - exercised only in minimal CPU installs + xops = None + + +class RMSNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-5) -> None: + super().__init__() + self.eps = float(eps) + self.weight = nn.Parameter(torch.ones(dim)) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + normalized = value.float() * torch.rsqrt( + value.float().square().mean(-1, keepdim=True) + self.eps + ) + return normalized.to(value.dtype) * self.weight.to(value.dtype) + + +class Attention(nn.Module): + def __init__(self, dim: int = 768, num_heads: int = 12) -> None: + super().__init__() + if dim % num_heads: + raise ValueError("attention width must be divisible by num_heads") + self.num_heads = int(num_heads) + self.head_dim = int(dim // num_heads) + self.q_proj = nn.Linear(dim, dim, bias=False) + self.k_proj = nn.Linear(dim, dim, bias=False) + self.v_proj = nn.Linear(dim, dim, bias=False) + self.proj = nn.Linear(dim, dim, bias=False) + self.q_norm = RMSNorm(self.head_dim) + self.k_norm = RMSNorm(self.head_dim) + + def forward(self, query: torch.Tensor, kv: torch.Tensor | None = None) -> torch.Tensor: + if kv is None: + kv = query + batch, query_tokens, channels = query.shape + key_tokens = kv.shape[1] + q = self.q_proj(query).reshape( + batch, query_tokens, self.num_heads, self.head_dim + ) + k = self.k_proj(kv).reshape(batch, key_tokens, self.num_heads, self.head_dim) + v = self.v_proj(kv).reshape(batch, key_tokens, self.num_heads, self.head_dim) + q = self.q_norm(q) + k = self.k_norm(k) + if xops is not None and query.is_cuda: + value = xops.memory_efficient_attention(q, k, v, p=0.0, op=None) + else: + value = F.scaled_dot_product_attention( + q.transpose(1, 2), + k.transpose(1, 2), + v.transpose(1, 2), + dropout_p=0.0, + ).transpose(1, 2) + value = value.reshape(batch, query_tokens, channels) + return self.proj(value) + + +__all__ = ["Attention", "RMSNorm"] diff --git a/worldcrafter/repencoder/repfeature/blocks.py b/worldcrafter/repencoder/repfeature/blocks.py new file mode 100644 index 0000000000000000000000000000000000000000..bb7886eecc7424632e24a0a2ed82c8d5d3d6c06d --- /dev/null +++ b/worldcrafter/repencoder/repfeature/blocks.py @@ -0,0 +1,68 @@ +from __future__ import annotations + +import torch +import torch.nn as nn + +from .attention import Attention + + +class Mlp(nn.Module): + def __init__(self, width: int = 768, ratio: int = 4) -> None: + super().__init__() + self.fc1 = nn.Linear(width, width * ratio, bias=False) + self.act = nn.GELU() + self.fc2 = nn.Linear(width * ratio, width, bias=False) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + return self.fc2(self.act(self.fc1(value))) + + +class BidirectionalBlock(nn.Module): + def __init__(self, width: int = 768, heads: int = 12) -> None: + super().__init__() + self.norm1_x = nn.LayerNorm(width, bias=False) + self.self_attn = Attention(width, heads) + self.norm2_x = nn.LayerNorm(width, bias=False) + self.cross_attn_x = Attention(width, heads) + self.norm3_x = nn.LayerNorm(width, bias=False) + self.mlp_x = Mlp(width) + self.norm1_rec = nn.LayerNorm(width, bias=False) + self.cross_attn_rec = Attention(width, heads) + self.norm2_rec = nn.LayerNorm(width, bias=False) + self.mlp_rec = Mlp(width) + + def forward( + self, target: torch.Tensor, representation: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + target = target + self.self_attn(self.norm1_x(target)) + target_norm = self.norm2_x(target) + representation_norm = self.norm1_rec(representation) + target = target + self.cross_attn_x(target_norm, kv=representation_norm) + representation = representation + self.cross_attn_rec( + representation_norm, kv=target_norm + ) + target = target + self.mlp_x(self.norm3_x(target)) + representation = representation + self.mlp_rec(self.norm2_rec(representation)) + return target, representation + + +class FinalTargetBlock(nn.Module): + def __init__(self, width: int = 768, heads: int = 12) -> None: + super().__init__() + self.norm1_x = nn.LayerNorm(width, bias=False) + self.self_attn = Attention(width, heads) + self.norm2_x = nn.LayerNorm(width, bias=False) + self.cross_attn_x = Attention(width, heads) + self.norm3_x = nn.LayerNorm(width, bias=False) + self.mlp_x = Mlp(width) + self.norm1_rec = nn.LayerNorm(width, bias=False) + + def forward(self, target: torch.Tensor, representation: torch.Tensor) -> torch.Tensor: + target = target + self.self_attn(self.norm1_x(target)) + target_norm = self.norm2_x(target) + representation_norm = self.norm1_rec(representation) + target = target + self.cross_attn_x(target_norm, kv=representation_norm) + return target + self.mlp_x(self.norm3_x(target)) + + +__all__ = ["BidirectionalBlock", "FinalTargetBlock"] diff --git a/worldcrafter/repencoder/repfeature/model.py b/worldcrafter/repencoder/repfeature/model.py new file mode 100644 index 0000000000000000000000000000000000000000..9a3958177eced8e67c1d22008b150f336ecf2d43 --- /dev/null +++ b/worldcrafter/repencoder/repfeature/model.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +import torch +import torch.nn as nn + +from ..config import RepEncoderConfig +from .blocks import BidirectionalBlock, FinalTargetBlock + + +class RepFeature(nn.Module): + def __init__(self, config: RepEncoderConfig | None = None) -> None: + super().__init__() + self.config = RepEncoderConfig() if config is None else config + self.target_embedding = nn.Conv2d( + 6, + self.config.scene_dim, + kernel_size=self.config.target_patch_size, + stride=self.config.target_patch_size, + bias=False, + ) + self.target_norm = nn.LayerNorm(self.config.scene_dim, bias=False) + self.target_register_tokens = nn.Parameter( + torch.zeros(1, 4, self.config.scene_dim) + ) + self.blocks = nn.ModuleList( + [ + BidirectionalBlock( + width=self.config.scene_dim, + heads=self.config.repfeature_heads, + ) + for _ in range(self.config.repfeature_depth - 1) + ] + ) + self.final_block = FinalTargetBlock( + width=self.config.scene_dim, + heads=self.config.repfeature_heads, + ) + + def forward( + self, + scene: torch.Tensor, + target_rays: torch.Tensor, + *, + target_microbatch: int = 4, + ) -> torch.Tensor: + if scene.ndim != 4 or tuple(scene.shape[1:]) != ( + self.config.source_views, + self.config.patch_tokens, + self.config.scene_dim, + ): + raise ValueError("scene must be [B,9,814,768]") + if target_rays.ndim != 5 or tuple(target_rays.shape[1:]) != ( + self.config.target_views, + 6, + self.config.ray_height, + self.config.ray_width, + ): + raise ValueError("target_rays must be [B,4,6,384,640]") + if scene.shape[0] != target_rays.shape[0]: + raise ValueError("scene and target_rays batch sizes differ") + target_microbatch = int(target_microbatch) + if target_microbatch <= 0: + raise ValueError("target_microbatch must be positive") + + batch = scene.shape[0] + flat_scene = scene.reshape( + batch, + self.config.source_views * self.config.patch_tokens, + self.config.scene_dim, + ) + pieces: list[torch.Tensor] = [] + for start in range(0, self.config.target_views, target_microbatch): + rays = target_rays[:, start : start + target_microbatch] + views = rays.shape[1] + flat_rays = rays.reshape( + batch * views, 6, self.config.ray_height, self.config.ray_width + ) + target = self.target_embedding(flat_rays).flatten(2).transpose(1, 2) + target = self.target_norm(target) + registers = self.target_register_tokens.expand(batch * views, -1, -1).to( + device=target.device, dtype=target.dtype + ) + target = torch.cat((registers, target), dim=1) + representation = ( + flat_scene[:, None] + .expand(batch, views, *flat_scene.shape[1:]) + .reshape(batch * views, *flat_scene.shape[1:]) + ) + for block in self.blocks: + target, representation = block(target, representation) + target = self.final_block(target, representation) + pieces.append( + target[:, 4:].reshape( + batch, + views, + self.config.latent_height * self.config.latent_width, + -1, + ) + ) + + joined = torch.cat(pieces, dim=1) + expected = ( + batch, + self.config.target_views, + self.config.latent_height * self.config.latent_width, + self.config.scene_dim, + ) + if tuple(joined.shape) != expected: + raise RuntimeError(f"repfeature output {tuple(joined.shape)} != {expected}") + return ( + joined.reshape( + batch, + self.config.target_views, + self.config.latent_height, + self.config.latent_width, + self.config.scene_dim, + ) + .permute(0, 4, 1, 2, 3) + .contiguous() + ) + + +__all__ = ["RepFeature"] diff --git a/worldcrafter/repencoder/trajectory_fov.py b/worldcrafter/repencoder/trajectory_fov.py new file mode 100644 index 0000000000000000000000000000000000000000..2a18442d53718427b9b89b56f6ff28509beb360c --- /dev/null +++ b/worldcrafter/repencoder/trajectory_fov.py @@ -0,0 +1,390 @@ +from __future__ import annotations + +from dataclasses import asdict, dataclass +import math +from typing import Sequence + +import torch + + +DEFAULT_HORIZONTAL_FOV_DEGREES = 100.0 +DEFAULT_VERTICAL_FOV_DEGREES = 71.13349068444832 +DEFAULT_NEAR_METERS = 0.1 +DEFAULT_FAR_METERS = 30.0 +FRUSTUM_SAMPLES_PER_AXIS = 10 + + +@dataclass(frozen=True) +class TrajectoryFovSelection: + selected_latents: tuple[int, ...] + selected_parent_chunks: tuple[int, ...] + selection_kinds: tuple[str, ...] + initial_fov_coverage_per_target: tuple[float, ...] + final_fov_coverage_per_target: tuple[float, ...] + initial_frontier_coverage_per_target: tuple[float, ...] + final_frontier_coverage_per_target: tuple[float, ...] + minimum_bridge_fraction: float + mean_bridge_fraction: float + disconnected_fill_count: int + step_diagnostics: tuple[dict[str, object], ...] + candidate_count: int + target_count: int + + def to_json(self) -> dict[str, object]: + return asdict(self) + + +def _validate_pose_batch(value: torch.Tensor, *, name: str) -> torch.Tensor: + pose = torch.as_tensor(value) + if pose.ndim != 3 or pose.shape[-2:] != (4, 4): + raise ValueError(f"{name} must be Nx4x4, got {tuple(pose.shape)}") + if not torch.is_floating_point(pose) or not torch.isfinite(pose).all(): + raise ValueError(f"{name} must contain finite floating-point poses") + return pose + + +def target_frustum_visibility_masks( + *, + target_c2w: torch.Tensor, + source_c2w: torch.Tensor, + horizontal_fov_degrees: float, + vertical_fov_degrees: float, + near_m: float, + far_m: float, + samples_per_axis: int = FRUSTUM_SAMPLES_PER_AXIS, +) -> torch.Tensor: + + targets = _validate_pose_batch(target_c2w, name="target_c2w") + sources = _validate_pose_batch(source_c2w, name="source_c2w").to( + device=targets.device, dtype=targets.dtype + ) + if int(samples_per_axis) != FRUSTUM_SAMPLES_PER_AXIS: + raise ValueError( + "trajectory FOV selection requires " + f"samples_per_axis={FRUSTUM_SAMPLES_PER_AXIS}" + ) + if not (math.isfinite(near_m) and math.isfinite(far_m) and 0.0 < near_m < far_m): + raise ValueError("near_m and far_m must be finite with 0 < near_m < far_m") + if not 0.0 < float(horizontal_fov_degrees) < 180.0: + raise ValueError("horizontal_fov_degrees must be in (0, 180)") + if not 0.0 < float(vertical_fov_degrees) < 180.0: + raise ValueError("vertical_fov_degrees must be in (0, 180)") + + device, dtype = targets.device, targets.dtype + z = torch.linspace(float(near_m), float(far_m), FRUSTUM_SAMPLES_PER_AXIS, device=device, dtype=dtype) + x = torch.linspace(-1.0, 1.0, FRUSTUM_SAMPLES_PER_AXIS, device=device, dtype=dtype) + y = torch.linspace(-1.0, 1.0, FRUSTUM_SAMPLES_PER_AXIS, device=device, dtype=dtype) + grid_x, grid_y, grid_z = torch.meshgrid(x, y, z, indexing="ij") + tan_h = math.tan(math.radians(float(horizontal_fov_degrees)) * 0.5) + tan_v = math.tan(math.radians(float(vertical_fov_degrees)) * 0.5) + target_points = torch.stack( + ( + grid_x.reshape(-1) * grid_z.reshape(-1) * tan_h, + grid_y.reshape(-1) * grid_z.reshape(-1) * tan_v, + grid_z.reshape(-1), + ), + dim=0, + ) + + points_world = ( + torch.bmm( + targets[:, :3, :3], + target_points.unsqueeze(0).expand(targets.shape[0], -1, -1), + ) + + targets[:, :3, 3:4] + ) + source_r_inv = sources[:, :3, :3].transpose(1, 2) + source_t_inv = -torch.bmm(source_r_inv, sources[:, :3, 3:4]) + points_in_source = torch.einsum("sij,tjp->stip", source_r_inv, points_world) + points_in_source = points_in_source + source_t_inv[:, None] + px = points_in_source[:, :, 0] + py = points_in_source[:, :, 1] + pz = points_in_source[:, :, 2] + yaw = torch.atan2(px, pz.clamp_min(1e-6)).abs() + pitch = torch.atan2( + py, torch.sqrt(px.square() + pz.square()).clamp_min(1e-6) + ).abs() + return ( + (pz >= float(near_m)) + & (pz <= float(far_m)) + & (yaw <= math.radians(float(horizontal_fov_degrees)) * 0.5) + & (pitch <= math.radians(float(vertical_fov_degrees)) * 0.5) + ) + + +def ordered_target_frontier_masks( + *, + anchor_c2w: torch.Tensor, + target_c2w: torch.Tensor, + horizontal_fov_degrees: float, + vertical_fov_degrees: float, + near_m: float, + far_m: float, + samples_per_axis: int = FRUSTUM_SAMPLES_PER_AXIS, +) -> tuple[torch.Tensor, torch.Tensor]: + + targets = _validate_pose_batch(target_c2w, name="target_c2w") + anchor = torch.as_tensor(anchor_c2w, device=targets.device, dtype=targets.dtype) + if anchor.shape != (4, 4) or not bool(torch.isfinite(anchor).all()): + raise ValueError("anchor_c2w must be one finite 4x4 pose") + prior_poses = torch.cat((anchor.unsqueeze(0), targets[:-1]), dim=0) + visibility = target_frustum_visibility_masks( + target_c2w=targets, + source_c2w=prior_poses, + horizontal_fov_degrees=horizontal_fov_degrees, + vertical_fov_degrees=vertical_fov_degrees, + near_m=near_m, + far_m=far_m, + samples_per_axis=samples_per_axis, + ) + prior_index = torch.arange(prior_poses.shape[0], device=targets.device)[:, None] + target_index = torch.arange(targets.shape[0], device=targets.device)[None, :] + chronological = prior_index <= target_index + overlap = (visibility & chronological.unsqueeze(-1)).any(dim=0) + return ~overlap, overlap + + +def _lexicographic_best( + rows: torch.Tensor, + *, + available: torch.Tensor, + stable_values: Sequence[int], +) -> int: + matrix = torch.as_tensor(rows).detach().cpu().to(torch.float64) + flags = torch.as_tensor(available).detach().cpu().tolist() + eligible = [index for index, flag in enumerate(flags) if bool(flag)] + if matrix.ndim != 2 or matrix.shape[0] != len(flags): + raise ValueError("rows and available have incompatible shapes") + if not eligible: + raise ValueError("no candidate is available") + return max( + eligible, + key=lambda index: ( + *(float(value) for value in matrix[index].tolist()), + -int(stable_values[index]), + ), + ) + + +def _coverage_ratio(mask: torch.Tensor) -> torch.Tensor: + return mask.float().mean(dim=-1) + + +def _frontier_ratio(mask: torch.Tensor, frontier: torch.Tensor) -> torch.Tensor: + cells = frontier.sum(dim=1).float() + covered = (mask & frontier).sum(dim=-1).float() + return torch.where( + cells > 0.0, + covered / cells.clamp_min(1.0), + torch.ones_like(cells), + ) + + +def _pick_best( + rows: torch.Tensor, + *, + available: torch.Tensor, + stable_values: Sequence[int], +) -> int: + return _lexicographic_best( + rows, + available=available, + stable_values=tuple(-int(value) for value in stable_values), + ) + + +@torch.inference_mode() +def select_trajectory_fov_history( + *, + controller_c2w: torch.Tensor, + candidate_latents: Sequence[int], + fixed_context_latents: Sequence[int], + target_c2w: torch.Tensor, + budget: int = 8, + horizontal_fov_degrees: float = DEFAULT_HORIZONTAL_FOV_DEGREES, + vertical_fov_degrees: float = DEFAULT_VERTICAL_FOV_DEGREES, + near_m: float = DEFAULT_NEAR_METERS, + far_m: float = DEFAULT_FAR_METERS, + frustum_samples_per_axis: int = FRUSTUM_SAMPLES_PER_AXIS, + latents_per_chunk: int = 9, +) -> TrajectoryFovSelection: + + poses = _validate_pose_batch(controller_c2w, name="controller_c2w") + targets = _validate_pose_batch(target_c2w, name="target_c2w").to( + device=poses.device, dtype=poses.dtype + ) + candidates = tuple(sorted({int(value) for value in candidate_latents})) + fixed = tuple(int(value) for value in fixed_context_latents) + if int(budget) <= 0: + raise ValueError("budget must be positive") + if len(fixed) != 1: + raise ValueError("trajectory FOV retrieval requires exactly one recent[-1] anchor") + if len(candidates) < int(budget): + raise ValueError(f"need at least {budget} candidates, got {len(candidates)}") + if set(candidates).intersection(fixed): + raise ValueError("candidate and fixed context latents must be disjoint") + if targets.shape[0] != 4: + raise ValueError(f"expected target slots 2/4/6/8, got {targets.shape[0]} poses") + if int(latents_per_chunk) <= 0: + raise ValueError("latents_per_chunk must be positive") + maximum_index = max((*candidates, *fixed)) + minimum_index = min((*candidates, *fixed)) + if minimum_index < 0: + raise ValueError("latent indices must be non-negative") + if maximum_index >= poses.shape[0]: + raise ValueError( + f"latent index {maximum_index} exceeds controller pose count {poses.shape[0]}" + ) + + combined = (*fixed, *candidates) + pose_index = torch.tensor(combined, device=poses.device, dtype=torch.long) + combined_c2w = poses.index_select(0, pose_index) + anchor_c2w = combined_c2w[0] + visibility = target_frustum_visibility_masks( + target_c2w=targets, + source_c2w=combined_c2w, + horizontal_fov_degrees=horizontal_fov_degrees, + vertical_fov_degrees=vertical_fov_degrees, + near_m=near_m, + far_m=far_m, + samples_per_axis=frustum_samples_per_axis, + ) + anchor_fov = visibility[0] + candidate_fov = visibility[1:] + frontier_fov, _ = ordered_target_frontier_masks( + anchor_c2w=anchor_c2w, + target_c2w=targets, + horizontal_fov_degrees=horizontal_fov_degrees, + vertical_fov_degrees=vertical_fov_degrees, + near_m=near_m, + far_m=far_m, + samples_per_axis=frustum_samples_per_axis, + ) + + available = torch.ones(len(candidates), device=poses.device, dtype=torch.bool) + covered = anchor_fov.clone() + selected_positions: list[int] = [] + selection_kinds: list[str] = [] + steps: list[dict[str, object]] = [] + bridge_values: list[float] = [] + disconnected_fill_count = 0 + + for slot in range(int(budget)): + resulting = covered.unsqueeze(0) | candidate_fov + full_ratio = resulting.float().mean(dim=2) + full_fairness = torch.sort(full_ratio, dim=1).values + frontier_ratio = _frontier_ratio(resulting, frontier_fov) + frontier_fairness = torch.sort(frontier_ratio, dim=1).values + + bridge_cells = (candidate_fov & covered.unsqueeze(0)).sum(dim=(1, 2)).float() + candidate_cells = candidate_fov.sum(dim=(1, 2)).float() + bridge_fraction = bridge_cells / candidate_cells.clamp_min(1.0) + connected = bridge_cells > 0.0 + + new_fov = candidate_fov & (~covered).unsqueeze(0) + new_fov_per_target = new_fov.sum(dim=2).float() + new_fov_total = new_fov_per_target.sum(dim=1) + new_frontier = new_fov & frontier_fov.unsqueeze(0) + new_frontier_per_target = new_frontier.sum(dim=2).float() + new_frontier_total = new_frontier_per_target.sum(dim=1) + + rows = torch.cat( + ( + full_fairness, + frontier_ratio[:, -1:], + frontier_fairness, + new_fov_total[:, None], + new_frontier_total[:, None], + ), + dim=1, + ) + eligible = available & connected + used_disconnected_fill = not bool(eligible.any()) + if used_disconnected_fill: + eligible = available + chosen = _pick_best(rows, available=eligible, stable_values=candidates) + + if used_disconnected_fill: + kind = "disconnected_fill" + disconnected_fill_count += 1 + elif int(new_frontier_total[chosen].item()) > 0: + kind = "connected_frontier_advance" + elif int(new_fov_total[chosen].item()) > 0: + kind = "connected_coverage_refinement" + else: + kind = "connected_saturated_fill" + + covered = resulting[chosen] + available[chosen] = False + selected_positions.append(chosen) + selection_kinds.append(kind) + bridge = float(bridge_fraction[chosen].item()) + bridge_values.append(bridge) + steps.append( + { + "slot": slot, + "latent": candidates[chosen], + "parent_chunk": candidates[chosen] // int(latents_per_chunk), + "kind": kind, + "bridge_fraction": bridge, + "raw_new_fov_cells_per_target": [ + int(value) for value in new_fov_per_target[chosen].cpu().tolist() + ], + "raw_new_frontier_cells_per_target": [ + int(value) for value in new_frontier_per_target[chosen].cpu().tolist() + ], + "absolute_fov_coverage_per_target": [ + float(value) for value in _coverage_ratio(covered).cpu().tolist() + ], + "frontier_coverage_per_target": [ + float(value) + for value in _frontier_ratio(covered, frontier_fov).cpu().tolist() + ], + } + ) + + selected = tuple(candidates[position] for position in selected_positions) + if len(selected) != int(budget) or len(set(selected)) != int(budget): + raise RuntimeError( + f"trajectory FOV retrieval did not produce {budget} unique sources: {selected}" + ) + return TrajectoryFovSelection( + selected_latents=selected, + selected_parent_chunks=tuple( + sorted({value // int(latents_per_chunk) for value in selected}) + ), + selection_kinds=tuple(selection_kinds), + initial_fov_coverage_per_target=tuple( + float(value) for value in _coverage_ratio(anchor_fov).cpu().tolist() + ), + final_fov_coverage_per_target=tuple( + float(value) for value in _coverage_ratio(covered).cpu().tolist() + ), + initial_frontier_coverage_per_target=tuple( + float(value) + for value in _frontier_ratio(anchor_fov, frontier_fov).cpu().tolist() + ), + final_frontier_coverage_per_target=tuple( + float(value) + for value in _frontier_ratio(covered, frontier_fov).cpu().tolist() + ), + minimum_bridge_fraction=min(bridge_values), + mean_bridge_fraction=sum(bridge_values) / len(bridge_values), + disconnected_fill_count=disconnected_fill_count, + step_diagnostics=tuple(steps), + candidate_count=len(candidates), + target_count=int(targets.shape[0]), + ) + + +__all__ = [ + "DEFAULT_FAR_METERS", + "DEFAULT_HORIZONTAL_FOV_DEGREES", + "DEFAULT_NEAR_METERS", + "DEFAULT_VERTICAL_FOV_DEGREES", + "FRUSTUM_SAMPLES_PER_AXIS", + "TrajectoryFovSelection", + "ordered_target_frontier_masks", + "select_trajectory_fov_history", + "target_frustum_visibility_masks", +] diff --git a/worldcrafter/repencoder/trajectory_memory_provider.py b/worldcrafter/repencoder/trajectory_memory_provider.py new file mode 100644 index 0000000000000000000000000000000000000000..f097d3ea09f86c9f90725221b5b6d20080ca1035 --- /dev/null +++ b/worldcrafter/repencoder/trajectory_memory_provider.py @@ -0,0 +1,521 @@ +from __future__ import annotations + +import math +from dataclasses import asdict, dataclass +from typing import Any, Mapping, Sequence + +import torch + +from .model import RepEncoder +from .trajectory_fov import ( + DEFAULT_FAR_METERS, + DEFAULT_HORIZONTAL_FOV_DEGREES, + DEFAULT_NEAR_METERS, + DEFAULT_VERTICAL_FOV_DEGREES, + FRUSTUM_SAMPLES_PER_AXIS, + select_trajectory_fov_history, +) + + +NUM_LATENT_FRAMES_PER_CHUNK = 9 +VAE_SCALE_FACTOR_TEMPORAL = 4 +WINDOW_NUM_FRAMES = 33 +TARGET_SLOTS = (2, 4, 6, 8) +RECENT_LOCAL_SLOT = 8 +EXPECTED_LATENT_CHW = (16, 48, 80) +HISTORY_SOURCE_BUDGET = 8 + + +@dataclass(frozen=True) +class RepEncoderInferenceProviderConfig: + seed: int = 0 + trajectory_fov_horizontal_fov_degrees: float = DEFAULT_HORIZONTAL_FOV_DEGREES + trajectory_fov_vertical_fov_degrees: float = DEFAULT_VERTICAL_FOV_DEGREES + trajectory_fov_near_m: float = DEFAULT_NEAR_METERS + trajectory_fov_far_m: float = DEFAULT_FAR_METERS + trajectory_fov_samples_per_axis: int = FRUSTUM_SAMPLES_PER_AXIS + near_zero_baseline_m: float = 0.05 + + def __post_init__(self) -> None: + if int(self.seed) < 0: + raise ValueError("seed must be non-negative") + for name in ( + "trajectory_fov_horizontal_fov_degrees", + "trajectory_fov_vertical_fov_degrees", + ): + value = float(getattr(self, name)) + if not math.isfinite(value) or not 0.0 < value < 180.0: + raise ValueError(f"{name} must be finite and in (0, 180)") + if int(self.trajectory_fov_samples_per_axis) != FRUSTUM_SAMPLES_PER_AXIS: + raise ValueError( + "trajectory FOV selection requires " + f"samples_per_axis={FRUSTUM_SAMPLES_PER_AXIS}" + ) + near_m = float(self.trajectory_fov_near_m) + far_m = float(self.trajectory_fov_far_m) + if not math.isfinite(near_m) or near_m <= 0.0: + raise ValueError("trajectory_fov_near_m must be finite and positive") + if not math.isfinite(far_m) or far_m <= near_m: + raise ValueError( + "trajectory_fov_far_m must be finite and greater than near" + ) + if ( + not math.isfinite(float(self.near_zero_baseline_m)) + or float(self.near_zero_baseline_m) < 0.0 + ): + raise ValueError("near_zero_baseline_m must be finite and non-negative") + + +@dataclass(frozen=True) +class RepEncoderInferenceBatchSelection: + batch_index: int + mode: str + retrieval_backend: str + source_indices: tuple[tuple[int, int], ...] + source_flat_indices: tuple[int, ...] + source_raw_frames: tuple[int, ...] + target_indices: tuple[tuple[int, int], ...] + target_raw_frames: tuple[int, ...] + selected_history_flat_indices: tuple[int, ...] + source_unique_parent_count: int + retrieval_diagnostics: Mapping[str, Any] + + def to_jsonable(self) -> dict[str, Any]: + value = asdict(self) + value["source_indices"] = [list(index) for index in self.source_indices] + value["target_indices"] = [list(index) for index in self.target_indices] + value["retrieval_diagnostics"] = dict(self.retrieval_diagnostics) + return value + + +@dataclass(frozen=True) +class RepEncoderInferenceRenderRecord: + chunk_index: int + pose_key: str + window_num_frames: int + target_slots: tuple[int, ...] + selections: tuple[RepEncoderInferenceBatchSelection, ...] + + def to_jsonable(self) -> dict[str, Any]: + return { + "chunk_index": self.chunk_index, + "pose_key": self.pose_key, + "window_num_frames": self.window_num_frames, + "target_slots": list(self.target_slots), + "selections": [selection.to_jsonable() for selection in self.selections], + } + + +def _as_global_metric_c2w( + camera_trajectory: Mapping[str, Any], *, batch_size: int, device: torch.device +) -> tuple[torch.Tensor, str]: + if not isinstance(camera_trajectory, Mapping): + raise TypeError("camera_trajectory must be a mapping") + pose_key = "c2w" + if camera_trajectory.get(pose_key) is None: + raise KeyError("RepEncoder inference requires camera_trajectory['c2w']") + pose = torch.as_tensor(camera_trajectory[pose_key]) + if pose.ndim == 3: + pose = pose.unsqueeze(0) + if pose.ndim != 4 or tuple(pose.shape[-2:]) not in {(3, 4), (4, 4)}: + raise ValueError( + f"{pose_key} must be [B,F,3,4] or [B,F,4,4], got {tuple(pose.shape)}" + ) + if not torch.is_floating_point(pose): + raise TypeError(f"{pose_key} must be floating point") + pose = pose.to(device=device, dtype=torch.float32) + if not torch.isfinite(pose).all(): + raise FloatingPointError(f"{pose_key} contains NaN or Inf") + if pose.shape[-2:] == (3, 4): + bottom = torch.zeros( + *pose.shape[:-2], 1, 4, device=pose.device, dtype=pose.dtype + ) + bottom[..., 0, 3] = 1.0 + pose = torch.cat((pose, bottom), dim=-2) + else: + expected_bottom = pose.new_tensor((0.0, 0.0, 0.0, 1.0)).expand( + *pose.shape[:-2], 4 + ) + if not torch.allclose(pose[..., 3, :], expected_bottom, atol=1e-5, rtol=0): + raise ValueError(f"{pose_key} has invalid homogeneous bottom rows") + if pose.shape[0] != int(batch_size): + raise ValueError( + f"{pose_key} batch {pose.shape[0]} does not match WorldCrafter latent batch {batch_size}" + ) + if int(batch_size) != 1: + raise ValueError("strict online trajectory-FOV retrieval supports batch_size=1") + if pose.shape[1] % WINDOW_NUM_FRAMES: + raise ValueError( + f"camera trajectory length must contain complete {WINDOW_NUM_FRAMES}-frame chunks, " + f"got {pose.shape[1]}" + ) + return pose, pose_key + + +def _anchor_raw_frames(num_chunks: int) -> torch.Tensor: + return torch.as_tensor( + [ + chunk * WINDOW_NUM_FRAMES + local * VAE_SCALE_FACTOR_TEMPORAL + for chunk in range(int(num_chunks)) + for local in range(NUM_LATENT_FRAMES_PER_CHUNK) + ], + dtype=torch.int64, + ).reshape(int(num_chunks), NUM_LATENT_FRAMES_PER_CHUNK) + + +class RepEncoderInferenceMemoryProvider: + def __init__( + self, + runtime: RepEncoder, + config: RepEncoderInferenceProviderConfig | Mapping[str, Any] | None = None, + *, + append_only: bool = False, + ) -> None: + if not callable(runtime): + raise TypeError("runtime must be a callable RepEncoder memory runtime") + if config is None: + config = RepEncoderInferenceProviderConfig() + elif isinstance(config, Mapping): + config = RepEncoderInferenceProviderConfig(**dict(config)) + if not isinstance(config, RepEncoderInferenceProviderConfig): + raise TypeError( + "config must be RepEncoderInferenceProviderConfig or a mapping" + ) + self.append_only = append_only + self._committed_pose = None + self.runtime = runtime + self.config = config + self.last_render_record: RepEncoderInferenceRenderRecord | None = None + self.render_records: list[RepEncoderInferenceRenderRecord] = [] + self._global_pose_latents: torch.Tensor | None = None + + def reset_sequence(self) -> None: + self._committed_pose = None + self._global_pose_latents = None + self.last_render_record = None + self.render_records.clear() + + def append_trajectory(self, pose: torch.Tensor, chunk_index: int) -> None: + """Commit one chunk without permitting changes to previously seen poses.""" + if not self.append_only: + raise RuntimeError("Trajectory append requires append_only=True") + pose, _ = _as_global_metric_c2w({"c2w": pose}, batch_size=1, device=pose.device) + previous = 0 if self._committed_pose is None else self._committed_pose.shape[1] + if ( + pose.shape[1] != (chunk_index + 1) * WINDOW_NUM_FRAMES + or previous != chunk_index * WINDOW_NUM_FRAMES + ): + raise ValueError("Only the next trajectory chunk may be appended") + if previous and not torch.equal(self._committed_pose, pose[:, :previous]): + raise RuntimeError("Committed trajectory prefix changed") + self._committed_pose = pose.detach().clone() + + def _validate_latents( + self, + generated_latents: torch.Tensor, + recent_latents: torch.Tensor, + *, + chunk_index: int, + ) -> None: + if ( + not isinstance(generated_latents, torch.Tensor) + or generated_latents.ndim != 5 + ): + raise TypeError("generated_latents must be a [B,16,T,48,80] torch tensor") + channels, latent_height, latent_width = EXPECTED_LATENT_CHW + expected_tail = ( + channels, + int(chunk_index) * NUM_LATENT_FRAMES_PER_CHUNK, + latent_height, + latent_width, + ) + if tuple(generated_latents.shape[1:]) != expected_tail: + raise ValueError( + "generated_latents do not match complete causal WorldCrafter history: " + f"expected [B,{expected_tail[0]},{expected_tail[1]},{expected_tail[2]}," + f"{expected_tail[3]}], got {tuple(generated_latents.shape)}" + ) + expected_recent = ( + generated_latents.shape[0], + generated_latents.shape[1], + 1, + generated_latents.shape[3], + generated_latents.shape[4], + ) + if ( + not isinstance(recent_latents, torch.Tensor) + or tuple(recent_latents.shape) != expected_recent + ): + raise ValueError( + f"recent_latents must be exact recent[-1] with shape {expected_recent}" + ) + if ( + recent_latents.device != generated_latents.device + or recent_latents.dtype != generated_latents.dtype + ): + raise ValueError( + "recent_latents must share generated_latents device and dtype" + ) + if not torch.is_floating_point(generated_latents): + raise TypeError( + "generated_latents must be floating point standardized Wan latents" + ) + runtime_device = getattr(self.runtime, "device", generated_latents.device) + if torch.device(runtime_device) != generated_latents.device: + raise ValueError( + "RepEncoder runtime and WorldCrafter latents must share one device; " + f"runtime={runtime_device}, worldcrafter={generated_latents.device}" + ) + + def _initialize_or_validate_pose( + self, *, pose: torch.Tensor, chunk_index: int + ) -> torch.Tensor: + total_chunks = int(pose.shape[1] // WINDOW_NUM_FRAMES) + anchor_raw_frames = _anchor_raw_frames(total_chunks) + anchor_indices = anchor_raw_frames.reshape(-1).to(device=pose.device) + global_pose_latents = pose[0].index_select(0, anchor_indices).contiguous() + if self.append_only: + if self._committed_pose is None or not torch.equal( + pose, self._committed_pose + ): + raise RuntimeError("Retrieval trajectory differs from committed poses") + if total_chunks != chunk_index + 1: + raise RuntimeError("Future trajectory chunks are not allowed") + self._global_pose_latents = global_pose_latents.detach().clone() + elif int(chunk_index) == 1: + self.reset_sequence() + self._global_pose_latents = global_pose_latents.detach().clone() + elif self._global_pose_latents is None: + self._global_pose_latents = global_pose_latents.detach().clone() + elif not torch.equal(self._global_pose_latents, global_pose_latents): + raise RuntimeError( + "global camera trajectory changed within one autoregressive sequence" + ) + return anchor_raw_frames + + def _select_history( + self, *, history_length: int, target4_c2w: torch.Tensor + ) -> tuple[tuple[int, ...], str, dict[str, Any]]: + if self._global_pose_latents is None: + raise RuntimeError("global retrieval poses were not initialized") + recent_index = int(history_length) - 1 + candidates = tuple(range(recent_index)) + if len(candidates) < HISTORY_SOURCE_BUDGET: + raise ValueError( + "insufficient strict history for source8: " + f"candidate_count={len(candidates)}, required={HISTORY_SOURCE_BUDGET}" + ) + if len(candidates) == HISTORY_SOURCE_BUDGET: + return ( + candidates, + "trajectory_fov_exact_budget_bootstrap", + { + "method": "trajectory_fov_a_to_b", + "selection_mode": "exact_budget_bootstrap_all_history", + "candidate_indices_global": list(candidates), + "fixed_context_indices_global": [recent_index], + "selected_indices_global": list(candidates), + "target_slots": list(TARGET_SLOTS), + }, + ) + + result = select_trajectory_fov_history( + controller_c2w=self._global_pose_latents, + candidate_latents=candidates, + fixed_context_latents=(recent_index,), + target_c2w=target4_c2w, + budget=HISTORY_SOURCE_BUDGET, + horizontal_fov_degrees=float( + self.config.trajectory_fov_horizontal_fov_degrees + ), + vertical_fov_degrees=float(self.config.trajectory_fov_vertical_fov_degrees), + near_m=float(self.config.trajectory_fov_near_m), + far_m=float(self.config.trajectory_fov_far_m), + frustum_samples_per_axis=int(self.config.trajectory_fov_samples_per_axis), + latents_per_chunk=NUM_LATENT_FRAMES_PER_CHUNK, + ) + diagnostics = result.to_json() + diagnostics.update( + { + "method": "trajectory_fov_a_to_b", + "selection_mode": "ordered_target_union_coverage", + "candidate_indices_global": list(candidates), + "fixed_context_indices_global": [recent_index], + "selected_indices_global": list(result.selected_latents), + "target_slots": list(TARGET_SLOTS), + "frustum_samples_per_axis": int( + self.config.trajectory_fov_samples_per_axis + ), + } + ) + return ( + tuple(int(value) for value in result.selected_latents), + "trajectory_fov_a_to_b_coverage", + diagnostics, + ) + + def render_memory( + self, + *, + generated_latents: torch.Tensor, + recent_latents: torch.Tensor, + camera_trajectory: Mapping[str, Any], + chunk_index: int, + num_latent_frames_per_chunk: int, + vae_scale_factor_temporal: int, + generator: torch.Generator | Sequence[torch.Generator] | None = None, + ) -> torch.Tensor: + chunk_index = int(chunk_index) + if chunk_index <= 0: + raise ValueError("RepEncoder inference requires chunk_index >= 1") + if int(num_latent_frames_per_chunk) != NUM_LATENT_FRAMES_PER_CHUNK: + raise ValueError( + "RepEncoder inference is fixed to nine latent frames per chunk" + ) + if int(vae_scale_factor_temporal) != VAE_SCALE_FACTOR_TEMPORAL: + raise ValueError( + "RepEncoder inference is fixed to Wan temporal scale factor 4" + ) + self._validate_latents( + generated_latents, recent_latents, chunk_index=chunk_index + ) + + pose, pose_key = _as_global_metric_c2w( + camera_trajectory, + batch_size=int(generated_latents.shape[0]), + device=generated_latents.device, + ) + anchor_raw_frames = self._initialize_or_validate_pose( + pose=pose, chunk_index=chunk_index + ) + if chunk_index >= anchor_raw_frames.shape[0]: + raise ValueError( + f"camera trajectory has only {anchor_raw_frames.shape[0]} complete chunks, " + f"cannot query chunk {chunk_index}" + ) + if self._global_pose_latents is None: + raise RuntimeError("global retrieval poses were not initialized") + + target_full_chunk_c2w = self._global_pose_latents[ + chunk_index + * NUM_LATENT_FRAMES_PER_CHUNK : (chunk_index + 1) + * NUM_LATENT_FRAMES_PER_CHUNK + ] + target4_c2w = target_full_chunk_c2w.index_select( + 0, + torch.as_tensor( + TARGET_SLOTS, device=target_full_chunk_c2w.device, dtype=torch.int64 + ), + ) + selected_history, mode, retrieval_diagnostics = self._select_history( + history_length=int(generated_latents.shape[2]), target4_c2w=target4_c2w + ) + if ( + len(selected_history) != HISTORY_SOURCE_BUDGET + or len(set(selected_history)) != HISTORY_SOURCE_BUDGET + ): + raise RuntimeError( + "trajectory-FOV retrieval did not return eight unique history views: " + f"{selected_history}" + ) + + del generator + recent_index = int(generated_latents.shape[2]) - 1 + source_flat_indices = (recent_index, *selected_history) + if len(set(source_flat_indices)) != NUM_LATENT_FRAMES_PER_CHUNK: + raise RuntimeError(f"source9 is not unique: {source_flat_indices}") + source_latents = ( + generated_latents.index_select( + 2, + torch.as_tensor( + source_flat_indices, + device=generated_latents.device, + dtype=torch.int64, + ), + ) + .permute(0, 2, 1, 3, 4) + .contiguous() + ) + + source_c2w_metric = ( + self._global_pose_latents.index_select( + 0, + torch.as_tensor( + source_flat_indices, + device=self._global_pose_latents.device, + dtype=torch.int64, + ), + ) + .unsqueeze(0) + .to(device=generated_latents.device, dtype=torch.float32) + ) + target_c2w_metric = target4_c2w.unsqueeze(0).to( + device=generated_latents.device, dtype=torch.float32 + ) + memory4 = self.runtime( + source_latents=source_latents, + source_c2w_metric=source_c2w_metric, + target_c2w_metric=target_c2w_metric, + near_zero_baseline_m=float(self.config.near_zero_baseline_m), + ) + expected_output = (int(generated_latents.shape[0]), 16, 4, 48, 80) + if ( + not isinstance(memory4, torch.Tensor) + or tuple(memory4.shape) != expected_output + ): + raise RuntimeError( + f"RepEncoder runtime must return standardized memory4 {expected_output}, " + f"got {None if not isinstance(memory4, torch.Tensor) else tuple(memory4.shape)}" + ) + if memory4.device != generated_latents.device: + raise ValueError( + "RepEncoder memory4 must remain on the WorldCrafter latent device" + ) + memory4 = memory4.to(dtype=generated_latents.dtype) + + source_indices = tuple( + (index // NUM_LATENT_FRAMES_PER_CHUNK, index % NUM_LATENT_FRAMES_PER_CHUNK) + for index in source_flat_indices + ) + target_indices = tuple((chunk_index, slot) for slot in TARGET_SLOTS) + selection = RepEncoderInferenceBatchSelection( + batch_index=0, + mode=mode, + retrieval_backend="trajectory_fov_a_to_b", + source_indices=source_indices, + source_flat_indices=source_flat_indices, + source_raw_frames=tuple( + int(anchor_raw_frames[chunk, local]) for chunk, local in source_indices + ), + target_indices=target_indices, + target_raw_frames=tuple( + int(anchor_raw_frames[chunk, local]) for chunk, local in target_indices + ), + selected_history_flat_indices=selected_history, + source_unique_parent_count=len({chunk for chunk, _ in source_indices}), + retrieval_diagnostics=retrieval_diagnostics, + ) + record = RepEncoderInferenceRenderRecord( + chunk_index=chunk_index, + pose_key=pose_key, + window_num_frames=WINDOW_NUM_FRAMES, + target_slots=TARGET_SLOTS, + selections=(selection,), + ) + self.last_render_record = record + self.render_records.append(record) + return memory4 + + +__all__ = [ + "HISTORY_SOURCE_BUDGET", + "RepEncoderInferenceBatchSelection", + "RepEncoderInferenceMemoryProvider", + "RepEncoderInferenceProviderConfig", + "RepEncoderInferenceRenderRecord", + "NUM_LATENT_FRAMES_PER_CHUNK", + "RECENT_LOCAL_SLOT", + "TARGET_SLOTS", + "VAE_SCALE_FACTOR_TEMPORAL", + "WINDOW_NUM_FRAMES", +] diff --git a/worldcrafter/repencoder/vggt/__init__.py b/worldcrafter/repencoder/vggt/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..f12f620e094b3a59da79ea4acc5a81881d42807f --- /dev/null +++ b/worldcrafter/repencoder/vggt/__init__.py @@ -0,0 +1,3 @@ +from .model import DinoTail, RepresentationBackbone + +__all__ = ["DinoTail", "RepresentationBackbone"] diff --git a/worldcrafter/repencoder/vggt/layers/__init__.py b/worldcrafter/repencoder/vggt/layers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..8f794cc0ae3c34f90c0d053fbf67b0ba801f6fde --- /dev/null +++ b/worldcrafter/repencoder/vggt/layers/__init__.py @@ -0,0 +1,11 @@ +from .attention import Attention, MemEffAttention +from .block import Block +from .rope import PositionGetter, RotaryPositionEmbedding2D + +__all__ = [ + "Attention", + "Block", + "MemEffAttention", + "PositionGetter", + "RotaryPositionEmbedding2D", +] diff --git a/worldcrafter/repencoder/vggt/layers/attention.py b/worldcrafter/repencoder/vggt/layers/attention.py new file mode 100644 index 0000000000000000000000000000000000000000..f369bcb8cd7f85022ca1ce43c543bbcd3fba33bc --- /dev/null +++ b/worldcrafter/repencoder/vggt/layers/attention.py @@ -0,0 +1,81 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# included in LICENSE.txt at the root of this repository. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/models/vision_transformer.py + + +import torch.nn.functional as F +from torch import nn, Tensor + + +class Attention(nn.Module): + def __init__( + self, + dim: int, + num_heads: int = 8, + qkv_bias: bool = True, + proj_bias: bool = True, + attn_drop: float = 0.0, + proj_drop: float = 0.0, + norm_layer: nn.Module = nn.LayerNorm, + qk_norm: bool = False, + fused_attn: bool = True, # use F.scaled_dot_product_attention or not + rope=None, + ) -> None: + super().__init__() + assert dim % num_heads == 0, "dim should be divisible by num_heads" + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.scale = self.head_dim**-0.5 + self.fused_attn = fused_attn + + self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() + self.k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(dim, dim, bias=proj_bias) + self.proj_drop = nn.Dropout(proj_drop) + self.rope = rope + + def forward(self, x: Tensor, pos=None) -> Tensor: + B, N, C = x.shape + qkv = ( + self.qkv(x) + .reshape(B, N, 3, self.num_heads, self.head_dim) + .permute(2, 0, 3, 1, 4) + ) + q, k, v = qkv.unbind(0) + q, k = self.q_norm(q), self.k_norm(k) + + if self.rope is not None: + q = self.rope(q, pos) + k = self.rope(k, pos) + + if self.fused_attn: + x = F.scaled_dot_product_attention( + q, k, v, dropout_p=self.attn_drop.p if self.training else 0.0 + ) + else: + q = q * self.scale + attn = q @ k.transpose(-2, -1) + attn = attn.softmax(dim=-1) + attn = self.attn_drop(attn) + x = attn @ v + + x = x.transpose(1, 2).reshape(B, N, C) + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class MemEffAttention(Attention): + def forward(self, x: Tensor, attn_bias=None, pos=None) -> Tensor: + if pos is not None or attn_bias is not None: + raise ValueError( + "This attention layer accepts dense tokens without positional bias" + ) + return super().forward(x) diff --git a/worldcrafter/repencoder/vggt/layers/block.py b/worldcrafter/repencoder/vggt/layers/block.py new file mode 100644 index 0000000000000000000000000000000000000000..c8239a3f13348287df931196afd6dd3f9bc20c7f --- /dev/null +++ b/worldcrafter/repencoder/vggt/layers/block.py @@ -0,0 +1,136 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# included in LICENSE.txt at the root of this repository. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/layers/patch_embed.py + +from typing import Callable + +import torch +from torch import nn, Tensor +from .attention import Attention +from .drop_path import DropPath +from .layer_scale import LayerScale +from .mlp import Mlp + + +class Block(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + mlp_ratio: float = 4.0, + qkv_bias: bool = True, + proj_bias: bool = True, + ffn_bias: bool = True, + drop: float = 0.0, + attn_drop: float = 0.0, + init_values=None, + drop_path: float = 0.0, + act_layer: Callable[..., nn.Module] = nn.GELU, + norm_layer: Callable[..., nn.Module] = nn.LayerNorm, + attn_class: Callable[..., nn.Module] = Attention, + ffn_layer: Callable[..., nn.Module] = Mlp, + qk_norm: bool = False, + fused_attn: bool = True, # use F.scaled_dot_product_attention or not + rope=None, + ) -> None: + super().__init__() + + self.norm1 = norm_layer(dim) + + self.attn = attn_class( + dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + proj_bias=proj_bias, + attn_drop=attn_drop, + proj_drop=drop, + qk_norm=qk_norm, + fused_attn=fused_attn, + rope=rope, + ) + + self.ls1 = ( + LayerScale(dim, init_values=init_values) if init_values else nn.Identity() + ) + self.drop_path1 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + + self.norm2 = norm_layer(dim) + mlp_hidden_dim = int(dim * mlp_ratio) + self.mlp = ffn_layer( + in_features=dim, + hidden_features=mlp_hidden_dim, + act_layer=act_layer, + drop=drop, + bias=ffn_bias, + ) + self.ls2 = ( + LayerScale(dim, init_values=init_values) if init_values else nn.Identity() + ) + self.drop_path2 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + + self.sample_drop_ratio = drop_path + + def forward(self, x: Tensor, pos=None) -> Tensor: + def attn_residual_func(x: Tensor, pos=None) -> Tensor: + return self.ls1(self.attn(self.norm1(x), pos=pos)) + + def ffn_residual_func(x: Tensor) -> Tensor: + return self.ls2(self.mlp(self.norm2(x))) + + if self.training and self.sample_drop_ratio > 0.1: + # the overhead is compensated only for a drop path rate larger than 0.1 + x = drop_add_residual_stochastic_depth( + x, + pos=pos, + residual_func=attn_residual_func, + sample_drop_ratio=self.sample_drop_ratio, + ) + x = drop_add_residual_stochastic_depth( + x, + residual_func=ffn_residual_func, + sample_drop_ratio=self.sample_drop_ratio, + ) + elif self.training and self.sample_drop_ratio > 0.0: + x = x + self.drop_path1(attn_residual_func(x, pos=pos)) + x = x + self.drop_path1(ffn_residual_func(x)) + else: + x = x + attn_residual_func(x, pos=pos) + x = x + ffn_residual_func(x) + return x + + +def drop_add_residual_stochastic_depth( + x: Tensor, + residual_func: Callable[[Tensor], Tensor], + sample_drop_ratio: float = 0.0, + pos=None, +) -> Tensor: + # 1) extract subset using permutation + b, n, d = x.shape + sample_subset_size = max(int(b * (1 - sample_drop_ratio)), 1) + brange = (torch.randperm(b, device=x.device))[:sample_subset_size] + x_subset = x[brange] + + # 2) apply residual_func to get residual + if pos is not None: + # if necessary, apply rope to the subset + pos = pos[brange] + residual = residual_func(x_subset, pos=pos) + else: + residual = residual_func(x_subset) + + x_flat = x.flatten(1) + residual = residual.flatten(1) + + residual_scale_factor = b / sample_subset_size + + # 3) add the residual + x_plus_residual = torch.index_add( + x_flat, 0, brange, residual.to(dtype=x.dtype), alpha=residual_scale_factor + ) + return x_plus_residual.view_as(x) diff --git a/worldcrafter/repencoder/vggt/layers/drop_path.py b/worldcrafter/repencoder/vggt/layers/drop_path.py new file mode 100644 index 0000000000000000000000000000000000000000..1354d399671871dcf14c1777097d2208f4f334c8 --- /dev/null +++ b/worldcrafter/repencoder/vggt/layers/drop_path.py @@ -0,0 +1,37 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the VGGT license found at +# https://github.com/facebookresearch/vggt/blob/main/LICENSE.txt + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/layers/drop.py + + +from torch import nn + + +def drop_path(x, drop_prob: float = 0.0, training: bool = False): + if drop_prob == 0.0 or not training: + return x + keep_prob = 1 - drop_prob + shape = (x.shape[0],) + (1,) * ( + x.ndim - 1 + ) # work with diff dim tensors, not just 2D ConvNets + random_tensor = x.new_empty(shape).bernoulli_(keep_prob) + if keep_prob > 0.0: + random_tensor.div_(keep_prob) + output = x * random_tensor + return output + + +class DropPath(nn.Module): + """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).""" + + def __init__(self, drop_prob=None): + super(DropPath, self).__init__() + self.drop_prob = drop_prob + + def forward(self, x): + return drop_path(x, self.drop_prob, self.training) diff --git a/worldcrafter/repencoder/vggt/layers/layer_scale.py b/worldcrafter/repencoder/vggt/layers/layer_scale.py new file mode 100644 index 0000000000000000000000000000000000000000..401bce901765e0293d2b147fc4f29c4c097d621c --- /dev/null +++ b/worldcrafter/repencoder/vggt/layers/layer_scale.py @@ -0,0 +1,23 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# included in LICENSE.txt at the root of this repository. + +# Modified from: https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/vision_transformer.py#L103-L110 + +from typing import Union + +import torch +from torch import nn, Tensor + + +class LayerScale(nn.Module): + def __init__( + self, dim: int, init_values: Union[float, Tensor] = 1e-5, inplace: bool = False + ) -> None: + super().__init__() + self.inplace = inplace + self.gamma = nn.Parameter(init_values * torch.ones(dim)) + + def forward(self, x: Tensor) -> Tensor: + return x.mul_(self.gamma) if self.inplace else x * self.gamma diff --git a/worldcrafter/repencoder/vggt/layers/mlp.py b/worldcrafter/repencoder/vggt/layers/mlp.py new file mode 100644 index 0000000000000000000000000000000000000000..d7e632bb34e55c2667c18a9c69078a56fe784849 --- /dev/null +++ b/worldcrafter/repencoder/vggt/layers/mlp.py @@ -0,0 +1,40 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# included in LICENSE.txt at the root of this repository. + +# References: +# https://github.com/facebookresearch/dino/blob/master/vision_transformer.py +# https://github.com/rwightman/pytorch-image-models/tree/master/timm/layers/mlp.py + + +from typing import Callable, Optional + +from torch import nn, Tensor + + +class Mlp(nn.Module): + def __init__( + self, + in_features: int, + hidden_features: Optional[int] = None, + out_features: Optional[int] = None, + act_layer: Callable[..., nn.Module] = nn.GELU, + drop: float = 0.0, + bias: bool = True, + ) -> None: + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + self.fc1 = nn.Linear(in_features, hidden_features, bias=bias) + self.act = act_layer() + self.fc2 = nn.Linear(hidden_features, out_features, bias=bias) + self.drop = nn.Dropout(drop) + + def forward(self, x: Tensor) -> Tensor: + x = self.fc1(x) + x = self.act(x) + x = self.drop(x) + x = self.fc2(x) + x = self.drop(x) + return x diff --git a/worldcrafter/repencoder/vggt/layers/rope.py b/worldcrafter/repencoder/vggt/layers/rope.py new file mode 100644 index 0000000000000000000000000000000000000000..6079378ca2a9fcaa558c75f705e6b9fff160f998 --- /dev/null +++ b/worldcrafter/repencoder/vggt/layers/rope.py @@ -0,0 +1,206 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# included in LICENSE.txt at the root of this repository. + + +# Implementation of 2D Rotary Position Embeddings (RoPE). + +# This module provides a clean implementation of 2D Rotary Position Embeddings, +# which extends the original RoPE concept to handle 2D spatial positions. + +# Inspired by: +# https://github.com/meta-llama/codellama/blob/main/llama/model.py +# https://github.com/naver-ai/rope-vit + + +from typing import Dict, Tuple + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class PositionGetter: + """Generates and caches 2D spatial positions for patches in a grid. + + This class efficiently manages the generation of spatial coordinates for patches + in a 2D grid, caching results to avoid redundant computations. + + Attributes: + position_cache: Dictionary storing precomputed position tensors for different + grid dimensions. + """ + + def __init__(self): + """Initializes the position generator with an empty cache.""" + self.position_cache: Dict[Tuple[int, int], torch.Tensor] = {} + + def __call__( + self, batch_size: int, height: int, width: int, device: torch.device + ) -> torch.Tensor: + """Generates spatial positions for a batch of patches. + + Args: + batch_size: Number of samples in the batch. + height: Height of the grid in patches. + width: Width of the grid in patches. + device: Target device for the position tensor. + + Returns: + Tensor of shape (batch_size, height*width, 2) containing y,x coordinates + for each position in the grid, repeated for each batch item. + """ + if (height, width) not in self.position_cache: + y_coords = torch.arange(height, device=device) + x_coords = torch.arange(width, device=device) + positions = torch.cartesian_prod(y_coords, x_coords) + self.position_cache[height, width] = positions + + cached_positions = self.position_cache[height, width] + return ( + cached_positions.view(1, height * width, 2) + .expand(batch_size, -1, -1) + .clone() + ) + + +class RotaryPositionEmbedding2D(nn.Module): + """2D Rotary Position Embedding implementation. + + This module applies rotary position embeddings to input tokens based on their + 2D spatial positions. It handles the position-dependent rotation of features + separately for vertical and horizontal dimensions. + + Args: + frequency: Base frequency for the position embeddings. Default: 100.0 + scaling_factor: Scaling factor for frequency computation. Default: 1.0 + + Attributes: + base_frequency: Base frequency for computing position embeddings. + scaling_factor: Factor to scale the computed frequencies. + frequency_cache: Cache for storing precomputed frequency components. + """ + + def __init__(self, frequency: float = 100.0, scaling_factor: float = 1.0): + """Initializes the 2D RoPE module.""" + super().__init__() + self.base_frequency = frequency + self.scaling_factor = scaling_factor + self.frequency_cache: Dict[Tuple, Tuple[torch.Tensor, torch.Tensor]] = {} + + def _compute_frequency_components( + self, dim: int, seq_len: int, device: torch.device, dtype: torch.dtype + ) -> Tuple[torch.Tensor, torch.Tensor]: + """Computes frequency components for rotary embeddings. + + Args: + dim: Feature dimension (must be even). + seq_len: Maximum sequence length. + device: Target device for computations. + dtype: Data type for the computed tensors. + + Returns: + Tuple of (cosine, sine) tensors for frequency components. + """ + cache_key = (dim, seq_len, device, dtype) + if cache_key not in self.frequency_cache: + # Compute frequency bands + exponents = torch.arange(0, dim, 2, device=device).float() / dim + inv_freq = 1.0 / (self.base_frequency**exponents) + + # Generate position-dependent frequencies + positions = torch.arange(seq_len, device=device, dtype=inv_freq.dtype) + angles = torch.einsum("i,j->ij", positions, inv_freq) + + # Compute and cache frequency components + angles = angles.to(dtype) + angles = torch.cat((angles, angles), dim=-1) + cos_components = angles.cos().to(dtype) + sin_components = angles.sin().to(dtype) + self.frequency_cache[cache_key] = (cos_components, sin_components) + + return self.frequency_cache[cache_key] + + @staticmethod + def _rotate_features(x: torch.Tensor) -> torch.Tensor: + """Performs feature rotation by splitting and recombining feature dimensions. + + Args: + x: Input tensor to rotate. + + Returns: + Rotated feature tensor. + """ + feature_dim = x.shape[-1] + x1, x2 = x[..., : feature_dim // 2], x[..., feature_dim // 2 :] + return torch.cat((-x2, x1), dim=-1) + + def _apply_1d_rope( + self, + tokens: torch.Tensor, + positions: torch.Tensor, + cos_comp: torch.Tensor, + sin_comp: torch.Tensor, + ) -> torch.Tensor: + """Applies 1D rotary position embeddings along one dimension. + + Args: + tokens: Input token features. + positions: Position indices. + cos_comp: Cosine components for rotation. + sin_comp: Sine components for rotation. + + Returns: + Tokens with applied rotary position embeddings. + """ + # Embed positions with frequency components + cos = F.embedding(positions, cos_comp)[:, None, :, :] + sin = F.embedding(positions, sin_comp)[:, None, :, :] + + # Apply rotation + return (tokens * cos) + (self._rotate_features(tokens) * sin) + + def forward(self, tokens: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: + """Applies 2D rotary position embeddings to input tokens. + + Args: + tokens: Input tensor of shape (batch_size, n_heads, n_tokens, dim). + The feature dimension (dim) must be divisible by 4. + positions: Position tensor of shape (batch_size, n_tokens, 2) containing + the y and x coordinates for each token. + + Returns: + Tensor of same shape as input with applied 2D rotary position embeddings. + + Raises: + AssertionError: If input dimensions are invalid or positions are malformed. + """ + # Validate inputs + assert tokens.size(-1) % 2 == 0, "Feature dimension must be even" + assert positions.ndim == 3 and positions.shape[-1] == 2, ( + "Positions must have shape (batch_size, n_tokens, 2)" + ) + + # Compute feature dimension for each spatial direction + feature_dim = tokens.size(-1) // 2 + + # Get frequency components + max_position = int(positions.max()) + 1 + cos_comp, sin_comp = self._compute_frequency_components( + feature_dim, max_position, tokens.device, tokens.dtype + ) + + # Split features for vertical and horizontal processing + vertical_features, horizontal_features = tokens.chunk(2, dim=-1) + + # Apply RoPE separately for each dimension + vertical_features = self._apply_1d_rope( + vertical_features, positions[..., 0], cos_comp, sin_comp + ) + horizontal_features = self._apply_1d_rope( + horizontal_features, positions[..., 1], cos_comp, sin_comp + ) + + # Combine processed features + return torch.cat((vertical_features, horizontal_features), dim=-1) diff --git a/worldcrafter/repencoder/vggt/model.py b/worldcrafter/repencoder/vggt/model.py new file mode 100644 index 0000000000000000000000000000000000000000..ca80638872ea99f98e525a5e5fcee9b44d9f72a8 --- /dev/null +++ b/worldcrafter/repencoder/vggt/model.py @@ -0,0 +1,184 @@ +from __future__ import annotations + +from functools import partial + +import torch +import torch.nn as nn + +from ..config import RepEncoderConfig +from .layers.attention import MemEffAttention +from .layers.block import Block +from .layers.rope import PositionGetter, RotaryPositionEmbedding2D + + +def _slice_expand_and_flatten(token: torch.Tensor, batch: int, views: int) -> torch.Tensor: + first = token[:, 0:1].expand(batch, 1, *token.shape[2:]) + others = token[:, 1:].expand(batch, views - 1, *token.shape[2:]) + return torch.cat((first, others), dim=1).reshape(batch * views, *token.shape[2:]) + + +class DinoTail(nn.Module): + def __init__(self, config: RepEncoderConfig | None = None) -> None: + super().__init__() + self.config = RepEncoderConfig() if config is None else config + self.latent_cls_token = nn.Parameter(torch.zeros(1, 1, self.config.dino_dim)) + self.latent_register_tokens = nn.Parameter(torch.zeros(1, 4, self.config.dino_dim)) + norm_layer = partial(nn.LayerNorm, eps=1e-6) + self.blocks = nn.ModuleList( + [ + Block( + dim=self.config.dino_dim, + num_heads=self.config.dino_heads, + mlp_ratio=4.0, + qkv_bias=True, + proj_bias=True, + ffn_bias=True, + init_values=1.0, + norm_layer=norm_layer, + attn_class=MemEffAttention, + qk_norm=False, + ) + for _ in range(self.config.dino_depth - self.config.dino_input_block) + ] + ) + self.norm = norm_layer(self.config.dino_dim) + + def forward(self, patch_tokens: torch.Tensor) -> torch.Tensor: + expected = ( + self.config.source_views, + self.config.patch_tokens, + self.config.dino_dim, + ) + if patch_tokens.ndim != 4 or tuple(patch_tokens.shape[1:]) != expected: + raise ValueError(f"DINO input must be [B,{expected}], got {tuple(patch_tokens.shape)}") + batch = patch_tokens.shape[0] + flat = patch_tokens.reshape( + batch * self.config.source_views, + self.config.patch_tokens, + self.config.dino_dim, + ) + prefix = torch.cat( + ( + self.latent_cls_token.expand(flat.shape[0], -1, -1), + self.latent_register_tokens.expand(flat.shape[0], -1, -1), + ), + dim=1, + ).to(device=flat.device, dtype=flat.dtype) + tokens = torch.cat((prefix, flat), dim=1) + for block in self.blocks: + tokens = block(tokens) + tokens = self.norm(tokens)[:, self.config.latent_special_tokens :] + return tokens.reshape(batch, *expected) + + +class RepresentationBackbone(nn.Module): + def __init__(self, config: RepEncoderConfig | None = None) -> None: + super().__init__() + self.config = RepEncoderConfig() if config is None else config + self.rope = RotaryPositionEmbedding2D(frequency=100) + self.position_getter = PositionGetter() + block_args = dict( + dim=self.config.dino_dim, + num_heads=self.config.vggt_heads, + mlp_ratio=4.0, + qkv_bias=True, + proj_bias=True, + ffn_bias=True, + init_values=0.01, + qk_norm=True, + rope=self.rope, + ) + self.frame_blocks = nn.ModuleList( + [Block(**block_args) for _ in range(self.config.vggt_depth)] + ) + self.global_blocks = nn.ModuleList( + [Block(**block_args) for _ in range(self.config.vggt_depth)] + ) + self.camera_token = nn.Parameter(torch.zeros(1, 2, 1, self.config.dino_dim)) + self.register_token = nn.Parameter( + torch.zeros(1, 2, self.config.vggt_register_tokens, self.config.dino_dim) + ) + self.camera_mlp = nn.Sequential( + nn.Linear(11, self.config.dino_dim, bias=True), + nn.SiLU(), + nn.Linear(self.config.dino_dim, self.config.dino_dim, bias=True), + ) + self.geo_feature_connector = nn.Linear( + self.config.dino_dim * 2, self.config.scene_dim, bias=True + ) + self.geo_feature_norm = nn.LayerNorm(self.config.scene_dim, bias=False) + + def forward( + self, + patch_tokens: torch.Tensor, + camera_tokens: torch.Tensor, + ) -> torch.Tensor: + batch, views, patches, channels = patch_tokens.shape + expected = ( + self.config.source_views, + self.config.patch_tokens, + self.config.dino_dim, + ) + if tuple(patch_tokens.shape[1:]) != expected: + raise ValueError(f"VGGT patch input must end in {expected}") + if tuple(camera_tokens.shape) != (batch, views, 11): + raise ValueError("camera_tokens must be [B,9,11]") + flat_patch = patch_tokens.reshape(batch * views, patches, channels) + projected_camera = self.camera_mlp(camera_tokens).unsqueeze(2) + camera = _slice_expand_and_flatten(self.camera_token, batch, views) + camera = camera + projected_camera.reshape(batch * views, 1, channels) + register = _slice_expand_and_flatten(self.register_token, batch, views) + tokens = torch.cat((camera, register, flat_patch), dim=1) + total_tokens = tokens.shape[1] + + pos = self.position_getter( + batch * views, + self.config.patch_height, + self.config.patch_width, + device=patch_tokens.device, + ) + pos = pos + 1 + special = torch.zeros( + batch * views, + self.config.latent_special_tokens, + 2, + device=patch_tokens.device, + dtype=pos.dtype, + ) + pos = torch.cat((special, pos), dim=1) + + frame_index = 0 + global_index = 0 + last_frame = None + last_global = None + for _ in range(self.config.vggt_depth): + if tokens.shape != (batch * views, total_tokens, channels): + tokens = tokens.reshape(batch, views, total_tokens, channels).reshape( + batch * views, total_tokens, channels + ) + frame_pos = pos.reshape(batch, views, total_tokens, 2).reshape( + batch * views, total_tokens, 2 + ) + tokens = self.frame_blocks[frame_index](tokens, pos=frame_pos) + frame_index += 1 + last_frame = tokens.reshape(batch, views, total_tokens, channels) + + tokens = tokens.reshape(batch, views, total_tokens, channels).reshape( + batch, views * total_tokens, channels + ) + global_pos = pos.reshape(batch, views, total_tokens, 2).reshape( + batch, views * total_tokens, 2 + ) + tokens = self.global_blocks[global_index](tokens, pos=global_pos) + global_index += 1 + last_global = tokens.reshape(batch, views, total_tokens, channels) + + if last_frame is None or last_global is None: + raise RuntimeError("VGGT representation path produced no output") + combined = torch.cat((last_frame, last_global), dim=-1) + patch = combined[:, :, self.config.latent_special_tokens :] + scene = self.geo_feature_connector(patch) + return self.geo_feature_norm(scene) + + +__all__ = ["DinoTail", "RepresentationBackbone"] diff --git a/worldcrafter/streaming.py b/worldcrafter/streaming.py new file mode 100644 index 0000000000000000000000000000000000000000..2ae1abb704bc624de7cc6bc6055652d6457bb14e --- /dev/null +++ b/worldcrafter/streaming.py @@ -0,0 +1,139 @@ +"""Persistent I2V sessions using the same Fast sampler as offline inference.""" + +import torch +from diffusers.utils.torch_utils import randn_tensor +from .diffusers.pipeline import _render_repencoder_memory_latents +from .fast.sampling import sample_fast + + +def stream_chunks( + pipe, *, image, prompt, generator, provider, max_chunks, cancel, on_step=None +): + """Yield readiness, then accept cumulative local/global poses for each chunk. + + Advance and close this generator under inference_mode on its owning thread. + The RNG calls and conditioning layout follow the offline I2V pipeline. + """ + if image is None or not 1 <= max_chunks <= 50: + raise ValueError("Interactive I2V requires an image and 1..50 chunks") + device = pipe._execution_device + pipe.fast_inference_mode = "i2v" + pipe._guidance_scale = 1.0 + pipe._attention_kwargs = None + pipe._interrupt = False + pipe._current_timestep = None + mean = torch.tensor( + pipe.vae.config.latents_mean, device=device, dtype=pipe.vae.dtype + ).view(1, 16, 1, 1, 1) + std = 1.0 / torch.tensor( + pipe.vae.config.latents_std, device=device, dtype=pipe.vae.dtype + ).view(1, 16, 1, 1, 1) + embeds, _ = pipe.encode_prompt( + prompt=prompt, + negative_prompt="", + do_classifier_free_guidance=False, + num_videos_per_prompt=1, + device=device, + max_sequence_length=512, + ) + embeds = embeds.to(pipe.transformer.dtype) + pixels = pipe.video_processor.preprocess(image, height=384, width=640) + prefix, fake = pipe.prepare_image_latents( + pixels, + latents_mean=mean, + latents_std=std, + num_latent_frames_per_chunk=9, + dtype=torch.float32, + device=device, + generator=generator, + ) + sigma = torch.rand(1, device=device, generator=generator) * (0.135 - 0.111) + 0.111 + prefix = ( + sigma * randn_tensor(prefix.shape, generator=generator, device=device) + + (1 - sigma) * prefix + ) + sigma = torch.rand(1, device=device, generator=generator) * (0.135 - 0.111) + 0.111 + fake = ( + sigma * randn_tensor(fake.shape, generator=generator, device=device) + + (1 - sigma) * fake + ) + history = torch.cat([torch.zeros(1, 16, 3, 48, 80, device=device), fake], dim=2) + generated = history[:, :, :0] + ids = torch.arange(17).split([1, 4, 2, 1, 9]) + short_ids = torch.cat([ids[0], ids[3]]).unsqueeze(0) + + def step_callback(*args): + cancel() + return on_step(*args) if on_step else args[-1] + + camera = yield {"ready": True} + for index in range(max_chunks): + cancel() + camera = dict(camera, c2w=camera["retrieval_pose"]) + memory = ( + generated.new_zeros(1, 16, 4, 48, 80) + if index == 0 + else _render_repencoder_memory_latents( + memory_provider=provider, + generated_latents=generated, + camera_trajectory=camera, + chunk_index=index, + num_latent_frames_per_chunk=9, + vae_scale_factor_temporal=4, + generator=generator, + ) + ) + mid, recent = history[:, :, -3:].split([2, 1], dim=2) + short = torch.cat([prefix, recent], dim=2) + latents = pipe.prepare_latents( + 1, + 16, + 384, + 640, + 33, + dtype=torch.float32, + device=device, + generator=generator, + latents=None, + ) + # Upsampling noise retains the sampler's independent default generator. + with pipe.progress_bar(total=6) as progress: + latents = sample_fast( + pipe, + latents=latents, + pyramid_num_stages=3, + prompt_embeds=embeds, + guidance_scale=1.0, + indices_hidden_states=ids[4].unsqueeze(0), + indices_latents_history_short=short_ids, + indices_latents_history_mid=ids[2].unsqueeze(0), + indices_latents_history_long=ids[1].unsqueeze(0), + latents_history_short=short, + latents_history_mid=mid, + latents_history_long=memory, + attention_kwargs={}, + camera_trajectory=camera, + num_latent_frames_per_chunk=9, + chunk_index=index, + camera_restart_each_chunk=False, + ucpe_pixel_center=True, + device=device, + transformer_dtype=pipe.transformer.dtype, + callback_on_step_end=step_callback, + callback_on_step_end_tensor_inputs=[], + progress_bar=progress, + ) + cancel() + generated = torch.cat([generated, latents], dim=2) + history = latents[:, :, -3:].clone() + decoded = pipe.vae.decode( + latents.to(pipe.vae.dtype) / std + mean, return_dict=False + )[0] + camera = yield dict( + chunk_index=index, + rgb=decoded.detach().cpu(), + latents=latents.detach().cpu(), + rng_state=generator.get_state().cpu(), + ) + del decoded + pipe._current_timestep = None diff --git a/worldcrafter/ucpe/__init__.py b/worldcrafter/ucpe/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..364b186082e043e0dec7810e9761c67df10ea85c --- /dev/null +++ b/worldcrafter/ucpe/__init__.py @@ -0,0 +1,13 @@ +from .bridge import ( + build_ucpe_attention_kwargs_for_chunk, + enable_ucpe_inference_sdpa_attention, + load_ucpe_camera_adapter_weights, + patch_worldcrafter_transformer_ucpe, +) + +__all__ = [ + "build_ucpe_attention_kwargs_for_chunk", + "enable_ucpe_inference_sdpa_attention", + "load_ucpe_camera_adapter_weights", + "patch_worldcrafter_transformer_ucpe", +] diff --git a/worldcrafter/ucpe/attention.py b/worldcrafter/ucpe/attention.py new file mode 100644 index 0000000000000000000000000000000000000000..09a573771ee4fd9ef58476177b30d67b1b183e80 --- /dev/null +++ b/worldcrafter/ucpe/attention.py @@ -0,0 +1,28 @@ +import torch +import torch.nn.functional as F + + +def flash_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + num_heads: int, + compatibility_mode: bool = False, +) -> torch.Tensor: + batch, tokens, channels = q.shape + head_dim = channels // num_heads + q = q.view(batch, tokens, num_heads, head_dim) + k = k.view(batch, tokens, num_heads, head_dim) + v = v.view(batch, tokens, num_heads, head_dim) + if compatibility_mode: + out = F.scaled_dot_product_attention( + q.transpose(1, 2), + k.transpose(1, 2), + v.transpose(1, 2), + is_causal=False, + ).transpose(1, 2) + else: + from ..kernels.attention_dispatch import attn_varlen_func + + out = attn_varlen_func(q, k, v) + return out.reshape(batch, tokens, channels) diff --git a/worldcrafter/ucpe/bridge.py b/worldcrafter/ucpe/bridge.py new file mode 100644 index 0000000000000000000000000000000000000000..a8635957f48eb07232351ca2fcfe533334953254 --- /dev/null +++ b/worldcrafter/ucpe/bridge.py @@ -0,0 +1,333 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any, TYPE_CHECKING + +import numpy as np +import torch +import torch.nn.functional as F +from einops import rearrange, repeat + + +from . import camera as ucpe_cc +from . import prope as prope_torch + +if TYPE_CHECKING: + from ..diffusers.transformer import WorldCrafterTransformer3DModel + + +UcpeSelfAttention = ucpe_cc.UcpeSelfAttention + + +def _flash_attention_sdpa( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + num_heads: int, + compatibility_mode: bool = False, +) -> torch.Tensor: + del compatibility_mode + batch, tokens, channels = q.shape + head_dim = channels // num_heads + q = q.view(batch, tokens, num_heads, head_dim).transpose(1, 2) + k = k.view(batch, tokens, num_heads, head_dim).transpose(1, 2) + v = v.view(batch, tokens, num_heads, head_dim).transpose(1, 2) + out = F.scaled_dot_product_attention(q, k, v, is_causal=False) + return out.transpose(1, 2).reshape(batch, tokens, channels) + + +def enable_ucpe_inference_sdpa_attention() -> None: + ucpe_cc.flash_attention = _flash_attention_sdpa + + +def _resolve_checkpoint_payload_path(checkpoint_path: str | Path) -> Path: + checkpoint_path = Path(checkpoint_path) + if checkpoint_path.is_dir(): + checkpoint_path = checkpoint_path / "camera_adapter.pth" + if not checkpoint_path.exists(): + raise FileNotFoundError(f"Camera adapter not found: {checkpoint_path}") + return checkpoint_path + + +def load_ucpe_checkpoint_state_dict(checkpoint_path: str | Path): + resolved_path = _resolve_checkpoint_payload_path(checkpoint_path) + state_obj = torch.load(resolved_path, map_location="cpu") + if not isinstance(state_obj, dict): + raise RuntimeError( + f"Checkpoint payload at {resolved_path} is expected to be a dict, got {type(state_obj)}" + ) + + for key in ("state_dict", "module"): + payload = state_obj.get(key) + if isinstance(payload, dict): + return payload, resolved_path + + return state_obj, resolved_path + + +def extract_ucpe_camera_adapter_state_dict(checkpoint_path: str | Path): + state_dict, resolved_path = load_ucpe_checkpoint_state_dict(checkpoint_path) + + mapped_state = {} + for key, value in state_dict.items(): + if ".cam_self_attn." in key and ( + key.startswith("pipe.dit.blocks.") or key.startswith("blocks.") + ): + mapped_state[key] = value + + if not mapped_state: + preview_keys = list(state_dict.keys())[:20] + raise RuntimeError( + "No UCPE camera adapter tensors found in checkpoint payload. " + f"checkpoint={checkpoint_path} resolved_payload={resolved_path} " + f"top_level_state_keys_preview={preview_keys}" + ) + return mapped_state, resolved_path + + +def _build_stage_rope_coeffs( + cameras: int, + patches_y: int, + patches_x: int, + head_dim: int, + freq_base: float, + freq_scale: float, + device: torch.device, + dtype: torch.dtype, +): + x_positions = torch.tile(torch.arange(patches_x, device=device), (patches_y * cameras,)) + y_positions = torch.tile( + torch.repeat_interleave(torch.arange(patches_y, device=device), patches_x), + (cameras,), + ) + coeffs_x = prope_torch._rope_precompute_coeffs( + x_positions, + freq_base=freq_base, + freq_scale=freq_scale, + feat_dim=head_dim // 4, + dtype=dtype, + ) + coeffs_y = prope_torch._rope_precompute_coeffs( + y_positions, + freq_base=freq_base, + freq_scale=freq_scale, + feat_dim=head_dim // 4, + dtype=dtype, + ) + return coeffs_x, coeffs_y + + +def patch_worldcrafter_transformer_ucpe( + transformer: WorldCrafterTransformer3DModel, + method: str, + height: int, + width: int, + attn_compress: int = 8, + adaptation_method: str = "parallel", + vae_scale_factor_spatial: int = 8, + attention_cls: type[UcpeSelfAttention] = UcpeSelfAttention, +): + if not any(key in method for key in ("gta", "prope", "relray")): + raise ValueError(f"Only UCPE attention-style methods are supported, got: {method}") + + patch_factor = vae_scale_factor_spatial * transformer.config.patch_size[1] + patches_x = width // patch_factor + patches_y = height // patch_factor + emb_dim = 3 if "absmap" in method else None + + for block in transformer.blocks: + num_heads = block.attn1.heads // attn_compress + if num_heads <= 0: + raise ValueError(f"attn_compress={attn_compress} is too large for heads={block.attn1.heads}") + hidden_dim = block.attn1.to_q.weight.shape[0] + block.cam_self_attn = attention_cls( + hidden_dim, + hidden_dim // attn_compress, + num_heads, + patches_x=patches_x, + patches_y=patches_y, + image_width=width, + image_height=height, + emb_dim=emb_dim, + adaptation_method=adaptation_method, + ) + + transformer.camera_condition = method + return ["cam_self_attn"] + + +def load_ucpe_camera_adapter_weights(transformer: WorldCrafterTransformer3DModel, checkpoint_path: str): + state_dict, resolved_path = extract_ucpe_camera_adapter_state_dict(checkpoint_path) + mapped_state = {} + for key, value in state_dict.items(): + if key.startswith("pipe.dit.blocks.") and ".cam_self_attn." in key: + mapped_state[key.replace("pipe.dit.", "", 1)] = value + elif key.startswith("blocks.") and ".cam_self_attn." in key: + mapped_state[key] = value + + expected_cam_keys = {key for key in transformer.state_dict().keys() if ".cam_self_attn." in key} + mapped_cam_keys = set(mapped_state.keys()) + missing_cam_keys = sorted(expected_cam_keys - mapped_cam_keys) + unexpected_cam_keys = sorted(mapped_cam_keys - expected_cam_keys) + + if not mapped_state: + raise RuntimeError( + f"No UCPE camera adapter tensors found in checkpoint: {checkpoint_path} " + f"(resolved_payload={resolved_path})" + ) + if missing_cam_keys or unexpected_cam_keys: + raise RuntimeError( + "UCPE camera adapter checkpoint does not exactly match the patched transformer. " + f"missing_cam_keys={len(missing_cam_keys)} unexpected_cam_keys={len(unexpected_cam_keys)} " + f"resolved_payload={resolved_path}" + ) + + load_info = transformer.load_state_dict(mapped_state, strict=False) + return { + "resolved_payload_path": str(resolved_path), + "loaded_tensor_keys": len(mapped_state), + "expected_tensor_keys": len(expected_cam_keys), + "missing_cam_keys": missing_cam_keys, + "unexpected_cam_keys": unexpected_cam_keys, + "missing_keys": list(load_info.missing_keys), + "unexpected_keys": list(load_info.unexpected_keys), + } + + +def _relative_pose_chunk( + global_c2w: Any, + *, + chunk_index: int, + window_num_frames: int, + device: torch.device, +) -> torch.Tensor | None: + if isinstance(global_c2w, torch.Tensor): + pose = global_c2w.detach().cpu().numpy() + else: + pose = np.asarray(global_c2w) + if pose.ndim == 3: + pose = pose[None] + if pose.ndim != 4 or pose.shape[0] != 1 or pose.shape[-2:] not in ((3, 4), (4, 4)): + raise ValueError( + "camera_trajectory['c2w'] must be [1,T,3,4] or [1,T,4,4], " + f"got {pose.shape}" + ) + start = int(chunk_index) * int(window_num_frames) + end = start + int(window_num_frames) + if pose.shape[1] < end: + return None + chunk = np.asarray(pose[0, start:end], dtype=np.float64) + homogeneous = np.zeros((window_num_frames, 4, 4), dtype=np.float64) + homogeneous[:, :3, :4] = chunk[:, :3, :4] + homogeneous[:, 3, 3] = 1.0 + relative = np.linalg.inv(homogeneous[0])[None] @ homogeneous + return torch.from_numpy(relative[:, :3, :4].astype(np.float32)).unsqueeze(0).to(device) + + +def build_ucpe_attention_kwargs_for_chunk( + transformer: WorldCrafterTransformer3DModel, + camera_trajectory: dict[str, Any] | None, + height: int, + width: int, + num_latent_frames_per_chunk: int, + chunk_index: int, + vae_scale_factor_temporal: int = 4, + token_grid_height: int | None = None, + token_grid_width: int | None = None, +): + if camera_trajectory is None or getattr(transformer, "camera_condition", "none") == "none": + return None + + x_fov = camera_trajectory["x_fov"] + xi = camera_trajectory["xi"] + method = transformer.camera_condition + + if x_fov.ndim == 0: + x_fov = x_fov.unsqueeze(0) + if xi.ndim == 0: + xi = xi.unsqueeze(0) + + window_num_frames = (num_latent_frames_per_chunk - 1) * vae_scale_factor_temporal + 1 + pose_chunk = _relative_pose_chunk( + camera_trajectory["c2w"], + chunk_index=chunk_index, + window_num_frames=window_num_frames, + device=x_fov.device, + ) + if pose_chunk is None: + return None + if pose_chunk.shape[-2:] == (4, 4): + pose_chunk = pose_chunk[..., :3, :4] + elif pose_chunk.shape[-2:] != (3, 4): + raise ValueError( + "pose_chunk is expected to be [B, T, 3, 4] or [B, T, 4, 4], " + f"got shape={tuple(pose_chunk.shape)}" + ) + pose_chunk = pose_chunk[:, ::vae_scale_factor_temporal].to(dtype=torch.float32) + c2w = torch.eye(4, device=pose_chunk.device, dtype=pose_chunk.dtype) + c2w = repeat(c2w, "... -> B T ...", B=pose_chunk.shape[0], T=pose_chunk.shape[1]).clone() + c2w[..., :3, :4] = pose_chunk + + if "gta" in method or "prope" in method: + raise NotImplementedError("Only relray_absmap is supported.") + + if "relray" not in method: + raise ValueError(f"Unsupported camera condition: {method}") + + attn = transformer.blocks[0].cam_self_attn + grid_h = token_grid_height if token_grid_height is not None else attn.patches_y + grid_w = token_grid_width if token_grid_width is not None else attn.patches_x + + d_cam = ucpe_cc.ucm_unproject_grid_fov( + x_fov=x_fov, + xi=xi, + height=grid_h, + width=grid_w, + device=pose_chunk.device, + dtype=pose_chunk.dtype, + ) + raymats = ucpe_cc.world_to_ray_mats(d_cam, c2w) + viewmats = rearrange(raymats, "B T H W ... -> B (T H W) ...") + + control_camera_dit_input = {"viewmats": viewmats} + + if token_grid_height is not None or token_grid_width is not None: + coeffs_x, coeffs_y = _build_stage_rope_coeffs( + cameras=pose_chunk.shape[1], + patches_y=grid_h, + patches_x=grid_w, + head_dim=attn.head_dim, + freq_base=attn.freq_base, + freq_scale=attn.freq_scale, + device=pose_chunk.device, + dtype=pose_chunk.dtype, + ) + control_camera_dit_input["coeffs_x"] = coeffs_x + control_camera_dit_input["coeffs_y"] = coeffs_y + + if "absmap" in method: + up_map, lat_map = ucpe_cc.compute_up_lat_map( + R=pose_chunk[..., :3, :3], + x_fov=x_fov, + xi=xi, + width=grid_w, + height=grid_h, + device=pose_chunk.device, + ) + cam_emb = torch.cat([up_map, lat_map], dim=-1) + cam_emb = rearrange(cam_emb, "B T H W C -> B (T H W) C") + control_camera_dit_input["cam_emb"] = cam_emb + + expected_token_count = pose_chunk.shape[1] * grid_h * grid_w + if control_camera_dit_input["viewmats"].shape[1] != expected_token_count: + raise ValueError( + "camera viewmats token count mismatch: " + f"expected {expected_token_count}, got {control_camera_dit_input['viewmats'].shape[1]}" + ) + if "cam_emb" in control_camera_dit_input and control_camera_dit_input["cam_emb"].shape[1] != expected_token_count: + raise ValueError( + "camera cam_emb token count mismatch: " + f"expected {expected_token_count}, got {control_camera_dit_input['cam_emb'].shape[1]}" + ) + + return {"camera_control_ucpe_input": control_camera_dit_input} diff --git a/worldcrafter/ucpe/camera.py b/worldcrafter/ucpe/camera.py new file mode 100644 index 0000000000000000000000000000000000000000..e9336ffa7f95c96119ce5f45a73a9d2416c20d5f --- /dev/null +++ b/worldcrafter/ucpe/camera.py @@ -0,0 +1,554 @@ +from functools import lru_cache + +import torch +from torch import nn +from .prope import PropeDotProductAttention +from .attention import flash_attention +from einops import rearrange, repeat, einsum +import torch.nn.functional as F + + +def compute_fx_from_fov_xi( + x_fov: torch.Tensor | float, + xi: torch.Tensor | float, + width: int, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """ + 根据水平视场角 (x_fov) 和 UCM 参数 (xi) 计算相机焦距 fx。 + + Args: + x_fov: float 或 [B] Tensor,水平视场角(单位:度) + xi: float 或 [B] Tensor,UCM 镜面参数 + width: 图像宽度(像素) + device: torch.device + dtype: torch.dtype + + Returns: + fx: [B] Tensor,焦距(像素单位) + """ + + # --- 转为 Tensor --- + def to_tensor_1d(x): + if torch.is_tensor(x): + return x.to(device=device, dtype=dtype).reshape(-1) + return torch.tensor([x], dtype=dtype, device=device) + + x_fov = to_tensor_1d(x_fov) + xi = to_tensor_1d(xi) + + # --- 自动广播 --- + B = max(x_fov.shape[0], xi.shape[0]) + x_fov = x_fov.view(-1).expand(B) + xi = xi.view(-1).expand(B) + + # --- 计算 fx --- + theta = torch.deg2rad(0.5 * x_fov) + eps = torch.finfo(dtype).eps + denom = torch.sin(theta).clamp_min(eps) + fx = (width * 0.5) * (torch.cos(theta) + xi) / denom + return fx + + +def project_ucm_points_fov(X, Y, Z, x_fov, xi, height, width): + fx = compute_fx_from_fov_xi(x_fov, xi, width, X.device, X.dtype) + return project_ucm_points(X, Y, Z, fx, fx, width / 2, height / 2, xi) + + +def project_ucm_points(X, Y, Z, fx, fy, cx, cy, xi): + def broadcast_param(param): + if not torch.is_tensor(param): + param = torch.tensor(param, device=X.device, dtype=X.dtype) + else: + param = param.to(device=X.device, dtype=X.dtype) + if param.ndim == 0: + return param + flat = param.reshape(-1) + if flat.numel() == 1: + return flat.view(1) + if X.ndim >= 1 and flat.numel() == X.shape[0]: + return flat.view(flat.shape[0], *([1] * (X.ndim - 1))) + return param + + fx = broadcast_param(fx) + fy = broadcast_param(fy) + cx = broadcast_param(cx) + cy = broadcast_param(cy) + xi = broadcast_param(xi) + radius = torch.sqrt(X * X + Y * Y + Z * Z) + denominator = Z + xi * radius + du = fx * (X / denominator) + cx + dv = fy * (Y / denominator) + cy + return du, dv + + +def _pixel_grid( + *, + height: int, + width: int, + batch: int, + dtype: torch.dtype, + device: torch.device, +) -> torch.Tensor: + xs = torch.linspace(0, width - 1, width, dtype=dtype, device=device) + ys = torch.linspace(0, height - 1, height, dtype=dtype, device=device) + ys, xs = torch.meshgrid([ys, xs], indexing="ij") + grid = torch.stack((xs, ys, torch.ones_like(xs)), dim=2) + return repeat(grid, "... -> b ...", b=batch) + + +@lru_cache(maxsize=128) +def _ucm_unproject_grid( + height: int, + width: int, + fx: float | torch.Tensor, + fy: float | torch.Tensor, + cx: float | torch.Tensor, + cy: float | torch.Tensor, + xi: float | torch.Tensor, + dtype: torch.dtype = torch.float32, + device: torch.device = torch.device("cpu"), + y_down: bool = True, +) -> torch.Tensor: + scalar_input = all(not torch.is_tensor(value) for value in (fx, fy, cx, cy, xi)) + + def tensor_1d(value): + if torch.is_tensor(value): + return value.to(device=device, dtype=dtype) + return torch.tensor([value], device=device, dtype=dtype) + + fx, fy, cx, cy, xi = map(tensor_1d, (fx, fy, cx, cy, xi)) + xs = torch.linspace(0, width - 1, width, dtype=dtype, device=device) + ys = torch.linspace(0, height - 1, height, dtype=dtype, device=device) + ys, xs = torch.meshgrid([ys, xs], indexing="ij") + grid = torch.stack((xs, ys, torch.ones_like(xs)), dim=2) + grid = repeat(grid, "... -> b ...", b=fx.shape[0]) + + x = (grid[..., 0] - cx[:, None, None]) / fx[:, None, None] + y = (grid[..., 1] - cy[:, None, None]) / fy[:, None, None] + if not y_down: + y = -y + r2 = x * x + y * y + alpha = xi[:, None, None] + torch.sqrt( + 1 + (1 - xi[:, None, None] * xi[:, None, None]) * r2 + ) + gamma = alpha / (1 + r2) + directions = torch.stack((gamma * x, gamma * y, gamma - xi[:, None, None]), dim=-1) + return directions[0] if scalar_input else directions + + +def ucm_unproject_grid_fov( + x_fov: float | torch.Tensor, + xi: float | torch.Tensor, + height: int, + width: int, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """ + 计算每个样本的相机方向向量 (UCM model, 用视场角定义)。 + 支持 float 或 [B] Tensor 的混合输入。 + - 若全为 float → 返回 [H, W, 3] + - 若任意为 [B] → 返回 [B, H, W, 3] + """ + if isinstance(device, str): + device = torch.device(device) + + is_batched = any( + torch.is_tensor(p) and p.reshape(-1).numel() > 1 for p in [x_fov, xi] + ) + + # --- 计算 fx, fy --- + fx = compute_fx_from_fov_xi(x_fov, xi, width, device, dtype) + fy = fx + xi_grid = ( + xi.to(device=device, dtype=dtype).reshape(-1) + if torch.is_tensor(xi) + else torch.tensor([xi], dtype=dtype, device=device) + ) + + d_cam = _ucm_unproject_grid( + height=height, + width=width, + fx=fx, + fy=fy, + cx=width / 2, + cy=height / 2, + xi=xi_grid, + dtype=dtype, + device=device, + y_down=True, + ) + + # --- 输出 shape 控制 --- + if not is_batched: + d_cam = d_cam[0] # [H, W, 3] + + return d_cam + + +def d_cam_to_angles(d_cam: torch.Tensor) -> torch.Tensor: + """ + 将方向向量 [x, y, z] 转换为 [azimuth, elevation]。 + 坐标系:z前,x右,y下(符合 UCM 投影输出) + + 输入: d_cam: [B, H, W, 3] + 输出: angles: [B, H, W, 2] — azimuth, elevation (单位: 弧度) + """ + d_unit = F.normalize(d_cam, dim=-1) # [B, H, W, 3] + + x = d_unit[..., 0] # right + y = d_unit[..., 1] # down + z = d_unit[..., 2] # forward + + # yaw / azimuth: angle in xz-plane + azimuth = torch.atan2(x, z) # ∈ [-π, π] + + # pitch / elevation: angle above xz-plane + elevation = -torch.asin(y) # y 向下 → elevation = -asin(y) + + return torch.stack([azimuth, elevation], dim=-1) # [B, H, W, 2] + + +def world_to_ray_mats( + d_cam: torch.Tensor, # [B, H, W, 3] + c2w: torch.Tensor, # [B, T, 4, 4] +) -> torch.Tensor: + """ + 构造每条 ray 的世界到 ray 局部坐标系的变换矩阵 world2ray。 + 坐标系定义: + - z: ray direction + - x: cam_y × ray_dir + - y: z × x + 返回: + raymats: [B, T, H, W, 4, 4] + """ + if d_cam.ndim == 3: + d_cam = d_cam.unsqueeze(0) + if c2w.ndim == 3: + c2w = c2w.unsqueeze(0) + if d_cam.ndim != 4 or d_cam.shape[-1] != 3: + raise ValueError( + f"d_cam must have shape [H,W,3] or [B,H,W,3], got {tuple(d_cam.shape)}" + ) + if c2w.ndim != 4 or c2w.shape[-2:] != (4, 4): + raise ValueError( + f"c2w must have shape [T,4,4] or [B,T,4,4], got {tuple(c2w.shape)}" + ) + if d_cam.shape[0] == 1 and c2w.shape[0] != 1: + d_cam = d_cam.expand(c2w.shape[0], -1, -1, -1) + elif c2w.shape[0] == 1 and d_cam.shape[0] != 1: + c2w = c2w.expand(d_cam.shape[0], -1, -1, -1) + elif d_cam.shape[0] != c2w.shape[0]: + raise ValueError( + f"d_cam and c2w batch mismatch: {d_cam.shape[0]} vs {c2w.shape[0]}" + ) + + B, H, W, _ = d_cam.shape + T = c2w.shape[1] + device = d_cam.device + dtype = d_cam.dtype + + # --- Expand ray dirs across frames --- + # [B,H,W,3] -> [B,T,H,W,3] + d_cam = repeat(d_cam, "b h w c -> b t h w c", t=T) + + # extract camera R,t + R_cam = c2w[..., :3, :3] # [B,T,3,3] + t_cam = c2w[..., :3, 3] # [B,T,3] + + # --- d_world: rotate ray directions into world --- + d_world = einsum(R_cam, d_cam, "b t i j, b t h w j -> b t h w i") + + # camera y-axis from each view + cam_y = R_cam[..., :, 1] # [B,T,3] + cam_y = repeat(cam_y, "b t c -> b t h w c", h=H, w=W) + + # === Construct orthonormal ray-local axes === + z_ray = F.normalize(d_world, dim=-1, eps=1e-6) + x_ray = torch.cross(cam_y, z_ray, dim=-1) + x_ray = F.normalize(x_ray, dim=-1, eps=1e-6) + y_ray = torch.cross(z_ray, x_ray, dim=-1) + y_ray = F.normalize(y_ray, dim=-1, eps=1e-6) + + # local->world rotation + R_l2w = torch.stack([x_ray, y_ray, z_ray], dim=-1) # [B,T,H,W,3,3] + + # world->local rotation (transpose) + R_w2l = rearrange(R_l2w, "b t h w i j -> b t h w j i") # ✅ + + # broadcast camera center + t_world = repeat(t_cam, "b t c -> b t h w c", h=H, w=W) + + # world->local translation + t_w2l = -einsum(R_w2l, t_world, "b t h w i j, b t h w j -> b t h w i") + + # assemble transform matrix + raymats = torch.zeros(B, T, H, W, 4, 4, device=device, dtype=dtype) + raymats[..., :3, :3] = R_w2l + raymats[..., :3, 3] = t_w2l + raymats[..., 3, 3] = 1.0 + + # NaN handling + mask = torch.isnan(d_world).any(-1) + raymats[mask] = torch.eye(4, device=device, dtype=dtype) + + return raymats + + +def compute_up_lat_map( + R: torch.Tensor, + x_fov: torch.Tensor, + xi: torch.Tensor, + height: int, + width: int, + device: torch.device = torch.device("cpu"), + delta: float = 0.1, +): + """ + 计算 up_map 和 lat_map。 + + Args: + R: [B, T, 3, 3] 相机 c2w 旋转矩阵 + x_fov: [B] 或 [B,T] 水平视场角(度) + xi: [B] 或 [B,T] UCM 参数 + height: int,图像/patch 高度 + width: int,图像/patch 宽度 + device: torch.device + delta: float,小旋转角度(弧度) + Returns: + up_map: [B, T, H, W, 2] 单位向量 map + lat_map: [B, T, H, W, 1] 纬度 map + """ + B, T, _, _ = R.shape + dtype = R.dtype + R = R.float() + + # Step1:生成每像素射线方向(相机坐标系) + d_cam = ucm_unproject_grid_fov( + x_fov=x_fov, + xi=xi, + height=height, + width=width, + device=device, + dtype=torch.float32, + ) # [B, H, W, 3] + if d_cam.ndim == 3: + d_cam = d_cam.unsqueeze(0) # [B, H, W, 3] + mask = d_cam.isnan().any(dim=-1, keepdim=True) # [B, H, W, 1] + + # Step2:从相机系旋转到世界系 + d_cam_exp = repeat(d_cam, "B H W C -> B T H W C", T=T) # [B, T, H, W, 3] + d_world = torch.einsum("btij,bthwj->bthwi", R, d_cam_exp) + d_world = d_world / torch.clamp_min(d_world.norm(dim=-1, keepdim=True), 1e-8) + + # Step3:计算纬度 map + Xw, Yw, Zw = d_world[..., 0], d_world[..., 1], d_world[..., 2] + lat_map = torch.atan2(-Yw, torch.sqrt(Xw**2 + Zw**2)).unsqueeze( + -1 + ) # [B, T, H, W, 1] + + # Step4:计算 up_map + v = d_world # 已归一化 + up_world = torch.tensor( + [0, -1, 0], device=device, dtype=torch.float32 + ) # 世界上方方向(+Y 向下设定) + k = torch.cross( + v, up_world.unsqueeze(0).unsqueeze(0).unsqueeze(0).expand_as(v), dim=-1 + ) + k = k / torch.clamp_min(k.norm(dim=-1, keepdim=True), 1e-8) + + delta = torch.tensor(delta, device=device, dtype=torch.float32) + cos_eps = torch.cos(delta) + sin_eps = torch.sin(delta) + # Rodrigues 公式旋转 v → v_rot + v_rot = ( + v * cos_eps + + torch.cross(k, v, dim=-1) * sin_eps + + k * (k * (v * 1).sum(dim=-1, keepdim=True)) * (1 - cos_eps) + ) + + dirs_cam = torch.einsum("btij,bthwj->bthwi", R.transpose(-1, -2), v_rot) + Xs, Ys, Zs = dirs_cam[..., 0], dirs_cam[..., 1], dirs_cam[..., 2] + + du, dv = project_ucm_points_fov( + Xs, + Ys, + Zs, + x_fov=x_fov.float(), + xi=xi.float(), + height=height, + width=width, + ) + grid = _pixel_grid( + height=height, + width=width, + batch=B, + dtype=torch.float32, + device=device, + ) # [B, H, W, 3] + grid_x = grid[..., 0].unsqueeze(1) # [B,1,H,W] + grid_y = grid[..., 1].unsqueeze(1) + + up_map = torch.stack((du - grid_x, dv - grid_y), dim=-1) # [B, T, H, W, 2] + up_map = up_map / torch.clamp_min(up_map.norm(dim=-1, keepdim=True), 1e-8) + + up_map = up_map.to(dtype=dtype) + lat_map = lat_map.to(dtype=dtype) + + # 扩 mask 到同 shape + mask_exp2 = mask.unsqueeze(1).expand(B, T, height, width, 1) + up_map = up_map.masked_fill(mask_exp2, 0.0) + lat_map = lat_map.masked_fill(mask_exp2, 0.0) + + return up_map, lat_map + + +class UcpeSelfAttention(nn.Module): + def __init__( + self, + dim: int, + attn_dim: int, + num_heads: int, + patches_x: int = 8, + patches_y: int = 8, + image_width: int = 128, + image_height: int = 128, + freq_base: float = 100.0, + freq_scale: float = 1.0, + precompute_coeffs: bool = True, + emb_dim: int | None = None, + adaptation_method: str = "parallel", + ): + super().__init__() + assert dim % num_heads == 0 + self.dim = dim + self.attn_dim = attn_dim + self.num_heads = num_heads + self.head_dim = attn_dim // num_heads + self.patches_x = patches_x + self.patches_y = patches_y + self.image_width = image_width + self.image_height = image_height + self.freq_base = freq_base + self.freq_scale = freq_scale + self.adaptation_method = adaptation_method + + self.q_proj = nn.Linear(dim, attn_dim) + self.k_proj = nn.Linear(dim, attn_dim) + self.v_proj = nn.Linear(dim, attn_dim) + self.out_proj = nn.Linear(attn_dim, dim) + if emb_dim is not None: + self.cam_encoder = nn.Linear(emb_dim, dim) + + nn.init.zeros_(self.out_proj.weight) + nn.init.zeros_(self.out_proj.bias) + + # 初始化 PRoPE attention 模块(带 precomputed coeffs) + self.prope_attn = PropeDotProductAttention( + head_dim=self.head_dim, + patches_x=patches_x, + patches_y=patches_y, + image_width=image_width, + image_height=image_height, + freq_base=freq_base, + freq_scale=freq_scale, + precompute_coeffs=precompute_coeffs, + ) + + def forward(self, x: torch.Tensor, control_camera_dit_input: dict): + """ + Args: + x: (B, T, D) — input tokens + control_camera_dit_input: dict with keys: + - viewmats: (B, N, 4, 4) + - K: (B, N, 3, 3) + """ + B, T, D = x.shape + N = control_camera_dit_input["viewmats"].shape[1] # number of cameras + H, W = self.patches_y, self.patches_x + assert ( + T == N * H * W or T == N + ), f"Expected token shape ({N}×{H}×{W} or {N}), got {T}" + + # Camera geometry intentionally stays fp32, while WorldCrafter hidden states + # and the UCPE adapter may independently be bf16/fp16 or fp32. Keep + # every Linear input in its weight dtype and restore the model hidden + # dtype at the attention/residual boundaries. + hidden_dtype = x.dtype + projection_dtype = self.q_proj.weight.dtype + projected_x = x.to(dtype=projection_dtype) + + if hasattr(self, "cam_encoder") and "cam_emb" in control_camera_dit_input: + cam_emb = control_camera_dit_input["cam_emb"].to( + dtype=self.cam_encoder.weight.dtype + ) + y = self.cam_encoder(cam_emb) + if y.shape[1] != T: + hw = T // cam_emb.shape[1] + y = repeat(y, "b f d -> b (f hw) d", hw=hw) + projected_x = projected_x + y.to(dtype=projection_dtype) + + # Project Q, K, V + q = ( + self.q_proj(projected_x) + .view(B, T, self.num_heads, self.head_dim) + .transpose(1, 2) + ) # [B, H, T, D_head] + k = ( + self.k_proj(projected_x) + .view(B, T, self.num_heads, self.head_dim) + .transpose(1, 2) + ) + v = ( + self.v_proj(projected_x) + .view(B, T, self.num_heads, self.head_dim) + .transpose(1, 2) + ) + + # Precompute camera-specific functions (only once per batch) + self.prope_attn._precompute_and_cache_apply_fns( + viewmats=control_camera_dit_input["viewmats"], + Ks=control_camera_dit_input.get("K", None), + coeffs_x=control_camera_dit_input.get("coeffs_x", None), + coeffs_y=control_camera_dit_input.get("coeffs_y", None), + ) + + # PRoPE matrices and coefficients are deliberately built in the FP32 + # camera-geometry dtype. Mixed-precision Linear/autocast outputs can be + # BF16/FP16 even when the adapter weights are FP32, so align + # Q/K/V before the einsum instead of relying on implicit promotion. + geometry_dtype = control_camera_dit_input["viewmats"].dtype + q = self.prope_attn._apply_to_q(q.to(dtype=geometry_dtype)) + k = self.prope_attn._apply_to_kv(k.to(dtype=geometry_dtype)) + v = self.prope_attn._apply_to_kv(v.to(dtype=geometry_dtype)) + + q = q.to(dtype=hidden_dtype) + k = k.to(dtype=hidden_dtype) + v = v.to(dtype=hidden_dtype) + + # Rearrange to [B, T, D] for flash_attention input + q = rearrange(q, "b h t d -> b t (h d)") + k = rearrange(k, "b h t d -> b t (h d)") + v = rearrange(v, "b h t d -> b t (h d)") + + # Fast attention (Flash/Sage/SDPA fallback) + out = flash_attention( + q, + k, + v, + num_heads=self.num_heads, + compatibility_mode=q.dtype not in (torch.float16, torch.bfloat16), + ) + + # reshape back + out = rearrange(out, "b t (h d) -> b h t d", h=self.num_heads) + + # Apply inverse transform for PRoPE + out = out.to(dtype=control_camera_dit_input["viewmats"].dtype) + out = self.prope_attn._apply_to_o(out) + + # Final projection + out = out.transpose(1, 2).reshape(B, T, -1).to(dtype=self.out_proj.weight.dtype) + return self.out_proj(out).to(dtype=hidden_dtype) diff --git a/worldcrafter/ucpe/prope.py b/worldcrafter/ucpe/prope.py new file mode 100644 index 0000000000000000000000000000000000000000..fc8287b2694ab6103c0079e9755579cf80ca50d6 --- /dev/null +++ b/worldcrafter/ucpe/prope.py @@ -0,0 +1,476 @@ +# MIT License +# +# Copyright (c) Authors of +# "Cameras as Relative Positional Encoding" https://arxiv.org/pdf/2507.10496 +# +# 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. + +# How to use PRoPE attention for self-attention: +# +# 1. Easiest way (fast): +# attn = PropeDotProductAttention(...) +# o = attn(q, k, v, viewmats, Ks) +# +# 2. More flexible way (fast): +# attn = PropeDotProductAttention(...) +# attn._precompute_and_cache_apply_fns(viewmats, Ks) +# q = attn._apply_to_q(q) +# k = attn._apply_to_kv(k) +# v = attn._apply_to_kv(v) +# o = F.scaled_dot_product_attention(q, k, v, **kwargs) +# o = attn._apply_to_o(o) +# +# 3. The most flexible way (but slower because repeated computation of RoPE coefficients): +# o = prope_dot_product_attention(q, k, v, ...) +# +# How to use PRoPE attention for cross-attention: +# +# attn_src = PropeDotProductAttention(...) +# attn_tgt = PropeDotProductAttention(...) +# attn_src._precompute_and_cache_apply_fns(viewmats_src, Ks_src) +# attn_tgt._precompute_and_cache_apply_fns(viewmats_tgt, Ks_tgt) +# q_src = attn_src._apply_to_q(q_src) +# k_tgt = attn_tgt._apply_to_kv(k_tgt) +# v_tgt = attn_tgt._apply_to_kv(v_tgt) +# o_src = F.scaled_dot_product_attention(q_src, k_tgt, v_tgt, **kwargs) +# o_src = attn_src._apply_to_o(o_src) + +from functools import partial +from typing import Callable, Optional, Tuple, List + +import torch +import torch.nn.functional as F + + +class PropeDotProductAttention(torch.nn.Module): + """PRoPE attention with precomputed RoPE coefficients.""" + + coeffs_x_0: torch.Tensor + coeffs_x_1: torch.Tensor + coeffs_y_0: torch.Tensor + coeffs_y_1: torch.Tensor + + def __init__( + self, + head_dim: int, + patches_x: int, + patches_y: int, + image_width: int, + image_height: int, + freq_base: float = 100.0, + freq_scale: float = 1.0, + precompute_coeffs: bool = True, + ): + super().__init__() + self.head_dim = head_dim + self.patches_x = patches_x + self.patches_y = patches_y + self.image_width = image_width + self.image_height = image_height + + if precompute_coeffs: + coeffs_x: Tuple[torch.Tensor, torch.Tensor] = _rope_precompute_coeffs( + torch.tile(torch.arange(patches_x), (patches_y,)), + freq_base=freq_base, + freq_scale=freq_scale, + feat_dim=head_dim // 4, + ) + coeffs_y: Tuple[torch.Tensor, torch.Tensor] = _rope_precompute_coeffs( + torch.repeat_interleave(torch.arange(patches_y), patches_x), + freq_base=freq_base, + freq_scale=freq_scale, + feat_dim=head_dim // 4, + ) + # Do not save coeffs to checkpoint as `cameras` might change during testing. + self.register_buffer("coeffs_x_0", coeffs_x[0], persistent=False) + self.register_buffer("coeffs_x_1", coeffs_x[1], persistent=False) + self.register_buffer("coeffs_y_0", coeffs_y[0], persistent=False) + self.register_buffer("coeffs_y_1", coeffs_y[1], persistent=False) + else: + self.coeffs_x_0 = None + self.coeffs_x_1 = None + self.coeffs_y_0 = None + self.coeffs_y_1 = None + + # override load_state_dict to not load coeffs if they exist (for backward compatibility) + def load_state_dict(self, state_dict, strict=True): + # remove coeffs from state_dict + state_dict.pop("coeffs_x_0", None) + state_dict.pop("coeffs_x_1", None) + state_dict.pop("coeffs_y_0", None) + state_dict.pop("coeffs_y_1", None) + super().load_state_dict(state_dict, strict) + + def forward( + self, + q: torch.Tensor, # (batch, num_heads, seqlen, head_dim) + k: torch.Tensor, # (batch, num_heads, seqlen, head_dim) + v: torch.Tensor, # (batch, num_heads, seqlen, head_dim) + viewmats: torch.Tensor, # (batch, cameras, 4, 4) + Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3) + **kwargs, + ) -> torch.Tensor: + return prope_dot_product_attention( + q, + k, + v, + viewmats=viewmats, + Ks=Ks, + patches_x=self.patches_x, + patches_y=self.patches_y, + image_width=self.image_width, + image_height=self.image_height, + coeffs_x=(self.coeffs_x_0, self.coeffs_x_1), + coeffs_y=(self.coeffs_y_0, self.coeffs_y_1), + **kwargs, + ) + + def _precompute_and_cache_apply_fns( + self, viewmats: torch.Tensor, Ks: Optional[torch.Tensor], + coeffs_x: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + coeffs_y: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + ): + (batch, cameras, _, _) = viewmats.shape + assert viewmats.shape == (batch, cameras, 4, 4) + assert Ks is None or Ks.shape == (batch, cameras, 3, 3) + + self.apply_fn_q, self.apply_fn_kv, self.apply_fn_o = _prepare_apply_fns( + head_dim=self.head_dim, + viewmats=viewmats, + Ks=Ks, + patches_x=self.patches_x, + patches_y=self.patches_y, + image_width=self.image_width, + image_height=self.image_height, + coeffs_x=(self.coeffs_x_0, self.coeffs_x_1) if coeffs_x is None else coeffs_x, + coeffs_y=(self.coeffs_y_0, self.coeffs_y_1) if coeffs_y is None else coeffs_y, + ) + + def _apply_to_q(self, q: torch.Tensor) -> torch.Tensor: + (batch, num_heads, seqlen, head_dim) = q.shape + assert head_dim == self.head_dim + assert q.shape == (batch, num_heads, seqlen, head_dim) + assert self.apply_fn_q is not None + return self.apply_fn_q(q) + + def _apply_to_kv(self, kv: torch.Tensor) -> torch.Tensor: + (batch, num_heads, seqlen, head_dim) = kv.shape + assert head_dim == self.head_dim + assert kv.shape == (batch, num_heads, seqlen, head_dim) + assert self.apply_fn_kv is not None + return self.apply_fn_kv(kv) + + def _apply_to_o(self, o: torch.Tensor) -> torch.Tensor: + (batch, num_heads, seqlen, head_dim) = o.shape + assert head_dim == self.head_dim + assert o.shape == (batch, num_heads, seqlen, head_dim) + assert self.apply_fn_o is not None + return self.apply_fn_o(o) + + +def prope_dot_product_attention( + q: torch.Tensor, # (batch, num_heads, seqlen, head_dim) + k: torch.Tensor, # (batch, num_heads, seqlen, head_dim) + v: torch.Tensor, # (batch, num_heads, seqlen, head_dim) + *, + viewmats: torch.Tensor, # (batch, cameras, 4, 4) + Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3) + patches_x: int, # How many patches wide is each image? + patches_y: int, # How many patches tall is each image? + image_width: int, # Width of the image. Used to normalize intrinsics. + image_height: int, # Height of the image. Used to normalize intrinsics. + coeffs_x: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + coeffs_y: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + **kwargs, +) -> torch.Tensor: + """Similar to torch.nn.functional.scaled_dot_product_attention, but applies PRoPE-style + positional encoding. + + Currently, we assume that the sequence length is equal to: + + cameras * patches_x * patches_y + + And token ordering allows the `(seqlen,)` axis to be reshaped into + `(cameras, patches_x, patches_y)`. + """ + # We're going to assume self-attention: all inputs are the same shape. + (batch, num_heads, seqlen, head_dim) = q.shape + cameras = viewmats.shape[1] + assert q.shape == k.shape == v.shape + assert viewmats.shape == (batch, cameras, 4, 4) + assert Ks is None or Ks.shape == (batch, cameras, 3, 3) + assert seqlen == cameras * patches_x * patches_y + + apply_fn_q, apply_fn_kv, apply_fn_o = _prepare_apply_fns( + head_dim=head_dim, + viewmats=viewmats, + Ks=Ks, + patches_x=patches_x, + patches_y=patches_y, + image_width=image_width, + image_height=image_height, + coeffs_x=coeffs_x, + coeffs_y=coeffs_y, + ) + + out = F.scaled_dot_product_attention( + query=apply_fn_q(q), + key=apply_fn_kv(k), + value=apply_fn_kv(v), + **kwargs, + ) + out = apply_fn_o(out) + assert out.shape == (batch, num_heads, seqlen, head_dim) + return out + + +def _prepare_apply_fns( + head_dim: int, # Q/K/V will have this last dimension + viewmats: torch.Tensor, # (batch, cameras, 4, 4) + Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3) + patches_x: int, # How many patches wide is each image? + patches_y: int, # How many patches tall is each image? + image_width: int, # Width of the image. Used to normalize intrinsics. + image_height: int, # Height of the image. Used to normalize intrinsics. + coeffs_x: Optional[torch.Tensor] = None, + coeffs_y: Optional[torch.Tensor] = None, +) -> Tuple[ + Callable[[torch.Tensor], torch.Tensor], + Callable[[torch.Tensor], torch.Tensor], + Callable[[torch.Tensor], torch.Tensor], +]: + """Prepare transforms for PRoPE-style positional encoding.""" + device = viewmats.device + (batch, cameras, _, _) = viewmats.shape + dtype = viewmats.dtype + + # Normalize camera intrinsics. + if Ks is not None: + Ks_norm = torch.zeros_like(Ks) + Ks_norm[..., 0, 0] = Ks[..., 0, 0] / image_width + Ks_norm[..., 1, 1] = Ks[..., 1, 1] / image_height + Ks_norm[..., 0, 2] = Ks[..., 0, 2] / image_width - 0.5 + Ks_norm[..., 1, 2] = Ks[..., 1, 2] / image_height - 0.5 + Ks_norm[..., 2, 2] = 1.0 + del Ks + + # Compute the camera projection matrices we use in PRoPE. + # - K is an `image<-camera` transform. + # - viewmats is a `camera<-world` transform. + # - P = lift(K) @ viewmats is an `image<-world` transform. + P = torch.einsum("...ij,...jk->...ik", _lift_K(Ks_norm), viewmats) + P_T = P.transpose(-1, -2) + P_inv = torch.einsum( + "...ij,...jk->...ik", + _invert_SE3(viewmats), + _lift_K(_invert_K(Ks_norm)), + ) + + else: + # GTA formula. P is `camera<-world` transform. + P = viewmats + P_T = P.transpose(-1, -2) + P_inv = _invert_SE3(viewmats) + + assert P.shape == P_inv.shape == (batch, cameras, 4, 4) + + # Precompute cos/sin terms for RoPE. We use tiles/repeats for 'row-major' + # broadcasting. + if coeffs_x is None: + coeffs_x = _rope_precompute_coeffs( + torch.tile(torch.arange(patches_x, device=device), (patches_y * cameras,)), + freq_base=100.0, + freq_scale=1.0, + feat_dim=head_dim // 4, + dtype=dtype, + ) + if coeffs_y is None: + coeffs_y = _rope_precompute_coeffs( + torch.tile( + torch.repeat_interleave( + torch.arange(patches_y, device=device), patches_x + ), + (cameras,), + ), + freq_base=100.0, + freq_scale=1.0, + feat_dim=head_dim // 4, + dtype=dtype, + ) + + # Block-diagonal transforms to the inputs and outputs of the attention operator. + assert head_dim % 4 == 0 + transforms_q = [ + (partial(_apply_tiled_projmat, matrix=P_T), head_dim // 2), + (partial(_rope_apply_coeffs, coeffs=coeffs_x), head_dim // 4), + (partial(_rope_apply_coeffs, coeffs=coeffs_y), head_dim // 4), + ] + transforms_kv = [ + (partial(_apply_tiled_projmat, matrix=P_inv), head_dim // 2), + (partial(_rope_apply_coeffs, coeffs=coeffs_x), head_dim // 4), + (partial(_rope_apply_coeffs, coeffs=coeffs_y), head_dim // 4), + ] + transforms_o = [ + (partial(_apply_tiled_projmat, matrix=P), head_dim // 2), + (partial(_rope_apply_coeffs, coeffs=coeffs_x, inverse=True), head_dim // 4), + (partial(_rope_apply_coeffs, coeffs=coeffs_y, inverse=True), head_dim // 4), + ] + + apply_fn_q = partial(_apply_block_diagonal, func_size_pairs=transforms_q) + apply_fn_kv = partial(_apply_block_diagonal, func_size_pairs=transforms_kv) + apply_fn_o = partial(_apply_block_diagonal, func_size_pairs=transforms_o) + return apply_fn_q, apply_fn_kv, apply_fn_o + + +def _apply_tiled_projmat( + feats: torch.Tensor, # (batch, num_heads, seqlen, feat_dim) + matrix: torch.Tensor, # (batch, cameras, D, D) or (batch, seqlen, D, D) +) -> torch.Tensor: + """Apply projection matrix to features.""" + # - seqlen => (cameras, patches_x * patches_y) + # - feat_dim => (feat_dim // 4, 4) + (batch, num_heads, seqlen, feat_dim) = feats.shape + D = matrix.shape[-1] + assert feat_dim % D == 0, f"feat_dim={feat_dim} must be divisible by D={D}" + + if matrix.shape[1] == seqlen: + # Per-ray projection: matrix shape [B, seqlen, D, D] + feats_ = feats.view(batch, num_heads, seqlen, feat_dim // D, D) + out = torch.einsum("btij,bntpj->bntpi", matrix, feats_) + return out.reshape(feats.shape) + + # Per-camera projection (original implementation) + cameras = matrix.shape[1] + assert seqlen > cameras and seqlen % cameras == 0 + assert matrix.shape == (batch, cameras, D, D) + assert feat_dim % D == 0 + return torch.einsum( + "bcij,bncpkj->bncpki", + matrix, + feats.reshape((batch, num_heads, cameras, -1, feat_dim // D, D)), + ).reshape(feats.shape) + + +def _rope_precompute_coeffs( + positions: torch.Tensor, # (seqlen,) + freq_base: float, + freq_scale: float, + feat_dim: int, + dtype: torch.dtype = torch.float32, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Precompute RoPE coefficients.""" + assert len(positions.shape) == 1 + assert feat_dim % 2 == 0 + num_freqs = feat_dim // 2 + freqs = freq_scale * ( + freq_base + ** ( + -torch.arange(num_freqs, device=positions.device)[None, None, None, :] + / num_freqs + ) + ) + angles = positions[None, None, :, None] * freqs + # Shape should be: `(batch, num_heads, seqlen, num_freqs)`; we're + # broadcasting across `batch` and `num_heads`. + assert angles.shape == (1, 1, positions.shape[0], num_freqs) + return torch.cos(angles).to(dtype), torch.sin(angles).to(dtype) + + +def _rope_apply_coeffs( + feats: torch.Tensor, # (batch, num_heads, seqlen, feat_dim) + coeffs: Tuple[torch.Tensor, torch.Tensor], + inverse: bool = False, +) -> torch.Tensor: + """Apply RoPE coefficients to features. We adopt a 'split' ordering + convention. (in contrast to 'interleaved')""" + cos, sin = coeffs + # We allow (cos, sin) to be either with shape (1, 1, seqlen, feat_dim // 2), + # or (1, 1, seqlen_per_image, feat_dim // 2) and we repeat it to + # match the shape of feats. + if cos.shape[2] != feats.shape[2]: + n_repeats = feats.shape[2] // cos.shape[2] + cos = cos.repeat(1, 1, n_repeats, 1) + sin = sin.repeat(1, 1, n_repeats, 1) + assert len(feats.shape) == len(cos.shape) == len(sin.shape) == 4 + assert cos.shape[-1] == sin.shape[-1] == feats.shape[-1] // 2 + x_in = feats[..., : feats.shape[-1] // 2] + y_in = feats[..., feats.shape[-1] // 2 :] + return torch.cat( + ( + [cos * x_in + sin * y_in, -sin * x_in + cos * y_in] + if not inverse + else [cos * x_in - sin * y_in, sin * x_in + cos * y_in] + ), + dim=-1, + ) + + +def _apply_block_diagonal( + feats: torch.Tensor, # (..., dim) + func_size_pairs: List[Tuple[Callable[[torch.Tensor], torch.Tensor], int]], +) -> torch.Tensor: + """Apply a block-diagonal function to an input array. + + Each function is specified as a tuple with form: + + ((Tensor) -> Tensor, int) + + Where the integer is the size of the input to the function. + """ + funcs, block_sizes = zip(*func_size_pairs) + assert feats.shape[-1] == sum(block_sizes) + x_blocks = torch.split(feats, block_sizes, dim=-1) + out = torch.cat( + [f(x_block) for f, x_block in zip(funcs, x_blocks)], + dim=-1, + ) + assert out.shape == feats.shape, "Input/output shapes should match." + return out + + +def _invert_SE3(transforms: torch.Tensor) -> torch.Tensor: + """Invert a 4x4 SE(3) matrix.""" + assert transforms.shape[-2:] == (4, 4) + Rinv = transforms[..., :3, :3].transpose(-1, -2) + out = torch.zeros_like(transforms) + out[..., :3, :3] = Rinv + out[..., :3, 3] = -torch.einsum("...ij,...j->...i", Rinv, transforms[..., :3, 3]) + out[..., 3, 3] = 1.0 + return out + + +def _lift_K(Ks: torch.Tensor) -> torch.Tensor: + """Lift 3x3 matrices to homogeneous 4x4 matrices.""" + assert Ks.shape[-2:] == (3, 3) + out = torch.zeros(Ks.shape[:-2] + (4, 4), device=Ks.device, dtype=Ks.dtype) + out[..., :3, :3] = Ks + out[..., 3, 3] = 1.0 + return out + + +def _invert_K(Ks: torch.Tensor) -> torch.Tensor: + """Invert 3x3 intrinsics matrices. Assumes no skew.""" + assert Ks.shape[-2:] == (3, 3) + out = torch.zeros_like(Ks) + out[..., 0, 0] = 1.0 / Ks[..., 0, 0] + out[..., 1, 1] = 1.0 / Ks[..., 1, 1] + out[..., 0, 2] = -Ks[..., 0, 2] / Ks[..., 0, 0] + out[..., 1, 2] = -Ks[..., 1, 2] / Ks[..., 1, 1] + out[..., 2, 2] = 1.0 + return out