diff --git a/.gitattributes b/.gitattributes index 0bf7a0cfe53ed8fd770d0144373fe39daf4c0bcf..734f186b39f745683ffc71e6367f0c539bd8304f 100644 --- a/.gitattributes +++ b/.gitattributes @@ -41,3 +41,18 @@ dino_model/facebookresearch_dino_main/.github/dino.gif filter=lfs diff=lfs merge dreamsim_model/facebookresearch_dino_main/.github/attention_maps.png filter=lfs diff=lfs merge=lfs -text dreamsim_model/facebookresearch_co-tracker_main/notebooks/demo.ipynb filter=lfs diff=lfs merge=lfs -text torch_hub/facebookresearch_dino_main/.github/dino.gif filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_dinov2_main/docs/ChannelAdaptiveDINO.png filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_dinov2_main/docs/Cell-DINO.png filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/teddy.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/assets/apple.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_dino_main/.github/attention_maps.png filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/notebooks/demo.ipynb filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/bear.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/cat.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/apple.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/backpack.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/pillow.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding-launch.mp4 filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/assets/teaser.png filter=lfs diff=lfs merge=lfs -text +torch_hub/facebookresearch_co-tracker_main/assets/bmx-bumps.gif filter=lfs diff=lfs merge=lfs -text diff --git a/torch_hub/checkpoints/alexnet-owt-7be5be79.pth b/torch_hub/checkpoints/alexnet-owt-7be5be79.pth new file mode 100644 index 0000000000000000000000000000000000000000..779e66b9c31e9a0b60a3f1530626d98f2d898c60 --- /dev/null +++ b/torch_hub/checkpoints/alexnet-owt-7be5be79.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7be5be791159472b1fbf3c69796f7cb30dca7ad8466c2df70058c37116cdee02 +size 244408911 diff --git a/torch_hub/checkpoints/convnext_tiny_1k_224_ema.pth b/torch_hub/checkpoints/convnext_tiny_1k_224_ema.pth new file mode 100644 index 0000000000000000000000000000000000000000..f2d606871a735782008abc9768eb2acb0bb8d229 --- /dev/null +++ b/torch_hub/checkpoints/convnext_tiny_1k_224_ema.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:14f3164e3ea6ac32ab3f574f528ce817696c9176fad4221e0a77a905a7360595 +size 114414741 diff --git a/torch_hub/checkpoints/cotracker2.pth b/torch_hub/checkpoints/cotracker2.pth new file mode 100644 index 0000000000000000000000000000000000000000..63543571e5ab3a388e59c41f53a5b46ab85bff26 --- /dev/null +++ b/torch_hub/checkpoints/cotracker2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:362f5274376d610dc987b6daf2c2fefe63e06e1835f4ec1a10d0a15c5a4eef4f +size 204396415 diff --git a/torch_hub/checkpoints/dino_vitbase16_pretrain.pth b/torch_hub/checkpoints/dino_vitbase16_pretrain.pth new file mode 100644 index 0000000000000000000000000000000000000000..a18201ca09ac579f341aca5514be8c0520e22727 --- /dev/null +++ b/torch_hub/checkpoints/dino_vitbase16_pretrain.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf34ad0f424b9029b593e8dc3ed553bf26e88bcba0d32bf3e62a6209cb64c85e +size 343242485 diff --git a/torch_hub/checkpoints/dinov2_vitb14_reg4_pretrain.pth b/torch_hub/checkpoints/dinov2_vitb14_reg4_pretrain.pth new file mode 100644 index 0000000000000000000000000000000000000000..a2eb1b912a37e79b6668a8e1acc5a9f5d5f921c5 --- /dev/null +++ b/torch_hub/checkpoints/dinov2_vitb14_reg4_pretrain.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:73182a088cf94833c94b1666d1c99e02fe87e2007bff57b564fb6206e25dba71 +size 346393545 diff --git a/torch_hub/checkpoints/pt_inception-2015-12-05-6726825d.pth b/torch_hub/checkpoints/pt_inception-2015-12-05-6726825d.pth new file mode 100644 index 0000000000000000000000000000000000000000..fab8037f95b5b88179f0ba378cea7d68c0ccfed3 --- /dev/null +++ b/torch_hub/checkpoints/pt_inception-2015-12-05-6726825d.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6726825d0af5f729cebd5821db510b11b1cfad8faad88a03f1befd49fb9129b2 +size 95628359 diff --git a/torch_hub/checkpoints/vgg16-397923af.pth b/torch_hub/checkpoints/vgg16-397923af.pth new file mode 100644 index 0000000000000000000000000000000000000000..9dfa9aa51ae961694af6b9dfa41bd8338de8272a --- /dev/null +++ b/torch_hub/checkpoints/vgg16-397923af.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:397923af8e79cdbb6a7127f12361acd7a2f83e06b05044ddf496e83de57a5bf0 +size 553433881 diff --git a/torch_hub/checkpoints/weights-inception-2015-12-05-6726825d.pth b/torch_hub/checkpoints/weights-inception-2015-12-05-6726825d.pth new file mode 100644 index 0000000000000000000000000000000000000000..fab8037f95b5b88179f0ba378cea7d68c0ccfed3 --- /dev/null +++ b/torch_hub/checkpoints/weights-inception-2015-12-05-6726825d.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6726825d0af5f729cebd5821db510b11b1cfad8faad88a03f1befd49fb9129b2 +size 95628359 diff --git a/torch_hub/facebookresearch_co-tracker_main/CODE_OF_CONDUCT.md b/torch_hub/facebookresearch_co-tracker_main/CODE_OF_CONDUCT.md new file mode 100644 index 0000000000000000000000000000000000000000..f913b6a55a6c5ab6e1224e11fc039c3d4c3b6283 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/CODE_OF_CONDUCT.md @@ -0,0 +1,80 @@ +# Code of Conduct + +## Our Pledge + +In the interest of fostering an open and welcoming environment, we as +contributors and maintainers pledge to make participation in our project and +our community a harassment-free experience for everyone, regardless of age, body +size, disability, ethnicity, sex characteristics, gender identity and expression, +level of experience, education, socio-economic status, nationality, personal +appearance, race, religion, or sexual identity and orientation. + +## Our Standards + +Examples of behavior that contributes to creating a positive environment +include: + +* Using welcoming and inclusive language +* Being respectful of differing viewpoints and experiences +* Gracefully accepting constructive criticism +* Focusing on what is best for the community +* Showing empathy towards other community members + +Examples of unacceptable behavior by participants include: + +* The use of sexualized language or imagery and unwelcome sexual attention or +advances +* Trolling, insulting/derogatory comments, and personal or political attacks +* Public or private harassment +* Publishing others' private information, such as a physical or electronic +address, without explicit permission +* Other conduct which could reasonably be considered inappropriate in a +professional setting + +## Our Responsibilities + +Project maintainers are responsible for clarifying the standards of acceptable +behavior and are expected to take appropriate and fair corrective action in +response to any instances of unacceptable behavior. + +Project maintainers have the right and responsibility to remove, edit, or +reject comments, commits, code, wiki edits, issues, and other contributions +that are not aligned to this Code of Conduct, or to ban temporarily or +permanently any contributor for other behaviors that they deem inappropriate, +threatening, offensive, or harmful. + +## Scope + +This Code of Conduct applies within all project spaces, and it also applies when +an individual is representing the project or its community in public spaces. +Examples of representing a project or community include using an official +project e-mail address, posting via an official social media account, or acting +as an appointed representative at an online or offline event. Representation of +a project may be further defined and clarified by project maintainers. + +This Code of Conduct also applies outside the project spaces when there is a +reasonable belief that an individual's behavior may have a negative impact on +the project or its community. + +## Enforcement + +Instances of abusive, harassing, or otherwise unacceptable behavior may be +reported by contacting the project team at . All +complaints will be reviewed and investigated and will result in a response that +is deemed necessary and appropriate to the circumstances. The project team is +obligated to maintain confidentiality with regard to the reporter of an incident. +Further details of specific enforcement policies may be posted separately. + +Project maintainers who do not follow or enforce the Code of Conduct in good +faith may face temporary or permanent repercussions as determined by other +members of the project's leadership. + +## Attribution + +This Code of Conduct is adapted from the [Contributor Covenant][homepage], version 1.4, +available at https://www.contributor-covenant.org/version/1/4/code-of-conduct.html + +[homepage]: https://www.contributor-covenant.org + +For answers to common questions about this code of conduct, see +https://www.contributor-covenant.org/faq \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/CONTRIBUTING.md b/torch_hub/facebookresearch_co-tracker_main/CONTRIBUTING.md new file mode 100644 index 0000000000000000000000000000000000000000..f3ed8c2929373655dfdc962d52978708a3cebbaf --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/CONTRIBUTING.md @@ -0,0 +1,28 @@ +# CoTracker +We want to make contributing to this project as easy and transparent as possible. + +## Pull Requests +We actively welcome your pull requests. + +1. Fork the repo and create your branch from `main`. +2. If you've changed APIs, update the documentation. +3. Make sure your code lints. +4. If you haven't already, complete the Contributor License Agreement ("CLA"). + +## Contributor License Agreement ("CLA") +In order to accept your pull request, we need you to submit a CLA. You only need +to do this once to work on any of Meta's open source projects. + +Complete your CLA here: + +## Issues +We use GitHub issues to track public bugs. Please ensure your description is +clear and has sufficient instructions to be able to reproduce the issue. + +Meta has a [bounty program](https://www.facebook.com/whitehat/) for the safe +disclosure of security bugs. In those cases, please go through the process +outlined on that page and do not file a public issue. + +## License +By contributing to CoTracker, you agree that your contributions will be licensed +under the LICENSE file in the root directory of this source tree. \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/LICENSE.md b/torch_hub/facebookresearch_co-tracker_main/LICENSE.md new file mode 100644 index 0000000000000000000000000000000000000000..e395ca3e2cdebf48a6375a3c1022d10caabba7db --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/LICENSE.md @@ -0,0 +1,399 @@ +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. \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/README.md b/torch_hub/facebookresearch_co-tracker_main/README.md new file mode 100644 index 0000000000000000000000000000000000000000..d7e05175054166c16aeaac5ddf937c5ffcc0f1b2 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/README.md @@ -0,0 +1,360 @@ +# CoTracker3: Simpler and Better Point Tracking by Pseudo-Labelling Real Videos + +**[Meta AI Research, GenAI](https://ai.facebook.com/research/)**; **[University of Oxford, VGG](https://www.robots.ox.ac.uk/~vgg/)** + +[Nikita Karaev](https://nikitakaraevv.github.io/), [Iurii Makarov](https://linkedin.com/in/lvoursl), [Jianyuan Wang](https://jytime.github.io/), [Ignacio Rocco](https://www.irocco.info/), [Benjamin Graham](https://ai.facebook.com/people/benjamin-graham/), [Natalia Neverova](https://nneverova.github.io/), [Andrea Vedaldi](https://www.robots.ox.ac.uk/~vedaldi/), [Christian Rupprecht](https://chrirupp.github.io/) + +### [Project Page](https://cotracker3.github.io/) | [Paper #1](https://arxiv.org/abs/2307.07635) | [Paper #2](https://arxiv.org/abs/2410.11831) | [X Thread](https://twitter.com/n_karaev/status/1742638906355470772) | [BibTeX](#citing-cotracker) + + + Open In Colab + + + Spaces + + + + +**CoTracker** is a fast transformer-based model that can track any point in a video. It brings to tracking some of the benefits of Optical Flow. + +CoTracker can track: + +- **Any pixel** in a video +- A **quasi-dense** set of pixels together +- Points can be manually selected or sampled on a grid in any video frame + +Try these tracking modes for yourself with our [Colab demo](https://colab.research.google.com/github/facebookresearch/co-tracker/blob/master/notebooks/demo.ipynb) or in the [Hugging Face Space 🤗](https://huggingface.co/spaces/facebook/cotracker). + +**Updates:** + +- [January 21, 2025] 📦 Kubric Dataset used for CoTracker3 now available! This dataset contains **6,000 high-resolution sequences** (512×512px, 120 frames) with slight camera motion, rendered using the Kubric engine. Check it out on [Hugging Face Dataset](https://huggingface.co/datasets/facebook/CoTracker3_Kubric). + +- [October 15, 2024] 📣 We're releasing CoTracker3! State-of-the-art point tracking with a lightweight architecture trained with 1000x less data than previous top-performing models. Code for baseline models and the pseudo-labeling pipeline are available in the repo, as well as model checkpoints. Check out our [paper](https://arxiv.org/abs/2410.11831) for more details. + +- [September 25, 2024] CoTracker2.1 is now available! This model has better performance on TAP-Vid benchmarks and follows the architecture of the original CoTracker. Try it out! + +- [June 14, 2024] We have released the code for [VGGSfM](https://github.com/facebookresearch/vggsfm), a model for recovering camera poses and 3D structure from any image sequences based on point tracking! VGGSfM is the first fully differentiable SfM framework that unlocks scalability and outperforms conventional SfM methods on standard benchmarks. + +- [December 27, 2023] CoTracker2 is now available! It can now track many more (up to **265*265**!) points jointly and it has a cleaner and more memory-efficient implementation. It also supports online processing. See the [updated paper](https://arxiv.org/abs/2307.07635) for more details. The old version remains available [here](https://github.com/facebookresearch/co-tracker/tree/8d364031971f6b3efec945dd15c468a183e58212). + +- [September 5, 2023] You can now run our Gradio demo [locally](./gradio_demo/app.py). + +## Quick start +The easiest way to use CoTracker is to load a pretrained model from `torch.hub`: + +### Offline mode: +```pip install imageio[ffmpeg]```, then: +```python +import torch +# Download the video +url = 'https://github.com/facebookresearch/co-tracker/raw/refs/heads/main/assets/apple.mp4' + +import imageio.v3 as iio +frames = iio.imread(url, plugin="FFMPEG") # plugin="pyav" + +device = 'cuda' +grid_size = 10 +video = torch.tensor(frames).permute(0, 3, 1, 2)[None].float().to(device) # B T C H W + +# Run Offline CoTracker: +cotracker = torch.hub.load("facebookresearch/co-tracker", "cotracker3_offline").to(device) +pred_tracks, pred_visibility = cotracker(video, grid_size=grid_size) # B T N 2, B T N 1 +``` +### Online mode: +```python +cotracker = torch.hub.load("facebookresearch/co-tracker", "cotracker3_online").to(device) + +# Run Online CoTracker, the same model with a different API: +# Initialize online processing +cotracker(video_chunk=video, is_first_step=True, grid_size=grid_size) + +# Process the video +for ind in range(0, video.shape[1] - cotracker.step, cotracker.step): + pred_tracks, pred_visibility = cotracker( + video_chunk=video[:, ind : ind + cotracker.step * 2] + ) # B T N 2, B T N 1 +``` +Online processing is more memory-efficient and allows for the processing of longer videos. However, in the example provided above, the video length is known! See [the online demo](./online_demo.py) for an example of tracking from an online stream with an unknown video length. + +### Visualize predicted tracks: +After [installing](#installation-instructions) CoTracker, you can visualize tracks with: +```python +from cotracker.utils.visualizer import Visualizer + +vis = Visualizer(save_dir="./saved_videos", pad_value=120, linewidth=3) +vis.visualize(video, pred_tracks, pred_visibility) +``` + +We offer a number of other ways to interact with CoTracker: + +1. Interactive Gradio demo: + - A demo is available in the [`facebook/cotracker` Hugging Face Space 🤗](https://huggingface.co/spaces/facebook/cotracker). + - You can use the gradio demo locally by running [`python -m gradio_demo.app`](./gradio_demo/app.py) after installing the required packages: `pip install -r gradio_demo/requirements.txt`. +2. Jupyter notebook: + - You can run the notebook in + [Google Colab](https://colab.research.google.com/github/facebookresearch/co-tracker/blob/master/notebooks/demo.ipynb). + - Or explore the notebook located at [`notebooks/demo.ipynb`](./notebooks/demo.ipynb). +2. You can [install](#installation-instructions) CoTracker _locally_ and then: + - Run an *offline* demo with 10 ⨉ 10 points sampled on a grid on the first frame of a video (results will be saved to `./saved_videos/demo.mp4`)): + + ```bash + python demo.py --grid_size 10 + ``` + - Run an *online* demo: + + ```bash + python online_demo.py + ``` + +A GPU is strongly recommended for using CoTracker locally. + + + + +## Installation Instructions +You can use a Pretrained Model via PyTorch Hub, as described above, or install CoTracker from this GitHub repo. +This is the best way if you need to run our local demo or evaluate/train CoTracker. + +Ensure you have both _PyTorch_ and _TorchVision_ installed on your system. Follow the instructions [here](https://pytorch.org/get-started/locally/) for the installation. +We strongly recommend installing both PyTorch and TorchVision with CUDA support, although for small tasks CoTracker can be run on CPU. + + + + +### Install a Development Version + +```bash +git clone https://github.com/facebookresearch/co-tracker +cd co-tracker +pip install -e . +pip install matplotlib flow_vis tqdm tensorboard +``` + +You can manually download all CoTracker3 checkpoints (baseline and scaled models, as well as single and sliding window architectures) from the links below and place them in the `checkpoints` folder as follows: + +```bash +mkdir -p checkpoints +cd checkpoints +# download the online (multi window) model +wget https://huggingface.co/facebook/cotracker3/resolve/main/scaled_online.pth +# download the offline (single window) model +wget https://huggingface.co/facebook/cotracker3/resolve/main/scaled_offline.pth +cd .. +``` +You can also download CoTracker3 checkpoints trained only on Kubric: +```bash +# download the online (sliding window) model +wget https://huggingface.co/facebook/cotracker3/resolve/main/baseline_online.pth +# download the offline (single window) model +wget https://huggingface.co/facebook/cotracker3/resolve/main/baseline_offline.pth +``` +For old checkpoints, see [this section](#previous-version). + +## Evaluation + +To reproduce the results presented in the paper, download the following datasets: + +- [TAP-Vid](https://github.com/deepmind/tapnet) +- [Dynamic Replica](https://dynamic-stereo.github.io/) + +And install the necessary dependencies: + +```bash +pip install hydra-core==1.1.0 mediapy +``` + +Then, execute the following command to evaluate the online model on TAP-Vid DAVIS: + +```bash +python ./cotracker/evaluation/evaluate.py --config-name eval_tapvid_davis_first exp_dir=./eval_outputs dataset_root=your/tapvid/path +``` +And the offline model: +```bash +python ./cotracker/evaluation/evaluate.py --config-name eval_tapvid_davis_first exp_dir=./eval_outputs dataset_root=/fsx-repligen/shared/datasets/tapvid offline_model=True window_len=60 checkpoint=./checkpoints/scaled_offline.pth +``` +We run evaluations jointly on all the target points at a time for faster inference. With such evaluations, the numbers are similar to those presented in the paper. If you want to reproduce the exact numbers from the paper, add the flag `single_point=True`. + +These are the numbers that you should be able to reproduce using the released checkpoint and the current version of the codebase: +| | Kinetics, $\delta_\text{avg}^\text{vis}$ | DAVIS, $\delta_\text{avg}^\text{vis}$ | RoboTAP, $\delta_\text{avg}^\text{vis}$ | RGB-S, $\delta_\text{avg}^\text{vis}$| +| :---: |:---: | :---: | :---: | :---: | +| CoTracker2, 27.12.23 | 61.8 | 74.6 | 69.6 | 73.4 | +| CoTracker2.1, 25.09.24 | 63 | 76.1 | 70.6 | 79.6 | +| CoTracker3 offline, 15.10.24 | 67.8 | **76.9** | 78.0 | **85.0** | +| CoTracker3 online, 15.10.24 | **68.3** | 76.7 | **78.8** | 82.7 | + + +## Training + +### Baseline +To train the CoTracker as described in our paper, you first need to generate annotations for [Google Kubric](https://github.com/google-research/kubric) MOVI-f dataset. +Instructions for annotation generation can be found [here](https://github.com/deepmind/tapnet). +You can also find a discussion on dataset generation in [this issue](https://github.com/facebookresearch/co-tracker/issues/8). + +Once you have the annotated dataset, you need to make sure you followed the steps for evaluation setup and install the training dependencies: + +```bash +pip install pip==24.0 +pip install pytorch_lightning==1.6.0 tensorboard opencv-python +``` + +Now you can launch training on Kubric. +Our model was trained for 50000 iterations on 32 GPUs (4 nodes with 8 GPUs). +Modify _dataset_root_ and _ckpt_path_ accordingly before running this command. For training on 4 nodes, add `--num_nodes 4`. + +Here is an example of how to launch training of the online model on Kubric: +```bash + python train_on_kubric.py --batch_size 1 --num_steps 50000 \ + --ckpt_path ./ --model_name cotracker_three --save_freq 200 --sequence_len 64 \ + --eval_datasets tapvid_davis_first tapvid_stacking --traj_per_sample 384 \ + --sliding_window_len 16 --train_datasets kubric --save_every_n_epoch 5 \ + --evaluate_every_n_epoch 5 --model_stride 4 --dataset_root ${path_to_your_dataset} \ + --num_nodes 4 --num_virtual_tracks 64 --mixed_precision --corr_radius 3 \ + --wdecay 0.0005 --linear_layer_for_vis_conf --validate_at_start --add_huber_loss +``` + +Training the offline model on Kubric: +```bash +python train_on_kubric.py --batch_size 1 --num_steps 50000 \ + --ckpt_path ./ --model_name cotracker_three --save_freq 200 --sequence_len 60 \ + --eval_datasets tapvid_davis_first tapvid_stacking --traj_per_sample 512 \ + --sliding_window_len 60 --train_datasets kubric --save_every_n_epoch 5 \ + --evaluate_every_n_epoch 5 --model_stride 4 --dataset_root ${path_to_your_dataset} \ + --num_nodes 4 --num_virtual_tracks 64 --mixed_precision --offline_model \ + --random_frame_rate --query_sampling_method random --corr_radius 3 \ + --wdecay 0.0005 --random_seq_len --linear_layer_for_vis_conf \ + --validate_at_start --add_huber_loss +``` + +### Fine-tuning with pseudo labels +In order to launch training with pseudo-labelling, you need to collect your own dataset of real videos. There is a sample class available in [`cotracker/datasets/real_dataset.py`](./cotracker/datasets/real_dataset.py) with keyword-based filtering that we used for training. Your class should implement loading a video and storing it in the `CoTrackerData` class as a field, while pseudo labels will be generated in `train_on_real_data.py`. + +You should have an existing Kubric-trained model for fine-tuning with pseudo labels. Here is an example of how you can launch fine-tuning of the online model: +```bash +python ./train_on_real_data.py --batch_size 1 --num_steps 15000 \ + --ckpt_path ./ --model_name cotracker_three --save_freq 200 --sequence_len 64 \ + --eval_datasets tapvid_stacking tapvid_davis_first --traj_per_sample 384 \ + --save_every_n_epoch 15 --evaluate_every_n_epoch 15 --model_stride 4 \ + --dataset_root ${path_to_your_dataset} --num_nodes 4 --real_data_splits 0 \ + --num_virtual_tracks 64 --mixed_precision --random_frame_rate \ + --restore_ckpt ./checkpoints/baseline_online.pth \ + --lr 0.00005 --real_data_filter_sift --validate_at_start \ + --sliding_window_len 16 --limit_samples 15000 + +``` +And the offline model: +```bash +python train_on_real_data.py --batch_size 1 --num_steps 15000 \ + --ckpt_path ./ --model_name cotracker_three --save_freq 200 --sequence_len 80 \ + --eval_datasets tapvid_stacking tapvid_davis_first --traj_per_sample 384 --save_every_n_epoch 15 \ + --evaluate_every_n_epoch 15 --model_stride 4 --dataset_root ${path_to_your_dataset} \ + --num_nodes 4 --real_data_splits 0 --num_virtual_tracks 64 --mixed_precision \ + --random_frame_rate --restore_ckpt ./checkpoints/baseline_offline.pth --lr 0.00005 \ + --real_data_filter_sift --validate_at_start --offline_model --limit_samples 15000 +``` + + + +## Development + +### Building the documentation + +To build CoTracker documentation, first install the dependencies: + +```bash +pip install sphinx +pip install sphinxcontrib-bibtex +``` + +Then you can use this command to generate the documentation in the `docs/_build/html` folder: + +```bash +make -C docs html +``` + + +## Previous versions +### CoTracker v2 +You could use CoTracker v2 with torch.hub in both offline and online modes. +#### Offline mode: +```pip install imageio[ffmpeg]```, then: +```python +import torch +# Download the video +url = 'https://github.com/facebookresearch/co-tracker/blob/main/assets/apple.mp4' + +import imageio.v3 as iio +frames = iio.imread(url, plugin="FFMPEG") # plugin="pyav" + +device = 'cuda' +grid_size = 10 +video = torch.tensor(frames).permute(0, 3, 1, 2)[None].float().to(device) # B T C H W + +# Run Offline CoTracker: +cotracker = torch.hub.load("facebookresearch/co-tracker", "cotracker2").to(device) +pred_tracks, pred_visibility = cotracker(video, grid_size=grid_size) # B T N 2, B T N 1 +``` +#### Online mode: +```python +cotracker = torch.hub.load("facebookresearch/co-tracker", "cotracker2_online").to(device) + +# Run Online CoTracker, the same model with a different API: +# Initialize online processing +cotracker(video_chunk=video, is_first_step=True, grid_size=grid_size) + +# Process the video +for ind in range(0, video.shape[1] - cotracker.step, cotracker.step): + pred_tracks, pred_visibility = cotracker( + video_chunk=video[:, ind : ind + cotracker.step * 2] + ) # B T N 2, B T N 1 +``` + +Checkpoint for v2 could be downloaded with the following command: +```bash +wget https://huggingface.co/facebook/cotracker/resolve/main/cotracker2.pth +``` + +### CoTracker v1 +It is directly available via pytorch hub: +```python +import torch +import einops +import timm +import tqdm + +cotracker = torch.hub.load("facebookresearch/co-tracker:v1.0", "cotracker_w8") +``` +The old version of the code is available [here](https://github.com/facebookresearch/co-tracker/tree/8d364031971f6b3efec945dd15c468a183e58212). +You can also download the corresponding checkpoints: +```bash +wget https://dl.fbaipublicfiles.com/cotracker/cotracker_stride_4_wind_8.pth +wget https://dl.fbaipublicfiles.com/cotracker/cotracker_stride_4_wind_12.pth +wget https://dl.fbaipublicfiles.com/cotracker/cotracker_stride_8_wind_16.pth +``` + +## License + +The majority of CoTracker is licensed under CC-BY-NC, however portions of the project are available under separate license terms: Particle Video Revisited is licensed under the MIT license, TAP-Vid and LocoTrack are licensed under the Apache 2.0 license. + +## Acknowledgments + +We would like to thank [PIPs](https://github.com/aharley/pips), [TAP-Vid](https://github.com/deepmind/tapnet), [LocoTrack](https://github.com/cvlab-kaist/locotrack) for publicly releasing their code and data. We also want to thank [Luke Melas-Kyriazi](https://lukemelas.github.io/) for proofreading the paper, [Jianyuan Wang](https://jytime.github.io/), [Roman Shapovalov](https://shapovalov.ro/) and [Adam W. Harley](https://adamharley.com/) for the insightful discussions. + +## Citing CoTracker + +If you find our repository useful, please consider giving it a star ⭐ and citing our research papers in your work: +```bibtex +@inproceedings{karaev23cotracker, + title = {CoTracker: It is Better to Track Together}, + author = {Nikita Karaev and Ignacio Rocco and Benjamin Graham and Natalia Neverova and Andrea Vedaldi and Christian Rupprecht}, + booktitle = {Proc. {ECCV}}, + year = {2024} +} +``` +```bibtex +@inproceedings{karaev24cotracker3, + title = {CoTracker3: Simpler and Better Point Tracking by Pseudo-Labelling Real Videos}, + author = {Nikita Karaev and Iurii Makarov and Jianyuan Wang and Natalia Neverova and Andrea Vedaldi and Christian Rupprecht}, + booktitle = {Proc. {arXiv:2410.11831}}, + year = {2024} +} +``` diff --git a/torch_hub/facebookresearch_co-tracker_main/__pycache__/hubconf.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/__pycache__/hubconf.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..0b314d27da7d4d42509ad07f08129b40b7eaed4d Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/__pycache__/hubconf.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/assets/apple.mp4 b/torch_hub/facebookresearch_co-tracker_main/assets/apple.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..80f158ffe2cc5625ef5df53e1ff423d2b83e989c --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/assets/apple.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7f48c5cfb1479e1dbc1df2373d5cad4f55c198bbdb379da0ece10087971542a +size 1219872 diff --git a/torch_hub/facebookresearch_co-tracker_main/assets/apple_mask.png b/torch_hub/facebookresearch_co-tracker_main/assets/apple_mask.png new file mode 100644 index 0000000000000000000000000000000000000000..76de57198180ddbf7bd4acd7fb2bd527830caaeb Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/assets/apple_mask.png differ diff --git a/torch_hub/facebookresearch_co-tracker_main/assets/bmx-bumps.gif b/torch_hub/facebookresearch_co-tracker_main/assets/bmx-bumps.gif new file mode 100644 index 0000000000000000000000000000000000000000..929f65515df76d4b09df231a6e010e624f825b52 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/assets/bmx-bumps.gif @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c96947f7db09bb39e9e2c365fdb34ff27860234fe8044fc07f271ba521f503ba +size 7670915 diff --git a/torch_hub/facebookresearch_co-tracker_main/assets/teaser.png b/torch_hub/facebookresearch_co-tracker_main/assets/teaser.png new file mode 100644 index 0000000000000000000000000000000000000000..49bbdaafe83b2c60b26f2d6eedd18a51504c31c6 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/assets/teaser.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1037b7035a9507bbb3a4afc120742cc2d1eb516f033d2a77bfbcb3501c4ccae7 +size 6030677 diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/__pycache__/__init__.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..0189c1f6fd9e08a6f202237c2ec53b5fcec22beb Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/__pycache__/__init__.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/__pycache__/predictor.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/__pycache__/predictor.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..279065415fe3bc043f3e51f46314b38f7c309f0b Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/__pycache__/predictor.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/dataclass_utils.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/dataclass_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..d171f94e6b02caa70c9156ae0e4ac6aac290432f --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/dataclass_utils.py @@ -0,0 +1,168 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + + +import json +import dataclasses +import numpy as np +from dataclasses import Field, MISSING +from typing import IO, TypeVar, Type, get_args, get_origin, Union, Any, Tuple + +_X = TypeVar("_X") + + +def load_dataclass(f: IO, cls: Type[_X], binary: bool = False) -> _X: + """ + Loads to a @dataclass or collection hierarchy including dataclasses + from a json recursively. + Call it like load_dataclass(f, typing.List[FrameAnnotationAnnotation]). + raises KeyError if json has keys not mapping to the dataclass fields. + + Args: + f: Either a path to a file, or a file opened for writing. + cls: The class of the loaded dataclass. + binary: Set to True if `f` is a file handle, else False. + """ + if binary: + asdict = json.loads(f.read().decode("utf8")) + else: + asdict = json.load(f) + + # in the list case, run a faster "vectorized" version + cls = get_args(cls)[0] + res = list(_dataclass_list_from_dict_list(asdict, cls)) + + return res + + +def _resolve_optional(type_: Any) -> Tuple[bool, Any]: + """Check whether `type_` is equivalent to `typing.Optional[T]` for some T.""" + if get_origin(type_) is Union: + args = get_args(type_) + if len(args) == 2 and args[1] == type(None): # noqa E721 + return True, args[0] + if type_ is Any: + return True, Any + + return False, type_ + + +def _unwrap_type(tp): + # strips Optional wrapper, if any + if get_origin(tp) is Union: + args = get_args(tp) + if len(args) == 2 and any(a is type(None) for a in args): # noqa: E721 + # this is typing.Optional + return args[0] if args[1] is type(None) else args[1] # noqa: E721 + return tp + + +def _get_dataclass_field_default(field: Field) -> Any: + if field.default_factory is not MISSING: + # pyre-fixme[29]: `Union[dataclasses._MISSING_TYPE, + # dataclasses._DefaultFactory[typing.Any]]` is not a function. + return field.default_factory() + elif field.default is not MISSING: + return field.default + else: + return None + + +def _dataclass_list_from_dict_list(dlist, typeannot): + """ + Vectorised version of `_dataclass_from_dict`. + The output should be equivalent to + `[_dataclass_from_dict(d, typeannot) for d in dlist]`. + + Args: + dlist: list of objects to convert. + typeannot: type of each of those objects. + Returns: + iterator or list over converted objects of the same length as `dlist`. + + Raises: + ValueError: it assumes the objects have None's in consistent places across + objects, otherwise it would ignore some values. This generally holds for + auto-generated annotations, but otherwise use `_dataclass_from_dict`. + """ + + cls = get_origin(typeannot) or typeannot + + if typeannot is Any: + return dlist + if all(obj is None for obj in dlist): # 1st recursion base: all None nodes + return dlist + if any(obj is None for obj in dlist): + # filter out Nones and recurse on the resulting list + idx_notnone = [(i, obj) for i, obj in enumerate(dlist) if obj is not None] + idx, notnone = zip(*idx_notnone) + converted = _dataclass_list_from_dict_list(notnone, typeannot) + res = [None] * len(dlist) + for i, obj in zip(idx, converted): + res[i] = obj + return res + + is_optional, contained_type = _resolve_optional(typeannot) + if is_optional: + return _dataclass_list_from_dict_list(dlist, contained_type) + + # otherwise, we dispatch by the type of the provided annotation to convert to + if issubclass(cls, tuple) and hasattr(cls, "_fields"): # namedtuple + # For namedtuple, call the function recursively on the lists of corresponding keys + types = cls.__annotations__.values() + dlist_T = zip(*dlist) + res_T = [ + _dataclass_list_from_dict_list(key_list, tp) + for key_list, tp in zip(dlist_T, types) + ] + return [cls(*converted_as_tuple) for converted_as_tuple in zip(*res_T)] + elif issubclass(cls, (list, tuple)): + # For list/tuple, call the function recursively on the lists of corresponding positions + types = get_args(typeannot) + if len(types) == 1: # probably List; replicate for all items + types = types * len(dlist[0]) + dlist_T = zip(*dlist) + res_T = ( + _dataclass_list_from_dict_list(pos_list, tp) + for pos_list, tp in zip(dlist_T, types) + ) + if issubclass(cls, tuple): + return list(zip(*res_T)) + else: + return [cls(converted_as_tuple) for converted_as_tuple in zip(*res_T)] + elif issubclass(cls, dict): + # For the dictionary, call the function recursively on concatenated keys and vertices + key_t, val_t = get_args(typeannot) + all_keys_res = _dataclass_list_from_dict_list( + [k for obj in dlist for k in obj.keys()], key_t + ) + all_vals_res = _dataclass_list_from_dict_list( + [k for obj in dlist for k in obj.values()], val_t + ) + indices = np.cumsum([len(obj) for obj in dlist]) + assert indices[-1] == len(all_keys_res) + + keys = np.split(list(all_keys_res), indices[:-1]) + all_vals_res_iter = iter(all_vals_res) + return [cls(zip(k, all_vals_res_iter)) for k in keys] + elif not dataclasses.is_dataclass(typeannot): + return dlist + + # dataclass node: 2nd recursion base; call the function recursively on the lists + # of the corresponding fields + assert dataclasses.is_dataclass(cls) + fieldtypes = { + f.name: (_unwrap_type(f.type), _get_dataclass_field_default(f)) + for f in dataclasses.fields(typeannot) + } + + # NOTE the default object is shared here + key_lists = ( + _dataclass_list_from_dict_list([obj.get(k, default) for obj in dlist], type_) + for k, (type_, default) in fieldtypes.items() + ) + transposed = zip(*key_lists) + return [cls(*vals_as_tuple) for vals_as_tuple in transposed] diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/dr_dataset.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/dr_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..354f644e2852ba9f6d87b55ca303abca2f19c306 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/dr_dataset.py @@ -0,0 +1,168 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + + +import os +import gzip +import torch +import numpy as np +import torch.utils.data as data +from collections import defaultdict +from dataclasses import dataclass +from typing import List, Optional, Any, Dict, Tuple + +from cotracker.datasets.utils import CoTrackerData +from cotracker.datasets.dataclass_utils import load_dataclass + + +@dataclass +class ImageAnnotation: + # path to jpg file, relative w.r.t. dataset_root + path: str + # H x W + size: Tuple[int, int] + + +@dataclass +class DynamicReplicaFrameAnnotation: + """A dataclass used to load annotations from json.""" + + # can be used to join with `SequenceAnnotation` + sequence_name: str + # 0-based, continuous frame number within sequence + frame_number: int + # timestamp in seconds from the video start + frame_timestamp: float + + image: ImageAnnotation + meta: Optional[Dict[str, Any]] = None + + camera_name: Optional[str] = None + trajectories: Optional[str] = None + + +class DynamicReplicaDataset(data.Dataset): + def __init__( + self, + root, + split="valid", + traj_per_sample=256, + crop_size=None, + sample_len=-1, + only_first_n_samples=-1, + rgbd_input=False, + ): + super(DynamicReplicaDataset, self).__init__() + self.root = root + self.sample_len = sample_len + self.split = split + self.traj_per_sample = traj_per_sample + self.rgbd_input = rgbd_input + self.crop_size = crop_size + frame_annotations_file = f"frame_annotations_{split}.jgz" + self.sample_list = [] + with gzip.open( + os.path.join(root, split, frame_annotations_file), "rt", encoding="utf8" + ) as zipfile: + frame_annots_list = load_dataclass( + zipfile, List[DynamicReplicaFrameAnnotation] + ) + seq_annot = defaultdict(list) + for frame_annot in frame_annots_list: + if frame_annot.camera_name == "left": + seq_annot[frame_annot.sequence_name].append(frame_annot) + + for seq_name in seq_annot.keys(): + seq_len = len(seq_annot[seq_name]) + + step = self.sample_len if self.sample_len > 0 else seq_len + counter = 0 + + for ref_idx in range(0, seq_len, step): + sample = seq_annot[seq_name][ref_idx : ref_idx + step] + self.sample_list.append(sample) + counter += 1 + if only_first_n_samples > 0 and counter >= only_first_n_samples: + break + + def __len__(self): + return len(self.sample_list) + + def crop(self, rgbs, trajs): + T, N, _ = trajs.shape + + S = len(rgbs) + H, W = rgbs[0].shape[:2] + assert S == T + + H_new = H + W_new = W + + # simple random crop + y0 = 0 if self.crop_size[0] >= H_new else (H_new - self.crop_size[0]) // 2 + x0 = 0 if self.crop_size[1] >= W_new else (W_new - self.crop_size[1]) // 2 + rgbs = [ + rgb[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]] + for rgb in rgbs + ] + + trajs[:, :, 0] -= x0 + trajs[:, :, 1] -= y0 + + return rgbs, trajs + + def __getitem__(self, index): + sample = self.sample_list[index] + T = len(sample) + rgbs, visibilities, traj_2d = [], [], [] + + H, W = sample[0].image.size + image_size = (H, W) + + for i in range(T): + traj_path = os.path.join( + self.root, self.split, sample[i].trajectories["path"] + ) + traj = torch.load(traj_path) + + visibilities.append(traj["verts_inds_vis"].numpy()) + + rgbs.append(traj["img"].numpy()) + traj_2d.append(traj["traj_2d"].numpy()[..., :2]) + + traj_2d = np.stack(traj_2d) + visibility = np.stack(visibilities) + T, N, D = traj_2d.shape + # subsample trajectories for augmentations + visible_inds_sampled = torch.randperm(N)[: self.traj_per_sample] + + traj_2d = traj_2d[:, visible_inds_sampled] + visibility = visibility[:, visible_inds_sampled] + + if self.crop_size is not None: + rgbs, traj_2d = self.crop(rgbs, traj_2d) + H, W, _ = rgbs[0].shape + image_size = self.crop_size + + visibility[traj_2d[:, :, 0] > image_size[1] - 1] = False + visibility[traj_2d[:, :, 0] < 0] = False + visibility[traj_2d[:, :, 1] > image_size[0] - 1] = False + visibility[traj_2d[:, :, 1] < 0] = False + + # filter out points that're visible for less than 10 frames + visible_inds_resampled = visibility.sum(0) > 10 + traj_2d = torch.from_numpy(traj_2d[:, visible_inds_resampled]) + visibility = torch.from_numpy(visibility[:, visible_inds_resampled]) + + rgbs = np.stack(rgbs, 0) + video = torch.from_numpy(rgbs).reshape(T, H, W, 3).permute(0, 3, 1, 2).float() + return CoTrackerData( + video=video, + trajectory=traj_2d, + visibility=visibility, + valid=torch.ones(T, N), + seq_name=sample[0].sequence_name, + ) diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/kubric_movif_dataset.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/kubric_movif_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..208196f2ddf71434c11d471745ed19e01af5c7f5 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/kubric_movif_dataset.py @@ -0,0 +1,542 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import torch +import cv2 + +import imageio +import numpy as np + +from cotracker.datasets.utils import CoTrackerData +from torchvision.transforms import ColorJitter, GaussianBlur +from PIL import Image +from cotracker.models.core.model_utils import smart_cat + + +class CoTrackerDataset(torch.utils.data.Dataset): + def __init__( + self, + data_root, + crop_size=(384, 512), + seq_len=24, + traj_per_sample=768, + sample_vis_last_frame=False, + use_augs=False, + ): + super(CoTrackerDataset, self).__init__() + np.random.seed(0) + torch.manual_seed(0) + self.data_root = data_root + self.seq_len = seq_len + self.traj_per_sample = traj_per_sample + self.sample_vis_last_frame = sample_vis_last_frame + self.use_augs = use_augs + self.crop_size = crop_size + # photometric augmentation + self.photo_aug = ColorJitter( + brightness=0.2, contrast=0.2, saturation=0.2, hue=0.25 / 3.14 + ) + self.blur_aug = GaussianBlur(11, sigma=(0.1, 2.0)) + + self.blur_aug_prob = 0.25 + self.color_aug_prob = 0.25 + + # occlusion augmentation + self.eraser_aug_prob = 0.5 + self.eraser_bounds = [2, 100] + self.eraser_max = 10 + + # occlusion augmentation + self.replace_aug_prob = 0.5 + self.replace_bounds = [2, 100] + self.replace_max = 10 + + # spatial augmentations + self.pad_bounds = [0, 100] + self.crop_size = crop_size + self.resize_lim = [0.25, 2.0] # sample resizes from here + self.resize_delta = 0.2 + self.max_crop_offset = 50 + + self.do_flip = True + self.h_flip_prob = 0.5 + self.v_flip_prob = 0.5 + + def getitem_helper(self, index): + return NotImplementedError + + def __getitem__(self, index): + gotit = False + + sample, gotit = self.getitem_helper(index) + if not gotit: + print("warning: sampling failed") + # fake sample, so we can still collate + sample = CoTrackerData( + video=torch.zeros( + (self.seq_len, 3, self.crop_size[0], self.crop_size[1]) + ), + trajectory=torch.zeros((self.seq_len, self.traj_per_sample, 2)), + visibility=torch.zeros((self.seq_len, self.traj_per_sample)), + valid=torch.zeros((self.seq_len, self.traj_per_sample)), + # dataset_name="kubric", + ) + + return sample, gotit + + def add_photometric_augs(self, rgbs, trajs, visibles, eraser=True, replace=True): + T, N, _ = trajs.shape + + S = len(rgbs) + H, W = rgbs[0].shape[:2] + assert S == T + + if eraser: + ############ eraser transform (per image after the first) ############ + rgbs = [rgb.astype(np.float32) for rgb in rgbs] + for i in range(1, S): + if np.random.rand() < self.eraser_aug_prob: + for _ in range( + np.random.randint(1, self.eraser_max + 1) + ): # number of times to occlude + xc = np.random.randint(0, W) + yc = np.random.randint(0, H) + dx = np.random.randint( + self.eraser_bounds[0], self.eraser_bounds[1] + ) + dy = np.random.randint( + self.eraser_bounds[0], self.eraser_bounds[1] + ) + x0 = np.clip(xc - dx / 2, 0, W - 1).round().astype(np.int32) + x1 = np.clip(xc + dx / 2, 0, W - 1).round().astype(np.int32) + y0 = np.clip(yc - dy / 2, 0, H - 1).round().astype(np.int32) + y1 = np.clip(yc + dy / 2, 0, H - 1).round().astype(np.int32) + + mean_color = np.mean( + rgbs[i][y0:y1, x0:x1, :].reshape(-1, 3), axis=0 + ) + rgbs[i][y0:y1, x0:x1, :] = mean_color + + occ_inds = np.logical_and( + np.logical_and(trajs[i, :, 0] >= x0, trajs[i, :, 0] < x1), + np.logical_and(trajs[i, :, 1] >= y0, trajs[i, :, 1] < y1), + ) + visibles[i, occ_inds] = 0 + rgbs = [rgb.astype(np.uint8) for rgb in rgbs] + + if replace: + rgbs_alt = [ + np.array(self.photo_aug(Image.fromarray(rgb)), dtype=np.uint8) + for rgb in rgbs + ] + rgbs_alt = [ + np.array(self.photo_aug(Image.fromarray(rgb)), dtype=np.uint8) + for rgb in rgbs_alt + ] + + ############ replace transform (per image after the first) ############ + rgbs = [rgb.astype(np.float32) for rgb in rgbs] + rgbs_alt = [rgb.astype(np.float32) for rgb in rgbs_alt] + for i in range(1, S): + if np.random.rand() < self.replace_aug_prob: + for _ in range( + np.random.randint(1, self.replace_max + 1) + ): # number of times to occlude + xc = np.random.randint(0, W) + yc = np.random.randint(0, H) + dx = np.random.randint( + self.replace_bounds[0], self.replace_bounds[1] + ) + dy = np.random.randint( + self.replace_bounds[0], self.replace_bounds[1] + ) + x0 = np.clip(xc - dx / 2, 0, W - 1).round().astype(np.int32) + x1 = np.clip(xc + dx / 2, 0, W - 1).round().astype(np.int32) + y0 = np.clip(yc - dy / 2, 0, H - 1).round().astype(np.int32) + y1 = np.clip(yc + dy / 2, 0, H - 1).round().astype(np.int32) + + wid = x1 - x0 + hei = y1 - y0 + y00 = np.random.randint(0, H - hei) + x00 = np.random.randint(0, W - wid) + fr = np.random.randint(0, S) + rep = rgbs_alt[fr][y00 : y00 + hei, x00 : x00 + wid, :] + rgbs[i][y0:y1, x0:x1, :] = rep + + occ_inds = np.logical_and( + np.logical_and(trajs[i, :, 0] >= x0, trajs[i, :, 0] < x1), + np.logical_and(trajs[i, :, 1] >= y0, trajs[i, :, 1] < y1), + ) + visibles[i, occ_inds] = 0 + rgbs = [rgb.astype(np.uint8) for rgb in rgbs] + + ############ photometric augmentation ############ + if np.random.rand() < self.color_aug_prob: + # random per-frame amount of aug + rgbs = [ + np.array(self.photo_aug(Image.fromarray(rgb)), dtype=np.uint8) + for rgb in rgbs + ] + + if np.random.rand() < self.blur_aug_prob: + # random per-frame amount of blur + rgbs = [ + np.array(self.blur_aug(Image.fromarray(rgb)), dtype=np.uint8) + for rgb in rgbs + ] + + return rgbs, trajs, visibles + + def add_spatial_augs(self, rgbs, trajs, visibles, crop_size): + T, N, __ = trajs.shape + + S = len(rgbs) + H, W = rgbs[0].shape[:2] + assert S == T + + rgbs = [rgb.astype(np.float32) for rgb in rgbs] + + ############ spatial transform ############ + + # padding + pad_x0 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1]) + pad_x1 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1]) + pad_y0 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1]) + pad_y1 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1]) + + rgbs = [ + np.pad(rgb, ((pad_y0, pad_y1), (pad_x0, pad_x1), (0, 0))) for rgb in rgbs + ] + trajs[:, :, 0] += pad_x0 + trajs[:, :, 1] += pad_y0 + H, W = rgbs[0].shape[:2] + + # scaling + stretching + scale = np.random.uniform(self.resize_lim[0], self.resize_lim[1]) + scale_x = scale + scale_y = scale + H_new = H + W_new = W + + scale_delta_x = 0.0 + scale_delta_y = 0.0 + + rgbs_scaled = [] + for s in range(S): + if s == 1: + scale_delta_x = np.random.uniform(-self.resize_delta, self.resize_delta) + scale_delta_y = np.random.uniform(-self.resize_delta, self.resize_delta) + elif s > 1: + scale_delta_x = ( + scale_delta_x * 0.8 + + np.random.uniform(-self.resize_delta, self.resize_delta) * 0.2 + ) + scale_delta_y = ( + scale_delta_y * 0.8 + + np.random.uniform(-self.resize_delta, self.resize_delta) * 0.2 + ) + scale_x = scale_x + scale_delta_x + scale_y = scale_y + scale_delta_y + + # bring h/w closer + scale_xy = (scale_x + scale_y) * 0.5 + scale_x = scale_x * 0.5 + scale_xy * 0.5 + scale_y = scale_y * 0.5 + scale_xy * 0.5 + + # don't get too crazy + scale_x = np.clip(scale_x, 0.2, 2.0) + scale_y = np.clip(scale_y, 0.2, 2.0) + + H_new = int(H * scale_y) + W_new = int(W * scale_x) + + # make it at least slightly bigger than the crop area, + # so that the random cropping can add diversity + H_new = np.clip(H_new, crop_size[0] + 10, None) + W_new = np.clip(W_new, crop_size[1] + 10, None) + # recompute scale in case we clipped + scale_x = (W_new - 1) / float(W - 1) + scale_y = (H_new - 1) / float(H - 1) + rgbs_scaled.append( + cv2.resize(rgbs[s], (W_new, H_new), interpolation=cv2.INTER_LINEAR) + ) + trajs[s, :, 0] *= scale_x + trajs[s, :, 1] *= scale_y + rgbs = rgbs_scaled + ok_inds = visibles[0, :] > 0 + vis_trajs = trajs[:, ok_inds] # S,?,2 + + if vis_trajs.shape[1] > 0: + mid_x = np.mean(vis_trajs[0, :, 0]) + mid_y = np.mean(vis_trajs[0, :, 1]) + else: + mid_y = crop_size[0] + mid_x = crop_size[1] + + x0 = int(mid_x - crop_size[1] // 2) + y0 = int(mid_y - crop_size[0] // 2) + + offset_x = 0 + offset_y = 0 + + for s in range(S): + # on each frame, shift a bit more + if s == 1: + offset_x = np.random.randint( + -self.max_crop_offset, self.max_crop_offset + ) + offset_y = np.random.randint( + -self.max_crop_offset, self.max_crop_offset + ) + elif s > 1: + offset_x = int( + offset_x * 0.8 + + np.random.randint(-self.max_crop_offset, self.max_crop_offset + 1) + * 0.2 + ) + offset_y = int( + offset_y * 0.8 + + np.random.randint(-self.max_crop_offset, self.max_crop_offset + 1) + * 0.2 + ) + x0 = x0 + offset_x + y0 = y0 + offset_y + + H_new, W_new = rgbs[s].shape[:2] + if H_new == crop_size[0]: + y0 = 0 + else: + y0 = min(max(0, y0), H_new - crop_size[0] - 1) + + if W_new == crop_size[1]: + x0 = 0 + else: + x0 = min(max(0, x0), W_new - crop_size[1] - 1) + + rgbs[s] = rgbs[s][y0 : y0 + crop_size[0], x0 : x0 + crop_size[1]] + trajs[s, :, 0] -= x0 + trajs[s, :, 1] -= y0 + + H_new = crop_size[0] + W_new = crop_size[1] + + # flip + h_flipped = False + v_flipped = False + if self.do_flip: + # h flip + if np.random.rand() < self.h_flip_prob: + h_flipped = True + rgbs = [rgb[:, ::-1] for rgb in rgbs] + # v flip + if np.random.rand() < self.v_flip_prob: + v_flipped = True + rgbs = [rgb[::-1] for rgb in rgbs] + if h_flipped: + trajs[:, :, 0] = W_new - trajs[:, :, 0] + if v_flipped: + trajs[:, :, 1] = H_new - trajs[:, :, 1] + return np.stack(rgbs), trajs + + def crop(self, rgbs, trajs, crop_size): + T, N, _ = trajs.shape + + S = len(rgbs) + H, W = rgbs[0].shape[:2] + assert S == T + + ############ spatial transform ############ + + H_new = H + W_new = W + + # simple random crop + y0 = 0 if crop_size[0] >= H_new else (H_new - crop_size[0]) // 2 + # np.random.randint(0, + x0 = 0 if crop_size[1] >= W_new else np.random.randint(0, W_new - crop_size[1]) + rgbs = [rgb[y0 : y0 + crop_size[0], x0 : x0 + crop_size[1]] for rgb in rgbs] + + trajs[:, :, 0] -= x0 + trajs[:, :, 1] -= y0 + + return np.stack(rgbs), trajs + + +class KubricMovifDataset(CoTrackerDataset): + def __init__( + self, + data_root, + crop_size=(384, 512), + seq_len=24, + traj_per_sample=768, + sample_vis_last_frame=False, + use_augs=False, + random_seq_len=False, + random_frame_rate=False, + random_number_traj=False, + split="train", + ): + super(KubricMovifDataset, self).__init__( + data_root=data_root, + crop_size=crop_size, + seq_len=seq_len, + traj_per_sample=traj_per_sample, + sample_vis_last_frame=sample_vis_last_frame, + use_augs=use_augs, + ) + self.random_seq_len = random_seq_len + self.random_frame_rate = random_frame_rate + self.random_number_traj = random_number_traj + self.pad_bounds = [0, 25] + self.resize_lim = [0.75, 1.25] # sample resizes from here + self.resize_delta = 0.05 + self.max_crop_offset = 15 + self.split = split + + self.seq_names = [ + fname + for fname in os.listdir(data_root) + if os.path.isdir(os.path.join(data_root, fname)) + ] + if self.split == "valid": + self.seq_names = self.seq_names[:30] + assert use_augs == False + + print("found %d unique videos in %s" % (len(self.seq_names), self.data_root)) + + def getitem_helper(self, index): + gotit = True + seq_name = self.seq_names[index] + npy_path = os.path.join(self.data_root, seq_name, seq_name + ".npy") + rgb_path = os.path.join(self.data_root, seq_name, "frames") + + img_paths = sorted(os.listdir(rgb_path)) + rgbs = [] + for i, img_path in enumerate(img_paths): + rgbs.append(imageio.v2.imread(os.path.join(rgb_path, img_path))) + + rgbs = np.stack(rgbs) + annot_dict = np.load(npy_path, allow_pickle=True).item() + traj_2d = annot_dict["coords"] + visibility = annot_dict["visibility"] + + frame_rate = 1 + final_num_traj = self.traj_per_sample + crop_size = self.crop_size + + # random crop + min_num_traj = 1 + assert self.traj_per_sample >= min_num_traj + if self.random_seq_len and self.random_number_traj: + final_num_traj = np.random.randint(min_num_traj, self.traj_per_sample) + alpha = final_num_traj / float(self.traj_per_sample) + seq_len = int(alpha * 10 + (1 - alpha) * self.seq_len) + seq_len = np.random.randint(seq_len - 2, seq_len + 2) + if self.random_frame_rate: + frame_rate = np.random.randint(1, int((120 / seq_len)) + 1) + elif self.random_number_traj: + final_num_traj = np.random.randint(min_num_traj, self.traj_per_sample) + alpha = final_num_traj / float(self.traj_per_sample) + seq_len = 8 * int(alpha * 2 + (1 - alpha) * self.seq_len // 8) + # seq_len = np.random.randint(seq_len , seq_len + 2) + if self.random_frame_rate: + frame_rate = np.random.randint(1, int((120 / seq_len)) + 1) + elif self.random_seq_len: + seq_len = np.random.randint(int(self.seq_len / 2), self.seq_len) + if self.random_frame_rate: + frame_rate = np.random.randint(1, int((120 / seq_len)) + 1) + else: + seq_len = self.seq_len + if self.random_frame_rate: + frame_rate = np.random.randint(1, int((120 / seq_len)) + 1) + + traj_2d = np.transpose(traj_2d, (1, 0, 2)) + visibility = np.transpose(np.logical_not(visibility), (1, 0)) + + no_augs = False + if seq_len < len(rgbs): + if seq_len * frame_rate < len(rgbs): + start_ind = np.random.choice(len(rgbs) - (seq_len * frame_rate), 1)[0] + else: + start_ind = 0 + rgbs = rgbs[start_ind : start_ind + seq_len * frame_rate : frame_rate] + traj_2d = traj_2d[start_ind : start_ind + seq_len * frame_rate : frame_rate] + visibility = visibility[ + start_ind : start_ind + seq_len * frame_rate : frame_rate + ] + + assert seq_len <= len(rgbs) + + if not no_augs: + if self.use_augs: + rgbs, traj_2d, visibility = self.add_photometric_augs( + rgbs, traj_2d, visibility, replace=False + ) + rgbs, traj_2d = self.add_spatial_augs( + rgbs, traj_2d, visibility, crop_size + ) + else: + rgbs, traj_2d = self.crop(rgbs, traj_2d, crop_size) + + visibility[traj_2d[:, :, 0] > crop_size[1] - 1] = False + visibility[traj_2d[:, :, 0] < 0] = False + visibility[traj_2d[:, :, 1] > crop_size[0] - 1] = False + visibility[traj_2d[:, :, 1] < 0] = False + + visibility = torch.from_numpy(visibility) + traj_2d = torch.from_numpy(traj_2d) + + crop_tensor = torch.tensor(crop_size).flip(0)[None, None] / 2.0 + close_pts_inds = torch.all( + torch.linalg.vector_norm(traj_2d[..., :2] - crop_tensor, dim=-1) < 1000.0, + dim=0, + ) + traj_2d = traj_2d[:, close_pts_inds] + visibility = visibility[:, close_pts_inds] + + visibile_pts_first_frame_inds = (visibility[0]).nonzero(as_tuple=False)[:, 0] + + visibile_pts_mid_frame_inds = (visibility[seq_len // 2]).nonzero( + as_tuple=False + )[:, 0] + visibile_pts_inds = torch.cat( + (visibile_pts_first_frame_inds, visibile_pts_mid_frame_inds), dim=0 + ) + if self.sample_vis_last_frame: + visibile_pts_last_frame_inds = (visibility[seq_len - 1]).nonzero( + as_tuple=False + )[:, 0] + visibile_pts_inds = torch.cat( + (visibile_pts_inds, visibile_pts_last_frame_inds), dim=0 + ) + point_inds = torch.randperm(len(visibile_pts_inds))[: self.traj_per_sample] + if len(point_inds) < self.traj_per_sample: + gotit = False + + visible_inds_sampled = visibile_pts_inds[point_inds] + + trajs = traj_2d[:, visible_inds_sampled].float() + visibles = visibility[:, visible_inds_sampled] + valids = torch.ones_like(visibles) + + trajs = trajs[:, :final_num_traj] + visibles = visibles[:, :final_num_traj] + valids = valids[:, :final_num_traj] + + rgbs = torch.from_numpy(rgbs).permute(0, 3, 1, 2).float() + + sample = CoTrackerData( + video=rgbs, + trajectory=trajs, + visibility=visibles, + valid=valids, + seq_name=seq_name, + ) + return sample, gotit + + def __len__(self): + return len(self.seq_names) diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/real_dataset.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/real_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..5a51e9d15e46a02a4361816ae6e77660edd62aad --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/real_dataset.py @@ -0,0 +1,282 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import torch +import json +import cv2 +import math +import imageio +import numpy as np + +from cotracker.datasets.utils import CoTrackerData +from torchvision.transforms import ColorJitter, GaussianBlur +from PIL import Image +from cotracker.models.core.model_utils import smart_cat +from torchvision.io import read_video +import torchvision +from cotracker.datasets.utils import collate_fn, collate_fn_train, dataclass_to_cuda_ +import torchvision.transforms.functional as F + + +class RealDataset(torch.utils.data.Dataset): + def __init__( + self, + crop_size=(384, 512), + seq_len=24, + traj_per_sample=768, + random_frame_rate=False, + random_seq_len=False, + data_splits=[0], + random_resize=False, + limit_samples=10000, + ): + super(RealDataset, self).__init__() + np.random.seed(0) + torch.manual_seed(0) + raise ValueError(f"This dataset wasn't released. You should collect your own dataset of real videos before training with this dataset class.") + + stopwords = set( + [ + "river", + "water", + "shore", + "lake", + "sea", + "ocean", + "silhouette", + "matte", + "online", + "virtual", + "meditation", + "artwork", + "drawing", + "animation", + "abstract", + "background", + "concept", + "cartoon", + "symbolic", + "painting", + "sketch", + "fireworks", + "fire", + "sky", + "darkness", + "timelapse", + "time-lapse", + "cgi", + "computer", + "computer-generated", + "drawing", + "draw", + "cgi", + "animate", + "cartoon", + "static", + "abstract", + "abstraction", + "3d", + "fandom", + "fantasy", + "graphics", + "cell", + "holographic", + "generated", + "generation" "telephoto", + "animated", + "disko", + "generate" "2d", + "3d", + "geometric", + "geometry", + "render", + "rendering", + "timelapse", + "slomo", + "slo", + "wallpaper", + "pattern", + "tile", + "generated", + "chroma", + "www", + "http", + "cannabis", + "loop", + "cycle", + "alpha", + "abstract", + "concept", + "digital", + "graphic", + "skies", + "fountain", + "train", + "rapid", + "fast", + "quick", + "vfx", + "effect", + ] + ) + + def no_stopwords_in_key(key, stopwords): + for s in stopwords: + if s in key.split(","): + return False + return True + + filelist_all = [] + + for part in data_splits: + filelist = np.load('YOUR FILELIST') + captions = np.load('YOUR CAPTIONS') + keywords = np.load('YOUR KEYWORDS') + + filtered_seqs_motion = [ + i + for i, key in enumerate(keywords) + if "motion" in key.split(",") + and ( + "man" in key.split(",") + or "woman" in key.split(",") + or "animal" in key.split(",") + or "child" in key.split(",") + ) + and no_stopwords_in_key(key, stopwords) + ] + print("filtered_seqs_motion", len(filtered_seqs_motion)) + filtered_seqs = filtered_seqs_motion + + print(f"filtered_seqs {part}", len(filtered_seqs)) + filelist_all = filelist_all + filelist[filtered_seqs].tolist() + + if len(filelist_all) > limit_samples: + break + + self.filelist = filelist_all[:limit_samples] + print(f"found {len(self.filelist)} unique videos") + self.traj_per_sample = traj_per_sample + self.crop_size = crop_size + self.seq_len = seq_len + self.random_frame_rate = random_frame_rate + self.random_resize = random_resize + self.random_seq_len = random_seq_len + + def crop(self, rgbs): + S = len(rgbs) + + H, W = rgbs.shape[2:] + + H_new = H + W_new = W + + # simple random crop + y0 = ( + 0 + if self.crop_size[0] >= H_new + else np.random.randint(0, H_new - self.crop_size[0]) + ) + x0 = ( + 0 + if self.crop_size[1] >= W_new + else np.random.randint(0, W_new - self.crop_size[1]) + ) + rgbs = [ + rgb[:, y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]] + for rgb in rgbs + ] + + return torch.stack(rgbs) + + def __getitem__(self, index): + gotit = False + + sample, gotit = self.getitem_helper(index) + if not gotit: + print("warning: sampling failed") + # fake sample, so we can still collate + sample = CoTrackerData( + video=torch.zeros( + (self.seq_len, 3, self.crop_size[0], self.crop_size[1]) + ), + trajectory=torch.ones(1, 1, 1, 2), + visibility=torch.ones(1, 1, 1), + valid=torch.ones(1, 1, 1), + ) + + return sample, gotit + + def sample_h_w(self): + area = np.random.uniform(0.6, 1) + a1 = np.random.uniform(area, 1) + a2 = np.random.uniform(area, 1) + h = (a1 + a2) / 2.0 + w = area / h + return h, w + + def getitem_helper(self, index): + gotit = True + video_path = self.filelist[index] + + rgbs, _, _ = read_video(str(video_path), output_format="TCHW", pts_unit="sec") + if rgbs.numel() == 0: + return None, False + seq_name = video_path + frame_rate = 1 + + if self.random_seq_len: + seq_len = np.random.randint(int(self.seq_len / 2), self.seq_len) + else: + seq_len = self.seq_len + + while len(rgbs) < seq_len: + rgbs = torch.cat([rgbs, rgbs.flip(0)]) + if seq_len < 8: + print("seq_len < 8, return NONE") + return None, False + if self.random_frame_rate: + max_frame_rate = min(4, int((len(rgbs) / seq_len))) + if max_frame_rate > 1: + frame_rate = np.random.randint(1, max_frame_rate) + + if seq_len * frame_rate < len(rgbs): + start_ind = np.random.choice(len(rgbs) - (seq_len * frame_rate), 1)[0] + else: + start_ind = 0 + rgbs = rgbs[start_ind : start_ind + seq_len * frame_rate : frame_rate] + + assert seq_len <= len(rgbs) + + if self.random_resize and np.random.rand() < 0.5: + video = [] + rgbs = rgbs.permute(0, 2, 3, 1).numpy() + + for i in range(len(rgbs)): + rgb = cv2.resize( + rgbs[i], + (self.crop_size[1], self.crop_size[0]), + interpolation=cv2.INTER_LINEAR, + ) + video.append(rgb) + video = torch.tensor(np.stack(video)).permute(0, 3, 1, 2) + + else: + video = self.crop(rgbs) + + sample = CoTrackerData( + video=video, + trajectory=torch.ones(seq_len, self.traj_per_sample, 2), + visibility=torch.ones(seq_len, self.traj_per_sample), + valid=torch.ones(seq_len, self.traj_per_sample), + seq_name=seq_name, + ) + + return sample, gotit + + def __len__(self): + return len(self.filelist) diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/tap_vid_datasets.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/tap_vid_datasets.py new file mode 100644 index 0000000000000000000000000000000000000000..1178741f2512a446458745d1753b4b93cb0e9ec0 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/tap_vid_datasets.py @@ -0,0 +1,244 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import io +import glob +import torch +import pickle +import numpy as np +import mediapy as media +import random +from PIL import Image +from typing import Mapping, Tuple, Union + +from cotracker.datasets.utils import CoTrackerData + +DatasetElement = Mapping[str, Mapping[str, Union[np.ndarray, str]]] + + +def resize_video(video: np.ndarray, output_size: Tuple[int, int]) -> np.ndarray: + """Resize a video to output_size.""" + # If you have a GPU, consider replacing this with a GPU-enabled resize op, + # such as a jitted jax.image.resize. It will make things faster. + return media.resize_video(video, output_size) + + +def sample_queries_first( + target_occluded: np.ndarray, + target_points: np.ndarray, + frames: np.ndarray, +) -> Mapping[str, np.ndarray]: + """Package a set of frames and tracks for use in TAPNet evaluations. + Given a set of frames and tracks with no query points, use the first + visible point in each track as the query. + Args: + target_occluded: Boolean occlusion flag, of shape [n_tracks, n_frames], + where True indicates occluded. + target_points: Position, of shape [n_tracks, n_frames, 2], where each point + is [x,y] scaled between 0 and 1. + frames: Video tensor, of shape [n_frames, height, width, 3]. Scaled between + -1 and 1. + Returns: + A dict with the keys: + video: Video tensor of shape [1, n_frames, height, width, 3] + query_points: Query points of shape [1, n_queries, 3] where + each point is [t, y, x] scaled to the range [-1, 1] + target_points: Target points of shape [1, n_queries, n_frames, 2] where + each point is [x, y] scaled to the range [-1, 1] + """ + valid = np.sum(~target_occluded, axis=1) > 0 + target_points = target_points[valid, :] + target_occluded = target_occluded[valid, :] + + query_points = [] + for i in range(target_points.shape[0]): + index = np.where(target_occluded[i] == 0)[0][0] + x, y = target_points[i, index, 0], target_points[i, index, 1] + query_points.append(np.array([index, y, x])) # [t, y, x] + query_points = np.stack(query_points, axis=0) + + return { + "video": frames[np.newaxis, ...], + "query_points": query_points[np.newaxis, ...], + "target_points": target_points[np.newaxis, ...], + "occluded": target_occluded[np.newaxis, ...], + } + + +def sample_queries_strided( + target_occluded: np.ndarray, + target_points: np.ndarray, + frames: np.ndarray, + query_stride: int = 5, +) -> Mapping[str, np.ndarray]: + """Package a set of frames and tracks for use in TAPNet evaluations. + + Given a set of frames and tracks with no query points, sample queries + strided every query_stride frames, ignoring points that are not visible + at the selected frames. + + Args: + target_occluded: Boolean occlusion flag, of shape [n_tracks, n_frames], + where True indicates occluded. + target_points: Position, of shape [n_tracks, n_frames, 2], where each point + is [x,y] scaled between 0 and 1. + frames: Video tensor, of shape [n_frames, height, width, 3]. Scaled between + -1 and 1. + query_stride: When sampling query points, search for un-occluded points + every query_stride frames and convert each one into a query. + + Returns: + A dict with the keys: + video: Video tensor of shape [1, n_frames, height, width, 3]. The video + has floats scaled to the range [-1, 1]. + query_points: Query points of shape [1, n_queries, 3] where + each point is [t, y, x] scaled to the range [-1, 1]. + target_points: Target points of shape [1, n_queries, n_frames, 2] where + each point is [x, y] scaled to the range [-1, 1]. + trackgroup: Index of the original track that each query point was + sampled from. This is useful for visualization. + """ + tracks = [] + occs = [] + queries = [] + trackgroups = [] + total = 0 + trackgroup = np.arange(target_occluded.shape[0]) + for i in range(0, target_occluded.shape[1], query_stride): + mask = target_occluded[:, i] == 0 + query = np.stack( + [ + i * np.ones(target_occluded.shape[0:1]), + target_points[:, i, 1], + target_points[:, i, 0], + ], + axis=-1, + ) + queries.append(query[mask]) + tracks.append(target_points[mask]) + occs.append(target_occluded[mask]) + trackgroups.append(trackgroup[mask]) + total += np.array(np.sum(target_occluded[:, i] == 0)) + + return { + "video": frames[np.newaxis, ...], + "query_points": np.concatenate(queries, axis=0)[np.newaxis, ...], + "target_points": np.concatenate(tracks, axis=0)[np.newaxis, ...], + "occluded": np.concatenate(occs, axis=0)[np.newaxis, ...], + "trackgroup": np.concatenate(trackgroups, axis=0)[np.newaxis, ...], + } + + +class TapVidDataset(torch.utils.data.Dataset): + def __init__( + self, + data_root, + dataset_type="davis", + resize_to=[256, 256], + queried_first=True, + fast_eval=False, + ): + local_random = random.Random() + local_random.seed(42) + self.fast_eval = fast_eval + self.dataset_type = dataset_type + self.resize_to = resize_to + self.queried_first = queried_first + if self.dataset_type == "kinetics": + all_paths = glob.glob(os.path.join(data_root, "*_of_0010.pkl")) + points_dataset = [] + for pickle_path in all_paths: + with open(pickle_path, "rb") as f: + data = pickle.load(f) + points_dataset = points_dataset + data + if fast_eval: + points_dataset = local_random.sample(points_dataset, 50) + self.points_dataset = points_dataset + + elif self.dataset_type == "robotap": + all_paths = glob.glob(os.path.join(data_root, "robotap_split*.pkl")) + points_dataset = None + for pickle_path in all_paths: + with open(pickle_path, "rb") as f: + data = pickle.load(f) + if points_dataset is None: + points_dataset = dict(data) + else: + points_dataset.update(data) + if fast_eval: + points_dataset_keys = local_random.sample( + sorted(points_dataset.keys()), 50 + ) + points_dataset = {k: points_dataset[k] for k in points_dataset_keys} + self.points_dataset = points_dataset + self.video_names = list(self.points_dataset.keys()) + else: + with open(data_root, "rb") as f: + self.points_dataset = pickle.load(f) + if self.dataset_type == "davis": + self.video_names = list(self.points_dataset.keys()) + elif self.dataset_type == "stacking": + # print("self.points_dataset", self.points_dataset) + self.video_names = [i for i in range(len(self.points_dataset))] + print("found %d unique videos in %s" % (len(self.points_dataset), data_root)) + + def __getitem__(self, index): + if self.dataset_type == "davis" or self.dataset_type == "robotap": + video_name = self.video_names[index] + else: + video_name = index + video = self.points_dataset[video_name] + frames = video["video"] + + if self.fast_eval and frames.shape[0] > 300: + return self.__getitem__((index + 1) % self.__len__()) + if isinstance(frames[0], bytes): + # TAP-Vid is stored and JPEG bytes rather than `np.ndarray`s. + def decode(frame): + byteio = io.BytesIO(frame) + img = Image.open(byteio) + return np.array(img) + + frames = np.array([decode(frame) for frame in frames]) + + target_points = self.points_dataset[video_name]["points"] + if self.resize_to is not None: + frames = resize_video(frames, self.resize_to) + target_points *= np.array( + [self.resize_to[1] - 1, self.resize_to[0] - 1] + ) # 1 should be mapped to resize_to-1 + else: + target_points *= np.array([frames.shape[2] - 1, frames.shape[1] - 1]) + + target_occ = self.points_dataset[video_name]["occluded"] + if self.queried_first: + converted = sample_queries_first(target_occ, target_points, frames) + else: + converted = sample_queries_strided(target_occ, target_points, frames) + assert converted["target_points"].shape[1] == converted["query_points"].shape[1] + + trajs = ( + torch.from_numpy(converted["target_points"])[0].permute(1, 0, 2).float() + ) # T, N, D + + rgbs = torch.from_numpy(frames).permute(0, 3, 1, 2).float() + visibles = torch.logical_not(torch.from_numpy(converted["occluded"]))[ + 0 + ].permute( + 1, 0 + ) # T, N + query_points = torch.from_numpy(converted["query_points"])[0] # T, N + return CoTrackerData( + rgbs, + trajs, + visibles, + seq_name=str(video_name), + query_points=query_points, + ) + + def __len__(self): + return len(self.points_dataset) diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/utils.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..eda3ade205f1eba6e9a801d548b53565ff84adff --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/datasets/utils.py @@ -0,0 +1,120 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + + +import torch +import dataclasses +import torch.nn.functional as F +from dataclasses import dataclass +from typing import Any, Optional, Dict + + +@dataclass(eq=False) +class CoTrackerData: + """ + Dataclass for storing video tracks data. + """ + + video: torch.Tensor # B, S, C, H, W + trajectory: torch.Tensor # B, S, N, 2 + visibility: torch.Tensor # B, S, N + # optional data + valid: Optional[torch.Tensor] = None # B, S, N + segmentation: Optional[torch.Tensor] = None # B, S, 1, H, W + seq_name: Optional[str] = None + query_points: Optional[torch.Tensor] = None # TapVID evaluation format + transforms: Optional[Dict[str, Any]] = None + aug_video: Optional[torch.Tensor] = None + + +def collate_fn(batch): + """ + Collate function for video tracks data. + """ + video = torch.stack([b.video for b in batch], dim=0) + trajectory = torch.stack([b.trajectory for b in batch], dim=0) + visibility = torch.stack([b.visibility for b in batch], dim=0) + query_points = segmentation = None + if batch[0].query_points is not None: + query_points = torch.stack([b.query_points for b in batch], dim=0) + if batch[0].segmentation is not None: + segmentation = torch.stack([b.segmentation for b in batch], dim=0) + seq_name = [b.seq_name for b in batch] + + return CoTrackerData( + video=video, + trajectory=trajectory, + visibility=visibility, + segmentation=segmentation, + seq_name=seq_name, + query_points=query_points, + ) + + +def collate_fn_train(batch): + """ + Collate function for video tracks data during training. + """ + gotit = [gotit for _, gotit in batch] + video = torch.stack([b.video for b, _ in batch], dim=0) + trajectory = torch.stack([b.trajectory for b, _ in batch], dim=0) + visibility = torch.stack([b.visibility for b, _ in batch], dim=0) + valid = torch.stack([b.valid for b, _ in batch], dim=0) + seq_name = [b.seq_name for b, _ in batch] + query_points = transforms = aug_video = None + if batch[0][0].query_points is not None: + query_points = torch.stack([b.query_points for b, _ in batch], dim=0) + + if batch[0][0].transforms is not None: + transforms = [b.transforms for b, _ in batch] + + if batch[0][0].aug_video is not None: + aug_video = torch.stack([b.aug_video for b, _ in batch], dim=0) + return ( + CoTrackerData( + video=video, + trajectory=trajectory, + visibility=visibility, + valid=valid, + seq_name=seq_name, + query_points=query_points, + aug_video=aug_video, + transforms=transforms, + ), + gotit, + ) + + +def try_to_cuda(t: Any) -> Any: + """ + Try to move the input variable `t` to a cuda device. + + Args: + t: Input. + + Returns: + t_cuda: `t` moved to a cuda device, if supported. + """ + try: + t = t.float().cuda() + except AttributeError: + pass + return t + + +def dataclass_to_cuda_(obj): + """ + Move all contents of a dataclass to cuda inplace if supported. + + Args: + batch: Input dataclass. + + Returns: + batch_cuda: `batch` moved to a cuda device, if supported. + """ + for f in dataclasses.fields(obj): + setattr(obj, f.name, try_to_cuda(getattr(obj, f.name))) + return obj diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_dynamic_replica.yaml b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_dynamic_replica.yaml new file mode 100644 index 0000000000000000000000000000000000000000..7d6fca91f30333b0ef9ff0e7392d481a3edcc270 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_dynamic_replica.yaml @@ -0,0 +1,6 @@ +defaults: + - default_config_eval +exp_dir: ./outputs/cotracker +dataset_name: dynamic_replica + + \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_davis_first.yaml b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_davis_first.yaml new file mode 100644 index 0000000000000000000000000000000000000000..d37a6c9cb8879c7e09ecd760eaa9fb767ec1d78f --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_davis_first.yaml @@ -0,0 +1,6 @@ +defaults: + - default_config_eval +exp_dir: ./outputs/cotracker +dataset_name: tapvid_davis_first + + \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_davis_strided.yaml b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_davis_strided.yaml new file mode 100644 index 0000000000000000000000000000000000000000..6e3cf3c1c1d7fe8ad0c5986af4d2ef973dbaa02f --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_davis_strided.yaml @@ -0,0 +1,6 @@ +defaults: + - default_config_eval +exp_dir: ./outputs/cotracker +dataset_name: tapvid_davis_strided + + \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_kinetics_first.yaml b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_kinetics_first.yaml new file mode 100644 index 0000000000000000000000000000000000000000..3be89144e1b635a72180532ef31a5512d6d4960f --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_kinetics_first.yaml @@ -0,0 +1,6 @@ +defaults: + - default_config_eval +exp_dir: ./outputs/cotracker +dataset_name: tapvid_kinetics_first + + \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_robotap_first.yaml b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_robotap_first.yaml new file mode 100644 index 0000000000000000000000000000000000000000..f259cdd604595d3dffe4aac056b1356e219d8e79 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_robotap_first.yaml @@ -0,0 +1,4 @@ +defaults: + - default_config_eval +exp_dir: ./outputs/cotracker +dataset_name: tapvid_robotap_first \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_stacking_first.yaml b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_stacking_first.yaml new file mode 100644 index 0000000000000000000000000000000000000000..ebde184297731e92f334ee12d80acb7acdc6c4a3 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_stacking_first.yaml @@ -0,0 +1,6 @@ +defaults: + - default_config_eval +exp_dir: ./outputs/cotracker +dataset_name: tapvid_stacking_first + + \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_stacking_strided.yaml b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_stacking_strided.yaml new file mode 100644 index 0000000000000000000000000000000000000000..7237a46901ea60cabacbc0d751e8dabdff1e0a77 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/configs/eval_tapvid_stacking_strided.yaml @@ -0,0 +1,6 @@ +defaults: + - default_config_eval +exp_dir: ./outputs/cotracker +dataset_name: tapvid_stacking_strided + + \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/eval_utils.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/eval_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..7002fa557eb4af487cf8536df87b297fd94ae236 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/eval_utils.py @@ -0,0 +1,138 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import numpy as np + +from typing import Iterable, Mapping, Tuple, Union + + +def compute_tapvid_metrics( + query_points: np.ndarray, + gt_occluded: np.ndarray, + gt_tracks: np.ndarray, + pred_occluded: np.ndarray, + pred_tracks: np.ndarray, + query_mode: str, +) -> Mapping[str, np.ndarray]: + """Computes TAP-Vid metrics (Jaccard, Pts. Within Thresh, Occ. Acc.) + See the TAP-Vid paper for details on the metric computation. All inputs are + given in raster coordinates. The first three arguments should be the direct + outputs of the reader: the 'query_points', 'occluded', and 'target_points'. + The paper metrics assume these are scaled relative to 256x256 images. + pred_occluded and pred_tracks are your algorithm's predictions. + This function takes a batch of inputs, and computes metrics separately for + each video. The metrics for the full benchmark are a simple mean of the + metrics across the full set of videos. These numbers are between 0 and 1, + but the paper multiplies them by 100 to ease reading. + Args: + query_points: The query points, an in the format [t, y, x]. Its size is + [b, n, 3], where b is the batch size and n is the number of queries + gt_occluded: A boolean array of shape [b, n, t], where t is the number + of frames. True indicates that the point is occluded. + gt_tracks: The target points, of shape [b, n, t, 2]. Each point is + in the format [x, y] + pred_occluded: A boolean array of predicted occlusions, in the same + format as gt_occluded. + pred_tracks: An array of track predictions from your algorithm, in the + same format as gt_tracks. + query_mode: Either 'first' or 'strided', depending on how queries are + sampled. If 'first', we assume the prior knowledge that all points + before the query point are occluded, and these are removed from the + evaluation. + Returns: + A dict with the following keys: + occlusion_accuracy: Accuracy at predicting occlusion. + pts_within_{x} for x in [1, 2, 4, 8, 16]: Fraction of points + predicted to be within the given pixel threshold, ignoring occlusion + prediction. + jaccard_{x} for x in [1, 2, 4, 8, 16]: Jaccard metric for the given + threshold + average_pts_within_thresh: average across pts_within_{x} + average_jaccard: average across jaccard_{x} + """ + + metrics = {} + # Fixed bug is described in: + # https://github.com/facebookresearch/co-tracker/issues/20 + eye = np.eye(gt_tracks.shape[2], dtype=np.int32) + + if query_mode == "first": + # evaluate frames after the query frame + query_frame_to_eval_frames = np.cumsum(eye, axis=1) - eye + elif query_mode == "strided": + # evaluate all frames except the query frame + query_frame_to_eval_frames = 1 - eye + else: + raise ValueError("Unknown query mode " + query_mode) + + query_frame = query_points[..., 0] + query_frame = np.round(query_frame).astype(np.int32) + evaluation_points = query_frame_to_eval_frames[query_frame] > 0 + + # Occlusion accuracy is simply how often the predicted occlusion equals the + # ground truth. + occ_acc = np.sum( + np.equal(pred_occluded, gt_occluded) & evaluation_points, + axis=(1, 2), + ) / np.sum(evaluation_points) + metrics["occlusion_accuracy"] = occ_acc + + # Next, convert the predictions and ground truth positions into pixel + # coordinates. + visible = np.logical_not(gt_occluded) + pred_visible = np.logical_not(pred_occluded) + all_frac_within = [] + all_jaccard = [] + for thresh in [1, 2, 4, 8, 16]: + # True positives are points that are within the threshold and where both + # the prediction and the ground truth are listed as visible. + within_dist = np.sum( + np.square(pred_tracks - gt_tracks), + axis=-1, + ) < np.square(thresh) + is_correct = np.logical_and(within_dist, visible) + + # Compute the frac_within_threshold, which is the fraction of points + # within the threshold among points that are visible in the ground truth, + # ignoring whether they're predicted to be visible. + count_correct = np.sum( + is_correct & evaluation_points, + axis=(1, 2), + ) + count_visible_points = np.sum(visible & evaluation_points, axis=(1, 2)) + frac_correct = count_correct / count_visible_points + metrics["pts_within_" + str(thresh)] = frac_correct + all_frac_within.append(frac_correct) + + true_positives = np.sum( + is_correct & pred_visible & evaluation_points, axis=(1, 2) + ) + + # The denominator of the jaccard metric is the true positives plus + # false positives plus false negatives. However, note that true positives + # plus false negatives is simply the number of points in the ground truth + # which is easier to compute than trying to compute all three quantities. + # Thus we just add the number of points in the ground truth to the number + # of false positives. + # + # False positives are simply points that are predicted to be visible, + # but the ground truth is not visible or too far from the prediction. + gt_positives = np.sum(visible & evaluation_points, axis=(1, 2)) + false_positives = (~visible) & pred_visible + false_positives = false_positives | ((~within_dist) & pred_visible) + false_positives = np.sum(false_positives & evaluation_points, axis=(1, 2)) + jaccard = true_positives / (gt_positives + false_positives) + metrics["jaccard_" + str(thresh)] = jaccard + all_jaccard.append(jaccard) + metrics["average_jaccard"] = np.mean( + np.stack(all_jaccard, axis=1), + axis=1, + ) + metrics["average_pts_within_thresh"] = np.mean( + np.stack(all_frac_within, axis=1), + axis=1, + ) + return metrics diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/evaluator.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/evaluator.py new file mode 100644 index 0000000000000000000000000000000000000000..7b31f9d6d0b4d337ae3d639fe7b6630c7cb8ab3d --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/core/evaluator.py @@ -0,0 +1,288 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +from collections import defaultdict +import os +from typing import Optional +import torch +from tqdm import tqdm +import numpy as np + +from torch.utils.tensorboard import SummaryWriter +from cotracker.datasets.utils import dataclass_to_cuda_ +from cotracker.utils.visualizer import Visualizer +from cotracker.models.core.model_utils import reduce_masked_mean +from cotracker.evaluation.core.eval_utils import compute_tapvid_metrics +from cotracker.predictor import CoTrackerOnlinePredictor +from cotracker.models.core.cotracker.cotracker3_offline import CoTrackerThreeOffline +from cotracker.models.core.cotracker.cotracker3_online import CoTrackerThreeOnline +import logging + + +class Evaluator: + """ + A class defining the CoTracker evaluator. + """ + + def __init__(self, exp_dir) -> None: + # Visualization + self.exp_dir = exp_dir + os.makedirs(exp_dir, exist_ok=True) + self.visualization_filepaths = defaultdict(lambda: defaultdict(list)) + self.visualize_dir = os.path.join(exp_dir, "visualisations") + + def compute_metrics(self, metrics, sample, pred_trajectory, dataset_name): + if isinstance(pred_trajectory, tuple): + pred_trajectory, pred_visibility = pred_trajectory + else: + pred_visibility = None + if "tapvid" in dataset_name: + B, T, N, D = sample.trajectory.shape + traj = sample.trajectory.clone() + thr = 0.6 + + if pred_visibility is None: + logging.warning("visibility is NONE") + pred_visibility = torch.zeros_like(sample.visibility) + + if not pred_visibility.dtype == torch.bool: + pred_visibility = pred_visibility > thr + + query_points = sample.query_points.clone().cpu().numpy() + + pred_visibility = pred_visibility[:, :, :N] + pred_trajectory = pred_trajectory[:, :, :N] + + gt_tracks = traj.permute(0, 2, 1, 3).cpu().numpy() + gt_occluded = ( + torch.logical_not(sample.visibility.clone().permute(0, 2, 1)) + .cpu() + .numpy() + ) + + pred_occluded = ( + torch.logical_not(pred_visibility.clone().permute(0, 2, 1)) + .cpu() + .numpy() + ) + pred_tracks = pred_trajectory.permute(0, 2, 1, 3).cpu().numpy() + + out_metrics = compute_tapvid_metrics( + query_points, + gt_occluded, + gt_tracks, + pred_occluded, + pred_tracks, + query_mode="strided" if "strided" in dataset_name else "first", + ) + + metrics[sample.seq_name[0]] = out_metrics + for metric_name in out_metrics.keys(): + if "avg" not in metrics: + metrics["avg"] = {} + metrics["avg"][metric_name] = np.mean( + [v[metric_name] for k, v in metrics.items() if k != "avg"] + ) + + logging.info(f"Metrics: {out_metrics}") + logging.info(f"avg: {metrics['avg']}") + print("metrics", out_metrics) + print("avg", metrics["avg"]) + elif dataset_name == "dynamic_replica" or dataset_name == "pointodyssey": + *_, N, _ = sample.trajectory.shape + B, T, N = sample.visibility.shape + H, W = sample.video.shape[-2:] + device = sample.video.device + + out_metrics = {} + + d_vis_sum = d_occ_sum = d_sum_all = 0.0 + thrs = [1, 2, 4, 8, 16] + sx_ = (W - 1) / 255.0 + sy_ = (H - 1) / 255.0 + sc_py = np.array([sx_, sy_]).reshape([1, 1, 2]) + sc_pt = torch.from_numpy(sc_py).float().to(device) + __, first_visible_inds = torch.max(sample.visibility, dim=1) + + frame_ids_tensor = torch.arange(T, device=device)[None, :, None].repeat( + B, 1, N + ) + start_tracking_mask = frame_ids_tensor > (first_visible_inds.unsqueeze(1)) + + for thr in thrs: + d_ = ( + torch.norm( + pred_trajectory[..., :2] / sc_pt + - sample.trajectory[..., :2] / sc_pt, + dim=-1, + ) + < thr + ).float() # B,S-1,N + d_occ = ( + reduce_masked_mean( + d_, (1 - sample.visibility) * start_tracking_mask + ).item() + * 100.0 + ) + d_occ_sum += d_occ + out_metrics[f"accuracy_occ_{thr}"] = d_occ + + d_vis = ( + reduce_masked_mean( + d_, sample.visibility * start_tracking_mask + ).item() + * 100.0 + ) + d_vis_sum += d_vis + out_metrics[f"accuracy_vis_{thr}"] = d_vis + + d_all = reduce_masked_mean(d_, start_tracking_mask).item() * 100.0 + d_sum_all += d_all + out_metrics[f"accuracy_{thr}"] = d_all + + d_occ_avg = d_occ_sum / len(thrs) + d_vis_avg = d_vis_sum / len(thrs) + d_all_avg = d_sum_all / len(thrs) + + sur_thr = 50 + dists = torch.norm( + pred_trajectory[..., :2] / sc_pt - sample.trajectory[..., :2] / sc_pt, + dim=-1, + ) # B,S,N + dist_ok = 1 - (dists > sur_thr).float() * sample.visibility # B,S,N + survival = torch.cumprod(dist_ok, dim=1) # B,S,N + out_metrics["survival"] = torch.mean(survival).item() * 100.0 + + out_metrics["accuracy_occ"] = d_occ_avg + out_metrics["accuracy_vis"] = d_vis_avg + out_metrics["accuracy"] = d_all_avg + + metrics[sample.seq_name[0]] = out_metrics + for metric_name in out_metrics.keys(): + if "avg" not in metrics: + metrics["avg"] = {} + metrics["avg"][metric_name] = float( + np.mean([v[metric_name] for k, v in metrics.items() if k != "avg"]) + ) + + logging.info(f"Metrics: {out_metrics}") + logging.info(f"avg: {metrics['avg']}") + print("metrics", out_metrics) + print("avg", metrics["avg"]) + + @torch.no_grad() + def evaluate_sequence( + self, + model, + test_dataloader: torch.utils.data.DataLoader, + dataset_name: str, + train_mode=False, + visualize_every: int = 50, + writer: Optional[SummaryWriter] = None, + step: Optional[int] = 0, + ): + metrics = {} + + vis = Visualizer( + save_dir=self.exp_dir, + fps=7, + ) + + for ind, sample in enumerate(tqdm(test_dataloader)): + if isinstance(sample, tuple): + sample, gotit = sample + if not all(gotit): + print("batch is None") + continue + if torch.cuda.is_available(): + dataclass_to_cuda_(sample) + device = torch.device("cuda") + else: + device = torch.device("cpu") + + if ( + not train_mode + and hasattr(model, "sequence_len") + and (sample.visibility[:, : model.sequence_len].sum() == 0) + ): + print(f"skipping batch {ind}") + continue + + if "tapvid" in dataset_name: + queries = sample.query_points.clone().float() + + queries = torch.stack( + [ + queries[:, :, 0], + queries[:, :, 2], + queries[:, :, 1], + ], + dim=2, + ).to(device) + else: + queries = torch.cat( + [ + torch.zeros_like(sample.trajectory[:, 0, :, :1]), + sample.trajectory[:, 0], + ], + dim=2, + ).to(device) + + if isinstance(model.model, CoTrackerThreeOnline): + online_model = CoTrackerOnlinePredictor(checkpoint=None) + online_model.model = model.model + online_model.step = model.model.window_len // 2 + online_model( + video_chunk=sample.video, + is_first_step=True, + queries=queries, + add_support_grid=False, + ) + # Process the video + for ind in range( + 0, sample.video.shape[1] - online_model.step, online_model.step + ): + pred_tracks, pred_visibility = online_model( + video_chunk=sample.video[:, ind : ind + online_model.step * 2], + add_support_grid=False, + grid_size=0, + ) # B T N 2, B T N 1 + pred_tracks = (pred_tracks, pred_visibility) + else: + pred_tracks = model(sample.video, queries) + + if "strided" in dataset_name: + inv_video = sample.video.flip(1).clone() + inv_queries = queries.clone() + inv_queries[:, :, 0] = inv_video.shape[1] - inv_queries[:, :, 0] - 1 + + pred_trj, pred_vsb = pred_tracks + inv_pred_trj, inv_pred_vsb = model(inv_video, inv_queries) + + inv_pred_trj = inv_pred_trj.flip(1) + inv_pred_vsb = inv_pred_vsb.flip(1) + + mask = pred_trj == 0 + + pred_trj[mask] = inv_pred_trj[mask] + pred_vsb[mask[:, :, :, 0]] = inv_pred_vsb[mask[:, :, :, 0]] + + pred_tracks = pred_trj, pred_vsb + + if dataset_name == "badja" or dataset_name == "fastcapture": + seq_name = sample.seq_name[0] + else: + seq_name = str(ind) + if ind % visualize_every == 0: + vis.visualize( + sample.video, + pred_tracks[0] if isinstance(pred_tracks, tuple) else pred_tracks, + filename=dataset_name + "_" + seq_name, + writer=writer, + step=step, + ) + self.compute_metrics(metrics, sample, pred_tracks, dataset_name) + return metrics diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/evaluate.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/evaluate.py new file mode 100644 index 0000000000000000000000000000000000000000..f8d25d2faa15dc5157b5498c7c2890246aac765e --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/evaluation/evaluate.py @@ -0,0 +1,190 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import json +import os +import hydra +import numpy as np +import torch + +from typing import Optional +from dataclasses import dataclass, field + +from omegaconf import OmegaConf + +from cotracker.datasets.utils import collate_fn +from cotracker.models.evaluation_predictor import EvaluationPredictor + +from cotracker.evaluation.core.evaluator import Evaluator +from cotracker.models.build_cotracker import build_cotracker + + +@dataclass(eq=False) +class DefaultConfig: + # Directory where all outputs of the experiment will be saved. + exp_dir: str = "./outputs" + + # Name of the dataset to be used for the evaluation. + dataset_name: str = "tapvid_davis_first" + # The root directory of the dataset. + dataset_root: str = "./" + + # Path to the pre-trained model checkpoint to be used for the evaluation. + # The default value is the path to a specific CoTracker model checkpoint. + checkpoint: str = "./checkpoints/scaled_online.pth" + # EvaluationPredictor parameters + # The size (N) of the support grid used in the predictor. + # The total number of points is (N*N). + grid_size: int = 5 + # The size (N) of the local support grid. + local_grid_size: int = 8 + num_uniformly_sampled_pts: int = 0 + sift_size: int = 0 + # A flag indicating whether to evaluate one ground truth point at a time. + single_point: bool = False + offline_model: bool = False + window_len: int = 16 + # The number of iterative updates for each sliding window. + n_iters: int = 6 + + seed: int = 0 + gpu_idx: int = 0 + local_extent: int = 50 + + v2: bool = False + + # Override hydra's working directory to current working dir, + # also disable storing the .hydra logs: + hydra: dict = field( + default_factory=lambda: { + "run": {"dir": "."}, + "output_subdir": None, + } + ) + + +def run_eval(cfg: DefaultConfig): + """ + The function evaluates CoTracker on a specified benchmark dataset based on a provided configuration. + + Args: + cfg (DefaultConfig): An instance of DefaultConfig class which includes: + - exp_dir (str): The directory path for the experiment. + - dataset_name (str): The name of the dataset to be used. + - dataset_root (str): The root directory of the dataset. + - checkpoint (str): The path to the CoTracker model's checkpoint. + - single_point (bool): A flag indicating whether to evaluate one ground truth point at a time. + - n_iters (int): The number of iterative updates for each sliding window. + - seed (int): The seed for setting the random state for reproducibility. + - gpu_idx (int): The index of the GPU to be used. + """ + # Creating the experiment directory if it doesn't exist + os.makedirs(cfg.exp_dir, exist_ok=True) + + # Saving the experiment configuration to a .yaml file in the experiment directory + cfg_file = os.path.join(cfg.exp_dir, "expconfig.yaml") + with open(cfg_file, "w") as f: + OmegaConf.save(config=cfg, f=f) + + evaluator = Evaluator(cfg.exp_dir) + cotracker_model = build_cotracker( + cfg.checkpoint, offline=cfg.offline_model, window_len=cfg.window_len, v2=cfg.v2 + ) + + # Creating the EvaluationPredictor object + predictor = EvaluationPredictor( + cotracker_model, + grid_size=cfg.grid_size, + local_grid_size=cfg.local_grid_size, + sift_size=cfg.sift_size, + single_point=cfg.single_point, + num_uniformly_sampled_pts=cfg.num_uniformly_sampled_pts, + n_iters=cfg.n_iters, + local_extent=cfg.local_extent, + interp_shape=(384, 512), + ) + + if torch.cuda.is_available(): + predictor.model = predictor.model.cuda() + + # Setting the random seeds + torch.manual_seed(cfg.seed) + np.random.seed(cfg.seed) + + # Constructing the specified dataset + curr_collate_fn = collate_fn + if "tapvid" in cfg.dataset_name: + from cotracker.datasets.tap_vid_datasets import TapVidDataset + + dataset_type = cfg.dataset_name.split("_")[1] + if dataset_type == "davis": + data_root = os.path.join( + cfg.dataset_root, "tapvid_davis", "tapvid_davis.pkl" + ) + elif dataset_type == "kinetics": + data_root = os.path.join(cfg.dataset_root, "tapvid_kinetics") + elif dataset_type == "robotap": + data_root = os.path.join(cfg.dataset_root, "tapvid_robotap") + elif dataset_type == "stacking": + data_root = os.path.join( + cfg.dataset_root, "tapvid_rgb_stacking", "tapvid_rgb_stacking.pkl" + ) + + test_dataset = TapVidDataset( + dataset_type=dataset_type, + data_root=data_root, + queried_first=not "strided" in cfg.dataset_name, + # resize_to=None, + ) + elif cfg.dataset_name == "dynamic_replica": + from cotracker.datasets.dr_dataset import DynamicReplicaDataset + + test_dataset = DynamicReplicaDataset( + cfg.dataset_root, sample_len=300, only_first_n_samples=1 + ) + + # Creating the DataLoader object + test_dataloader = torch.utils.data.DataLoader( + test_dataset, + batch_size=1, + shuffle=False, + num_workers=1, + collate_fn=curr_collate_fn, + ) + + # Timing and conducting the evaluation + import time + + start = time.time() + evaluate_result = evaluator.evaluate_sequence( + predictor, test_dataloader, dataset_name=cfg.dataset_name + ) + end = time.time() + print(end - start) + + # Saving the evaluation results to a .json file + evaluate_result = evaluate_result["avg"] + print("evaluate_result", evaluate_result) + result_file = os.path.join(cfg.exp_dir, f"result_eval_.json") + evaluate_result["time"] = end - start + print(f"Dumping eval results to {result_file}.") + with open(result_file, "w") as f: + json.dump(evaluate_result, f) + + +cs = hydra.core.config_store.ConfigStore.instance() +cs.store(name="default_config_eval", node=DefaultConfig) + + +@hydra.main(config_path="./configs/", config_name="default_config_eval") +def evaluate(cfg: DefaultConfig) -> None: + os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" + os.environ["CUDA_VISIBLE_DEVICES"] = str(cfg.gpu_idx) + run_eval(cfg) + + +if __name__ == "__main__": + evaluate() diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__pycache__/__init__.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..d6503fb165a993c64297f177ca0d103a2334ac9a Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__pycache__/__init__.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__pycache__/build_cotracker.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__pycache__/build_cotracker.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..954af9168faa2bc084b114473c1488c777ee74ec Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/__pycache__/build_cotracker.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/bootstap_predictor.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/bootstap_predictor.py new file mode 100644 index 0000000000000000000000000000000000000000..6a16a4952ef2bc886f696f64f811cddfc4248ef6 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/bootstap_predictor.py @@ -0,0 +1,65 @@ +import torch +import torch.nn.functional as F + +import sys + +import matplotlib.pyplot as plt +import mediapy as media +import numpy as np +from tapnet.torch.tapir_model import TAPIR + + +def postprocess_occlusions(occlusions, expected_dist): + visibles = (1 - F.sigmoid(occlusions)) * (1 - F.sigmoid(expected_dist)) > 0.5 + return visibles + + +class TAPIRPredictor(torch.nn.Module): + def __init__(self, bootstap=False, model=None): + super().__init__() + self.interp_shape = (256, 256) + if model is None: + if bootstap: + checkpoint = "./tapnet/bootstapir_checkpoint.pt" + model = TAPIR(pyramid_level=1, extra_convs=True) + else: + checkpoint = "./tapnet/tapir_checkpoint_panning.pt" + model = TAPIR(pyramid_level=0, extra_convs=False) + model.load_state_dict(torch.load(checkpoint)) + self.model = model.eval().to("cuda") + + def forward(self, rgbs, queries=None, grid_size=0, iters=6, eval_depth=False): + B, T, C, H, W = rgbs.shape + rgbs_ = rgbs.reshape(B * T, C, H, W) + rgbs_ = F.interpolate(rgbs_, tuple(self.interp_shape), mode="bilinear") + rgbs_ = rgbs_.reshape(B, T, 3, self.interp_shape[0], self.interp_shape[1]) + rgbs_ = rgbs_[0].permute(0, 2, 3, 1) + rgbs_ = (rgbs_ / 255.0) * 2 - 1 + + if queries is not None: + queries = queries.clone().float() + B, N, D = queries.shape + assert D == 3 + assert B == 1 + queries[:, :, 1] *= self.interp_shape[1] / W + queries[:, :, 2] *= self.interp_shape[0] / H + queries = torch.stack( + [queries[..., 0], queries[..., 2], queries[..., 1]], dim=-1 + ) + + outputs = self.model(video=rgbs_[None], query_points=queries) + tracks, occlusions, expected_dist = ( + outputs["tracks"], + outputs["occlusion"][0], + outputs["expected_dist"][0], + ) + visibility = postprocess_occlusions(occlusions, expected_dist)[None].permute( + 0, 2, 1 + ) + + tracks = tracks.permute(0, 2, 1, 3) + + tracks[:, :, :, 0] *= W / float(self.interp_shape[1]) + tracks[:, :, :, 1] *= H / float(self.interp_shape[0]) + + return tracks, visibility diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/build_cotracker.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/build_cotracker.py new file mode 100644 index 0000000000000000000000000000000000000000..8710a836e27fb250552fa0873abc778e819751e6 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/build_cotracker.py @@ -0,0 +1,45 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch + +from cotracker.models.core.cotracker.cotracker import CoTracker2 +from cotracker.models.core.cotracker.cotracker3_offline import CoTrackerThreeOffline +from cotracker.models.core.cotracker.cotracker3_online import CoTrackerThreeOnline + + +def build_cotracker( + checkpoint: str, +): + if checkpoint is None: + return build_cotracker() + model_name = checkpoint.split("/")[-1].split(".")[0] + if model_name == "cotracker": + return build_cotracker(checkpoint=checkpoint) + else: + raise ValueError(f"Unknown model name {model_name}") + + +def build_cotracker(checkpoint=None, offline=True, window_len=16, v2=False): + if v2: + cotracker = CoTracker2(stride=4, window_len=window_len) + else: + if offline: + cotracker = CoTrackerThreeOffline( + stride=4, corr_radius=3, window_len=window_len + ) + else: + cotracker = CoTrackerThreeOnline( + stride=4, corr_radius=3, window_len=window_len + ) + + if checkpoint is not None: + with open(checkpoint, "rb") as f: + state_dict = torch.load(f, map_location="cpu") + if "model" in state_dict: + state_dict = state_dict["model"] + cotracker.load_state_dict(state_dict) + return cotracker diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/__init__.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..4d143798f287d4bad3a93eed1e237bcafb0f0560 Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/__init__.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/embeddings.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/embeddings.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..f723d2a2b209e3da1590e4f1fcc8f78cfb1ff99b Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/embeddings.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/model_utils.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/model_utils.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..c3ae0e24448c53d2d4fcdd4b6a1133e999c5d276 Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/__pycache__/model_utils.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/__init__.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..480ca4e7b88e931cfe4548923186a45edc69ebce Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/__init__.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/blocks.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/blocks.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..325a4ae445438618ac0909eb70178bc1fa211dfc Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/blocks.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..fd67a133d9e9580c3b3375f8f02558f72f316acc Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker3_offline.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker3_offline.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..785e988a3e7415f532c34708dfe5955b557bdc83 Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker3_offline.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker3_online.cpython-310.pyc b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker3_online.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..dc5ea3ee664353e9f98cb360f2bba93e7bc91c60 Binary files /dev/null and b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/__pycache__/cotracker3_online.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/blocks.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/blocks.py new file mode 100644 index 0000000000000000000000000000000000000000..12dd7992e8ff80a81a58d6a1b3bb25edcd975669 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/blocks.py @@ -0,0 +1,438 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch +import torch.nn as nn +import torch.nn.functional as F +from functools import partial +from typing import Callable +import collections +from torch import Tensor +from itertools import repeat + +from cotracker.models.core.model_utils import bilinear_sampler + + +# From PyTorch internals +def _ntuple(n): + def parse(x): + if isinstance(x, collections.abc.Iterable) and not isinstance(x, str): + return tuple(x) + return tuple(repeat(x, n)) + + return parse + + +def exists(val): + return val is not None + + +def default(val, d): + return val if exists(val) else d + + +to_2tuple = _ntuple(2) + + +class Mlp(nn.Module): + """MLP as used in Vision Transformer, MLP-Mixer and related networks""" + + def __init__( + self, + in_features, + hidden_features=None, + out_features=None, + act_layer=nn.GELU, + norm_layer=None, + bias=True, + drop=0.0, + use_conv=False, + ): + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + bias = to_2tuple(bias) + drop_probs = to_2tuple(drop) + linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear + + self.fc1 = linear_layer(in_features, hidden_features, bias=bias[0]) + self.act = act_layer() + self.drop1 = nn.Dropout(drop_probs[0]) + self.norm = ( + norm_layer(hidden_features) if norm_layer is not None else nn.Identity() + ) + self.fc2 = linear_layer(hidden_features, out_features, bias=bias[1]) + self.drop2 = nn.Dropout(drop_probs[1]) + + def forward(self, x): + x = self.fc1(x) + x = self.act(x) + x = self.drop1(x) + x = self.fc2(x) + x = self.drop2(x) + return x + + +class ResidualBlock(nn.Module): + def __init__(self, in_planes, planes, norm_fn="group", stride=1): + super(ResidualBlock, self).__init__() + + self.conv1 = nn.Conv2d( + in_planes, + planes, + kernel_size=3, + padding=1, + stride=stride, + padding_mode="zeros", + ) + self.conv2 = nn.Conv2d( + planes, planes, kernel_size=3, padding=1, padding_mode="zeros" + ) + self.relu = nn.ReLU(inplace=True) + + num_groups = planes // 8 + + if norm_fn == "group": + self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes) + self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes) + if not stride == 1: + self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes) + + elif norm_fn == "batch": + self.norm1 = nn.BatchNorm2d(planes) + self.norm2 = nn.BatchNorm2d(planes) + if not stride == 1: + self.norm3 = nn.BatchNorm2d(planes) + + elif norm_fn == "instance": + self.norm1 = nn.InstanceNorm2d(planes) + self.norm2 = nn.InstanceNorm2d(planes) + if not stride == 1: + self.norm3 = nn.InstanceNorm2d(planes) + + elif norm_fn == "none": + self.norm1 = nn.Sequential() + self.norm2 = nn.Sequential() + if not stride == 1: + self.norm3 = nn.Sequential() + + if stride == 1: + self.downsample = None + + else: + self.downsample = nn.Sequential( + nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm3 + ) + + def forward(self, x): + y = x + y = self.relu(self.norm1(self.conv1(y))) + y = self.relu(self.norm2(self.conv2(y))) + + if self.downsample is not None: + x = self.downsample(x) + + return self.relu(x + y) + + +class BasicEncoder(nn.Module): + def __init__(self, input_dim=3, output_dim=128, stride=4): + super(BasicEncoder, self).__init__() + self.stride = stride + self.norm_fn = "instance" + self.in_planes = output_dim // 2 + self.norm1 = nn.InstanceNorm2d(self.in_planes) + self.norm2 = nn.InstanceNorm2d(output_dim * 2) + + self.conv1 = nn.Conv2d( + input_dim, + self.in_planes, + kernel_size=7, + stride=2, + padding=3, + padding_mode="zeros", + ) + self.relu1 = nn.ReLU(inplace=True) + self.layer1 = self._make_layer(output_dim // 2, stride=1) + self.layer2 = self._make_layer(output_dim // 4 * 3, stride=2) + self.layer3 = self._make_layer(output_dim, stride=2) + self.layer4 = self._make_layer(output_dim, stride=2) + + self.conv2 = nn.Conv2d( + output_dim * 3 + output_dim // 4, + output_dim * 2, + kernel_size=3, + padding=1, + padding_mode="zeros", + ) + self.relu2 = nn.ReLU(inplace=True) + self.conv3 = nn.Conv2d(output_dim * 2, output_dim, kernel_size=1) + for m in self.modules(): + if isinstance(m, nn.Conv2d): + nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu") + elif isinstance(m, (nn.InstanceNorm2d)): + if m.weight is not None: + nn.init.constant_(m.weight, 1) + if m.bias is not None: + nn.init.constant_(m.bias, 0) + + def _make_layer(self, dim, stride=1): + layer1 = ResidualBlock(self.in_planes, dim, self.norm_fn, stride=stride) + layer2 = ResidualBlock(dim, dim, self.norm_fn, stride=1) + layers = (layer1, layer2) + + self.in_planes = dim + return nn.Sequential(*layers) + + def forward(self, x): + _, _, H, W = x.shape + + x = self.conv1(x) + x = self.norm1(x) + x = self.relu1(x) + + a = self.layer1(x) + b = self.layer2(a) + c = self.layer3(b) + d = self.layer4(c) + + def _bilinear_intepolate(x): + return F.interpolate( + x, + (H // self.stride, W // self.stride), + mode="bilinear", + align_corners=True, + ) + + a = _bilinear_intepolate(a) + b = _bilinear_intepolate(b) + c = _bilinear_intepolate(c) + d = _bilinear_intepolate(d) + + x = self.conv2(torch.cat([a, b, c, d], dim=1)) + x = self.norm2(x) + x = self.relu2(x) + x = self.conv3(x) + return x + + +class EfficientCorrBlock: + def __init__( + self, + fmaps, + num_levels=4, + radius=4, + padding_mode="zeros", + ): + B, S, C, H, W = fmaps.shape + self.padding_mode = padding_mode + self.num_levels = num_levels + self.radius = radius + self.fmaps_pyramid = [] + + self.fmaps_pyramid.append(fmaps) + for i in range(self.num_levels - 1): + fmaps_ = fmaps.reshape(B * S, C, H, W) + fmaps_ = F.avg_pool2d(fmaps_, 2, stride=2) + _, _, H, W = fmaps_.shape + fmaps = fmaps_.reshape(B, S, C, H, W) + self.fmaps_pyramid.append(fmaps) + + def sample(self, coords, target): + r = self.radius + device = coords.device + B, S, N, D = coords.shape + assert D == 2 + + target = target.permute(0, 1, 3, 2).unsqueeze(-1) + + out_pyramid = [] + for i in range(self.num_levels): + pyramid = self.fmaps_pyramid[i] + C, H, W = pyramid.shape[2:] + centroid_lvl = ( + torch.cat( + [torch.zeros_like(coords[..., :1], device=device), coords], dim=-1 + ).reshape(B * S, N, 1, 1, 3) + / 2**i + ) + + dx = torch.linspace(-r, r, 2 * r + 1, device=device) + dy = torch.linspace(-r, r, 2 * r + 1, device=device) + + xgrid, ygrid = torch.meshgrid(dy, dx, indexing="ij") + zgrid = torch.zeros_like(xgrid, device=device) + delta = torch.stack([zgrid, xgrid, ygrid], axis=-1) + delta_lvl = delta.view(1, 1, 2 * r + 1, 2 * r + 1, 3) + coords_lvl = centroid_lvl + delta_lvl + pyramid_sample = bilinear_sampler( + pyramid.reshape(B * S, C, 1, H, W), coords_lvl + ) + + corr = torch.sum(target * pyramid_sample.reshape(B, S, C, N, -1), dim=2) + corr = corr / torch.sqrt(torch.tensor(C).float()) + out_pyramid.append(corr) + + out = torch.cat(out_pyramid, dim=-1) # B, S, N, LRR*2 + out = out.permute(0, 2, 1, 3).contiguous().view(B * N, S, -1).float() + return out + + +class CorrBlock: + def __init__( + self, + fmaps, + num_levels=4, + radius=4, + multiple_track_feats=False, + padding_mode="zeros", + ): + B, S, C, H, W = fmaps.shape + self.S, self.C, self.H, self.W = S, C, H, W + self.padding_mode = padding_mode + self.num_levels = num_levels + self.radius = radius + self.fmaps_pyramid = [] + self.multiple_track_feats = multiple_track_feats + + self.fmaps_pyramid.append(fmaps) + for i in range(self.num_levels - 1): + fmaps_ = fmaps.reshape(B * S, C, H, W) + fmaps_ = F.avg_pool2d(fmaps_, 2, stride=2) + _, _, H, W = fmaps_.shape + fmaps = fmaps_.reshape(B, S, C, H, W) + self.fmaps_pyramid.append(fmaps) + + def sample(self, coords): + r = self.radius + B, S, N, D = coords.shape + assert D == 2 + + H, W = self.H, self.W + out_pyramid = [] + for i in range(self.num_levels): + corrs = self.corrs_pyramid[i] # B, S, N, H, W + *_, H, W = corrs.shape + + dx = torch.linspace(-r, r, 2 * r + 1) + dy = torch.linspace(-r, r, 2 * r + 1) + delta = torch.stack(torch.meshgrid(dy, dx, indexing="ij"), axis=-1).to( + coords.device + ) + + centroid_lvl = coords.reshape(B * S * N, 1, 1, 2) / 2**i + delta_lvl = delta.view(1, 2 * r + 1, 2 * r + 1, 2) + coords_lvl = centroid_lvl + delta_lvl + + corrs = bilinear_sampler( + corrs.reshape(B * S * N, 1, H, W), + coords_lvl, + padding_mode=self.padding_mode, + ) + corrs = corrs.view(B, S, N, -1) + out_pyramid.append(corrs) + + out = torch.cat(out_pyramid, dim=-1) # B, S, N, LRR*2 + out = out.permute(0, 2, 1, 3).contiguous().view(B * N, S, -1).float() + return out + + def corr(self, targets): + B, S, N, C = targets.shape + if self.multiple_track_feats: + targets_split = targets.split(C // self.num_levels, dim=-1) + B, S, N, C = targets_split[0].shape + + assert C == self.C + assert S == self.S + + fmap1 = targets + + self.corrs_pyramid = [] + for i, fmaps in enumerate(self.fmaps_pyramid): + *_, H, W = fmaps.shape + fmap2s = fmaps.view(B, S, C, H * W) # B S C H W -> B S C (H W) + if self.multiple_track_feats: + fmap1 = targets_split[i] + corrs = torch.matmul(fmap1, fmap2s) + corrs = corrs.view(B, S, N, H, W) # B S N (H W) -> B S N H W + corrs = corrs / torch.sqrt(torch.tensor(C).float()) + self.corrs_pyramid.append(corrs) + + +class Attention(nn.Module): + def __init__( + self, query_dim, context_dim=None, num_heads=8, dim_head=48, qkv_bias=False + ): + super().__init__() + inner_dim = dim_head * num_heads + context_dim = default(context_dim, query_dim) + self.scale = dim_head**-0.5 + self.heads = num_heads + + self.to_q = nn.Linear(query_dim, inner_dim, bias=qkv_bias) + self.to_kv = nn.Linear(context_dim, inner_dim * 2, bias=qkv_bias) + self.to_out = nn.Linear(inner_dim, query_dim) + + def forward(self, x, context=None, attn_bias=None): + B, N1, C = x.shape + h = self.heads + + q = self.to_q(x).reshape(B, N1, h, C // h).permute(0, 2, 1, 3) + context = default(context, x) + k, v = self.to_kv(context).chunk(2, dim=-1) + + N2 = context.shape[1] + k = k.reshape(B, N2, h, C // h).permute(0, 2, 1, 3) + v = v.reshape(B, N2, h, C // h).permute(0, 2, 1, 3) + + sim = (q @ k.transpose(-2, -1)) * self.scale + + if attn_bias is not None: + sim = sim + attn_bias + attn = sim.softmax(dim=-1) + + x = (attn @ v).transpose(1, 2).reshape(B, N1, C) + return self.to_out(x) + + +class AttnBlock(nn.Module): + def __init__( + self, + hidden_size, + num_heads, + attn_class: Callable[..., nn.Module] = Attention, + mlp_ratio=4.0, + **block_kwargs + ): + super().__init__() + self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.attn = attn_class( + hidden_size, num_heads=num_heads, qkv_bias=True, **block_kwargs + ) + + self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + mlp_hidden_dim = int(hidden_size * mlp_ratio) + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = Mlp( + in_features=hidden_size, + hidden_features=mlp_hidden_dim, + act_layer=approx_gelu, + drop=0, + ) + + def forward(self, x, mask=None): + attn_bias = mask + if mask is not None: + mask = ( + (mask[:, None] * mask[:, :, None]) + .unsqueeze(1) + .expand(-1, self.attn.num_heads, -1, -1) + ) + max_neg_value = -torch.finfo(x.dtype).max + attn_bias = (~mask) * max_neg_value + x = x + self.attn(self.norm1(x), attn_bias=attn_bias) + x = x + self.mlp(self.norm2(x)) + return x diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker.py new file mode 100644 index 0000000000000000000000000000000000000000..5f23f9c856c05f6060312c25054b90c317bfdba7 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker.py @@ -0,0 +1,577 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from cotracker.models.core.model_utils import sample_features4d, sample_features5d +from cotracker.models.core.embeddings import ( + get_2d_embedding, + get_1d_sincos_pos_embed_from_grid, + get_2d_sincos_pos_embed, +) + +from cotracker.models.core.cotracker.blocks import ( + Mlp, + BasicEncoder, + AttnBlock, + CorrBlock, + Attention, +) + +torch.manual_seed(0) + + +class CoTracker2(nn.Module): + def __init__( + self, + window_len=8, + stride=4, + add_space_attn=True, + num_virtual_tracks=64, + model_resolution=(384, 512), + ): + super(CoTracker2, self).__init__() + self.window_len = window_len + self.stride = stride + self.hidden_dim = 256 + self.latent_dim = 128 + self.add_space_attn = add_space_attn + self.fnet = BasicEncoder(output_dim=self.latent_dim) + self.num_virtual_tracks = num_virtual_tracks + self.model_resolution = model_resolution + self.input_dim = 456 + self.updateformer = EfficientUpdateFormer( + space_depth=6, + time_depth=6, + input_dim=self.input_dim, + hidden_size=384, + output_dim=self.latent_dim + 2, + mlp_ratio=4.0, + add_space_attn=add_space_attn, + num_virtual_tracks=num_virtual_tracks, + ) + + time_grid = torch.linspace(0, window_len - 1, window_len).reshape( + 1, window_len, 1 + ) + + self.register_buffer( + "time_emb", get_1d_sincos_pos_embed_from_grid(self.input_dim, time_grid[0]) + ) + + self.register_buffer( + "pos_emb", + get_2d_sincos_pos_embed( + embed_dim=self.input_dim, + grid_size=( + model_resolution[0] // stride, + model_resolution[1] // stride, + ), + ), + ) + self.norm = nn.GroupNorm(1, self.latent_dim) + self.track_feat_updater = nn.Sequential( + nn.Linear(self.latent_dim, self.latent_dim), + nn.GELU(), + ) + self.vis_predictor = nn.Sequential( + nn.Linear(self.latent_dim, 1), + ) + + def forward_window( + self, + fmaps, + coords, + track_feat=None, + vis=None, + track_mask=None, + attention_mask=None, + iters=4, + ): + # B = batch size + # S = number of frames in the window) + # N = number of tracks + # C = channels of a point feature vector + # E = positional embedding size + # LRR = local receptive field radius + # D = dimension of the transformer input tokens + + # track_feat = B S N C + # vis = B S N 1 + # track_mask = B S N 1 + # attention_mask = B S N + + B, S_init, N, __ = track_mask.shape + B, S, *_ = fmaps.shape + + track_mask = F.pad(track_mask, (0, 0, 0, 0, 0, S - S_init), "constant") + track_mask_vis = ( + torch.cat([track_mask, vis], dim=-1) + .permute(0, 2, 1, 3) + .reshape(B * N, S, 2) + ) + + corr_block = CorrBlock( + fmaps, + num_levels=4, + radius=3, + padding_mode="border", + ) + + sampled_pos_emb = ( + sample_features4d(self.pos_emb.repeat(B, 1, 1, 1), coords[:, 0]) + .reshape(B * N, self.input_dim) + .unsqueeze(1) + ) # B E N -> (B N) 1 E + + coord_preds = [] + for __ in range(iters): + coords = coords.detach() # B S N 2 + corr_block.corr(track_feat) + + # Sample correlation features around each point + fcorrs = corr_block.sample(coords) # (B N) S LRR + + # Get the flow embeddings + flows = (coords - coords[:, 0:1]).permute(0, 2, 1, 3).reshape(B * N, S, 2) + flow_emb = get_2d_embedding(flows, 64, cat_coords=True) # N S E + + track_feat_ = track_feat.permute(0, 2, 1, 3).reshape( + B * N, S, self.latent_dim + ) + + transformer_input = torch.cat( + [flow_emb, fcorrs, track_feat_, track_mask_vis], dim=2 + ) + x = transformer_input + sampled_pos_emb + self.time_emb + x = x.view(B, N, S, -1) # (B N) S D -> B N S D + + delta = self.updateformer( + x, + attention_mask.reshape(B * S, N), # B S N -> (B S) N + ) + + delta_coords = delta[..., :2].permute(0, 2, 1, 3) + coords = coords + delta_coords + coord_preds.append(coords * self.stride) + + delta_feats_ = delta[..., 2:].reshape(B * N * S, self.latent_dim) + track_feat_ = track_feat.permute(0, 2, 1, 3).reshape( + B * N * S, self.latent_dim + ) + track_feat_ = self.track_feat_updater(self.norm(delta_feats_)) + track_feat_ + track_feat = track_feat_.reshape(B, N, S, self.latent_dim).permute( + 0, 2, 1, 3 + ) # (B N S) C -> B S N C + + vis_pred = self.vis_predictor(track_feat).reshape(B, S, N) + return coord_preds, vis_pred + + def get_track_feat(self, fmaps, queried_frames, queried_coords): + sample_frames = queried_frames[:, None, :, None] + sample_coords = torch.cat( + [ + sample_frames, + queried_coords[:, None], + ], + dim=-1, + ) + sample_track_feats = sample_features5d(fmaps, sample_coords) + return sample_track_feats + + def init_video_online_processing(self): + self.online_ind = 0 + self.online_track_feat = None + self.online_coords_predicted = None + self.online_vis_predicted = None + + def forward(self, video, queries, iters=4, is_train=False, is_online=False): + """Predict tracks + + Args: + video (FloatTensor[B, T, 3]): input videos. + queries (FloatTensor[B, N, 3]): point queries. + iters (int, optional): number of updates. Defaults to 4. + is_train (bool, optional): enables training mode. Defaults to False. + is_online (bool, optional): enables online mode. Defaults to False. Before enabling, call model.init_video_online_processing(). + + Returns: + - coords_predicted (FloatTensor[B, T, N, 2]): + - vis_predicted (FloatTensor[B, T, N]): + - train_data: `None` if `is_train` is false, otherwise: + - all_vis_predictions (List[FloatTensor[B, S, N, 1]]): + - all_coords_predictions (List[FloatTensor[B, S, N, 2]]): + - mask (BoolTensor[B, T, N]): + """ + B, T, C, H, W = video.shape + B, N, __ = queries.shape + S = self.window_len + device = queries.device + + # B = batch size + # S = number of frames in the window of the padded video + # S_trimmed = actual number of frames in the window + # N = number of tracks + # C = color channels (3 for RGB) + # E = positional embedding size + # LRR = local receptive field radius + # D = dimension of the transformer input tokens + + # video = B T C H W + # queries = B N 3 + # coords_init = B S N 2 + # vis_init = B S N 1 + + assert S >= 2 # A tracker needs at least two frames to track something + if is_online: + assert T <= S, "Online mode: video chunk must be <= window size." + assert ( + self.online_ind is not None + ), "Call model.init_video_online_processing() first." + assert not is_train, "Training not supported in online mode." + step = S // 2 # How much the sliding window moves at every step + video = 2 * (video / 255.0) - 1.0 + + # The first channel is the frame number + # The rest are the coordinates of points we want to track + queried_frames = queries[:, :, 0].long() + + queried_coords = queries[..., 1:] + queried_coords = queried_coords / self.stride + + # We store our predictions here + coords_predicted = torch.zeros((B, T, N, 2), device=device) + vis_predicted = torch.zeros((B, T, N), device=device) + if is_online: + if self.online_coords_predicted is None: + # Init online predictions with zeros + self.online_coords_predicted = coords_predicted + self.online_vis_predicted = vis_predicted + else: + # Pad online predictions with zeros for the current window + pad = min(step, T - step) + coords_predicted = F.pad( + self.online_coords_predicted, (0, 0, 0, 0, 0, pad), "constant" + ) + vis_predicted = F.pad( + self.online_vis_predicted, (0, 0, 0, pad), "constant" + ) + all_coords_predictions, all_vis_predictions = [], [] + + # Pad the video so that an integer number of sliding windows fit into it + # TODO: we may drop this requirement because the transformer should not care + # TODO: pad the features instead of the video + pad = ( + S - T if is_online else (S - T % S) % S + ) # We don't want to pad if T % S == 0 + video = video.reshape(B, 1, T, C * H * W) + video_pad = video[:, :, -1:].repeat(1, 1, pad, 1) + video = torch.cat([video, video_pad], dim=2) + + # Compute convolutional features for the video or for the current chunk in case of online mode + fmaps = self.fnet(video.reshape(-1, C, H, W)).reshape( + B, -1, self.latent_dim, H // self.stride, W // self.stride + ) + + # We compute track features + track_feat = self.get_track_feat( + fmaps, + queried_frames - self.online_ind if is_online else queried_frames, + queried_coords, + ).repeat(1, S, 1, 1) + if is_online: + # We update track features for the current window + sample_frames = queried_frames[:, None, :, None] # B 1 N 1 + left = 0 if self.online_ind == 0 else self.online_ind + step + right = self.online_ind + S + sample_mask = (sample_frames >= left) & (sample_frames < right) + if self.online_track_feat is None: + self.online_track_feat = torch.zeros_like(track_feat, device=device) + self.online_track_feat += track_feat * sample_mask + track_feat = self.online_track_feat.clone() + # We process ((num_windows - 1) * step + S) frames in total, so there are + # (ceil((T - S) / step) + 1) windows + num_windows = (T - S + step - 1) // step + 1 + # We process only the current video chunk in the online mode + indices = [self.online_ind] if is_online else range(0, step * num_windows, step) + + coords_init = queried_coords.reshape(B, 1, N, 2).expand(B, S, N, 2).float() + vis_init = torch.ones((B, S, N, 1), device=device).float() * 10 + for ind in indices: + # We copy over coords and vis for tracks that are queried + # by the end of the previous window, which is ind + overlap + if ind > 0: + overlap = S - step + copy_over = (queried_frames < ind + overlap)[ + :, None, :, None + ] # B 1 N 1 + coords_prev = torch.nn.functional.pad( + coords_predicted[:, ind : ind + overlap] / self.stride, + (0, 0, 0, 0, 0, step), + "replicate", + ) # B S N 2 + vis_prev = torch.nn.functional.pad( + vis_predicted[:, ind : ind + overlap, :, None].clone(), + (0, 0, 0, 0, 0, step), + "replicate", + ) # B S N 1 + coords_init = torch.where( + copy_over.expand_as(coords_init), coords_prev, coords_init + ) + vis_init = torch.where( + copy_over.expand_as(vis_init), vis_prev, vis_init + ) + + # The attention mask is 1 for the spatio-temporal points within + # a track which is updated in the current window + attention_mask = ( + (queried_frames < ind + S).reshape(B, 1, N).repeat(1, S, 1) + ) # B S N + + # The track mask is 1 for the spatio-temporal points that actually + # need updating: only after begin queried, and not if contained + # in a previous window + track_mask = ( + queried_frames[:, None, :, None] + <= torch.arange(ind, ind + S, device=device)[None, :, None, None] + ).contiguous() # B S N 1 + + if ind > 0: + track_mask[:, :overlap, :, :] = False + + # Predict the coordinates and visibility for the current window + coords, vis = self.forward_window( + fmaps=fmaps if is_online else fmaps[:, ind : ind + S], + coords=coords_init, + track_feat=attention_mask.unsqueeze(-1) * track_feat, + vis=vis_init, + track_mask=track_mask, + attention_mask=attention_mask, + iters=iters, + ) + + S_trimmed = ( + T if is_online else min(T - ind, S) + ) # accounts for last window duration + coords_predicted[:, ind : ind + S] = coords[-1][:, :S_trimmed] + vis_predicted[:, ind : ind + S] = vis[:, :S_trimmed] + if is_train: + all_coords_predictions.append( + [coord[:, :S_trimmed] for coord in coords] + ) + all_vis_predictions.append(torch.sigmoid(vis[:, :S_trimmed])) + + if is_online: + self.online_ind += step + self.online_coords_predicted = coords_predicted + self.online_vis_predicted = vis_predicted + vis_predicted = torch.sigmoid(vis_predicted) + + if is_train: + mask = ( + queried_frames[:, None] + <= torch.arange(0, T, device=device)[None, :, None] + ) + train_data = (all_coords_predictions, all_vis_predictions, mask) + else: + train_data = None + + return coords_predicted, vis_predicted, train_data + + +class EfficientUpdateFormer(nn.Module): + """ + Transformer model that updates track estimates. + """ + + def __init__( + self, + space_depth=6, + time_depth=6, + input_dim=320, + hidden_size=384, + num_heads=8, + output_dim=130, + mlp_ratio=4.0, + num_virtual_tracks=64, + add_space_attn=True, + linear_layer_for_vis_conf=False, + ): + super().__init__() + self.out_channels = 2 + self.num_heads = num_heads + self.hidden_size = hidden_size + self.input_transform = torch.nn.Linear(input_dim, hidden_size, bias=True) + if linear_layer_for_vis_conf: + self.flow_head = torch.nn.Linear(hidden_size, output_dim - 2, bias=True) + self.vis_conf_head = torch.nn.Linear(hidden_size, 2, bias=True) + else: + self.flow_head = torch.nn.Linear(hidden_size, output_dim, bias=True) + self.num_virtual_tracks = num_virtual_tracks + self.virual_tracks = nn.Parameter( + torch.randn(1, num_virtual_tracks, 1, hidden_size) + ) + self.add_space_attn = add_space_attn + self.linear_layer_for_vis_conf = linear_layer_for_vis_conf + self.time_blocks = nn.ModuleList( + [ + AttnBlock( + hidden_size, + num_heads, + mlp_ratio=mlp_ratio, + attn_class=Attention, + ) + for _ in range(time_depth) + ] + ) + + if add_space_attn: + self.space_virtual_blocks = nn.ModuleList( + [ + AttnBlock( + hidden_size, + num_heads, + mlp_ratio=mlp_ratio, + attn_class=Attention, + ) + for _ in range(space_depth) + ] + ) + self.space_point2virtual_blocks = nn.ModuleList( + [ + CrossAttnBlock( + hidden_size, hidden_size, num_heads, mlp_ratio=mlp_ratio + ) + for _ in range(space_depth) + ] + ) + self.space_virtual2point_blocks = nn.ModuleList( + [ + CrossAttnBlock( + hidden_size, hidden_size, num_heads, mlp_ratio=mlp_ratio + ) + for _ in range(space_depth) + ] + ) + assert len(self.time_blocks) >= len(self.space_virtual2point_blocks) + self.initialize_weights() + + def initialize_weights(self): + def _basic_init(module): + if isinstance(module, nn.Linear): + torch.nn.init.xavier_uniform_(module.weight) + if module.bias is not None: + nn.init.constant_(module.bias, 0) + torch.nn.init.trunc_normal_(self.flow_head.weight, std=0.001) + if self.linear_layer_for_vis_conf: + torch.nn.init.trunc_normal_(self.vis_conf_head.weight, std=0.001) + + def _trunc_init(module): + """ViT weight initialization, original timm impl (for reproducibility)""" + if isinstance(module, nn.Linear): + torch.nn.init.trunc_normal_(module.weight, std=0.02) + if module.bias is not None: + nn.init.zeros_(module.bias) + + self.apply(_basic_init) + + def forward(self, input_tensor, mask=None, add_space_attn=True): + tokens = self.input_transform(input_tensor) + + B, _, T, _ = tokens.shape + virtual_tokens = self.virual_tracks.repeat(B, 1, T, 1) + tokens = torch.cat([tokens, virtual_tokens], dim=1) + + _, N, _, _ = tokens.shape + j = 0 + layers = [] + for i in range(len(self.time_blocks)): + time_tokens = tokens.contiguous().view(B * N, T, -1) # B N T C -> (B N) T C + time_tokens = self.time_blocks[i](time_tokens) + + tokens = time_tokens.view(B, N, T, -1) # (B N) T C -> B N T C + if ( + add_space_attn + and hasattr(self, "space_virtual_blocks") + and (i % (len(self.time_blocks) // len(self.space_virtual_blocks)) == 0) + ): + space_tokens = ( + tokens.permute(0, 2, 1, 3).contiguous().view(B * T, N, -1) + ) # B N T C -> (B T) N C + + point_tokens = space_tokens[:, : N - self.num_virtual_tracks] + virtual_tokens = space_tokens[:, N - self.num_virtual_tracks :] + + virtual_tokens = self.space_virtual2point_blocks[j]( + virtual_tokens, point_tokens, mask=mask + ) + + virtual_tokens = self.space_virtual_blocks[j](virtual_tokens) + point_tokens = self.space_point2virtual_blocks[j]( + point_tokens, virtual_tokens, mask=mask + ) + + space_tokens = torch.cat([point_tokens, virtual_tokens], dim=1) + tokens = space_tokens.view(B, T, N, -1).permute( + 0, 2, 1, 3 + ) # (B T) N C -> B N T C + j += 1 + tokens = tokens[:, : N - self.num_virtual_tracks] + + flow = self.flow_head(tokens) + if self.linear_layer_for_vis_conf: + vis_conf = self.vis_conf_head(tokens) + flow = torch.cat([flow, vis_conf], dim=-1) + + return flow + + +class CrossAttnBlock(nn.Module): + def __init__( + self, hidden_size, context_dim, num_heads=1, mlp_ratio=4.0, **block_kwargs + ): + super().__init__() + self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.norm_context = nn.LayerNorm(hidden_size) + self.cross_attn = Attention( + hidden_size, + context_dim=context_dim, + num_heads=num_heads, + qkv_bias=True, + **block_kwargs + ) + + self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + mlp_hidden_dim = int(hidden_size * mlp_ratio) + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = Mlp( + in_features=hidden_size, + hidden_features=mlp_hidden_dim, + act_layer=approx_gelu, + drop=0, + ) + + def forward(self, x, context, mask=None): + attn_bias = None + if mask is not None: + if mask.shape[1] == x.shape[1]: + mask = mask[:, None, :, None].expand( + -1, self.cross_attn.heads, -1, context.shape[1] + ) + else: + mask = mask[:, None, None].expand( + -1, self.cross_attn.heads, x.shape[1], -1 + ) + + max_neg_value = -torch.finfo(x.dtype).max + attn_bias = (~mask) * max_neg_value + x = x + self.cross_attn( + self.norm1(x), context=self.norm_context(context), attn_bias=attn_bias + ) + x = x + self.mlp(self.norm2(x)) + return x diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker3_offline.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker3_offline.py new file mode 100644 index 0000000000000000000000000000000000000000..01990373612a7440ac6147acf5efd7f24e4dd5e2 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker3_offline.py @@ -0,0 +1,233 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch +import torch.nn as nn +import torch.nn.functional as F +from cotracker.models.core.cotracker.cotracker3_online import CoTrackerThreeBase, posenc + +torch.manual_seed(0) + + +class CoTrackerThreeOffline(CoTrackerThreeBase): + def __init__(self, **args): + super(CoTrackerThreeOffline, self).__init__(**args) + + def forward( + self, + video, + queries, + iters=4, + is_train=False, + add_space_attn=True, + fmaps_chunk_size=200, + ): + """Predict tracks + + Args: + video (FloatTensor[B, T, 3]): input videos. + queries (FloatTensor[B, N, 3]): point queries. + iters (int, optional): number of updates. Defaults to 4. + is_train (bool, optional): enables training mode. Defaults to False. + Returns: + - coords_predicted (FloatTensor[B, T, N, 2]): + - vis_predicted (FloatTensor[B, T, N]): + - train_data: `None` if `is_train` is false, otherwise: + - all_vis_predictions (List[FloatTensor[B, S, N, 1]]): + - all_coords_predictions (List[FloatTensor[B, S, N, 2]]): + - mask (BoolTensor[B, T, N]): + """ + + B, T, C, H, W = video.shape + device = queries.device + assert H % self.stride == 0 and W % self.stride == 0 + + B, N, __ = queries.shape + # B = batch size + # S_trimmed = actual number of frames in the window + # N = number of tracks + # C = color channels (3 for RGB) + # E = positional embedding size + # LRR = local receptive field radius + # D = dimension of the transformer input tokens + + # video = B T C H W + # queries = B N 3 + # coords_init = B T N 2 + # vis_init = B T N 1 + + assert T >= 1 # A tracker needs at least two frames to track something + + video = 2 * (video / 255.0) - 1.0 + dtype = video.dtype + queried_frames = queries[:, :, 0].long() + + queried_coords = queries[..., 1:3] + queried_coords = queried_coords / self.stride + + # We store our predictions here + all_coords_predictions, all_vis_predictions, all_confidence_predictions = ( + [], + [], + [], + ) + C_ = C + H4, W4 = H // self.stride, W // self.stride + # Compute convolutional features for the video or for the current chunk in case of online mode + + if T > fmaps_chunk_size: + fmaps = [] + for t in range(0, T, fmaps_chunk_size): + video_chunk = video[:, t : t + fmaps_chunk_size] + fmaps_chunk = self.fnet(video_chunk.reshape(-1, C_, H, W)) + T_chunk = video_chunk.shape[1] + C_chunk, H_chunk, W_chunk = fmaps_chunk.shape[1:] + fmaps.append(fmaps_chunk.reshape(B, T_chunk, C_chunk, H_chunk, W_chunk)) + fmaps = torch.cat(fmaps, dim=1).reshape(-1, C_chunk, H_chunk, W_chunk) + else: + fmaps = self.fnet(video.reshape(-1, C_, H, W)) + fmaps = fmaps.permute(0, 2, 3, 1) + fmaps = fmaps / torch.sqrt( + torch.maximum( + torch.sum(torch.square(fmaps), axis=-1, keepdims=True), + torch.tensor(1e-12, device=fmaps.device), + ) + ) + fmaps = fmaps.permute(0, 3, 1, 2).reshape( + B, -1, self.latent_dim, H // self.stride, W // self.stride + ) + fmaps = fmaps.to(dtype) + + # We compute track features + fmaps_pyramid = [] + track_feat_pyramid = [] + track_feat_support_pyramid = [] + fmaps_pyramid.append(fmaps) + for i in range(self.corr_levels - 1): + fmaps_ = fmaps.reshape( + B * T, self.latent_dim, fmaps.shape[-2], fmaps.shape[-1] + ) + fmaps_ = F.avg_pool2d(fmaps_, 2, stride=2) + fmaps = fmaps_.reshape( + B, T, self.latent_dim, fmaps_.shape[-2], fmaps_.shape[-1] + ) + fmaps_pyramid.append(fmaps) + + for i in range(self.corr_levels): + track_feat, track_feat_support = self.get_track_feat( + fmaps_pyramid[i], + queried_frames, + queried_coords / 2**i, + support_radius=self.corr_radius, + ) + track_feat_pyramid.append(track_feat.repeat(1, T, 1, 1)) + track_feat_support_pyramid.append(track_feat_support.unsqueeze(1)) + + D_coords = 2 + + coord_preds, vis_preds, confidence_preds = [], [], [] + + vis = torch.zeros((B, T, N), device=device).float() + confidence = torch.zeros((B, T, N), device=device).float() + coords = queried_coords.reshape(B, 1, N, 2).expand(B, T, N, 2).float() + + r = 2 * self.corr_radius + 1 + + for it in range(iters): + coords = coords.detach() # B T N 2 + coords_init = coords.view(B * T, N, 2) + corr_embs = [] + corr_feats = [] + for i in range(self.corr_levels): + corr_feat = self.get_correlation_feat( + fmaps_pyramid[i], coords_init / 2**i + ) + track_feat_support = ( + track_feat_support_pyramid[i] + .view(B, 1, r, r, N, self.latent_dim) + .squeeze(1) + .permute(0, 3, 1, 2, 4) + ) + corr_volume = torch.einsum( + "btnhwc,bnijc->btnhwij", corr_feat, track_feat_support + ) + corr_emb = self.corr_mlp(corr_volume.reshape(B * T * N, r * r * r * r)) + corr_embs.append(corr_emb) + corr_embs = torch.cat(corr_embs, dim=-1) + corr_embs = corr_embs.view(B, T, N, corr_embs.shape[-1]) + + transformer_input = [vis[..., None], confidence[..., None], corr_embs] + + rel_coords_forward = coords[:, :-1] - coords[:, 1:] + rel_coords_backward = coords[:, 1:] - coords[:, :-1] + + rel_coords_forward = torch.nn.functional.pad( + rel_coords_forward, (0, 0, 0, 0, 0, 1) + ) + rel_coords_backward = torch.nn.functional.pad( + rel_coords_backward, (0, 0, 0, 0, 1, 0) + ) + scale = ( + torch.tensor( + [self.model_resolution[1], self.model_resolution[0]], + device=coords.device, + ) + / self.stride + ) + rel_coords_forward = rel_coords_forward / scale + rel_coords_backward = rel_coords_backward / scale + + rel_pos_emb_input = posenc( + torch.cat([rel_coords_forward, rel_coords_backward], dim=-1), + min_deg=0, + max_deg=10, + ) # batch, num_points, num_frames, 84 + transformer_input.append(rel_pos_emb_input) + + x = ( + torch.cat(transformer_input, dim=-1) + .permute(0, 2, 1, 3) + .reshape(B * N, T, -1) + ) + + x = x + self.interpolate_time_embed(x, T) + x = x.view(B, N, T, -1) # (B N) T D -> B N T D + + delta = self.updateformer( + x, + add_space_attn=add_space_attn, + ) + + delta_coords = delta[..., :D_coords].permute(0, 2, 1, 3) + delta_vis = delta[..., D_coords].permute(0, 2, 1) + delta_confidence = delta[..., D_coords + 1].permute(0, 2, 1) + + vis = vis + delta_vis + confidence = confidence + delta_confidence + + coords = coords + delta_coords + coords_append = coords.clone() + coords_append[..., :2] = coords_append[..., :2] * float(self.stride) + coord_preds.append(coords_append) + vis_preds.append(torch.sigmoid(vis)) + confidence_preds.append(torch.sigmoid(confidence)) + + if is_train: + all_coords_predictions.append([coord[..., :2] for coord in coord_preds]) + all_vis_predictions.append(vis_preds) + all_confidence_predictions.append(confidence_preds) + + if is_train: + train_data = ( + all_coords_predictions, + all_vis_predictions, + all_confidence_predictions, + torch.ones_like(vis_preds[-1], device=vis_preds[-1].device), + ) + else: + train_data = None + + return coord_preds[-1][..., :2], vis_preds[-1], confidence_preds[-1], train_data diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker3_online.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker3_online.py new file mode 100644 index 0000000000000000000000000000000000000000..8748a4f7eb59b63e717914cc2d1b54af0ad891ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/cotracker3_online.py @@ -0,0 +1,541 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch +import torch.nn as nn +import torch.nn.functional as F +from cotracker.models.core.model_utils import sample_features5d, bilinear_sampler +from cotracker.models.core.embeddings import get_1d_sincos_pos_embed_from_grid + +from cotracker.models.core.cotracker.blocks import Mlp, BasicEncoder +from cotracker.models.core.cotracker.cotracker import EfficientUpdateFormer + +torch.manual_seed(0) + + +def posenc(x, min_deg, max_deg): + """Cat x with a positional encoding of x with scales 2^[min_deg, max_deg-1]. + Instead of computing [sin(x), cos(x)], we use the trig identity + cos(x) = sin(x + pi/2) and do one vectorized call to sin([x, x+pi/2]). + Args: + x: torch.Tensor, variables to be encoded. Note that x should be in [-pi, pi]. + min_deg: int, the minimum (inclusive) degree of the encoding. + max_deg: int, the maximum (exclusive) degree of the encoding. + legacy_posenc_order: bool, keep the same ordering as the original tf code. + Returns: + encoded: torch.Tensor, encoded variables. + """ + if min_deg == max_deg: + return x + scales = torch.tensor( + [2**i for i in range(min_deg, max_deg)], dtype=x.dtype, device=x.device + ) + + xb = (x[..., None, :] * scales[:, None]).reshape(list(x.shape[:-1]) + [-1]) + four_feat = torch.sin(torch.cat([xb, xb + 0.5 * torch.pi], dim=-1)) + return torch.cat([x] + [four_feat], dim=-1) + + +class CoTrackerThreeBase(nn.Module): + def __init__( + self, + window_len=8, + stride=4, + corr_radius=3, + corr_levels=4, + num_virtual_tracks=64, + model_resolution=(384, 512), + add_space_attn=True, + linear_layer_for_vis_conf=True, + ): + super(CoTrackerThreeBase, self).__init__() + self.window_len = window_len + self.stride = stride + self.corr_radius = corr_radius + self.corr_levels = corr_levels + self.hidden_dim = 256 + self.latent_dim = 128 + + self.linear_layer_for_vis_conf = linear_layer_for_vis_conf + self.fnet = BasicEncoder(input_dim=3, output_dim=self.latent_dim, stride=stride) + + highres_dim = 128 + lowres_dim = 256 + + self.num_virtual_tracks = num_virtual_tracks + self.model_resolution = model_resolution + + self.input_dim = 1110 + + self.updateformer = EfficientUpdateFormer( + space_depth=3, + time_depth=3, + input_dim=self.input_dim, + hidden_size=384, + output_dim=4, + mlp_ratio=4.0, + num_virtual_tracks=num_virtual_tracks, + add_space_attn=add_space_attn, + linear_layer_for_vis_conf=linear_layer_for_vis_conf, + ) + self.corr_mlp = Mlp(in_features=49 * 49, hidden_features=384, out_features=256) + + time_grid = torch.linspace(0, window_len - 1, window_len).reshape( + 1, window_len, 1 + ) + + self.register_buffer( + "time_emb", get_1d_sincos_pos_embed_from_grid(self.input_dim, time_grid[0]) + ) + + def get_support_points(self, coords, r, reshape_back=True): + B, _, N, _ = coords.shape + device = coords.device + centroid_lvl = coords.reshape(B, N, 1, 1, 3) + + dx = torch.linspace(-r, r, 2 * r + 1, device=device) + dy = torch.linspace(-r, r, 2 * r + 1, device=device) + + xgrid, ygrid = torch.meshgrid(dy, dx, indexing="ij") + zgrid = torch.zeros_like(xgrid, device=device) + delta = torch.stack([zgrid, xgrid, ygrid], axis=-1) + delta_lvl = delta.view(1, 1, 2 * r + 1, 2 * r + 1, 3) + coords_lvl = centroid_lvl + delta_lvl + + if reshape_back: + return coords_lvl.reshape(B, N, (2 * r + 1) ** 2, 3).permute(0, 2, 1, 3) + else: + return coords_lvl + + def get_track_feat(self, fmaps, queried_frames, queried_coords, support_radius=0): + + sample_frames = queried_frames[:, None, :, None] + sample_coords = torch.cat( + [ + sample_frames, + queried_coords[:, None], + ], + dim=-1, + ) + support_points = self.get_support_points(sample_coords, support_radius) + support_track_feats = sample_features5d(fmaps, support_points) + return ( + support_track_feats[:, None, support_track_feats.shape[1] // 2], + support_track_feats, + ) + + def get_correlation_feat(self, fmaps, queried_coords): + B, T, D, H_, W_ = fmaps.shape + N = queried_coords.shape[1] + r = self.corr_radius + sample_coords = torch.cat( + [torch.zeros_like(queried_coords[..., :1]), queried_coords], dim=-1 + )[:, None] + support_points = self.get_support_points(sample_coords, r, reshape_back=False) + correlation_feat = bilinear_sampler( + fmaps.reshape(B * T, D, 1, H_, W_), support_points + ) + return correlation_feat.view(B, T, D, N, (2 * r + 1), (2 * r + 1)).permute( + 0, 1, 3, 4, 5, 2 + ) + + def interpolate_time_embed(self, x, t): + previous_dtype = x.dtype + T = self.time_emb.shape[1] + + if t == T: + return self.time_emb + + time_emb = self.time_emb.float() + time_emb = F.interpolate( + time_emb.permute(0, 2, 1), size=t, mode="linear" + ).permute(0, 2, 1) + return time_emb.to(previous_dtype) + + +class CoTrackerThreeOnline(CoTrackerThreeBase): + def __init__(self, **args): + super(CoTrackerThreeOnline, self).__init__(**args) + + def init_video_online_processing(self): + self.online_ind = 0 + self.online_track_feat = [None] * self.corr_levels + self.online_track_support = [None] * self.corr_levels + self.online_coords_predicted = None + self.online_vis_predicted = None + self.online_conf_predicted = None + + def forward_window( + self, + fmaps_pyramid, + coords, + track_feat_support_pyramid, + vis=None, + conf=None, + attention_mask=None, + iters=4, + add_space_attn=False, + ): + B, S, *_ = fmaps_pyramid[0].shape + N = coords.shape[2] + r = 2 * self.corr_radius + 1 + + coord_preds, vis_preds, conf_preds = [], [], [] + for it in range(iters): + coords = coords.detach() # B T N 2 + coords_init = coords.view(B * S, N, 2) + corr_embs = [] + corr_feats = [] + for i in range(self.corr_levels): + corr_feat = self.get_correlation_feat( + fmaps_pyramid[i], coords_init / 2**i + ) + track_feat_support = ( + track_feat_support_pyramid[i] + .view(B, 1, r, r, N, self.latent_dim) + .squeeze(1) + .permute(0, 3, 1, 2, 4) + ) + corr_volume = torch.einsum( + "btnhwc,bnijc->btnhwij", corr_feat, track_feat_support + ) + corr_emb = self.corr_mlp(corr_volume.reshape(B * S * N, r * r * r * r)) + + corr_embs.append(corr_emb) + + corr_embs = torch.cat(corr_embs, dim=-1) + corr_embs = corr_embs.view(B, S, N, corr_embs.shape[-1]) + + transformer_input = [vis, conf, corr_embs] + + rel_coords_forward = coords[:, :-1] - coords[:, 1:] + rel_coords_backward = coords[:, 1:] - coords[:, :-1] + + rel_coords_forward = torch.nn.functional.pad( + rel_coords_forward, (0, 0, 0, 0, 0, 1) + ) + rel_coords_backward = torch.nn.functional.pad( + rel_coords_backward, (0, 0, 0, 0, 1, 0) + ) + + scale = ( + torch.tensor( + [self.model_resolution[1], self.model_resolution[0]], + device=coords.device, + ) + / self.stride + ) + rel_coords_forward = rel_coords_forward / scale + rel_coords_backward = rel_coords_backward / scale + + rel_pos_emb_input = posenc( + torch.cat([rel_coords_forward, rel_coords_backward], dim=-1), + min_deg=0, + max_deg=10, + ) # batch, num_points, num_frames, 84 + transformer_input.append(rel_pos_emb_input) + + x = ( + torch.cat(transformer_input, dim=-1) + .permute(0, 2, 1, 3) + .reshape(B * N, S, -1) + ) + + x = x + self.interpolate_time_embed(x, S) + x = x.view(B, N, S, -1) # (B N) T D -> B N T D + + delta = self.updateformer(x, add_space_attn=add_space_attn) + + delta_coords = delta[..., :2].permute(0, 2, 1, 3) + delta_vis = delta[..., 2:3].permute(0, 2, 1, 3) + delta_conf = delta[..., 3:].permute(0, 2, 1, 3) + + vis = vis + delta_vis + conf = conf + delta_conf + + coords = coords + delta_coords + coord_preds.append(coords[..., :2] * float(self.stride)) + + vis_preds.append(vis[..., 0]) + conf_preds.append(conf[..., 0]) + return coord_preds, vis_preds, conf_preds + + def forward( + self, + video, + queries, + iters=4, + is_train=False, + add_space_attn=True, + fmaps_chunk_size=200, + is_online=False, + ): + """Predict tracks + + Args: + video (FloatTensor[B, T, 3]): input videos. + queries (FloatTensor[B, N, 3]): point queries. + iters (int, optional): number of updates. Defaults to 4. + is_train (bool, optional): enables training mode. Defaults to False. + Returns: + - coords_predicted (FloatTensor[B, T, N, 2]): + - vis_predicted (FloatTensor[B, T, N]): + - train_data: `None` if `is_train` is false, otherwise: + - all_vis_predictions (List[FloatTensor[B, S, N, 1]]): + - all_coords_predictions (List[FloatTensor[B, S, N, 2]]): + - mask (BoolTensor[B, T, N]): + """ + + B, T, C, H, W = video.shape + device = queries.device + assert H % self.stride == 0 and W % self.stride == 0 + + B, N, __ = queries.shape + # B = batch size + # S_trimmed = actual number of frames in the window + # N = number of tracks + # C = color channels (3 for RGB) + # E = positional embedding size + # LRR = local receptive field radius + # D = dimension of the transformer input tokens + + # video = B T C H W + # queries = B N 3 + # coords_init = B T N 2 + # vis_init = B T N 1 + S = self.window_len + assert S >= 2 # A tracker needs at least two frames to track something + if is_online: + assert T <= S, "Online mode: video chunk must be <= window size." + assert ( + self.online_ind is not None + ), "Call model.init_video_online_processing() first." + assert not is_train, "Training not supported in online mode." + + step = S // 2 # How much the sliding window moves at every step + + video = 2 * (video / 255.0) - 1.0 + pad = ( + S - T if is_online else (S - T % S) % S + ) # We don't want to pad if T % S == 0 + video = video.reshape(B, 1, T, C * H * W) + if pad > 0: + padding_tensor = video[:, :, -1:, :].expand(B, 1, pad, C * H * W) + video = torch.cat([video, padding_tensor], dim=2) + video = video.reshape(B, -1, C, H, W) + T_pad = video.shape[1] + # The first channel is the frame number + # The rest are the coordinates of points we want to track + dtype = video.dtype + queried_frames = queries[:, :, 0].long() + + queried_coords = queries[..., 1:3] + queried_coords = queried_coords / self.stride + + # We store our predictions here + coords_predicted = torch.zeros((B, T, N, 2), device=device) + vis_predicted = torch.zeros((B, T, N), device=device) + conf_predicted = torch.zeros((B, T, N), device=device) + + if is_online: + if self.online_coords_predicted is None: + # Init online predictions with zeros + self.online_coords_predicted = coords_predicted + self.online_vis_predicted = vis_predicted + self.online_conf_predicted = conf_predicted + else: + # Pad online predictions with zeros for the current window + pad = min(step, T - step) + coords_predicted = F.pad( + self.online_coords_predicted, (0, 0, 0, 0, 0, pad), "constant" + ) + vis_predicted = F.pad( + self.online_vis_predicted, (0, 0, 0, pad), "constant" + ) + conf_predicted = F.pad( + self.online_conf_predicted, (0, 0, 0, pad), "constant" + ) + + # We store our predictions here + all_coords_predictions, all_vis_predictions, all_confidence_predictions = ( + [], + [], + [], + ) + + C_ = C + H4, W4 = H // self.stride, W // self.stride + + # Compute convolutional features for the video or for the current chunk in case of online mode + if (not is_train) and (T > fmaps_chunk_size): + fmaps = [] + for t in range(0, T, fmaps_chunk_size): + video_chunk = video[:, t : t + fmaps_chunk_size] + fmaps_chunk = self.fnet(video_chunk.reshape(-1, C_, H, W)) + T_chunk = video_chunk.shape[1] + C_chunk, H_chunk, W_chunk = fmaps_chunk.shape[1:] + fmaps.append(fmaps_chunk.reshape(B, T_chunk, C_chunk, H_chunk, W_chunk)) + fmaps = torch.cat(fmaps, dim=1).reshape(-1, C_chunk, H_chunk, W_chunk) + else: + fmaps = self.fnet(video.reshape(-1, C_, H, W)) + fmaps = fmaps.permute(0, 2, 3, 1) + fmaps = fmaps / torch.sqrt( + torch.maximum( + torch.sum(torch.square(fmaps), axis=-1, keepdims=True), + torch.tensor(1e-12, device=fmaps.device), + ) + ) + fmaps = fmaps.permute(0, 3, 1, 2).reshape( + B, -1, self.latent_dim, H // self.stride, W // self.stride + ) + fmaps = fmaps.to(dtype) + + # We compute track features + fmaps_pyramid = [] + track_feat_pyramid = [] + track_feat_support_pyramid = [] + fmaps_pyramid.append(fmaps) + for i in range(self.corr_levels - 1): + fmaps_ = fmaps.reshape( + B * T_pad, self.latent_dim, fmaps.shape[-2], fmaps.shape[-1] + ) + fmaps_ = F.avg_pool2d(fmaps_, 2, stride=2) + fmaps = fmaps_.reshape( + B, T_pad, self.latent_dim, fmaps_.shape[-2], fmaps_.shape[-1] + ) + fmaps_pyramid.append(fmaps) + if is_online: + sample_frames = queried_frames[:, None, :, None] # B 1 N 1 + left = 0 if self.online_ind == 0 else self.online_ind + step + right = self.online_ind + S + sample_mask = (sample_frames >= left) & (sample_frames < right) + + for i in range(self.corr_levels): + track_feat, track_feat_support = self.get_track_feat( + fmaps_pyramid[i], + queried_frames - self.online_ind if is_online else queried_frames, + queried_coords / 2**i, + support_radius=self.corr_radius, + ) + + if is_online: + if self.online_track_feat[i] is None: + self.online_track_feat[i] = torch.zeros_like( + track_feat, device=device + ) + self.online_track_support[i] = torch.zeros_like( + track_feat_support, device=device + ) + + self.online_track_feat[i] += track_feat * sample_mask + self.online_track_support[i] += track_feat_support * sample_mask + track_feat_pyramid.append( + self.online_track_feat[i].repeat(1, T_pad, 1, 1) + ) + track_feat_support_pyramid.append( + self.online_track_support[i].unsqueeze(1) + ) + else: + track_feat_pyramid.append(track_feat.repeat(1, T_pad, 1, 1)) + track_feat_support_pyramid.append(track_feat_support.unsqueeze(1)) + + D_coords = 2 + coord_preds, vis_preds, confidence_preds = [], [], [] + + vis_init = torch.zeros((B, S, N, 1), device=device).float() + conf_init = torch.zeros((B, S, N, 1), device=device).float() + coords_init = queried_coords.reshape(B, 1, N, 2).expand(B, S, N, 2).float() + + num_windows = (T - S + step - 1) // step + 1 + # We process only the current video chunk in the online mode + indices = [self.online_ind] if is_online else range(0, step * num_windows, step) + + for ind in indices: + if ind > 0: + overlap = S - step + copy_over = (queried_frames < ind + overlap)[ + :, None, :, None + ] # B 1 N 1 + coords_prev = coords_predicted[:, ind : ind + overlap] / self.stride + padding_tensor = coords_prev[:, -1:, :, :].expand(-1, step, -1, -1) + coords_prev = torch.cat([coords_prev, padding_tensor], dim=1) + + vis_prev = vis_predicted[:, ind : ind + overlap, :, None].clone() + padding_tensor = vis_prev[:, -1:, :, :].expand(-1, step, -1, -1) + vis_prev = torch.cat([vis_prev, padding_tensor], dim=1) + + conf_prev = conf_predicted[:, ind : ind + overlap, :, None].clone() + padding_tensor = conf_prev[:, -1:, :, :].expand(-1, step, -1, -1) + conf_prev = torch.cat([conf_prev, padding_tensor], dim=1) + + coords_init = torch.where( + copy_over.expand_as(coords_init), coords_prev, coords_init + ) + vis_init = torch.where( + copy_over.expand_as(vis_init), vis_prev, vis_init + ) + conf_init = torch.where( + copy_over.expand_as(conf_init), conf_prev, conf_init + ) + + attention_mask = (queried_frames < ind + S).reshape(B, 1, N) # B S N + # import ipdb; ipdb.set_trace() + coords, viss, confs = self.forward_window( + fmaps_pyramid=( + fmaps_pyramid + if is_online + else [fmap[:, ind : ind + S] for fmap in fmaps_pyramid] + ), + coords=coords_init, + track_feat_support_pyramid=[ + attention_mask[:, None, :, :, None] * tfeat + for tfeat in track_feat_support_pyramid + ], + vis=vis_init, + conf=conf_init, + attention_mask=attention_mask.repeat(1, S, 1), + iters=iters, + add_space_attn=add_space_attn, + ) + S_trimmed = ( + T if is_online else min(T - ind, S) + ) # accounts for last window duration + coords_predicted[:, ind : ind + S] = coords[-1][:, :S_trimmed] + vis_predicted[:, ind : ind + S] = viss[-1][:, :S_trimmed] + conf_predicted[:, ind : ind + S] = confs[-1][:, :S_trimmed] + if is_train: + all_coords_predictions.append( + [coord[:, :S_trimmed] for coord in coords] + ) + all_vis_predictions.append( + [torch.sigmoid(vis[:, :S_trimmed]) for vis in viss] + ) + all_confidence_predictions.append( + [torch.sigmoid(conf[:, :S_trimmed]) for conf in confs] + ) + if is_online: + self.online_ind += step + self.online_coords_predicted = coords_predicted + self.online_vis_predicted = vis_predicted + self.online_conf_predicted = conf_predicted + vis_predicted = torch.sigmoid(vis_predicted) + conf_predicted = torch.sigmoid(conf_predicted) + + if is_train: + valid_mask = ( + queried_frames[:, None] + <= torch.arange(0, T, device=device)[None, :, None] + ) + train_data = ( + all_coords_predictions, + all_vis_predictions, + all_confidence_predictions, + valid_mask, + ) + else: + train_data = None + + return coords_predicted, vis_predicted, conf_predicted, train_data diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/losses.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/losses.py new file mode 100644 index 0000000000000000000000000000000000000000..e9286ec8951404fb20c509cfd6ee08b0a32ab66d --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/cotracker/losses.py @@ -0,0 +1,118 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch +import torch.nn.functional as F +from cotracker.models.core.model_utils import reduce_masked_mean +import torch.nn as nn +from typing import List + + +def sequence_loss( + flow_preds, + flow_gt, + valids, + vis=None, + gamma=0.8, + add_huber_loss=False, + loss_only_for_visible=False, +): + """Loss function defined over sequence of flow predictions""" + total_flow_loss = 0.0 + for j in range(len(flow_gt)): + B, S, N, D = flow_gt[j].shape + B, S2, N = valids[j].shape + assert S == S2 + n_predictions = len(flow_preds[j]) + flow_loss = 0.0 + for i in range(n_predictions): + i_weight = gamma ** (n_predictions - i - 1) + flow_pred = flow_preds[j][i] + if add_huber_loss: + i_loss = huber_loss(flow_pred, flow_gt[j], delta=6.0) + else: + i_loss = (flow_pred - flow_gt[j]).abs() # B, S, N, 2 + i_loss = torch.mean(i_loss, dim=3) # B, S, N + valid_ = valids[j].clone() + if loss_only_for_visible: + valid_ = valid_ * vis[j] + flow_loss += i_weight * reduce_masked_mean(i_loss, valid_) + flow_loss = flow_loss / n_predictions + total_flow_loss += flow_loss + return total_flow_loss / len(flow_gt) + + +def huber_loss(x, y, delta=1.0): + """Calculate element-wise Huber loss between x and y""" + diff = x - y + abs_diff = diff.abs() + flag = (abs_diff <= delta).float() + return flag * 0.5 * diff**2 + (1 - flag) * delta * (abs_diff - 0.5 * delta) + + +def sequence_BCE_loss(vis_preds, vis_gts): + total_bce_loss = 0.0 + for j in range(len(vis_preds)): + n_predictions = len(vis_preds[j]) + bce_loss = 0.0 + for i in range(n_predictions): + vis_loss = F.binary_cross_entropy(vis_preds[j][i], vis_gts[j]) + bce_loss += vis_loss + bce_loss = bce_loss / n_predictions + total_bce_loss += bce_loss + return total_bce_loss / len(vis_preds) + + +def sequence_prob_loss( + tracks: torch.Tensor, + confidence: torch.Tensor, + target_points: torch.Tensor, + visibility: torch.Tensor, + expected_dist_thresh: float = 12.0, +): + """Loss for classifying if a point is within pixel threshold of its target.""" + # Points with an error larger than 12 pixels are likely to be useless; marking + # them as occluded will actually improve Jaccard metrics and give + # qualitatively better results. + total_logprob_loss = 0.0 + for j in range(len(tracks)): + n_predictions = len(tracks[j]) + logprob_loss = 0.0 + for i in range(n_predictions): + err = torch.sum((tracks[j][i].detach() - target_points[j]) ** 2, dim=-1) + valid = (err <= expected_dist_thresh**2).float() + logprob = F.binary_cross_entropy(confidence[j][i], valid, reduction="none") + logprob *= visibility[j] + logprob = torch.mean(logprob, dim=[1, 2]) + logprob_loss += logprob + logprob_loss = logprob_loss / n_predictions + total_logprob_loss += logprob_loss + return total_logprob_loss / len(tracks) + + +def masked_mean(data: torch.Tensor, mask: torch.Tensor | None, dim: List[int]): + if mask is None: + return data.mean(dim=dim, keepdim=True) + mask = mask.float() + mask_sum = torch.sum(mask, dim=dim, keepdim=True) + mask_mean = torch.sum(data * mask, dim=dim, keepdim=True) / torch.clamp( + mask_sum, min=1.0 + ) + return mask_mean + + +def masked_mean_var(data: torch.Tensor, mask: torch.Tensor, dim: List[int]): + if mask is None: + return data.mean(dim=dim, keepdim=True), data.var(dim=dim, keepdim=True) + mask = mask.float() + mask_sum = torch.sum(mask, dim=dim, keepdim=True) + mask_mean = torch.sum(data * mask, dim=dim, keepdim=True) / torch.clamp( + mask_sum, min=1.0 + ) + mask_var = torch.sum( + mask * (data - mask_mean) ** 2, dim=dim, keepdim=True + ) / torch.clamp(mask_sum, min=1.0) + return mask_mean.squeeze(dim), mask_var.squeeze(dim) diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/embeddings.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/embeddings.py new file mode 100644 index 0000000000000000000000000000000000000000..897cd5d9f41121a9692281a719a2d24914293318 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/embeddings.py @@ -0,0 +1,120 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +from typing import Tuple, Union +import torch + + +def get_2d_sincos_pos_embed( + embed_dim: int, grid_size: Union[int, Tuple[int, int]] +) -> torch.Tensor: + """ + This function initializes a grid and generates a 2D positional embedding using sine and cosine functions. + It is a wrapper of get_2d_sincos_pos_embed_from_grid. + Args: + - embed_dim: The embedding dimension. + - grid_size: The grid size. + Returns: + - pos_embed: The generated 2D positional embedding. + """ + if isinstance(grid_size, tuple): + grid_size_h, grid_size_w = grid_size + else: + grid_size_h = grid_size_w = grid_size + grid_h = torch.arange(grid_size_h, dtype=torch.float) + grid_w = torch.arange(grid_size_w, dtype=torch.float) + grid = torch.meshgrid(grid_w, grid_h, indexing="xy") + grid = torch.stack(grid, dim=0) + grid = grid.reshape([2, 1, grid_size_h, grid_size_w]) + pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid) + return pos_embed.reshape(1, grid_size_h, grid_size_w, -1).permute(0, 3, 1, 2) + + +def get_2d_sincos_pos_embed_from_grid( + embed_dim: int, grid: torch.Tensor +) -> torch.Tensor: + """ + This function generates a 2D positional embedding from a given grid using sine and cosine functions. + + Args: + - embed_dim: The embedding dimension. + - grid: The grid to generate the embedding from. + + Returns: + - emb: The generated 2D positional embedding. + """ + assert embed_dim % 2 == 0 + + # use half of dimensions to encode grid_h + emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2) + emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2) + + emb = torch.cat([emb_h, emb_w], dim=2) # (H*W, D) + return emb + + +def get_1d_sincos_pos_embed_from_grid( + embed_dim: int, pos: torch.Tensor +) -> torch.Tensor: + """ + This function generates a 1D positional embedding from a given grid using sine and cosine functions. + + Args: + - embed_dim: The embedding dimension. + - pos: The position to generate the embedding from. + + Returns: + - emb: The generated 1D positional embedding. + """ + assert embed_dim % 2 == 0 + omega = torch.arange(embed_dim // 2, dtype=torch.double) + omega /= embed_dim / 2.0 + omega = 1.0 / 10000**omega # (D/2,) + + pos = pos.reshape(-1) # (M,) + out = torch.einsum("m,d->md", pos, omega) # (M, D/2), outer product + + emb_sin = torch.sin(out) # (M, D/2) + emb_cos = torch.cos(out) # (M, D/2) + + emb = torch.cat([emb_sin, emb_cos], dim=1) # (M, D) + return emb[None].float() + + +def get_2d_embedding(xy: torch.Tensor, C: int, cat_coords: bool = True) -> torch.Tensor: + """ + This function generates a 2D positional embedding from given coordinates using sine and cosine functions. + + Args: + - xy: The coordinates to generate the embedding from. + - C: The size of the embedding. + - cat_coords: A flag to indicate whether to concatenate the original coordinates to the embedding. + + Returns: + - pe: The generated 2D positional embedding. + """ + B, N, D = xy.shape + assert D == 2 + + x = xy[:, :, 0:1] + y = xy[:, :, 1:2] + div_term = ( + torch.arange(0, C, 2, device=xy.device, dtype=torch.float32) * (1000.0 / C) + ).reshape(1, 1, int(C / 2)) + + pe_x = torch.zeros(B, N, C, device=xy.device, dtype=torch.float32) + pe_y = torch.zeros(B, N, C, device=xy.device, dtype=torch.float32) + + pe_x[:, :, 0::2] = torch.sin(x * div_term) + pe_x[:, :, 1::2] = torch.cos(x * div_term) + + pe_y[:, :, 0::2] = torch.sin(y * div_term) + pe_y[:, :, 1::2] = torch.cos(y * div_term) + + pe = torch.cat([pe_x, pe_y], dim=2) # (B, N, C*3) + if cat_coords: + pe = torch.cat([xy, pe], dim=2) # (B, N, C*3+3) + return pe diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/model_utils.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/model_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..cad72771417e1aa6d6687f40f5f4175fbe3862a1 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/core/model_utils.py @@ -0,0 +1,426 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import numpy as np +import random +import torch +import torch.nn.functional as F +from typing import Optional, Tuple + +EPS = 1e-6 + + +def smart_cat(tensor1, tensor2, dim): + if tensor1 is None: + return tensor2 + return torch.cat([tensor1, tensor2], dim=dim) + + +def get_uniformly_sampled_pts( + size: int, + num_frames: int, + extent: Tuple[float, ...], + device: Optional[torch.device] = torch.device("cpu"), +): + time_points = torch.randint(low=0, high=num_frames, size=(size, 1), device=device) + space_points = torch.rand(size, 2, device=device) * torch.tensor( + [extent[1], extent[0]], device=device + ) + points = torch.cat((time_points, space_points), dim=1) + return points[None] + + +def get_superpoint_sampled_pts( + video, + size: int, + num_frames: int, + extent: Tuple[float, ...], + device: Optional[torch.device] = torch.device("cpu"), +): + extractor = SuperPoint(max_num_keypoints=48).eval().cuda() + points = list() + for _ in range(8): + frame_num = random.randint(0, int(num_frames * 0.25)) + key_points = extractor.extract( + video[0, frame_num, :, :, :] / 255.0, resize=None + )["keypoints"] + frame_tensor = torch.full((1, key_points.shape[1], 1), frame_num).cuda() + points.append(torch.cat([frame_tensor.cuda(), key_points], dim=2)) + return torch.cat(points, dim=1)[:, :size, :] + + +def get_sift_sampled_pts( + video, + size: int, + num_frames: int, + extent: Tuple[float, ...], + device: Optional[torch.device] = torch.device("cpu"), + num_sampled_frames: int = 8, + sampling_length_percent: float = 0.25, +): + import cv2 + # assert size == 384, "hardcoded for experiment" + sift = cv2.SIFT_create(nfeatures=size // num_sampled_frames) + points = list() + for _ in range(num_sampled_frames): + frame_num = random.randint(0, int(num_frames * sampling_length_percent)) + key_points, _ = sift.detectAndCompute( + video[0, frame_num, :, :, :] + .cpu() + .permute(1, 2, 0) + .numpy() + .astype(np.uint8), + None, + ) + for kp in key_points: + points.append([frame_num, int(kp.pt[0]), int(kp.pt[1])]) + return torch.tensor(points[:size], device=device)[None] + + +def get_points_on_a_grid( + size: int, + extent: Tuple[float, ...], + center: Optional[Tuple[float, ...]] = None, + device: Optional[torch.device] = torch.device("cpu"), +): + r"""Get a grid of points covering a rectangular region + + `get_points_on_a_grid(size, extent)` generates a :attr:`size` by + :attr:`size` grid fo points distributed to cover a rectangular area + specified by `extent`. + + The `extent` is a pair of integer :math:`(H,W)` specifying the height + and width of the rectangle. + + Optionally, the :attr:`center` can be specified as a pair :math:`(c_y,c_x)` + specifying the vertical and horizontal center coordinates. The center + defaults to the middle of the extent. + + Points are distributed uniformly within the rectangle leaving a margin + :math:`m=W/64` from the border. + + It returns a :math:`(1, \text{size} \times \text{size}, 2)` tensor of + points :math:`P_{ij}=(x_i, y_i)` where + + .. math:: + P_{ij} = \left( + c_x + m -\frac{W}{2} + \frac{W - 2m}{\text{size} - 1}\, j,~ + c_y + m -\frac{H}{2} + \frac{H - 2m}{\text{size} - 1}\, i + \right) + + Points are returned in row-major order. + + Args: + size (int): grid size. + extent (tuple): height and with of the grid extent. + center (tuple, optional): grid center. + device (str, optional): Defaults to `"cpu"`. + + Returns: + Tensor: grid. + """ + if size == 1: + return torch.tensor([extent[1] / 2, extent[0] / 2], device=device)[None, None] + + if center is None: + center = [extent[0] / 2, extent[1] / 2] + + margin = extent[1] / 64 + range_y = (margin - extent[0] / 2 + center[0], extent[0] / 2 + center[0] - margin) + range_x = (margin - extent[1] / 2 + center[1], extent[1] / 2 + center[1] - margin) + grid_y, grid_x = torch.meshgrid( + torch.linspace(*range_y, size, device=device), + torch.linspace(*range_x, size, device=device), + indexing="ij", + ) + return torch.stack([grid_x, grid_y], dim=-1).reshape(1, -1, 2) + + +def reduce_masked_mean(input, mask, dim=None, keepdim=False): + r"""Masked mean + + `reduce_masked_mean(x, mask)` computes the mean of a tensor :attr:`input` + over a mask :attr:`mask`, returning + + .. math:: + \text{output} = + \frac + {\sum_{i=1}^N \text{input}_i \cdot \text{mask}_i} + {\epsilon + \sum_{i=1}^N \text{mask}_i} + + where :math:`N` is the number of elements in :attr:`input` and + :attr:`mask`, and :math:`\epsilon` is a small constant to avoid + division by zero. + + `reduced_masked_mean(x, mask, dim)` computes the mean of a tensor + :attr:`input` over a mask :attr:`mask` along a dimension :attr:`dim`. + Optionally, the dimension can be kept in the output by setting + :attr:`keepdim` to `True`. Tensor :attr:`mask` must be broadcastable to + the same dimension as :attr:`input`. + + The interface is similar to `torch.mean()`. + + Args: + inout (Tensor): input tensor. + mask (Tensor): mask. + dim (int, optional): Dimension to sum over. Defaults to None. + keepdim (bool, optional): Keep the summed dimension. Defaults to False. + + Returns: + Tensor: mean tensor. + """ + + mask = mask.expand_as(input) + + prod = input * mask + + if dim is None: + numer = torch.sum(prod) + denom = torch.sum(mask) + else: + numer = torch.sum(prod, dim=dim, keepdim=keepdim) + denom = torch.sum(mask, dim=dim, keepdim=keepdim) + + mean = numer / (EPS + denom) + return mean + + +def bilinear_sampler(input, coords, align_corners=True, padding_mode="border"): + r"""Sample a tensor using bilinear interpolation + + `bilinear_sampler(input, coords)` samples a tensor :attr:`input` at + coordinates :attr:`coords` using bilinear interpolation. It is the same + as `torch.nn.functional.grid_sample()` but with a different coordinate + convention. + + The input tensor is assumed to be of shape :math:`(B, C, H, W)`, where + :math:`B` is the batch size, :math:`C` is the number of channels, + :math:`H` is the height of the image, and :math:`W` is the width of the + image. The tensor :attr:`coords` of shape :math:`(B, H_o, W_o, 2)` is + interpreted as an array of 2D point coordinates :math:`(x_i,y_i)`. + + Alternatively, the input tensor can be of size :math:`(B, C, T, H, W)`, + in which case sample points are triplets :math:`(t_i,x_i,y_i)`. Note + that in this case the order of the components is slightly different + from `grid_sample()`, which would expect :math:`(x_i,y_i,t_i)`. + + If `align_corners` is `True`, the coordinate :math:`x` is assumed to be + in the range :math:`[0,W-1]`, with 0 corresponding to the center of the + left-most image pixel :math:`W-1` to the center of the right-most + pixel. + + If `align_corners` is `False`, the coordinate :math:`x` is assumed to + be in the range :math:`[0,W]`, with 0 corresponding to the left edge of + the left-most pixel :math:`W` to the right edge of the right-most + pixel. + + Similar conventions apply to the :math:`y` for the range + :math:`[0,H-1]` and :math:`[0,H]` and to :math:`t` for the range + :math:`[0,T-1]` and :math:`[0,T]`. + + Args: + input (Tensor): batch of input images. + coords (Tensor): batch of coordinates. + align_corners (bool, optional): Coordinate convention. Defaults to `True`. + padding_mode (str, optional): Padding mode. Defaults to `"border"`. + + Returns: + Tensor: sampled points. + """ + + sizes = input.shape[2:] + + assert len(sizes) in [2, 3] + + if len(sizes) == 3: + # t x y -> x y t to match dimensions T H W in grid_sample + coords = coords[..., [1, 2, 0]] + + if align_corners: + coords = coords * torch.tensor( + [2 / max(size - 1, 1) for size in reversed(sizes)], device=coords.device + ) + else: + coords = coords * torch.tensor( + [2 / size for size in reversed(sizes)], device=coords.device + ) + + coords -= 1 + + return F.grid_sample( + input, coords, align_corners=align_corners, padding_mode=padding_mode + ) + + +def sample_features4d(input, coords): + r"""Sample spatial features + + `sample_features4d(input, coords)` samples the spatial features + :attr:`input` represented by a 4D tensor :math:`(B, C, H, W)`. + + The field is sampled at coordinates :attr:`coords` using bilinear + interpolation. :attr:`coords` is assumed to be of shape :math:`(B, R, + 3)`, where each sample has the format :math:`(x_i, y_i)`. This uses the + same convention as :func:`bilinear_sampler` with `align_corners=True`. + + The output tensor has one feature per point, and has shape :math:`(B, + R, C)`. + + Args: + input (Tensor): spatial features. + coords (Tensor): points. + + Returns: + Tensor: sampled features. + """ + + B, _, _, _ = input.shape + + # B R 2 -> B R 1 2 + coords = coords.unsqueeze(2) + + # B C R 1 + feats = bilinear_sampler(input, coords) + + return feats.permute(0, 2, 1, 3).view( + B, -1, feats.shape[1] * feats.shape[3] + ) # B C R 1 -> B R C + + +def sample_features5d(input, coords): + r"""Sample spatio-temporal features + + `sample_features5d(input, coords)` works in the same way as + :func:`sample_features4d` but for spatio-temporal features and points: + :attr:`input` is a 5D tensor :math:`(B, T, C, H, W)`, :attr:`coords` is + a :math:`(B, R1, R2, 3)` tensor of spatio-temporal point :math:`(t_i, + x_i, y_i)`. The output tensor has shape :math:`(B, R1, R2, C)`. + + Args: + input (Tensor): spatio-temporal features. + coords (Tensor): spatio-temporal points. + + Returns: + Tensor: sampled features. + """ + + B, T, _, _, _ = input.shape + + # B T C H W -> B C T H W + input = input.permute(0, 2, 1, 3, 4) + + # B R1 R2 3 -> B R1 R2 1 3 + coords = coords.unsqueeze(3) + + # B C R1 R2 1 + feats = bilinear_sampler(input, coords) + + return feats.permute(0, 2, 3, 1, 4).view( + B, feats.shape[2], feats.shape[3], feats.shape[1] + ) # B C R1 R2 1 -> B R1 R2 C + + +def get_grid( + height, + width, + shape=None, + dtype="torch", + device="cpu", + align_corners=True, + normalize=True, +): + H, W = height, width + S = shape if shape else [] + if align_corners: + x = torch.linspace(0, 1, W, device=device) + y = torch.linspace(0, 1, H, device=device) + if not normalize: + x = x * (W - 1) + y = y * (H - 1) + else: + x = torch.linspace(0.5 / W, 1.0 - 0.5 / W, W, device=device) + y = torch.linspace(0.5 / H, 1.0 - 0.5 / H, H, device=device) + if not normalize: + x = x * W + y = y * H + x_view, y_view, exp = [1 for _ in S] + [1, -1], [1 for _ in S] + [-1, 1], S + [H, W] + x = x.view(*x_view).expand(*exp) + y = y.view(*y_view).expand(*exp) + grid = torch.stack([x, y], dim=-1) + if dtype == "numpy": + grid = grid.numpy() + return grid + + +def bilinear_sampler(input, coords, align_corners=True, padding_mode="border"): + r"""Sample a tensor using bilinear interpolation + + `bilinear_sampler(input, coords)` samples a tensor :attr:`input` at + coordinates :attr:`coords` using bilinear interpolation. It is the same + as `torch.nn.functional.grid_sample()` but with a different coordinate + convention. + + The input tensor is assumed to be of shape :math:`(B, C, H, W)`, where + :math:`B` is the batch size, :math:`C` is the number of channels, + :math:`H` is the height of the image, and :math:`W` is the width of the + image. The tensor :attr:`coords` of shape :math:`(B, H_o, W_o, 2)` is + interpreted as an array of 2D point coordinates :math:`(x_i,y_i)`. + + Alternatively, the input tensor can be of size :math:`(B, C, T, H, W)`, + in which case sample points are triplets :math:`(t_i,x_i,y_i)`. Note + that in this case the order of the components is slightly different + from `grid_sample()`, which would expect :math:`(x_i,y_i,t_i)`. + + If `align_corners` is `True`, the coordinate :math:`x` is assumed to be + in the range :math:`[0,W-1]`, with 0 corresponding to the center of the + left-most image pixel :math:`W-1` to the center of the right-most + pixel. + + If `align_corners` is `False`, the coordinate :math:`x` is assumed to + be in the range :math:`[0,W]`, with 0 corresponding to the left edge of + the left-most pixel :math:`W` to the right edge of the right-most + pixel. + + Similar conventions apply to the :math:`y` for the range + :math:`[0,H-1]` and :math:`[0,H]` and to :math:`t` for the range + :math:`[0,T-1]` and :math:`[0,T]`. + + Args: + input (Tensor): batch of input images. + coords (Tensor): batch of coordinates. + align_corners (bool, optional): Coordinate convention. Defaults to `True`. + padding_mode (str, optional): Padding mode. Defaults to `"border"`. + + Returns: + Tensor: sampled points. + """ + + sizes = input.shape[2:] + + assert len(sizes) in [2, 3] + + if len(sizes) == 3: + # t x y -> x y t to match dimensions T H W in grid_sample + coords = coords[..., [1, 2, 0]] + + if align_corners: + coords = coords * torch.tensor( + [2 / max(size - 1, 1) for size in reversed(sizes)], device=coords.device + ) + else: + coords = coords * torch.tensor( + [2 / size for size in reversed(sizes)], device=coords.device + ) + + coords -= 1 + + return F.grid_sample( + input, coords, align_corners=align_corners, padding_mode=padding_mode + ) + + +def round_to_multiple_of_4(n): + return round(n / 4) * 4 diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/models/evaluation_predictor.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/evaluation_predictor.py new file mode 100644 index 0000000000000000000000000000000000000000..71c84e0687ed8b6f2bad694da19147eef238eb98 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/models/evaluation_predictor.py @@ -0,0 +1,199 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch +import torch.nn.functional as F +from typing import Tuple + +from cotracker.models.core.cotracker.cotracker3_offline import CoTrackerThreeOffline +from cotracker.models.core.model_utils import ( + get_points_on_a_grid, + get_uniformly_sampled_pts, + get_sift_sampled_pts, +) +import numpy as np +import sys + +from torchvision.transforms import Compose +from tqdm import tqdm +from cotracker.models.core.model_utils import bilinear_sampler + + +class EvaluationPredictor(torch.nn.Module): + def __init__( + self, + cotracker_model: CoTrackerThreeOffline, + interp_shape: Tuple[int, int] = (384, 512), + grid_size: int = 5, + local_grid_size: int = 8, + single_point: bool = True, + sift_size: int = 0, + num_uniformly_sampled_pts: int = 0, + n_iters: int = 6, + local_extent: int = 50, + ) -> None: + super(EvaluationPredictor, self).__init__() + self.grid_size = grid_size + self.local_grid_size = local_grid_size + self.sift_size = sift_size + self.single_point = single_point + self.interp_shape = interp_shape + self.n_iters = n_iters + self.num_uniformly_sampled_pts = num_uniformly_sampled_pts + self.model = cotracker_model + self.local_extent = local_extent + self.model.eval() + + def forward(self, video, queries): + queries = queries.clone() + B, T, C, H, W = video.shape + B, N, D = queries.shape + + assert D == 3 + assert B == 1 + interp_shape = self.interp_shape + + video = video.reshape(B * T, C, H, W) + video = F.interpolate( + video, tuple(interp_shape), mode="bilinear", align_corners=True + ) + video = video.reshape(B, T, 3, interp_shape[0], interp_shape[1]) + + device = video.device + + queries[:, :, 1] *= (interp_shape[1] - 1) / (W - 1) + queries[:, :, 2] *= (interp_shape[0] - 1) / (H - 1) + + if self.single_point: + traj_e = torch.zeros((B, T, N, 2), device=device) + vis_e = torch.zeros((B, T, N), device=device) + conf_e = torch.zeros((B, T, N), device=device) + + for pind in range((N)): + query = queries[:, pind : pind + 1] + t = query[0, 0, 0].long() + start_ind = 0 + traj_e_pind, vis_e_pind, conf_e_pind = self._process_one_point( + video[:,start_ind:], query + ) + traj_e[:, start_ind:, pind : pind + 1] = traj_e_pind[:, :, :1] + vis_e[:, start_ind:, pind : pind + 1] = vis_e_pind[:, :, :1] + conf_e[:, start_ind:, pind : pind + 1] = conf_e_pind[:, :, :1] + else: + if self.grid_size > 0: + xy = get_points_on_a_grid(self.grid_size, video.shape[3:]) + xy = torch.cat([torch.zeros_like(xy[:, :, :1]), xy], dim=2).to( + device + ) # + queries = torch.cat([queries, xy], dim=1) # + + if self.num_uniformly_sampled_pts > 0: + xy = get_uniformly_sampled_pts( + self.num_uniformly_sampled_pts, + video.shape[1], + video.shape[3:], + device=device, + ) + queries = torch.cat([queries, xy], dim=1) # + + sift_size = self.sift_size + if sift_size > 0: + xy = get_sift_sampled_pts(video, sift_size, T, [H, W], device=device) + if xy.shape[1] == sift_size: + queries = torch.cat([queries, xy], dim=1) # + else: + sift_size = 0 + + preds = self.model(video=video, queries=queries, iters=self.n_iters) + traj_e, vis_e = preds[0], preds[1] + conf_e = None + if len(preds) > 3: + conf_e = preds[2] + if ( + sift_size > 0 + or self.grid_size > 0 + or self.num_uniformly_sampled_pts > 0 + ): + traj_e = traj_e[ + :, + :, + : -self.grid_size**2 - sift_size - self.num_uniformly_sampled_pts, + ] + vis_e = vis_e[ + :, + :, + : -self.grid_size**2 - sift_size - self.num_uniformly_sampled_pts, + ] + if conf_e is not None: + conf_e = conf_e[ + :, + :, + : -self.grid_size**2 + - sift_size + - self.num_uniformly_sampled_pts, + ] + + traj_e[:, :, :, 0] *= (W - 1) / float(interp_shape[1] - 1) + traj_e[:, :, :, 1] *= (H - 1) / float(interp_shape[0] - 1) + if conf_e is not None: + vis_e = vis_e * conf_e + + return traj_e, vis_e + + def _process_one_point(self, video, query): + t = query[0, 0, 0].long() + B, T, C, H, W = video.shape + device = query.device + if self.local_grid_size > 0: + xy_target = get_points_on_a_grid( + self.local_grid_size, + (self.local_extent, self.local_extent), + [query[0, 0, 2].item(), query[0, 0, 1].item()], + ) + + xy_target = torch.cat( + [torch.zeros_like(xy_target[:, :, :1]), xy_target], dim=2 + ).to( + device + ) # + query = torch.cat([query, xy_target], dim=1) # + + if self.grid_size > 0: + xy = get_points_on_a_grid(self.grid_size, video.shape[3:]) + xy = torch.cat([torch.zeros_like(xy[:, :, :1]), xy], dim=2).to(device) # + query = torch.cat([query, xy], dim=1) # + + sift_size = self.sift_size + if sift_size > 0: + xy = get_sift_sampled_pts(video, sift_size, T, [H, W], device=device) + sift_size = xy.shape[1] + if sift_size > 0: + query = torch.cat([query, xy], dim=1) # + + num_uniformly_sampled_pts = self.sift_size - sift_size + if num_uniformly_sampled_pts > 0: + xy2 = get_uniformly_sampled_pts( + num_uniformly_sampled_pts, + video.shape[1], + video.shape[3:], + device=device, + ) + query = torch.cat([query, xy2], dim=1) # + + if self.num_uniformly_sampled_pts > 0: + xy = get_uniformly_sampled_pts( + self.num_uniformly_sampled_pts, + video.shape[1], + video.shape[3:], + device=device, + ) + query = torch.cat([query, xy], dim=1) # + + traj_e_pind, vis_e_pind, conf_e_pind, __ = self.model( + video=video, queries=query, iters=self.n_iters + ) + + return traj_e_pind[..., :2], vis_e_pind, conf_e_pind diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/predictor.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/predictor.py new file mode 100644 index 0000000000000000000000000000000000000000..da7e7aba1c4ad988a111c2c81bb492de966a375e --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/predictor.py @@ -0,0 +1,309 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch +import torch.nn.functional as F + +from cotracker.models.core.model_utils import smart_cat, get_points_on_a_grid +from cotracker.models.build_cotracker import build_cotracker + + +class CoTrackerPredictor(torch.nn.Module): + def __init__( + self, + checkpoint="./checkpoints/scaled_offline.pth", + offline=True, + v2=False, + window_len=60, + ): + super().__init__() + self.v2 = v2 + self.support_grid_size = 6 + model = build_cotracker( + checkpoint, + v2=v2, + offline=offline, + window_len=window_len, + ) + self.interp_shape = model.model_resolution + self.model = model + self.model.eval() + + @torch.no_grad() + def forward( + self, + video, # (B, T, 3, H, W) + # input prompt types: + # - None. Dense tracks are computed in this case. You can adjust *query_frame* to compute tracks starting from a specific frame. + # *backward_tracking=True* will compute tracks in both directions. + # - queries. Queried points of shape (B, N, 3) in format (t, x, y) for frame index and pixel coordinates. + # - grid_size. Grid of N*N points from the first frame. if segm_mask is provided, then computed only for the mask. + # You can adjust *query_frame* and *backward_tracking* for the regular grid in the same way as for dense tracks. + queries: torch.Tensor = None, + segm_mask: torch.Tensor = None, # Segmentation mask of shape (B, 1, H, W) + grid_size: int = 0, + grid_query_frame: int = 0, # only for dense and regular grid tracks + backward_tracking: bool = False, + ): + if queries is None and grid_size == 0: + tracks, visibilities = self._compute_dense_tracks( + video, + grid_query_frame=grid_query_frame, + backward_tracking=backward_tracking, + ) + else: + tracks, visibilities = self._compute_sparse_tracks( + video, + queries, + segm_mask, + grid_size, + add_support_grid=(grid_size == 0 or segm_mask is not None), + grid_query_frame=grid_query_frame, + backward_tracking=backward_tracking, + ) + + return tracks, visibilities + + def _compute_dense_tracks( + self, video, grid_query_frame, grid_size=80, backward_tracking=False + ): + *_, H, W = video.shape + grid_step = W // grid_size + grid_width = W // grid_step + grid_height = H // grid_step + tracks = visibilities = None + grid_pts = torch.zeros((video.shape[0], grid_width * grid_height, 3)).to(video.device) + grid_pts[:, :, 0] = grid_query_frame + for offset in range(grid_step * grid_step): + print(f"step {offset} / {grid_step * grid_step}") + ox = offset % grid_step + oy = offset // grid_step + grid_pts[:, :, 1] = ( + torch.arange(grid_width).repeat(grid_height) * grid_step + ox + ) + grid_pts[:, :, 2] = ( + torch.arange(grid_height).repeat_interleave(grid_width) * grid_step + oy + ) + tracks_step, visibilities_step = self._compute_sparse_tracks( + video=video, + queries=grid_pts, + backward_tracking=backward_tracking, + ) + tracks = smart_cat(tracks, tracks_step, dim=2) + visibilities = smart_cat(visibilities, visibilities_step, dim=2) + + return tracks, visibilities + + def _compute_sparse_tracks( + self, + video, + queries, + segm_mask=None, + grid_size=0, + add_support_grid=False, + grid_query_frame=0, + backward_tracking=False, + ): + B, T, C, H, W = video.shape + + video = video.reshape(B * T, C, H, W) + video = F.interpolate( + video, tuple(self.interp_shape), mode="bilinear", align_corners=True + ) + video = video.reshape(B, T, 3, self.interp_shape[0], self.interp_shape[1]) + + if queries is not None: + B, N, D = queries.shape + assert D == 3 + queries = queries.clone() + queries[:, :, 1:] *= queries.new_tensor( + [ + (self.interp_shape[1] - 1) / (W - 1), + (self.interp_shape[0] - 1) / (H - 1), + ] + ) + elif grid_size > 0: + grid_pts = get_points_on_a_grid( + grid_size, self.interp_shape, device=video.device + ) + if segm_mask is not None: + segm_mask = F.interpolate( + segm_mask, tuple(self.interp_shape), mode="nearest" + ) + point_mask = segm_mask[0, 0][ + (grid_pts[0, :, 1]).round().long().cpu(), + (grid_pts[0, :, 0]).round().long().cpu(), + ].bool() + grid_pts = grid_pts[:, point_mask] + + queries = torch.cat( + [torch.ones_like(grid_pts[:, :, :1]) * grid_query_frame, grid_pts], + dim=2, + ).repeat(B, 1, 1) + + if add_support_grid: + grid_pts = get_points_on_a_grid( + self.support_grid_size, self.interp_shape, device=video.device + ) + grid_pts = torch.cat( + [torch.zeros_like(grid_pts[:, :, :1]), grid_pts], dim=2 + ) + grid_pts = grid_pts.repeat(B, 1, 1) + queries = torch.cat([queries, grid_pts], dim=1) + + tracks, visibilities, *_ = self.model.forward( + video=video, queries=queries, iters=6 + ) + + if backward_tracking: + tracks, visibilities = self._compute_backward_tracks( + video, queries, tracks, visibilities + ) + if add_support_grid: + queries[:, -self.support_grid_size**2 :, 0] = T - 1 + if add_support_grid: + tracks = tracks[:, :, : -self.support_grid_size**2] + visibilities = visibilities[:, :, : -self.support_grid_size**2] + thr = 0.9 + visibilities = visibilities > thr + + # correct query-point predictions + # see https://github.com/facebookresearch/co-tracker/issues/28 + + # TODO: batchify + for i in range(len(queries)): + queries_t = queries[i, : tracks.size(2), 0].to(torch.int64) + arange = torch.arange(0, len(queries_t)) + + # overwrite the predictions with the query points + tracks[i, queries_t, arange] = queries[i, : tracks.size(2), 1:] + + # correct visibilities, the query points should be visible + visibilities[i, queries_t, arange] = True + + tracks *= tracks.new_tensor( + [(W - 1) / (self.interp_shape[1] - 1), (H - 1) / (self.interp_shape[0] - 1)] + ) + return tracks, visibilities + + def _compute_backward_tracks(self, video, queries, tracks, visibilities): + inv_video = video.flip(1).clone() + inv_queries = queries.clone() + inv_queries[:, :, 0] = inv_video.shape[1] - inv_queries[:, :, 0] - 1 + + inv_tracks, inv_visibilities, *_ = self.model( + video=inv_video, queries=inv_queries, iters=6 + ) + + inv_tracks = inv_tracks.flip(1) + inv_visibilities = inv_visibilities.flip(1) + arange = torch.arange(video.shape[1], device=queries.device)[None, :, None] + + mask = (arange < queries[:, None, :, 0]).unsqueeze(-1).repeat(1, 1, 1, 2) + + tracks[mask] = inv_tracks[mask] + visibilities[mask[:, :, :, 0]] = inv_visibilities[mask[:, :, :, 0]] + return tracks, visibilities + + +class CoTrackerOnlinePredictor(torch.nn.Module): + def __init__( + self, + checkpoint="./checkpoints/scaled_online.pth", + offline=False, + v2=False, + window_len=16, + ): + super().__init__() + self.v2 = v2 + self.support_grid_size = 6 + model = build_cotracker(checkpoint, v2=v2, offline=False, window_len=window_len) + self.interp_shape = model.model_resolution + self.step = model.window_len // 2 + self.model = model + self.model.eval() + + @torch.no_grad() + def forward( + self, + video_chunk, + is_first_step: bool = False, + queries: torch.Tensor = None, + grid_size: int = 5, + grid_query_frame: int = 0, + add_support_grid=False, + ): + B, T, C, H, W = video_chunk.shape + # Initialize online video processing and save queried points + # This needs to be done before processing *each new video* + if is_first_step: + self.model.init_video_online_processing() + if queries is not None: + B, N, D = queries.shape + self.N = N + assert D == 3 + queries = queries.clone() + queries[:, :, 1:] *= queries.new_tensor( + [ + (self.interp_shape[1] - 1) / (W - 1), + (self.interp_shape[0] - 1) / (H - 1), + ] + ) + if add_support_grid: + grid_pts = get_points_on_a_grid( + self.support_grid_size, self.interp_shape, device=video_chunk.device + ) + grid_pts = torch.cat( + [torch.zeros_like(grid_pts[:, :, :1]), grid_pts], dim=2 + ) + queries = torch.cat([queries, grid_pts], dim=1) + elif grid_size > 0: + grid_pts = get_points_on_a_grid( + grid_size, self.interp_shape, device=video_chunk.device + ) + self.N = grid_size**2 + queries = torch.cat( + [torch.ones_like(grid_pts[:, :, :1]) * grid_query_frame, grid_pts], + dim=2, + ) + + self.queries = queries + return (None, None) + + video_chunk = video_chunk.reshape(B * T, C, H, W) + video_chunk = F.interpolate( + video_chunk, tuple(self.interp_shape), mode="bilinear", align_corners=True + ) + video_chunk = video_chunk.reshape( + B, T, 3, self.interp_shape[0], self.interp_shape[1] + ) + if self.v2: + tracks, visibilities, __ = self.model( + video=video_chunk, queries=self.queries, iters=6, is_online=True + ) + else: + tracks, visibilities, confidence, __ = self.model( + video=video_chunk, queries=self.queries, iters=6, is_online=True + ) + if add_support_grid: + tracks = tracks[:,:,:self.N] + visibilities = visibilities[:,:,:self.N] + if not self.v2: + confidence = confidence[:,:,:self.N] + + if not self.v2: + visibilities = visibilities * confidence + thr = 0.6 + return ( + tracks + * tracks.new_tensor( + [ + (W - 1) / (self.interp_shape[1] - 1), + (H - 1) / (self.interp_shape[0] - 1), + ] + ), + visibilities > thr, + ) diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/__init__.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5277f46157403e47fd830fc519144b97ef69d4ae --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/train_utils.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/train_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..5463295cfff6bee5508a1238beed117ef0072a06 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/train_utils.py @@ -0,0 +1,255 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import sys +import torch +import signal +import socket +from torch.utils.data import ConcatDataset +from cotracker.datasets.utils import collate_fn, collate_fn_train +from torch.utils.tensorboard import SummaryWriter +from cotracker.datasets.dr_dataset import DynamicReplicaDataset +from cotracker.models.evaluation_predictor import EvaluationPredictor + + +# define the handler function +# for training on a slurm cluster +def sig_handler(signum, frame): + print("caught signal", signum) + print(socket.gethostname(), "USR1 signal caught.") + # do other stuff to cleanup here + print("requeuing job " + os.environ["SLURM_JOB_ID"]) + os.system("scontrol requeue " + os.environ["SLURM_JOB_ID"]) + sys.exit(-1) + + +def term_handler(signum, frame): + print("bypassing sigterm", flush=True) + + +def get_eval_dataloader(dataset_root, ds_name): + from cotracker.datasets.tap_vid_datasets import TapVidDataset + + collate_fn_local = collate_fn + if ds_name == "dynamic_replica": + from cotracker.datasets.dr_dataset import DynamicReplicaDataset + + eval_dataset = DynamicReplicaDataset( + root=os.path.join(dataset_root, "dynamic_replica"), + sample_len=300, + only_first_n_samples=1, + rgbd_input=False, + ) + elif ds_name == "tapvid_davis_first": + data_root = os.path.join(dataset_root, "tapvid/tapvid_davis/tapvid_davis.pkl") + eval_dataset = TapVidDataset( + dataset_type="davis", data_root=data_root, queried_first=True + ) + elif ds_name == "tapvid_davis_strided": + data_root = os.path.join(dataset_root, "tapvid/tapvid_davis/tapvid_davis.pkl") + eval_dataset = TapVidDataset( + dataset_type="davis", data_root=data_root, queried_first=False + ) + elif ds_name == "tapvid_kinetics_first": + eval_dataset = TapVidDataset( + dataset_type="kinetics", + data_root=os.path.join(dataset_root, "tapvid", "tapvid_kinetics"), + ) + elif ds_name == "tapvid_stacking": + eval_dataset = TapVidDataset( + dataset_type="stacking", + data_root=os.path.join( + dataset_root, "tapvid", "tapvid_rgb_stacking", "tapvid_rgb_stacking.pkl" + ), + ) + elif ds_name == "tapvid_robotap": + eval_dataset = TapVidDataset( + dataset_type="robotap", + data_root=os.path.join(dataset_root, "tapvid", "tapvid_robotap"), + ) + elif ds_name == "kubric": + from cotracker.datasets.kubric_movif_dataset import KubricMovifDataset + + eval_dataset = KubricMovifDataset( + data_root=os.path.join( + args.dataset_root, "kubric/kubric_movi_f_120_frames_dense/movi_f" + ), + traj_per_sample=1024, + use_augs=False, + split="valid", + sample_vis_1st_frame=True, + ) + collate_fn_local = collate_fn_train + eval_dataloader_dr = torch.utils.data.DataLoader( + eval_dataset, + batch_size=1, + shuffle=False, + num_workers=1, + collate_fn=collate_fn_local, + ) + return eval_dataloader_dr + + +def get_train_dataset(args): + dataset = None + if "kubric" in args.train_datasets: + from cotracker.datasets import kubric_movif_dataset + + kubric = kubric_movif_dataset.KubricMovifDataset( + data_root=os.path.join( + args.dataset_root, "kubric/kubric_movi_f_120_frames_dense/movi_f" + ), + crop_size=args.crop_size, + seq_len=args.sequence_len, + traj_per_sample=args.traj_per_sample, + sample_vis_last_frame=args.query_sampling_method is not None + and ("random" in args.query_sampling_method), + use_augs=not args.dont_use_augs, + random_seq_len=args.random_seq_len, + random_frame_rate=args.random_frame_rate, + random_number_traj=args.random_number_traj, + ) + + if dataset is None: + dataset = ConcatDataset(4 * [kubric]) + else: + dataset = ConcatDataset(4 * [kubric] + [dataset]) + print("add kubric to train", len(dataset)) + + if "dr" in args.train_datasets: + dr = DynamicReplicaDataset( + root=os.path.join(args.dataset_root, "dynamic_replica"), + sample_len=args.sequence_len, + split="train", + traj_per_sample=args.traj_per_sample, + crop_size=args.crop_size, + ) + if dataset is None: + dataset = dr + else: + dataset = ConcatDataset([dr] + [dataset]) + + return dataset + + +def run_test_eval(evaluator, model, dataloaders, writer, step, query_random=False): + model.eval() + for ds_name, dataloader in dataloaders: + visualize_every = 1 + grid_size = 5 + num_uniformly_sampled_pts = 0 + if ds_name == "dynamic_replica": + visualize_every = 8 + grid_size = 0 + elif ds_name == "kubric": + visualize_every = 5 + grid_size = 0 + elif "davis" in ds_name or "tapvid_stacking" in ds_name: + visualize_every = 5 + elif "robotap" in ds_name: + visualize_every = 20 + elif "kinetics" in ds_name: + visualize_every = 50 + if query_random: + grid_size = 0 + num_uniformly_sampled_pts = 100 + + predictor = EvaluationPredictor( + model.module.module, + grid_size=grid_size, + local_grid_size=0, + single_point=False, + num_uniformly_sampled_pts=num_uniformly_sampled_pts, + n_iters=6, + ) + + if torch.cuda.is_available(): + predictor.model = predictor.model.cuda() + + metrics = evaluator.evaluate_sequence( + model=predictor, + test_dataloader=dataloader, + dataset_name=ds_name, + train_mode=True, + writer=writer, + step=step, + visualize_every=visualize_every, + ) + + if ds_name == "dynamic_replica" or ds_name == "kubric": + metrics = { + f"{ds_name}_avg_{k}": v + for k, v in metrics["avg"].items() + if not ("1" in k or "2" in k or "4" in k or "8" in k) + } + + if "tapvid" in ds_name: + metrics = { + f"{ds_name}_avg_OA": metrics["avg"]["occlusion_accuracy"], + f"{ds_name}_avg_delta": metrics["avg"]["average_pts_within_thresh"], + f"{ds_name}_avg_Jaccard": metrics["avg"]["average_jaccard"], + } + + writer.add_scalars(f"Eval_{ds_name}", metrics, step) + + +class Logger: + SUM_FREQ = 100 + + def __init__(self, model, scheduler, ckpt_path): + self.model = model + self.scheduler = scheduler + self.ckpt_path = ckpt_path + self.total_steps = 0 + self.running_loss = {} + self.writer = SummaryWriter(log_dir=os.path.join(ckpt_path, "runs")) + + def _print_training_status(self): + metrics_data = [ + self.running_loss[k] / Logger.SUM_FREQ + for k in sorted(self.running_loss.keys()) + ] + training_str = "[{:6d}] ".format(self.total_steps + 1) + metrics_str = ("{:10.4f}, " * len(metrics_data)).format(*metrics_data) + + # print the training status + logging.info( + f"Training Metrics ({self.total_steps}): {training_str + metrics_str}" + ) + + if self.writer is None: + self.writer = SummaryWriter(log_dir=os.path.join(self.ckpt_path, "runs")) + + for k in self.running_loss: + self.writer.add_scalar( + k, self.running_loss[k] / Logger.SUM_FREQ, self.total_steps + ) + self.running_loss[k] = 0.0 + + def push(self, metrics, task): + self.total_steps += 1 + + for key in metrics: + task_key = str(key) + "_" + task + if task_key not in self.running_loss: + self.running_loss[task_key] = 0.0 + + self.running_loss[task_key] += metrics[key] + + if self.total_steps % Logger.SUM_FREQ == Logger.SUM_FREQ - 1: + self._print_training_status() + self.running_loss = {} + + def write_dict(self, results): + if self.writer is None: + self.writer = SummaryWriter(log_dir=os.path.join(self.ckpt_path, "runs")) + + for key in results: + self.writer.add_scalar(key, results[key], self.total_steps) + + def close(self): + self.writer.close() diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/visualizer.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/visualizer.py new file mode 100644 index 0000000000000000000000000000000000000000..fbe008ab57f410fcf1f7ba1ff66dfd2f77d777f0 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/utils/visualizer.py @@ -0,0 +1,363 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import numpy as np +import imageio +import torch + +from matplotlib import cm +import torch.nn.functional as F +import torchvision.transforms as transforms +import matplotlib.pyplot as plt +from PIL import Image, ImageDraw + + +def read_video_from_path(path): + try: + reader = imageio.get_reader(path) + except Exception as e: + print("Error opening video file: ", e) + return None + frames = [] + for i, im in enumerate(reader): + frames.append(np.array(im)) + return np.stack(frames) + + +def draw_circle(rgb, coord, radius, color=(255, 0, 0), visible=True, color_alpha=None): + # Create a draw object + draw = ImageDraw.Draw(rgb) + # Calculate the bounding box of the circle + left_up_point = (coord[0] - radius, coord[1] - radius) + right_down_point = (coord[0] + radius, coord[1] + radius) + # Draw the circle + color = tuple(list(color) + [color_alpha if color_alpha is not None else 255]) + + draw.ellipse( + [left_up_point, right_down_point], + fill=tuple(color) if visible else None, + outline=tuple(color), + ) + return rgb + + +def draw_line(rgb, coord_y, coord_x, color, linewidth): + draw = ImageDraw.Draw(rgb) + draw.line( + (coord_y[0], coord_y[1], coord_x[0], coord_x[1]), + fill=tuple(color), + width=linewidth, + ) + return rgb + + +def add_weighted(rgb, alpha, original, beta, gamma): + return (rgb * alpha + original * beta + gamma).astype("uint8") + + +class Visualizer: + def __init__( + self, + save_dir: str = "./results", + grayscale: bool = False, + pad_value: int = 0, + fps: int = 10, + mode: str = "rainbow", # 'cool', 'optical_flow' + linewidth: int = 2, + show_first_frame: int = 10, + tracks_leave_trace: int = 0, # -1 for infinite + ): + self.mode = mode + self.save_dir = save_dir + if mode == "rainbow": + self.color_map = cm.get_cmap("gist_rainbow") + elif mode == "cool": + self.color_map = cm.get_cmap(mode) + self.show_first_frame = show_first_frame + self.grayscale = grayscale + self.tracks_leave_trace = tracks_leave_trace + self.pad_value = pad_value + self.linewidth = linewidth + self.fps = fps + + def visualize( + self, + video: torch.Tensor, # (B,T,C,H,W) + tracks: torch.Tensor, # (B,T,N,2) + visibility: torch.Tensor = None, # (B, T, N, 1) bool + gt_tracks: torch.Tensor = None, # (B,T,N,2) + segm_mask: torch.Tensor = None, # (B,1,H,W) + filename: str = "video", + writer=None, # tensorboard Summary Writer, used for visualization during training + step: int = 0, + query_frame=0, + save_video: bool = True, + compensate_for_camera_motion: bool = False, + opacity: float = 1.0, + ): + if compensate_for_camera_motion: + assert segm_mask is not None + if segm_mask is not None: + coords = tracks[0, query_frame].round().long() + segm_mask = segm_mask[0, query_frame][coords[:, 1], coords[:, 0]].long() + + video = F.pad( + video, + (self.pad_value, self.pad_value, self.pad_value, self.pad_value), + "constant", + 255, + ) + color_alpha = int(opacity * 255) + tracks = tracks + self.pad_value + + if self.grayscale: + transform = transforms.Grayscale() + video = transform(video) + video = video.repeat(1, 1, 3, 1, 1) + + res_video = self.draw_tracks_on_video( + video=video, + tracks=tracks, + visibility=visibility, + segm_mask=segm_mask, + gt_tracks=gt_tracks, + query_frame=query_frame, + compensate_for_camera_motion=compensate_for_camera_motion, + color_alpha=color_alpha, + ) + if save_video: + self.save_video(res_video, filename=filename, writer=writer, step=step) + return res_video + + def save_video(self, video, filename, writer=None, step=0): + if writer is not None: + writer.add_video( + filename, + video.to(torch.uint8), + global_step=step, + fps=self.fps, + ) + else: + os.makedirs(self.save_dir, exist_ok=True) + wide_list = list(video.unbind(1)) + wide_list = [wide[0].permute(1, 2, 0).cpu().numpy() for wide in wide_list] + + # Prepare the video file path + save_path = os.path.join(self.save_dir, f"{filename}.mp4") + + # Create a writer object + video_writer = imageio.get_writer(save_path, fps=self.fps) + + # Write frames to the video file + for frame in wide_list[2:-1]: + video_writer.append_data(frame) + + video_writer.close() + + print(f"Video saved to {save_path}") + + def draw_tracks_on_video( + self, + video: torch.Tensor, + tracks: torch.Tensor, + visibility: torch.Tensor = None, + segm_mask: torch.Tensor = None, + gt_tracks=None, + query_frame=0, + compensate_for_camera_motion=False, + color_alpha: int = 255, + ): + B, T, C, H, W = video.shape + _, _, N, D = tracks.shape + + assert D == 2 + assert C == 3 + video = video[0].permute(0, 2, 3, 1).byte().detach().cpu().numpy() # S, H, W, C + tracks = tracks[0].long().detach().cpu().numpy() # S, N, 2 + if gt_tracks is not None: + gt_tracks = gt_tracks[0].detach().cpu().numpy() + + res_video = [] + + # process input video + for rgb in video: + res_video.append(rgb.copy()) + vector_colors = np.zeros((T, N, 3)) + + if self.mode == "optical_flow": + import flow_vis + + vector_colors = flow_vis.flow_to_color(tracks - tracks[query_frame][None]) + elif segm_mask is None: + if self.mode == "rainbow": + y_min, y_max = ( + tracks[query_frame, :, 1].min(), + tracks[query_frame, :, 1].max(), + ) + norm = plt.Normalize(y_min, y_max) + for n in range(N): + if isinstance(query_frame, torch.Tensor): + query_frame_ = query_frame[n] + else: + query_frame_ = query_frame + color = self.color_map(norm(tracks[query_frame_, n, 1])) + color = np.array(color[:3])[None] * 255 + vector_colors[:, n] = np.repeat(color, T, axis=0) + else: + # color changes with time + for t in range(T): + color = np.array(self.color_map(t / T)[:3])[None] * 255 + vector_colors[t] = np.repeat(color, N, axis=0) + else: + if self.mode == "rainbow": + vector_colors[:, segm_mask <= 0, :] = 255 + + y_min, y_max = ( + tracks[0, segm_mask > 0, 1].min(), + tracks[0, segm_mask > 0, 1].max(), + ) + norm = plt.Normalize(y_min, y_max) + for n in range(N): + if segm_mask[n] > 0: + color = self.color_map(norm(tracks[0, n, 1])) + color = np.array(color[:3])[None] * 255 + vector_colors[:, n] = np.repeat(color, T, axis=0) + + else: + # color changes with segm class + segm_mask = segm_mask.cpu() + color = np.zeros((segm_mask.shape[0], 3), dtype=np.float32) + color[segm_mask > 0] = np.array(self.color_map(1.0)[:3]) * 255.0 + color[segm_mask <= 0] = np.array(self.color_map(0.0)[:3]) * 255.0 + vector_colors = np.repeat(color[None], T, axis=0) + + # draw tracks + if self.tracks_leave_trace != 0: + for t in range(query_frame + 1, T): + first_ind = ( + max(0, t - self.tracks_leave_trace) + if self.tracks_leave_trace >= 0 + else 0 + ) + curr_tracks = tracks[first_ind : t + 1] + curr_colors = vector_colors[first_ind : t + 1] + if compensate_for_camera_motion: + diff = ( + tracks[first_ind : t + 1, segm_mask <= 0] + - tracks[t : t + 1, segm_mask <= 0] + ).mean(1)[:, None] + + curr_tracks = curr_tracks - diff + curr_tracks = curr_tracks[:, segm_mask > 0] + curr_colors = curr_colors[:, segm_mask > 0] + + res_video[t] = self._draw_pred_tracks( + res_video[t], + curr_tracks, + curr_colors, + ) + if gt_tracks is not None: + res_video[t] = self._draw_gt_tracks( + res_video[t], gt_tracks[first_ind : t + 1] + ) + + # draw points + for t in range(T): + img = Image.fromarray(np.uint8(res_video[t])) + for i in range(N): + coord = (tracks[t, i, 0], tracks[t, i, 1]) + visibile = True + if visibility is not None: + visibile = visibility[0, t, i] + if coord[0] != 0 and coord[1] != 0: + if not compensate_for_camera_motion or ( + compensate_for_camera_motion and segm_mask[i] > 0 + ): + img = draw_circle( + img, + coord=coord, + radius=int(self.linewidth * 2), + color=vector_colors[t, i].astype(int), + visible=visibile, + color_alpha=color_alpha, + ) + res_video[t] = np.array(img) + + # construct the final rgb sequence + if self.show_first_frame > 0: + res_video = [res_video[0]] * self.show_first_frame + res_video[1:] + return torch.from_numpy(np.stack(res_video)).permute(0, 3, 1, 2)[None].byte() + + def _draw_pred_tracks( + self, + rgb: np.ndarray, # H x W x 3 + tracks: np.ndarray, # T x 2 + vector_colors: np.ndarray, + alpha: float = 0.5, + ): + T, N, _ = tracks.shape + rgb = Image.fromarray(np.uint8(rgb)) + for s in range(T - 1): + vector_color = vector_colors[s] + original = rgb.copy() + alpha = (s / T) ** 2 + for i in range(N): + coord_y = (int(tracks[s, i, 0]), int(tracks[s, i, 1])) + coord_x = (int(tracks[s + 1, i, 0]), int(tracks[s + 1, i, 1])) + if coord_y[0] != 0 and coord_y[1] != 0: + rgb = draw_line( + rgb, + coord_y, + coord_x, + vector_color[i].astype(int), + self.linewidth, + ) + if self.tracks_leave_trace > 0: + rgb = Image.fromarray( + np.uint8( + add_weighted( + np.array(rgb), alpha, np.array(original), 1 - alpha, 0 + ) + ) + ) + rgb = np.array(rgb) + return rgb + + def _draw_gt_tracks( + self, + rgb: np.ndarray, # H x W x 3, + gt_tracks: np.ndarray, # T x 2 + ): + T, N, _ = gt_tracks.shape + color = np.array((211, 0, 0)) + rgb = Image.fromarray(np.uint8(rgb)) + for t in range(T): + for i in range(N): + gt_tracks = gt_tracks[t][i] + # draw a red cross + if gt_tracks[0] > 0 and gt_tracks[1] > 0: + length = self.linewidth * 3 + coord_y = (int(gt_tracks[0]) + length, int(gt_tracks[1]) + length) + coord_x = (int(gt_tracks[0]) - length, int(gt_tracks[1]) - length) + rgb = draw_line( + rgb, + coord_y, + coord_x, + color, + self.linewidth, + ) + coord_y = (int(gt_tracks[0]) - length, int(gt_tracks[1]) + length) + coord_x = (int(gt_tracks[0]) + length, int(gt_tracks[1]) - length) + rgb = draw_line( + rgb, + coord_y, + coord_x, + color, + self.linewidth, + ) + rgb = np.array(rgb) + return rgb diff --git a/torch_hub/facebookresearch_co-tracker_main/cotracker/version.py b/torch_hub/facebookresearch_co-tracker_main/cotracker/version.py new file mode 100644 index 0000000000000000000000000000000000000000..e550747c2a6dd5d6bb2ec61889417359e8125945 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/cotracker/version.py @@ -0,0 +1,8 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + + +__version__ = "3.0.0" diff --git a/torch_hub/facebookresearch_co-tracker_main/demo.py b/torch_hub/facebookresearch_co-tracker_main/demo.py new file mode 100644 index 0000000000000000000000000000000000000000..702a0322ef4ce8b9b15b80be76a57d05a467a2c6 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/demo.py @@ -0,0 +1,109 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import torch +import argparse +import numpy as np + +from PIL import Image +from cotracker.utils.visualizer import Visualizer, read_video_from_path +from cotracker.predictor import CoTrackerPredictor + +DEFAULT_DEVICE = ( + "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu" +) + +# if DEFAULT_DEVICE == "mps": +# os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1" + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument( + "--video_path", + default="./assets/apple.mp4", + help="path to a video", + ) + parser.add_argument( + "--mask_path", + default="./assets/apple_mask.png", + help="path to a segmentation mask", + ) + parser.add_argument( + "--checkpoint", + # default="./checkpoints/cotracker.pth", + default=None, + help="CoTracker model parameters", + ) + parser.add_argument("--grid_size", type=int, default=10, help="Regular grid size") + parser.add_argument( + "--grid_query_frame", + type=int, + default=0, + help="Compute dense and grid tracks starting from this frame", + ) + parser.add_argument( + "--backward_tracking", + action="store_true", + help="Compute tracks in both directions, not only forward", + ) + parser.add_argument( + "--use_v2_model", + action="store_true", + help="Pass it if you wish to use CoTracker2, CoTracker++ is the default now", + ) + parser.add_argument( + "--offline", + action="store_true", + help="Pass it if you would like to use the offline model, in case of online don't pass it", + ) + + args = parser.parse_args() + + # load the input video frame by frame + video = read_video_from_path(args.video_path) + video = torch.from_numpy(video).permute(0, 3, 1, 2)[None].float() + segm_mask = np.array(Image.open(os.path.join(args.mask_path))) + segm_mask = torch.from_numpy(segm_mask)[None, None] + + if args.checkpoint is not None: + if args.use_v2_model: + model = CoTrackerPredictor(checkpoint=args.checkpoint, v2=args.use_v2_model) + else: + if args.offline: + window_len = 60 + else: + window_len = 16 + model = CoTrackerPredictor( + checkpoint=args.checkpoint, + v2=args.use_v2_model, + offline=args.offline, + window_len=window_len, + ) + else: + model = torch.hub.load("facebookresearch/co-tracker", "cotracker3_offline") + + model = model.to(DEFAULT_DEVICE) + video = video.to(DEFAULT_DEVICE) + + pred_tracks, pred_visibility = model( + video, + grid_size=args.grid_size, + grid_query_frame=args.grid_query_frame, + backward_tracking=args.backward_tracking, + # segm_mask=segm_mask + ) + print("computed") + + # save a video with predicted tracks + seq_name = args.video_path.split("/")[-1] + vis = Visualizer(save_dir="./saved_videos", pad_value=120, linewidth=3) + vis.visualize( + video, + pred_tracks, + pred_visibility, + query_frame=0 if args.backward_tracking else args.grid_query_frame, + ) diff --git a/torch_hub/facebookresearch_co-tracker_main/docs/Makefile b/torch_hub/facebookresearch_co-tracker_main/docs/Makefile new file mode 100644 index 0000000000000000000000000000000000000000..b5d2aef658d769b18a40e127804ffd65b8ba6882 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/docs/Makefile @@ -0,0 +1,13 @@ +SPHINXOPTS ?= +SPHINXBUILD ?= sphinx-build +SOURCEDIR = source +BUILDDIR = _build +O = -a + +help: + @$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) + +.PHONY: help Makefile + +%: Makefile + @$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/docs/source/apis/models.rst b/torch_hub/facebookresearch_co-tracker_main/docs/source/apis/models.rst new file mode 100644 index 0000000000000000000000000000000000000000..0f1f3b9f680eebe19ad3a29bf4e62ad04efabb79 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/docs/source/apis/models.rst @@ -0,0 +1,14 @@ +Models +====== + +CoTracker models: + +.. currentmodule:: cotracker.models + +Model Utils +----------- + +.. automodule:: cotracker.models.core.model_utils + :members: + :undoc-members: + :show-inheritance: \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/docs/source/apis/utils.rst b/torch_hub/facebookresearch_co-tracker_main/docs/source/apis/utils.rst new file mode 100644 index 0000000000000000000000000000000000000000..00130831c916e5107b8301aa912dbf05e2d261cc --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/docs/source/apis/utils.rst @@ -0,0 +1,11 @@ +Utils +===== + +CoTracker utilizes the following utilities: + +.. currentmodule:: cotracker + +.. automodule:: cotracker.utils.visualizer + :members: + :undoc-members: + :show-inheritance: \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/docs/source/conf.py b/torch_hub/facebookresearch_co-tracker_main/docs/source/conf.py new file mode 100644 index 0000000000000000000000000000000000000000..743bcd34604e60cb8cd2cb87bdd83c18101c93ff --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/docs/source/conf.py @@ -0,0 +1,39 @@ +__version__ = None +exec(open("../../cotracker/version.py", "r").read()) + +project = "CoTracker" +copyright = "2023-24, Meta Platforms, Inc. and affiliates" +author = "Meta Platforms" +release = __version__ + +extensions = [ + "sphinx.ext.napoleon", + "sphinx.ext.duration", + "sphinx.ext.doctest", + "sphinx.ext.autodoc", + "sphinx.ext.autosummary", + "sphinx.ext.intersphinx", + "sphinxcontrib.bibtex", +] + +intersphinx_mapping = { + "python": ("https://docs.python.org/3/", None), + "sphinx": ("https://www.sphinx-doc.org/en/master/", None), +} +intersphinx_disabled_domains = ["std"] + +# templates_path = ["_templates"] +html_theme = "alabaster" + +# Ignore >>> when copying code +copybutton_prompt_text = r">>> |\.\.\. " +copybutton_prompt_is_regexp = True + +# -- Options for EPUB output +epub_show_urls = "footnote" + +# typehints +autodoc_typehints = "description" + +# citations +bibtex_bibfiles = ["references.bib"] diff --git a/torch_hub/facebookresearch_co-tracker_main/docs/source/index.rst b/torch_hub/facebookresearch_co-tracker_main/docs/source/index.rst new file mode 100644 index 0000000000000000000000000000000000000000..2f7a6d06ed2df2cff1090f4e3b8b86cafef8c181 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/docs/source/index.rst @@ -0,0 +1,29 @@ +gsplat +=================================== + +.. image:: ../../assets/bmx-bumps.gif + :width: 800 + :alt: Example of cotracker in action + +Overview +-------- + +*CoTracker* is an open-source tracker :cite:p:`karaev2023cotracker`. + +Links +----- + +.. toctree:: + :glob: + :maxdepth: 1 + :caption: Python API + + apis/* + + +Citations +--------- + +.. bibliography:: + :style: unsrt + :filter: docname in docnames diff --git a/torch_hub/facebookresearch_co-tracker_main/docs/source/references.bib b/torch_hub/facebookresearch_co-tracker_main/docs/source/references.bib new file mode 100644 index 0000000000000000000000000000000000000000..27e046e1a416354e47ef00bfdb22a08cbfbd4608 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/docs/source/references.bib @@ -0,0 +1,6 @@ +@article{karaev2023cotracker, + title = {CoTracker: It is Better to Track Together}, + author = {Nikita Karaev and Ignacio Rocco and Benjamin Graham and Natalia Neverova and Andrea Vedaldi and Christian Rupprecht}, + journal = {arXiv:2307.07635}, + year = {2023} +} diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/app.py b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/app.py new file mode 100644 index 0000000000000000000000000000000000000000..5c6779cc1afcfcc7b9c09d4bddf64985625e975e --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/app.py @@ -0,0 +1,614 @@ +# This Gradio demo code is from https://github.com/cvlab-kaist/locotrack/blob/main/demo/demo.py +# We updated it to work with CoTracker3 models. We thank authors of LocoTrack +# for such an amazing Gradio demo. + +import os +import sys +import uuid + +import gradio as gr +import mediapy +import numpy as np +import cv2 +import matplotlib +import torch +import colorsys +import random +from typing import List, Optional, Sequence, Tuple + +import numpy as np + + +# Generate random colormaps for visualizing different points. +def get_colors(num_colors: int) -> List[Tuple[int, int, int]]: + """Gets colormap for points.""" + colors = [] + for i in np.arange(0.0, 360.0, 360.0 / num_colors): + hue = i / 360.0 + lightness = (50 + np.random.rand() * 10) / 100.0 + saturation = (90 + np.random.rand() * 10) / 100.0 + color = colorsys.hls_to_rgb(hue, lightness, saturation) + colors.append( + (int(color[0] * 255), int(color[1] * 255), int(color[2] * 255)) + ) + random.shuffle(colors) + return colors + +def get_points_on_a_grid( + size: int, + extent: Tuple[float, ...], + center: Optional[Tuple[float, ...]] = None, + device: Optional[torch.device] = torch.device("cpu"), +): + r"""Get a grid of points covering a rectangular region + + `get_points_on_a_grid(size, extent)` generates a :attr:`size` by + :attr:`size` grid fo points distributed to cover a rectangular area + specified by `extent`. + + The `extent` is a pair of integer :math:`(H,W)` specifying the height + and width of the rectangle. + + Optionally, the :attr:`center` can be specified as a pair :math:`(c_y,c_x)` + specifying the vertical and horizontal center coordinates. The center + defaults to the middle of the extent. + + Points are distributed uniformly within the rectangle leaving a margin + :math:`m=W/64` from the border. + + It returns a :math:`(1, \text{size} \times \text{size}, 2)` tensor of + points :math:`P_{ij}=(x_i, y_i)` where + + .. math:: + P_{ij} = \left( + c_x + m -\frac{W}{2} + \frac{W - 2m}{\text{size} - 1}\, j,~ + c_y + m -\frac{H}{2} + \frac{H - 2m}{\text{size} - 1}\, i + \right) + + Points are returned in row-major order. + + Args: + size (int): grid size. + extent (tuple): height and with of the grid extent. + center (tuple, optional): grid center. + device (str, optional): Defaults to `"cpu"`. + + Returns: + Tensor: grid. + """ + if size == 1: + return torch.tensor([extent[1] / 2, extent[0] / 2], device=device)[None, None] + + if center is None: + center = [extent[0] / 2, extent[1] / 2] + + margin = extent[1] / 64 + range_y = (margin - extent[0] / 2 + center[0], extent[0] / 2 + center[0] - margin) + range_x = (margin - extent[1] / 2 + center[1], extent[1] / 2 + center[1] - margin) + grid_y, grid_x = torch.meshgrid( + torch.linspace(*range_y, size, device=device), + torch.linspace(*range_x, size, device=device), + indexing="ij", + ) + return torch.stack([grid_x, grid_y], dim=-1).reshape(1, -1, 2) + +def paint_point_track( + frames: np.ndarray, + point_tracks: np.ndarray, + visibles: np.ndarray, + colormap: Optional[List[Tuple[int, int, int]]] = None, +) -> np.ndarray: + """Converts a sequence of points to color code video. + + Args: + frames: [num_frames, height, width, 3], np.uint8, [0, 255] + point_tracks: [num_points, num_frames, 2], np.float32, [0, width / height] + visibles: [num_points, num_frames], bool + colormap: colormap for points, each point has a different RGB color. + + Returns: + video: [num_frames, height, width, 3], np.uint8, [0, 255] + """ + num_points, num_frames = point_tracks.shape[0:2] + if colormap is None: + colormap = get_colors(num_colors=num_points) + height, width = frames.shape[1:3] + dot_size_as_fraction_of_min_edge = 0.015 + radius = int(round(min(height, width) * dot_size_as_fraction_of_min_edge)) + diam = radius * 2 + 1 + quadratic_y = np.square(np.arange(diam)[:, np.newaxis] - radius - 1) + quadratic_x = np.square(np.arange(diam)[np.newaxis, :] - radius - 1) + icon = (quadratic_y + quadratic_x) - (radius**2) / 2.0 + sharpness = 0.15 + icon = np.clip(icon / (radius * 2 * sharpness), 0, 1) + icon = 1 - icon[:, :, np.newaxis] + icon1 = np.pad(icon, [(0, 1), (0, 1), (0, 0)]) + icon2 = np.pad(icon, [(1, 0), (0, 1), (0, 0)]) + icon3 = np.pad(icon, [(0, 1), (1, 0), (0, 0)]) + icon4 = np.pad(icon, [(1, 0), (1, 0), (0, 0)]) + + video = frames.copy() + for t in range(num_frames): + # Pad so that points that extend outside the image frame don't crash us + image = np.pad( + video[t], + [ + (radius + 1, radius + 1), + (radius + 1, radius + 1), + (0, 0), + ], + ) + for i in range(num_points): + # The icon is centered at the center of a pixel, but the input coordinates + # are raster coordinates. Therefore, to render a point at (1,1) (which + # lies on the corner between four pixels), we need 1/4 of the icon placed + # centered on the 0'th row, 0'th column, etc. We need to subtract + # 0.5 to make the fractional position come out right. + x, y = point_tracks[i, t, :] + 0.5 + x = min(max(x, 0.0), width) + y = min(max(y, 0.0), height) + + if visibles[i, t]: + x1, y1 = np.floor(x).astype(np.int32), np.floor(y).astype(np.int32) + x2, y2 = x1 + 1, y1 + 1 + + # bilinear interpolation + patch = ( + icon1 * (x2 - x) * (y2 - y) + + icon2 * (x2 - x) * (y - y1) + + icon3 * (x - x1) * (y2 - y) + + icon4 * (x - x1) * (y - y1) + ) + x_ub = x1 + 2 * radius + 2 + y_ub = y1 + 2 * radius + 2 + image[y1:y_ub, x1:x_ub, :] = (1 - patch) * image[ + y1:y_ub, x1:x_ub, : + ] + patch * np.array(colormap[i])[np.newaxis, np.newaxis, :] + + # Remove the pad + video[t] = image[ + radius + 1 : -radius - 1, radius + 1 : -radius - 1 + ].astype(np.uint8) + return video + + +PREVIEW_WIDTH = 768 # Width of the preview video +VIDEO_INPUT_RESO = (384, 512) # Resolution of the input video +POINT_SIZE = 4 # Size of the query point in the preview video +FRAME_LIMIT = 300 # Limit the number of frames to process + + +def get_point(frame_num, video_queried_preview, query_points, query_points_color, query_count, evt: gr.SelectData): + print(f"You selected {(evt.index[0], evt.index[1], frame_num)}") + + current_frame = video_queried_preview[int(frame_num)] + + # Get the mouse click + query_points[int(frame_num)].append((evt.index[0], evt.index[1], frame_num)) + + # Choose the color for the point from matplotlib colormap + color = matplotlib.colormaps.get_cmap("gist_rainbow")(query_count % 20 / 20) + color = (int(color[0] * 255), int(color[1] * 255), int(color[2] * 255)) + # print(f"Color: {color}") + query_points_color[int(frame_num)].append(color) + + # Draw the point on the frame + x, y = evt.index + current_frame_draw = cv2.circle(current_frame, (x, y), POINT_SIZE, color, -1) + + # Update the frame + video_queried_preview[int(frame_num)] = current_frame_draw + + # Update the query count + query_count += 1 + return ( + current_frame_draw, # Updated frame for preview + video_queried_preview, # Updated preview video + query_points, # Updated query points + query_points_color, # Updated query points color + query_count # Updated query count + ) + + +def undo_point(frame_num, video_preview, video_queried_preview, query_points, query_points_color, query_count): + if len(query_points[int(frame_num)]) == 0: + return ( + video_queried_preview[int(frame_num)], + video_queried_preview, + query_points, + query_points_color, + query_count + ) + + # Get the last point + query_points[int(frame_num)].pop(-1) + query_points_color[int(frame_num)].pop(-1) + + # Redraw the frame + current_frame_draw = video_preview[int(frame_num)].copy() + for point, color in zip(query_points[int(frame_num)], query_points_color[int(frame_num)]): + x, y, _ = point + current_frame_draw = cv2.circle(current_frame_draw, (x, y), POINT_SIZE, color, -1) + + # Update the query count + query_count -= 1 + + # Update the frame + video_queried_preview[int(frame_num)] = current_frame_draw + return ( + current_frame_draw, # Updated frame for preview + video_queried_preview, # Updated preview video + query_points, # Updated query points + query_points_color, # Updated query points color + query_count # Updated query count + ) + + +def clear_frame_fn(frame_num, video_preview, video_queried_preview, query_points, query_points_color, query_count): + query_count -= len(query_points[int(frame_num)]) + + query_points[int(frame_num)] = [] + query_points_color[int(frame_num)] = [] + + video_queried_preview[int(frame_num)] = video_preview[int(frame_num)].copy() + + return ( + video_preview[int(frame_num)], # Set the preview frame to the original frame + video_queried_preview, + query_points, # Cleared query points + query_points_color, # Cleared query points color + query_count # New query count + ) + + + +def clear_all_fn(frame_num, video_preview): + return ( + video_preview[int(frame_num)], + video_preview.copy(), + [[] for _ in range(len(video_preview))], + [[] for _ in range(len(video_preview))], + 0 + ) + + +def choose_frame(frame_num, video_preview_array): + return video_preview_array[int(frame_num)] + + +def preprocess_video_input(video_path): + video_arr = mediapy.read_video(video_path) + video_fps = video_arr.metadata.fps + num_frames = video_arr.shape[0] + if num_frames > FRAME_LIMIT: + gr.Warning(f"The video is too long. Only the first {FRAME_LIMIT} frames will be used.", duration=5) + video_arr = video_arr[:FRAME_LIMIT] + num_frames = FRAME_LIMIT + + # Resize to preview size for faster processing, width = PREVIEW_WIDTH + height, width = video_arr.shape[1:3] + new_height, new_width = int(PREVIEW_WIDTH * height / width), PREVIEW_WIDTH + + preview_video = mediapy.resize_video(video_arr, (new_height, new_width)) + input_video = mediapy.resize_video(video_arr, VIDEO_INPUT_RESO) + + preview_video = np.array(preview_video) + input_video = np.array(input_video) + + interactive = True + + return ( + video_arr, # Original video + preview_video, # Original preview video, resized for faster processing + preview_video.copy(), # Copy of preview video for visualization + input_video, # Resized video input for model + # None, # video_feature, # Extracted feature + video_fps, # Set the video FPS + gr.update(open=False), # Close the video input drawer + # tracking_mode, # Set the tracking mode + preview_video[0], # Set the preview frame to the first frame + gr.update(minimum=0, maximum=num_frames - 1, value=0, interactive=interactive), # Set slider interactive + [[] for _ in range(num_frames)], # Set query_points to empty + [[] for _ in range(num_frames)], # Set query_points_color to empty + [[] for _ in range(num_frames)], + 0, # Set query count to 0 + gr.update(interactive=interactive), # Make the buttons interactive + gr.update(interactive=interactive), + gr.update(interactive=interactive), + gr.update(interactive=True), + ) + + +def track( + video_preview, + video_input, + video_fps, + query_points, + query_points_color, + query_count, +): + tracking_mode = 'selected' + if query_count == 0: + tracking_mode='grid' + + device = "cuda" if torch.cuda.is_available() else "cpu" + dtype = torch.float if device == "cuda" else torch.float + + # Convert query points to tensor, normalize to input resolution + if tracking_mode!='grid': + query_points_tensor = [] + for frame_points in query_points: + query_points_tensor.extend(frame_points) + + query_points_tensor = torch.tensor(query_points_tensor).float() + query_points_tensor *= torch.tensor([ + VIDEO_INPUT_RESO[1], VIDEO_INPUT_RESO[0], 1 + ]) / torch.tensor([ + [video_preview.shape[2], video_preview.shape[1], 1] + ]) + query_points_tensor = query_points_tensor[None].flip(-1).to(device, dtype) # xyt -> tyx + query_points_tensor = query_points_tensor[:, :, [0, 2, 1]] # tyx -> txy + + video_input = torch.tensor(video_input).unsqueeze(0).to(device, dtype) + + model = torch.hub.load("facebookresearch/co-tracker", "cotracker3_online") + model = model.to(device) + + video_input = video_input.permute(0, 1, 4, 2, 3) + if tracking_mode=='grid': + xy = get_points_on_a_grid(15, video_input.shape[3:], device=device) + queries = torch.cat([torch.zeros_like(xy[:, :, :1]), xy], dim=2).to(device) # + add_support_grid=False + cmap = matplotlib.colormaps.get_cmap("gist_rainbow") + query_points_color = [[]] + query_count = queries.shape[1] + for i in range(query_count): + # Choose the color for the point from matplotlib colormap + color = cmap(i / float(query_count)) + color = (int(color[0] * 255), int(color[1] * 255), int(color[2] * 255)) + query_points_color[0].append(color) + + else: + queries = query_points_tensor + add_support_grid=True + + model(video_chunk=video_input, is_first_step=True, grid_size=0, queries=queries, add_support_grid=add_support_grid) + # + for ind in range(0, video_input.shape[1] - model.step, model.step): + pred_tracks, pred_visibility = model( + video_chunk=video_input[:, ind : ind + model.step * 2], + grid_size=0, + queries=queries, + add_support_grid=add_support_grid + ) # B T N 2, B T N 1 + tracks = (pred_tracks * torch.tensor([video_preview.shape[2], video_preview.shape[1]]).to(device) / torch.tensor([VIDEO_INPUT_RESO[1], VIDEO_INPUT_RESO[0]]).to(device))[0].permute(1, 0, 2).cpu().numpy() + pred_occ = pred_visibility[0].permute(1, 0).cpu().numpy() + + # make color array + colors = [] + for frame_colors in query_points_color: + colors.extend(frame_colors) + colors = np.array(colors) + + painted_video = paint_point_track(video_preview,tracks,pred_occ,colors) + + # save video + video_file_name = uuid.uuid4().hex + ".mp4" + video_path = os.path.join(os.path.dirname(__file__), "tmp") + video_file_path = os.path.join(video_path, video_file_name) + os.makedirs(video_path, exist_ok=True) + + mediapy.write_video(video_file_path, painted_video, fps=video_fps) + + return video_file_path + + +with gr.Blocks() as demo: + video = gr.State() + video_queried_preview = gr.State() + video_preview = gr.State() + video_input = gr.State() + video_fps = gr.State(24) + + query_points = gr.State([]) + query_points_color = gr.State([]) + is_tracked_query = gr.State([]) + query_count = gr.State(0) + + gr.Markdown("# 🎨 CoTracker3: Simpler and Better Point Tracking by Pseudo-Labelling Real Videos") + gr.Markdown("
\ +

Welcome to CoTracker! This space demonstrates point (pixel) tracking in videos. \ + The model tracks points on a grid or points selected by you.

\ +

To get started, simply upload your .mp4 video or click on one of the example videos to load them. The shorter the video, the faster the processing. We recommend submitting short videos of length 2-7 seconds.

\ +

After you uploaded a video, please click \"Submit\" and then click \"Track\" for grid tracking or specify points you want to track before clicking. Enjoy the results!

\ +

For more details, check out our GitHub Repo ⭐. We thank the authors of LocoTrack for their interactive demo.

\ +
" + ) + + + gr.Markdown("## First step: upload your video or select an example video, and click submit.") + with gr.Row(): + + + with gr.Accordion("Your video input", open=True) as video_in_drawer: + video_in = gr.Video(label="Video Input", format="mp4") + submit = gr.Button("Submit", scale=0) + + import os + apple = os.path.join(os.path.dirname(__file__), "videos", "apple.mp4") + bear = os.path.join(os.path.dirname(__file__), "videos", "bear.mp4") + paragliding_launch = os.path.join( + os.path.dirname(__file__), "videos", "paragliding-launch.mp4" + ) + paragliding = os.path.join(os.path.dirname(__file__), "videos", "paragliding.mp4") + cat = os.path.join(os.path.dirname(__file__), "videos", "cat.mp4") + pillow = os.path.join(os.path.dirname(__file__), "videos", "pillow.mp4") + teddy = os.path.join(os.path.dirname(__file__), "videos", "teddy.mp4") + backpack = os.path.join(os.path.dirname(__file__), "videos", "backpack.mp4") + + + gr.Examples(examples=[bear, apple, paragliding, paragliding_launch, cat, pillow, teddy, backpack], + inputs = [ + video_in + ], + ) + + + gr.Markdown("## Second step: Simply click \"Track\" to track a grid of points or select query points on the video before clicking") + with gr.Row(): + with gr.Column(): + with gr.Row(): + query_frames = gr.Slider( + minimum=0, maximum=100, value=0, step=1, label="Choose Frame", interactive=False) + with gr.Row(): + undo = gr.Button("Undo", interactive=False) + clear_frame = gr.Button("Clear Frame", interactive=False) + clear_all = gr.Button("Clear All", interactive=False) + + with gr.Row(): + current_frame = gr.Image( + label="Click to add query points", + type="numpy", + interactive=False + ) + + with gr.Row(): + track_button = gr.Button("Track", interactive=False) + + with gr.Column(): + output_video = gr.Video( + label="Output Video", + interactive=False, + autoplay=True, + loop=True, + ) + + + + submit.click( + fn = preprocess_video_input, + inputs = [video_in], + outputs = [ + video, + video_preview, + video_queried_preview, + video_input, + video_fps, + video_in_drawer, + current_frame, + query_frames, + query_points, + query_points_color, + is_tracked_query, + query_count, + undo, + clear_frame, + clear_all, + track_button, + ], + queue = False + ) + + query_frames.change( + fn = choose_frame, + inputs = [query_frames, video_queried_preview], + outputs = [ + current_frame, + ], + queue = False + ) + + current_frame.select( + fn = get_point, + inputs = [ + query_frames, + video_queried_preview, + query_points, + query_points_color, + query_count, + ], + outputs = [ + current_frame, + video_queried_preview, + query_points, + query_points_color, + query_count + ], + queue = False + ) + + undo.click( + fn = undo_point, + inputs = [ + query_frames, + video_preview, + video_queried_preview, + query_points, + query_points_color, + query_count + ], + outputs = [ + current_frame, + video_queried_preview, + query_points, + query_points_color, + query_count + ], + queue = False + ) + + clear_frame.click( + fn = clear_frame_fn, + inputs = [ + query_frames, + video_preview, + video_queried_preview, + query_points, + query_points_color, + query_count + ], + outputs = [ + current_frame, + video_queried_preview, + query_points, + query_points_color, + query_count + ], + queue = False + ) + + clear_all.click( + fn = clear_all_fn, + inputs = [ + query_frames, + video_preview, + ], + outputs = [ + current_frame, + video_queried_preview, + query_points, + query_points_color, + query_count + ], + queue = False + ) + + + track_button.click( + fn = track, + inputs = [ + video_preview, + video_input, + video_fps, + query_points, + query_points_color, + query_count, + ], + outputs = [ + output_video, + ], + queue = True, + ) + + +demo.launch(show_api=False, show_error=True, debug=True, share=True) \ No newline at end of file diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/requirements.txt b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..465be7fd5c439c122d54e5c5ab2c7666d247b396 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/requirements.txt @@ -0,0 +1,12 @@ +torch==1.13.0 +torchvision==0.14.0 +matplotlib==3.7.5 +moviepy==1.0.3 +flow_vis +gradio +imageio[ffmpeg] +opencv-python +imutils==0.5.4 +mediapy==1.2.2 +numpy +git+https://github.com/facebookresearch/co-tracker.git diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/apple.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/apple.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..80f158ffe2cc5625ef5df53e1ff423d2b83e989c --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/apple.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7f48c5cfb1479e1dbc1df2373d5cad4f55c198bbdb379da0ece10087971542a +size 1219872 diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/backpack.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/backpack.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..907644929bcfa0485684069f4090aef73a663f48 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/backpack.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4b5ac6b2285ffb48e3a740e419e38c781df9c963589a5fd894e5b4e13dd6a8b8 +size 1208738 diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/bear.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/bear.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..c2e971804f6878f237f0f8696a742a456a69dcb3 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/bear.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eeffab1780be601b19b2097be81a8c2d4fa2b624ac1028be0a32191d25acca0f +size 893943 diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/cat.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/cat.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..4b1d8aca22ac111ba49c3348eb3fe5352cab5642 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/cat.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a41ab94e89fd0f44f829d545ab6e604437146608357266709e0cc726378d1872 +size 253409 diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding-launch.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding-launch.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..c2f81a429e947455b1e694a864b8d685b1adf561 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding-launch.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4cdc5c670b372cc93763a4cbed99a79080ff98185b9386cea90920d71c14d31f +size 829010 diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..1e557367d8304833b400b6e39533a6acd2a1fe7c --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/paragliding.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d5cc8498d96ed551b2cc0b7086070f0abb82855b845f22770ca0993ce968ceaa +size 437963 diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/pillow.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/pillow.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..4bb67dd98c42a50cb734ef7ecbdea3e13c59af01 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/pillow.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8f05818f586d7b0796fcd4714ea4be489c93701598cadc86ce7973fc24655fee +size 1407147 diff --git a/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/teddy.mp4 b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/teddy.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..4b6617ccc58f7b6bf7c6d23bfe827d57ae84071d --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/gradio_demo/videos/teddy.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:720503173c3b23b1d3d3fefa0e930558f944f0562e6a7b3c23810fc7046b39c7 +size 1337504 diff --git a/torch_hub/facebookresearch_co-tracker_main/hubconf.py b/torch_hub/facebookresearch_co-tracker_main/hubconf.py new file mode 100644 index 0000000000000000000000000000000000000000..e8365279946241259b1b3ce5fdd7454043faf043 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/hubconf.py @@ -0,0 +1,119 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch + +_COTRACKER2_URL = ( + "https://huggingface.co/facebook/cotracker/resolve/main/cotracker2.pth" +) +_COTRACKER2v1_URL = ( + "https://huggingface.co/facebook/cotracker/resolve/main/cotracker2v1.pth" +) +_COTRACKER3_SCALED_OFFLINE_URL = ( + "https://huggingface.co/facebook/cotracker3/resolve/main/scaled_offline.pth" +) +_COTRACKER3_SCALED_ONLINE_URL = ( + "https://huggingface.co/facebook/cotracker3/resolve/main/scaled_online.pth" +) + + +def _make_cotracker_predictor( + *, pretrained: bool = True, online=False, version="3", **kwargs +): + if online: + from cotracker.predictor import CoTrackerOnlinePredictor + + if version == "2": + predictor = CoTrackerOnlinePredictor(checkpoint=None, window_len=8, v2=True) + elif version == "2.1": + predictor = CoTrackerOnlinePredictor( + checkpoint=None, window_len=16, v2=True + ) + elif version == "3": + predictor = CoTrackerOnlinePredictor( + checkpoint=None, window_len=16, v2=False + ) + else: + from cotracker.predictor import CoTrackerPredictor + + if version == "2": + predictor = CoTrackerPredictor(checkpoint=None, window_len=8, v2=True) + elif version == "2.1": + predictor = CoTrackerPredictor(checkpoint=None, window_len=16, v2=True) + elif version == "3": + predictor = CoTrackerPredictor(checkpoint=None, window_len=60, v2=False) + if pretrained: + if version == "2": + state_dict = torch.hub.load_state_dict_from_url( + _COTRACKER2_URL, map_location="cpu" + ) + elif version == "2.1": + state_dict = torch.hub.load_state_dict_from_url( + _COTRACKER2v1_URL, map_location="cpu" + ) + elif version == "3": + if online: + state_dict = torch.hub.load_state_dict_from_url( + _COTRACKER3_SCALED_ONLINE_URL, map_location="cpu" + ) + else: + state_dict = torch.hub.load_state_dict_from_url( + _COTRACKER3_SCALED_OFFLINE_URL, map_location="cpu" + ) + else: + raise Exception("Provided version does not exist") + predictor.model.load_state_dict(state_dict) + return predictor + + +def cotracker2(*, pretrained: bool = True, **kwargs): + """ + CoTracker2 with stride 4 and window length 8. Can track up to 265*265 points jointly. + """ + return _make_cotracker_predictor(pretrained=pretrained, online=False, version="2", **kwargs) + + +def cotracker2_online(*, pretrained: bool = True, **kwargs): + """ + Online CoTracker2 with stride 4 and window length 8. Can track up to 265*265 points jointly. + """ + return _make_cotracker_predictor(pretrained=pretrained, online=True, version="2", **kwargs) + + +def cotracker2v1(*, pretrained: bool = True, **kwargs): + """ + CoTracker2 with stride 4 and window length 16. + """ + return _make_cotracker_predictor( + pretrained=pretrained, online=False, version="2.1", **kwargs + ) + + +def cotracker2v1_online(*, pretrained: bool = True, **kwargs): + """ + Online CoTracker2 with stride 4 and window length 16. + """ + return _make_cotracker_predictor( + pretrained=pretrained, online=True, version="2.1", **kwargs + ) + + +def cotracker3_offline(*, pretrained: bool = True, **kwargs): + """ + Scaled offline CoTracker3 with stride 4 and window length 16. + """ + return _make_cotracker_predictor( + pretrained=pretrained, online=False, version="3", **kwargs + ) + + +def cotracker3_online(*, pretrained: bool = True, **kwargs): + """ + Scaled online CoTracker3 with stride 4 and window length 16. + """ + return _make_cotracker_predictor( + pretrained=pretrained, online=True, version="3", **kwargs + ) diff --git a/torch_hub/facebookresearch_co-tracker_main/launch_training_kubric_offline.sh b/torch_hub/facebookresearch_co-tracker_main/launch_training_kubric_offline.sh new file mode 100644 index 0000000000000000000000000000000000000000..bd38865aa03d6d98218b05da54f7e4dc99ba6d87 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/launch_training_kubric_offline.sh @@ -0,0 +1,27 @@ +#!/bin/bash + +EXP_DIR=$1 +EXP_NAME=$2 +DATE=$3 +DATASET_ROOT=$4 +NUM_STEPS=$5 + + +echo `which python` + +mkdir -p ${EXP_DIR}/${DATE}_${EXP_NAME}/logs/; +mkdir ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3; +find . \( -name "*.sh" -o -name "*.py" \) -type f -exec cp --parents {} ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 \; + +export PYTHONPATH=`(cd ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 && pwd)`:`pwd`:$PYTHONPATH +sbatch --comment=${EXP_NAME} --partition=learn --account=repligen --qos=repligen --time=39:00:00 --gpus-per-node=8 --nodes=4 --ntasks-per-node=8 \ +--job-name=${EXP_NAME} --cpus-per-task=10 --signal=USR1@60 --open-mode=append \ +--output=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.out \ +--error=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.err \ +--wrap="srun --label python ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3/train_on_kubric.py --batch_size 1 \ +--num_steps ${NUM_STEPS} --ckpt_path ${EXP_DIR}/${DATE}_${EXP_NAME} --model_name cotracker_three \ +--save_freq 200 --sequence_len 60 --eval_datasets tapvid_davis_first tapvid_stacking \ +--traj_per_sample 512 --sliding_window_len 60 --train_datasets kubric \ +--save_every_n_epoch 5 --evaluate_every_n_epoch 5 --model_stride 4 --dataset_root ${DATASET_ROOT} --num_nodes 4 \ +--num_virtual_tracks 64 --mixed_precision --offline_model --random_frame_rate --query_sampling_method random \ +--corr_radius 3 --wdecay 0.0005 --random_seq_len --linear_layer_for_vis_conf --validate_at_start --add_huber_loss" diff --git a/torch_hub/facebookresearch_co-tracker_main/launch_training_kubric_online.sh b/torch_hub/facebookresearch_co-tracker_main/launch_training_kubric_online.sh new file mode 100644 index 0000000000000000000000000000000000000000..3e2bab6215b52dcee3971de88a3b546624803297 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/launch_training_kubric_online.sh @@ -0,0 +1,27 @@ +#!/bin/bash + +EXP_DIR=$1 +EXP_NAME=$2 +DATE=$3 +DATASET_ROOT=$4 +NUM_STEPS=$5 + + +echo `which python` + +mkdir -p ${EXP_DIR}/${DATE}_${EXP_NAME}/logs/; +mkdir ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3; +find . \( -name "*.sh" -o -name "*.py" \) -type f -exec cp --parents {} ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 \; + +export PYTHONPATH=`(cd ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 && pwd)`:`pwd`:$PYTHONPATH +sbatch --comment=${EXP_NAME} --partition=learn --account=repligen --qos=repligen --time=39:00:00 --gpus-per-node=8 --nodes=4 --ntasks-per-node=8 \ +--job-name=${EXP_NAME} --cpus-per-task=10 --signal=USR1@60 --open-mode=append \ +--output=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.out \ +--error=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.err \ +--wrap="srun --label python ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3/train_on_kubric.py --batch_size 1 \ +--num_steps ${NUM_STEPS} --ckpt_path ${EXP_DIR}/${DATE}_${EXP_NAME} --model_name cotracker_three \ +--save_freq 200 --sequence_len 64 --eval_datasets tapvid_davis_first tapvid_stacking \ +--traj_per_sample 384 --sliding_window_len 16 --train_datasets kubric \ +--save_every_n_epoch 5 --evaluate_every_n_epoch 5 --model_stride 4 --dataset_root ${DATASET_ROOT} --num_nodes 4 \ +--num_virtual_tracks 64 --mixed_precision \ +--corr_radius 3 --wdecay 0.0005 --linear_layer_for_vis_conf --validate_at_start --add_huber_loss" diff --git a/torch_hub/facebookresearch_co-tracker_main/launch_training_scaling_offline.sh b/torch_hub/facebookresearch_co-tracker_main/launch_training_scaling_offline.sh new file mode 100644 index 0000000000000000000000000000000000000000..01d13ef9d9fda9ce503d2a0cfa4474d0b8d24a5f --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/launch_training_scaling_offline.sh @@ -0,0 +1,29 @@ +#!/bin/bash + +EXP_DIR=$1 +EXP_NAME=$2 +DATE=$3 +DATASET_ROOT=$4 +NUM_STEPS=$5 + + +echo `which python` + +mkdir -p ${EXP_DIR}/${DATE}_${EXP_NAME}/logs/; +mkdir ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3; +find . \( -name "*.sh" -o -name "*.py" \) -type f -exec cp --parents {} ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 \; + +export PYTHONPATH=`(cd ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 && pwd)`:`pwd`:$PYTHONPATH +sbatch --comment=${EXP_NAME} --partition=learn --account=repligen --qos=repligen --time=120:00:00 --gpus-per-node=8 --nodes=4 --ntasks-per-node=8 \ +--job-name=${EXP_NAME} --cpus-per-task=10 --signal=USR1@60 --open-mode=append \ +--output=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.out \ +--error=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.err \ +--wrap="srun --label python ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3/train_on_real_data.py --batch_size 1 \ +--num_steps ${NUM_STEPS} --ckpt_path ${EXP_DIR}/${DATE}_${EXP_NAME} --model_name cotracker_three \ +--save_freq 200 --sequence_len 80 --eval_datasets tapvid_stacking tapvid_davis_first \ +--traj_per_sample 384 --save_every_n_epoch 15 --evaluate_every_n_epoch 15 --model_stride 4 --dataset_root ${DATASET_ROOT} --num_nodes 4 --real_data_splits 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 \ +--num_virtual_tracks 64 --mixed_precision --random_frame_rate \ +--restore_ckpt ./checkpoints/baseline_offline.pth --lr 0.00005 \ +--real_data_filter_sift --validate_at_start --offline_model --limit_samples 10000" +# 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 + diff --git a/torch_hub/facebookresearch_co-tracker_main/launch_training_scaling_online.sh b/torch_hub/facebookresearch_co-tracker_main/launch_training_scaling_online.sh new file mode 100644 index 0000000000000000000000000000000000000000..d31e8d5701ae7596b44a7e1ec82d212c396ffcd1 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/launch_training_scaling_online.sh @@ -0,0 +1,28 @@ +#!/bin/bash + +EXP_DIR=$1 +EXP_NAME=$2 +DATE=$3 +DATASET_ROOT=$4 +NUM_STEPS=$5 + + +echo `which python` + +mkdir -p ${EXP_DIR}/${DATE}_${EXP_NAME}/logs/; +mkdir ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3; +find . \( -name "*.sh" -o -name "*.py" \) -type f -exec cp --parents {} ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 \; + +export PYTHONPATH=`(cd ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3 && pwd)`:`pwd`:$PYTHONPATH +sbatch --comment=${EXP_NAME} --partition=learn --account=repligen --qos=repligen --time=120:00:00 --gpus-per-node=8 --nodes=4 --ntasks-per-node=8 \ +--job-name=${EXP_NAME} --cpus-per-task=10 --signal=USR1@60 --open-mode=append \ +--output=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.out \ +--error=${EXP_DIR}/${DATE}_${EXP_NAME}/logs/%j_%x_%A_%a_%N.err \ +--wrap="srun --label python ${EXP_DIR}/${DATE}_${EXP_NAME}/cotracker3/train_on_real_data.py --batch_size 1 \ +--num_steps ${NUM_STEPS} --ckpt_path ${EXP_DIR}/${DATE}_${EXP_NAME} --model_name cotracker_three \ +--save_freq 200 --sequence_len 64 --eval_datasets tapvid_stacking tapvid_davis_first \ +--traj_per_sample 384 --save_every_n_epoch 15 --evaluate_every_n_epoch 15 --model_stride 4 --dataset_root ${DATASET_ROOT} --num_nodes 4 --real_data_splits 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 \ +--num_virtual_tracks 64 --mixed_precision --random_frame_rate \ +--restore_ckpt ./checkpoints/baseline_online.pth --lr 0.00005 \ +--real_data_filter_sift --validate_at_start --sliding_window_len 16 --limit_samples 10000" +# 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 diff --git a/torch_hub/facebookresearch_co-tracker_main/notebooks/demo.ipynb b/torch_hub/facebookresearch_co-tracker_main/notebooks/demo.ipynb new file mode 100644 index 0000000000000000000000000000000000000000..b41c701cdd0c43eeb6ed1eff39e2510b62216ad7 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/notebooks/demo.ipynb @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2d8bd0b51945974f9f73770a7b100ce7256efcd0408aed53f235443ff0f48c5d +size 14500914 diff --git a/torch_hub/facebookresearch_co-tracker_main/online_demo.py b/torch_hub/facebookresearch_co-tracker_main/online_demo.py new file mode 100644 index 0000000000000000000000000000000000000000..5514f86f1b578d32833d0f955b93ee6558dfc779 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/online_demo.py @@ -0,0 +1,104 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import torch +import argparse +import imageio.v3 as iio +import numpy as np + +from cotracker.utils.visualizer import Visualizer +from cotracker.predictor import CoTrackerOnlinePredictor + + +DEFAULT_DEVICE = ( + "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu" +) + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument( + "--video_path", + default="./assets/apple.mp4", + help="path to a video", + ) + parser.add_argument( + "--checkpoint", + default=None, + help="CoTracker model parameters", + ) + parser.add_argument("--grid_size", type=int, default=10, help="Regular grid size") + parser.add_argument( + "--grid_query_frame", + type=int, + default=0, + help="Compute dense and grid tracks starting from this frame", + ) + + args = parser.parse_args() + + if not os.path.isfile(args.video_path): + raise ValueError("Video file does not exist") + + if args.checkpoint is not None: + model = CoTrackerOnlinePredictor(checkpoint=args.checkpoint) + else: + model = torch.hub.load("facebookresearch/co-tracker", "cotracker3_online") + model = model.to(DEFAULT_DEVICE) + + window_frames = [] + + def _process_step(window_frames, is_first_step, grid_size, grid_query_frame): + video_chunk = ( + torch.tensor( + np.stack(window_frames[-model.step * 2 :]), device=DEFAULT_DEVICE + ) + .float() + .permute(0, 3, 1, 2)[None] + ) # (1, T, 3, H, W) + return model( + video_chunk, + is_first_step=is_first_step, + grid_size=grid_size, + grid_query_frame=grid_query_frame, + ) + + # Iterating over video frames, processing one window at a time: + is_first_step = True + for i, frame in enumerate( + iio.imiter( + args.video_path, + plugin="FFMPEG", + ) + ): + if i % model.step == 0 and i != 0: + pred_tracks, pred_visibility = _process_step( + window_frames, + is_first_step, + grid_size=args.grid_size, + grid_query_frame=args.grid_query_frame, + ) + is_first_step = False + window_frames.append(frame) + # Processing the final video frames in case video length is not a multiple of model.step + pred_tracks, pred_visibility = _process_step( + window_frames[-(i % model.step) - model.step - 1 :], + is_first_step, + grid_size=args.grid_size, + grid_query_frame=args.grid_query_frame, + ) + + print("Tracks are computed") + + # save a video with predicted tracks + seq_name = args.video_path.split("/")[-1] + video = torch.tensor(np.stack(window_frames), device=DEFAULT_DEVICE).permute( + 0, 3, 1, 2 + )[None] + vis = Visualizer(save_dir="./saved_videos", pad_value=120, linewidth=3) + vis.visualize( + video, pred_tracks, pred_visibility, query_frame=args.grid_query_frame + ) diff --git a/torch_hub/facebookresearch_co-tracker_main/setup.py b/torch_hub/facebookresearch_co-tracker_main/setup.py new file mode 100644 index 0000000000000000000000000000000000000000..a7aae5f0df37891ec89037a731d3acb029df7a39 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/setup.py @@ -0,0 +1,18 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +from setuptools import find_packages, setup + +setup( + name="cotracker", + version="3.0", + install_requires=[], + packages=find_packages(exclude="notebooks"), + extras_require={ + "all": ["matplotlib"], + "dev": ["flake8", "black"], + }, +) diff --git a/torch_hub/facebookresearch_co-tracker_main/tests/test_bilinear_sample.py b/torch_hub/facebookresearch_co-tracker_main/tests/test_bilinear_sample.py new file mode 100644 index 0000000000000000000000000000000000000000..29e5322a0f007a5c4f23f38e509092b8f00e9714 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/tests/test_bilinear_sample.py @@ -0,0 +1,51 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + + +import torch +import unittest + +from cotracker.models.core.model_utils import bilinear_sampler + + +class TestBilinearSampler(unittest.TestCase): + # Sample from an image (4d) + def _test4d(self, align_corners): + H, W = 4, 5 + # Construct a grid to obtain indentity sampling + input = torch.randn(H * W).view(1, 1, H, W).float() + coords = torch.meshgrid(torch.arange(H), torch.arange(W)) + coords = torch.stack(coords[::-1], dim=-1).float()[None] + if not align_corners: + coords = coords + 0.5 + sampled_input = bilinear_sampler(input, coords, align_corners=align_corners) + torch.testing.assert_close(input, sampled_input) + + # Sample from a video (5d) + def _test5d(self, align_corners): + T, H, W = 3, 4, 5 + # Construct a grid to obtain indentity sampling + input = torch.randn(H * W).view(1, 1, H, W).float() + input = torch.stack([input, input + 1, input + 2], dim=2) + coords = torch.meshgrid(torch.arange(T), torch.arange(W), torch.arange(H)) + coords = torch.stack(coords, dim=-1).float().permute(0, 2, 1, 3)[None] + + if not align_corners: + coords = coords + 0.5 + sampled_input = bilinear_sampler(input, coords, align_corners=align_corners) + torch.testing.assert_close(input, sampled_input) + + def test4d(self): + self._test4d(align_corners=True) + self._test4d(align_corners=False) + + def test5d(self): + self._test5d(align_corners=True) + self._test5d(align_corners=False) + + +# run the test +unittest.main() diff --git a/torch_hub/facebookresearch_co-tracker_main/train_on_kubric.py b/torch_hub/facebookresearch_co-tracker_main/train_on_kubric.py new file mode 100644 index 0000000000000000000000000000000000000000..e293cb26b28ee2eb8ace0c9a59b0170425292c19 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/train_on_kubric.py @@ -0,0 +1,706 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import random +import torch +import signal +import socket +import sys +import json +import torch.nn.functional as F +import numpy as np +import argparse +import logging +from pathlib import Path +from tqdm import tqdm +import torch.optim as optim + +from torch.cuda.amp import GradScaler +from pytorch_lightning.lite import LightningLite + +from cotracker.models.core.cotracker.cotracker3_offline import CoTrackerThreeOffline +from cotracker.models.core.cotracker.cotracker3_online import CoTrackerThreeOnline + +from cotracker.utils.visualizer import Visualizer +from cotracker.models.core.model_utils import get_uniformly_sampled_pts +from cotracker.evaluation.core.evaluator import Evaluator +from cotracker.datasets.utils import collate_fn, collate_fn_train, dataclass_to_cuda_ +from cotracker.models.core.cotracker.losses import ( + sequence_loss, + sequence_BCE_loss, + sequence_prob_loss, +) +from cotracker.utils.train_utils import ( + Logger, + get_eval_dataloader, + get_train_dataset, + sig_handler, + term_handler, + run_test_eval, +) + + +def fetch_optimizer(args, model): + """Create the optimizer and learning rate scheduler""" + mlp_params = sum( + p.numel() + for name, p in model.named_parameters() + if p.requires_grad and "corr_mlp" in name + ) + print(f"Total number of MlP parameters: {mlp_params}") + + mlp_params = sum( + p.numel() + for name, p in model.named_parameters() + if p.requires_grad and "cmdtop" in name + ) + print(f"Total number of cmdtop parameters: {mlp_params}") + + total_params = sum(p.numel() for p in model.parameters() if p.requires_grad) + print(f"Total number of parameters: {total_params}") + optimizer = optim.AdamW( + model.parameters(), lr=args.lr, weight_decay=args.wdecay, eps=1e-8 + ) + scheduler = optim.lr_scheduler.OneCycleLR( + optimizer, + args.lr, + args.num_steps + 100, + pct_start=0.05, + cycle_momentum=False, + anneal_strategy="cos", + ) + return optimizer, scheduler + + +def forward_batch(batch, model, args): + video = batch.video + trajs_g = batch.trajectory + vis_g = batch.visibility + valids = batch.valid + + B, T, C, H, W = video.shape + assert C == 3 + B, T, N, D = trajs_g.shape + device = video.device + + __, first_positive_inds = torch.max(vis_g, dim=1) + + if args.query_sampling_method == "random": + assert B == 1 + true_indices = torch.nonzero(vis_g[0]) + # Group the indices by the first column (N) + grouped_indices = true_indices[:, 1].unique() + # Initialize an empty tensor to hold the sampled points + sampled_points = torch.empty((B, N, D)) + indices = torch.empty((B, N, 1)) + # For each unique N + for n in grouped_indices: + # Get the T indices where visibilities[0, :, n] is True + t_indices = true_indices[true_indices[:, 1] == n, 0] + + # Select a random index from t_indices + random_index = t_indices[torch.randint(0, len(t_indices), (1,))] + + # Use this random index to sample a point from the trajectories tensor + sampled_points[0, n] = trajs_g[0, random_index, n] + indices[0, n] = random_index.float() + # model.window_len = vis_g.shape[1] + queries = torch.cat([indices, sampled_points], dim=2) + else: + # We want to make sure that during training the model sees visible points + # that it does not need to track just yet: they are visible but queried from a later frame + N_rand = N // 4 + # inds of visible points in the 1st frame + nonzero_inds = [ + [torch.nonzero(vis_g[b, :, i]) for i in range(N)] for b in range(B) + ] + + for b in range(B): + rand_vis_inds = torch.cat( + [ + nonzero_row[torch.randint(len(nonzero_row), size=(1,))] + for nonzero_row in nonzero_inds[b] + ], + dim=1, + ) + first_positive_inds[b] = torch.cat( + [rand_vis_inds[:, :N_rand], first_positive_inds[b : b + 1, N_rand:]], + dim=1, + ) + + ind_array_ = torch.arange(T, device=device) + ind_array_ = ind_array_[None, :, None].repeat(B, 1, N) + assert torch.allclose( + vis_g[ind_array_ == first_positive_inds[:, None, :]], + torch.ones(1, device=device), + ) + gather = torch.gather( + trajs_g, 1, first_positive_inds[:, :, None, None].repeat(1, 1, N, D) + ) + xys = torch.diagonal(gather, dim1=1, dim2=2).permute(0, 2, 1) + + queries = torch.cat([first_positive_inds[:, :, None], xys[:, :, :2]], dim=2) + + assert B == 1 + + if ( + torch.isnan(queries).any() + or torch.isnan(trajs_g).any() + or queries.abs().max() > 1500 + ): + print("failed_sample") + print("queries time", queries[..., 0]) + print("queries ", queries[..., 1:]) + queries = torch.ones_like(queries).to(queries.device).float() + print("new queries", queries) + valids = torch.zeros_like(valids).to(valids.device).float() + print("new valids", valids) + + model_output = model( + video=video, queries=queries[..., :3], iters=args.train_iters, is_train=True + ) + + tracks, visibility, confidence, train_data = model_output + coord_predictions, vis_predictions, confidence_predicitons, valid_mask = train_data + + vis_gts = [] + invis_gts = [] + traj_gts = [] + valids_gts = [] + + if args.offline_model: + S = T + seq_len = (S // 2) + 1 + else: + S = args.sliding_window_len + seq_len = T + + for ind in range(0, seq_len - S // 2, S // 2): + vis_gts.append(vis_g[:, ind : ind + S]) + invis_gts.append(1 - vis_g[:, ind : ind + S]) + traj_gts.append(trajs_g[:, ind : ind + S, :, :2]) + val = valids[:, ind : ind + S] + if not args.offline_model: + val = val * valid_mask[:, ind : ind + S] + valids_gts.append(val) + + seq_loss_visible = sequence_loss( + coord_predictions, + traj_gts, + valids_gts, + vis=vis_gts, + gamma=0.8, + add_huber_loss=args.add_huber_loss, + loss_only_for_visible=True, + ) + confidence_loss = sequence_prob_loss( + coord_predictions, confidence_predicitons, traj_gts, vis_gts + ) + vis_loss = sequence_BCE_loss(vis_predictions, vis_gts) + + output = {"flow": {"predictions": tracks[0].detach()}} + output["flow"]["loss"] = seq_loss_visible.mean() * 0.05 + output["flow"]["queries"] = queries.clone() + + if not args.train_only_on_visible: + seq_loss_invisible = sequence_loss( + coord_predictions, + traj_gts, + valids_gts, + vis=invis_gts, + gamma=0.8, + add_huber_loss=False, + loss_only_for_visible=True, + ) + output["flow_invisible"] = {"loss": seq_loss_invisible.mean() * 0.01} + output["visibility"] = { + "loss": vis_loss.mean(), + "predictions": visibility[0].detach(), + } + output["confidence"] = { + "loss": confidence_loss.mean(), + } + return output + + +class Lite(LightningLite): + def run(self, args): + def seed_everything(seed: int): + random.seed(seed) + os.environ["PYTHONHASHSEED"] = str(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False + + seed_everything(42) + + def seed_worker(worker_id): + worker_seed = torch.initial_seed() % 2**32 + np.random.seed(worker_seed + worker_id) + random.seed(worker_seed + worker_id) + + g = torch.Generator() + g.manual_seed(42) + if self.global_rank == 0: + eval_dataloaders = [] + for ds_name in args.eval_datasets: + eval_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + if not args.debug: + final_dataloaders = [dl for dl in eval_dataloaders] + + ds_name = "dynamic_replica" + final_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + + ds_name = "tapvid_robotap" + final_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + + ds_name = "tapvid_kinetics_first" + final_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + + evaluator = Evaluator(args.ckpt_path) + + visualizer = Visualizer( + save_dir=args.ckpt_path, + pad_value=180, + fps=1, + show_first_frame=0, + tracks_leave_trace=0, + ) + + if args.model_name == "cotracker_three": + if args.offline_model: + model = CoTrackerThreeOffline( + stride=args.model_stride, + corr_radius=args.corr_radius, + corr_levels=args.corr_levels, + window_len=args.sliding_window_len, + num_virtual_tracks=args.num_virtual_tracks, + model_resolution=args.crop_size, + linear_layer_for_vis_conf=args.linear_layer_for_vis_conf, + ) + else: + model = CoTrackerThreeOnline( + stride=args.model_stride, + corr_radius=args.corr_radius, + corr_levels=args.corr_levels, + window_len=args.sliding_window_len, + num_virtual_tracks=args.num_virtual_tracks, + model_resolution=args.crop_size, + linear_layer_for_vis_conf=args.linear_layer_for_vis_conf, + ) + else: + raise ValueError(f"Model {args.model_name} doesn't exist") + + with open(args.ckpt_path + "/meta.json", "w") as file: + json.dump(vars(args), file, sort_keys=True, indent=4) + + model.cuda() + + train_dataset = get_train_dataset(args) + train_loader = torch.utils.data.DataLoader( + train_dataset, + batch_size=args.batch_size, + shuffle=True, + num_workers=args.num_workers, + worker_init_fn=seed_worker, + generator=g, + pin_memory=True, + collate_fn=collate_fn_train, + drop_last=True, + ) + train_loader = self.setup_dataloaders(train_loader, move_to_device=False) + print("LEN TRAIN LOADER", len(train_loader)) + optimizer, scheduler = fetch_optimizer(args, model) + + total_steps = 0 + if self.global_rank == 0: + logger = Logger(model, scheduler, ckpt_path=args.ckpt_path) + + folder_ckpts = [ + f + for f in os.listdir(args.ckpt_path) + if not os.path.isdir(f) and f.endswith(".pth") and not "final" in f + ] + if len(folder_ckpts) > 0: + ckpt_path = sorted(folder_ckpts)[-1] + ckpt = self.load(os.path.join(args.ckpt_path, ckpt_path)) + logging.info(f"Loading checkpoint {ckpt_path}") + if "model" in ckpt: + model.load_state_dict(ckpt["model"]) + else: + model.load_state_dict(ckpt) + if "optimizer" in ckpt: + logging.info("Load optimizer") + optimizer.load_state_dict(ckpt["optimizer"]) + if "scheduler" in ckpt: + logging.info("Load scheduler") + scheduler.load_state_dict(ckpt["scheduler"]) + if "total_steps" in ckpt: + total_steps = ckpt["total_steps"] + logging.info(f"Load total_steps {total_steps}") + + elif args.restore_ckpt is not None: + assert args.restore_ckpt.endswith(".pth") or args.restore_ckpt.endswith( + ".pt" + ) + logging.info("Loading checkpoint...") + + strict = False + state_dict = self.load(args.restore_ckpt) + if "model" in state_dict: + state_dict = state_dict["model"] + state_dict = { + k: v + for k, v in state_dict.items() + if "time_emb" not in k and "pos_emb" not in k + } + if list(state_dict.keys())[0].startswith("module."): + state_dict = { + k.replace("module.", ""): v for k, v in state_dict.items() + } + model.load_state_dict(state_dict, strict=strict) + + logging.info(f"Done loading checkpoint") + model, optimizer = self.setup(model, optimizer, move_to_device=False) + model.train() + + save_freq = args.save_freq + scaler = GradScaler(enabled=False) + + should_keep_training = True + global_batch_num = 0 + epoch = -1 + + while should_keep_training: + epoch += 1 + for i_batch, batch in enumerate(tqdm(train_loader)): + batch, gotit = batch + if not all(gotit): + print("batch is None") + continue + + dataclass_to_cuda_(batch) + + optimizer.zero_grad(set_to_none=True) + + assert model.training + + output = forward_batch(batch, model, args) + + loss = 0 + for k, v in output.items(): + if "loss" in v: + loss += v["loss"] + + if self.global_rank == 0: + for k, v in output.items(): + if "loss" in v: + logger.writer.add_scalar( + f"live_{k}_loss", v["loss"].item(), total_steps + ) + if "metrics" in v: + logger.push(v["metrics"], k) + if total_steps % save_freq == save_freq - 1: + visualizer.visualize( + video=batch.video.clone(), + tracks=batch.trajectory.clone()[..., :2], + visibility=batch.visibility.clone(), + filename="train_gt_traj_0", + writer=logger.writer, + step=total_steps, + ) + + visualizer.visualize( + video=batch.video.clone(), + tracks=output["flow"]["predictions"][None], + visibility=output["visibility"]["predictions"][None] > 0.8, + filename="train_pred_traj_0", + writer=logger.writer, + step=total_steps, + ) + + if len(output) > 1: + logger.writer.add_scalar( + f"live_total_loss", loss.item(), total_steps + ) + logger.writer.add_scalar( + f"learning_rate", optimizer.param_groups[0]["lr"], total_steps + ) + global_batch_num += 1 + + self.barrier() + self.backward(scaler.scale(loss)) + + scaler.unscale_(optimizer) + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) + + scaler.step(optimizer) + scheduler.step() + scaler.update() + total_steps += 1 + if self.global_rank == 0: + if (i_batch >= len(train_loader) - 1) or ( + total_steps == 1 and args.validate_at_start + ): + if (epoch + 1) % args.save_every_n_epoch == 0: + ckpt_iter = "0" * (6 - len(str(total_steps))) + str( + total_steps + ) + save_path = Path( + f"{args.ckpt_path}/model_{args.model_name}_{ckpt_iter}.pth" + ) + + save_dict = { + "model": model.module.module.state_dict(), + "optimizer": optimizer.state_dict(), + "scheduler": scheduler.state_dict(), + "total_steps": total_steps, + } + + logging.info(f"Saving file {save_path}") + self.save(save_dict, save_path) + + if (epoch + 1) % args.evaluate_every_n_epoch == 0 or ( + args.validate_at_start and epoch == 0 + ): + run_test_eval( + evaluator, + model, + eval_dataloaders, + logger.writer, + total_steps, + query_random=( + args.query_sampling_method is not None + and "random" in args.query_sampling_method + ), + ) + model.train() + torch.cuda.empty_cache() + + self.barrier() + if total_steps > args.num_steps: + should_keep_training = False + break + + if self.global_rank == 0: + print("FINISHED TRAINING") + + PATH = f"{args.ckpt_path}/{args.model_name}_final.pth" + torch.save(model.module.module.state_dict(), PATH) + run_test_eval( + evaluator, + model, + final_dataloaders, + logger.writer, + total_steps, + query_random=( + args.query_sampling_method is not None + and "random" in args.query_sampling_method + ), + ) + logger.close() + + +if __name__ == "__main__": + signal.signal(signal.SIGUSR1, sig_handler) + signal.signal(signal.SIGTERM, term_handler) + parser = argparse.ArgumentParser() + parser.add_argument("--model_name", default="cotracker_three", help="model name") + parser.add_argument("--restore_ckpt", help="path to restore a checkpoint") + parser.add_argument("--ckpt_path", help="path to save checkpoints") + parser.add_argument( + "--batch_size", type=int, default=4, help="batch size used during training." + ) + parser.add_argument("--num_nodes", type=int, default=1) + parser.add_argument( + "--num_workers", type=int, default=10, help="number of dataloader workers" + ) + + parser.add_argument( + "--mixed_precision", action="store_true", help="use mixed precision" + ) + parser.add_argument("--lr", type=float, default=0.0005, help="max learning rate.") + parser.add_argument( + "--wdecay", type=float, default=0.00001, help="Weight decay in optimizer." + ) + parser.add_argument( + "--num_steps", type=int, default=200000, help="length of training schedule." + ) + parser.add_argument( + "--evaluate_every_n_epoch", + type=int, + default=1, + help="evaluate during training after every n epochs, after every epoch by default", + ) + parser.add_argument( + "--save_every_n_epoch", + type=int, + default=1, + help="save checkpoints during training after every n epochs, after every epoch by default", + ) + parser.add_argument( + "--validate_at_start", + action="store_true", + help="whether to run evaluation before training starts", + ) + parser.add_argument( + "--save_freq", + type=int, + default=100, + help="frequency of trajectory visualization during training", + ) + parser.add_argument( + "--traj_per_sample", + type=int, + default=768, + help="the number of trajectories to sample for training", + ) + parser.add_argument( + "--dataset_root", type=str, help="path lo all the datasets (train and eval)" + ) + + parser.add_argument( + "--train_iters", + type=int, + default=4, + help="number of updates to the disparity field in each forward pass.", + ) + parser.add_argument( + "--sequence_len", type=int, default=8, help="train sequence length" + ) + parser.add_argument( + "--eval_datasets", + nargs="+", + default=["tapvid_davis_first"], + help="what datasets to use for evaluation", + ) + parser.add_argument( + "--train_datasets", + nargs="+", + default=["kubric"], + help="what datasets to use for evaluation", + ) + parser.add_argument( + "--random_frame_rate", + action="store_true", + help="remove space attention from CoTracker", + ) + parser.add_argument( + "--num_virtual_tracks", + type=int, + default=None, + help="stride of the CoTracker feature network", + ) + parser.add_argument( + "--dont_use_augs", + action="store_true", + help="don't apply augmentations during training", + ) + parser.add_argument( + "--offline_model", + action="store_true", + help="only sample trajectories with points visible on the first frame", + ) + parser.add_argument( + "--sliding_window_len", + type=int, + default=16, + help="length of the CoTracker sliding window", + ) + parser.add_argument( + "--model_stride", + type=int, + default=4, + help="stride of the CoTracker feature network", + ) + parser.add_argument( + "--corr_radius", + type=int, + default=3, + help="stride of the CoTracker feature network", + ) + parser.add_argument( + "--corr_levels", + type=int, + default=4, + help="stride of the CoTracker feature network", + ) + parser.add_argument( + "--crop_size", + type=int, + nargs="+", + default=[384, 512], + help="crop videos to this resolution during training", + ) + parser.add_argument( + "--eval_max_seq_len", + type=int, + default=1000, + help="maximum length of evaluation videos", + ) + parser.add_argument( + "--query_sampling_method", + type=str, + help="path lo all the datasets (train and eval)", + ) + parser.add_argument( + "--random_number_traj", + action="store_true", + help="only sample trajectories with points visible on the first frame", + ) + parser.add_argument( + "--add_huber_loss", + action="store_true", + help="only sample trajectories with points visible on the first frame", + ) + parser.add_argument( + "--debug", + action="store_true", + help="only sample trajectories with points visible on the first frame", + ) + parser.add_argument( + "--random_seq_len", + action="store_true", + help="only sample trajectories with points visible on the first frame", + ) + parser.add_argument( + "--linear_layer_for_vis_conf", + action="store_true", + help="stride of the CoTracker feature network", + ) + parser.add_argument( + "--train_only_on_visible", + action="store_true", + help="stride of the CoTracker feature network", + ) + + args = parser.parse_args() + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(levelname)-8s [%(filename)s:%(lineno)d] %(message)s", + ) + + Path(args.ckpt_path).mkdir(exist_ok=True, parents=True) + from pytorch_lightning.strategies import DDPStrategy + + Lite( + strategy=DDPStrategy(find_unused_parameters=False), + devices="auto", + accelerator="gpu", + precision="bf16" if args.mixed_precision else 32, + num_nodes=args.num_nodes, + ).run(args) diff --git a/torch_hub/facebookresearch_co-tracker_main/train_on_real_data.py b/torch_hub/facebookresearch_co-tracker_main/train_on_real_data.py new file mode 100644 index 0000000000000000000000000000000000000000..72fc56519f5a92fb73d883bc4b5a435ee38385e1 --- /dev/null +++ b/torch_hub/facebookresearch_co-tracker_main/train_on_real_data.py @@ -0,0 +1,849 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import os +import random +import torch +import signal +import socket +import sys +import json + +import numpy as np +import argparse +import logging +from pathlib import Path +from tqdm import tqdm +import torch.optim as optim +import torchvision +from torch.utils.data import DataLoader +from torch.cuda.amp import GradScaler + +from pytorch_lightning.lite import LightningLite + +from cotracker.models.bootstap_predictor import TAPIRPredictor +from cotracker.models.core.cotracker.cotracker import CoTracker2 +from cotracker.models.core.cotracker.cotracker3_offline import CoTrackerThreeOffline +from cotracker.models.core.cotracker.cotracker3_online import CoTrackerThreeOnline + +from cotracker.utils.visualizer import Visualizer + +from cotracker.evaluation.core.evaluator import Evaluator +from cotracker.datasets.utils import collate_fn, collate_fn_train, dataclass_to_cuda_ +from cotracker.models.core.model_utils import ( + get_uniformly_sampled_pts, + get_points_on_a_grid, + get_sift_sampled_pts, + get_superpoint_sampled_pts, +) +from cotracker.models.core.cotracker.losses import sequence_loss +from cotracker.models.build_cotracker import build_cotracker +from cotracker.utils.train_utils import ( + Logger, + get_eval_dataloader, + sig_handler, + term_handler, + run_test_eval, +) + + +def fetch_optimizer(args, model): + """Create the optimizer and learning rate scheduler""" + total_params = sum(p.numel() for p in model.parameters() if p.requires_grad) + print(f"Total number of parameters: {total_params}") + for name, param in model.named_parameters(): + if "vis_conf_head" in name: + param.requires_grad = False + + optimizer = optim.AdamW( + model.parameters(), lr=args.lr, weight_decay=args.wdecay, eps=1e-8 + ) + scheduler = optim.lr_scheduler.OneCycleLR( + optimizer, + args.lr, + args.num_steps + 100, + pct_start=0.0, + cycle_momentum=False, + anneal_strategy="cos", + ) + return optimizer, scheduler + + +def forward_batch(batch, model, args, teacher_models): + video = batch.video + trajs_g = batch.trajectory + vis_g = batch.visibility + valids = batch.valid + B, T, C, H, W = video.shape + assert C == 3 + B, T, N, D = trajs_g.shape + device = video.device + failed_sample = False + if args.real_data_filter_sift: + queries = get_sift_sampled_pts(video, N, T, [H, W], device=device) + + if queries.shape[1] < N: + logging.warning( + f"SIFT wasn't able to extract enough features: {queries.shape[1]}" + ) + failed_sample = True + queries = get_uniformly_sampled_pts(N, T, [H, W], device=device) + elif args.real_data_filter_superpoint: + queries = get_superpoint_sampled_pts(video, N, T, [H, W], device=device) + + if queries.shape[1] < N: + logging.warning("SuperPoint wasn't able to extract enough features") + failed_sample = True + queries = get_uniformly_sampled_pts(N, T, [H, W], device=device) + else: + queries = get_uniformly_sampled_pts(N, T, [H, W], device=device) + # Inference with additional points sampled on a regular grid usually makes predictions better. + # So we sample these points and discard them thereafter + + teacher_model_ind = random.choice(range(len(teacher_models))) + + teacher_model_type, teacher_model = teacher_models[teacher_model_ind] + uniform_size = grid_size = sift_size = 0 + queries_cat = queries.clone() + if "online" in teacher_model_type: + grid_size = args.train_grid_size + sift_size = args.train_sift_size + if grid_size > 0: + xy = get_points_on_a_grid(grid_size, [H, W], device=device) + xy = torch.cat([torch.zeros_like(xy[:, :, :1]), xy], dim=2) # + queries_cat = torch.cat([queries_cat, xy], dim=1) # + + if sift_size > 0: + xy = get_sift_sampled_pts(video, sift_size, T, [H, W], device=device) + if xy.shape[1] == sift_size: + queries_cat = torch.cat([queries_cat, xy], dim=1) # + else: + sift_size = 0 + elif "offline" in teacher_model_type: + uniform_size = 100 + if uniform_size > 0: + xy = get_uniformly_sampled_pts(uniform_size, T, [H, W], device=device) + queries_cat = torch.cat([queries_cat, xy], dim=1) # + elif teacher_model_type == "tapir": + pass + else: + raise ValueError(f"Model type {teacher_model_type} doesn't exist") + + if "cotracker_three" in teacher_model_type: + with torch.no_grad(): + ( + trajs_g, + vis_g, + confidence, + __, + ) = teacher_model(video, queries_cat) + else: + with torch.no_grad(): + trajs_g, vis_g, *_ = teacher_model(video, queries_cat) + confidence = torch.ones_like(vis_g) + + # discarding additional points + if sift_size > 0 or grid_size > 0 or uniform_size > 0: + trajs_g = trajs_g[:, :, : -(grid_size**2) - sift_size - uniform_size] + vis_g = vis_g[:, :, : -(grid_size**2) - sift_size - uniform_size] + confidence = confidence[:, :, : -(grid_size**2) - sift_size - uniform_size] + + vis_g = vis_g > 0.9 + + batch.trajectory = trajs_g + batch.visibility = vis_g + + if args.model_name == "cotracker_three": + if ( + torch.isnan(queries).any() + or torch.isnan(trajs_g).any() + or queries.abs().max() > 1500 + ): + logging.warning("failed_sample") + queries = torch.ones_like(queries).to(queries.device).float() + valids = torch.zeros_like(valids).to(valids.device).float() + + tracks, visibility, confidence, train_data = model( + video=video, queries=queries, iters=args.train_iters, is_train=True + ) + coord_predictions, vis_predictions, confidence_predicitons, valid_mask = ( + train_data + ) + + if failed_sample: + valid_mask = torch.zeros_like(vis_g) + logging.warning("Making mask zero for failed sample") + + vis_gts = [] + invis_gts = [] + traj_gts = [] + valids_gts = [] + if args.offline_model: + S = T + seq_len = (S // 2) + 1 + else: + S = args.sliding_window_len + seq_len = T + for ind in range(0, seq_len - S // 2, S // 2): + vis_gts.append(vis_g[:, ind : ind + S].float()) + invis_gts.append(1 - vis_g[:, ind : ind + S].float()) + traj_gts.append(trajs_g[:, ind : ind + S, :, :2]) + valids_gts.append(valids[:, ind : ind + S] * valid_mask[:, ind : ind + S]) + + seq_loss = sequence_loss( + coord_predictions, + traj_gts, + valids_gts, + vis=vis_gts, + gamma=0.8, + add_huber_loss=True, + loss_only_for_visible=True, + ) + + output = { + "flow": {"predictions": (tracks[0].detach() * valid_mask[..., None])[0]} + } + output["flow"]["loss"] = seq_loss.mean() * 0.05 + output["flow"]["queries"] = queries.clone() + output["flow"]["query_frame"] = queries[0, :, 0].cpu().int() + + output["visibility"] = { + "predictions": visibility[0].detach(), + } + if not (teacher_model_type == "tapir" or args.train_only_visible_points): + seq_loss_invisible = sequence_loss( + coord_predictions, + traj_gts, + valids_gts, + vis=invis_gts, + gamma=0.8, + add_huber_loss=False, + loss_only_for_visible=True, + ) + output["flow_invisible"] = {"loss": seq_loss_invisible.mean() * 0.01} + + return output + else: + predictions, visibility, train_data = model( + video=video, queries=queries, iters=args.train_iters, is_train=True + ) + coord_predictions, vis_predictions, valid_mask = train_data + + if failed_sample: + valid_mask = torch.zeros_like(valid_mask) + logging.warning("Making mask zero for failed sample") + + vis_gts = [] + traj_gts = [] + valids_gts = [] + delta = 6 + S = args.sliding_window_len + pred_ind = 0 + + for ind in range(0, args.sequence_len - S // 2, S // 2): + vis_gts.append(vis_g[:, ind : ind + S]) + traj_gts.append(trajs_g[:, ind : ind + S]) + if ( + teacher_model_type == "tapir" + or teacher_model_type == "online_cotracker_three" + or args.train_only_visible_points + ): + valids_gts.append( + valids[:, ind : ind + S] + * valid_mask[:, ind : ind + S] + * vis_g[:, ind : ind + S] + > 0.9 + ) + else: + valids_gts.append( + valids[:, ind : ind + S] * valid_mask[:, ind : ind + S] + ) + pred_ind += 1 + + seq_loss = sequence_loss( + coord_predictions, + traj_gts, + vis_gts, + valids_gts, + gamma=0.8, + loss_only_for_visible_pts=False, + ) + + batch.trajectory = batch.trajectory * valid_mask[..., None] + output = {"flow": {}} + + output["flow"]["predictions"] = (predictions.detach() * valid_mask[..., None])[ + 0 + ] + output["flow"]["loss"] = seq_loss.mean() + output["flow"]["query_frame"] = queries[0, :, 0].cpu().int() + output["visibility"] = { + "predictions": visibility[0].detach(), + } + return output + + +class Lite(LightningLite): + def run(self, args): + def seed_everything(seed: int): + random.seed(seed) + os.environ["PYTHONHASHSEED"] = str(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False + + seed_everything(0) + + def seed_worker(worker_id): + worker_seed = torch.initial_seed() % 2**32 + np.random.seed(worker_seed) + random.seed(worker_seed) + + g = torch.Generator() + g.manual_seed(0) + + if self.global_rank == 0: + eval_dataloaders = [] + for ds_name in args.eval_datasets: + eval_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + if not args.debug: + final_dataloaders = [dl for dl in eval_dataloaders] + ds_name = "tapvid_kinetics_first" + final_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + + ds_name = "tapvid_robotap" + final_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + + ds_name = "dynamic_replica" + final_dataloaders.append( + (ds_name, get_eval_dataloader(args.dataset_root, ds_name)) + ) + evaluator = Evaluator(args.ckpt_path) + + visualizer = Visualizer( + save_dir=args.ckpt_path, + pad_value=180, + fps=1, + show_first_frame=0, + tracks_leave_trace=0, + ) + + if args.model_name == "cotracker": + model = CoTracker2( + stride=args.model_stride, + window_len=args.sliding_window_len, + num_virtual_tracks=args.num_virtual_tracks, + model_resolution=args.crop_size, + ) + elif args.model_name == "cotracker_three": + if args.offline_model: + model = CoTrackerThreeOffline( + stride=4, + corr_radius=3, + window_len=60, + model_resolution=(384, 512), + linear_layer_for_vis_conf=True, + ) + else: + model = CoTrackerThreeOnline( + stride=4, + corr_radius=3, + window_len=16, + model_resolution=(384, 512), + linear_layer_for_vis_conf=True, + ) + else: + raise ValueError(f"Model {args.model_name} doesn't exist") + + with open(args.ckpt_path + "/meta.json", "w") as file: + json.dump(vars(args), file, sort_keys=True, indent=4) + + model.cuda() + teacher_models = [] + from cotracker.datasets import real_dataset + + train_dataset = real_dataset.RealDataset( + crop_size=args.crop_size, + seq_len=args.sequence_len, + traj_per_sample=args.traj_per_sample, + random_frame_rate=args.random_frame_rate, + random_seq_len=args.offline_model, + data_splits=args.real_data_splits, + random_resize=False, + limit_samples=args.limit_samples, + ) + + if args.model_name == "cotracker": + teacher_model_online = ( + build_cotracker( + window_len=args.sliding_window_len, checkpoint=args.restore_ckpt + ) + .cuda() + .eval() + ) + teacher_models.append(("online", teacher_model_online)) + elif args.model_name == "cotracker_three": + teacher_model_online = ( + build_cotracker( + window_len=16, + offline=False, + checkpoint="./checkpoints/cotracker2v1.pth", + v2=True, + ) + .cuda() + .eval() + ) + teacher_models.append(("online", teacher_model_online)) + else: + raise ValueError(f"Model {args.model_name} doesn't exist") + + online_checkpoint = "./checkpoints/baseline_online.pth" + if args.model_name == "cotracker_three" and not args.offline_model: + online_checkpoint = args.restore_ckpt + print("online_checkpoint", online_checkpoint) + teacher_model_online_cot_three = ( + build_cotracker(checkpoint=online_checkpoint, offline=False, window_len=16) + .cuda() + .eval() + ) + teacher_models.append( + ("online_cotracker_three", teacher_model_online_cot_three) + ) + + offline_checkpoint = "./checkpoints/baseline_offline.pth" + if args.model_name == "cotracker_three" and args.offline_model: + offline_checkpoint = args.restore_ckpt + + teacher_model_offline_cot_three = ( + build_cotracker(checkpoint=offline_checkpoint, offline=True, window_len=60) + .cuda() + .eval() + ) + teacher_models.append( + ("offline_cotracker_three", teacher_model_offline_cot_three) + ) + + teacher_model_tapir = TAPIRPredictor() + teacher_models.append(("tapir", teacher_model_tapir)) + + train_loader = DataLoader( + train_dataset, + batch_size=args.batch_size, + shuffle=True, + num_workers=args.num_workers, + worker_init_fn=seed_worker, + generator=g, + pin_memory=True, + collate_fn=collate_fn_train, + drop_last=True, + ) + + train_loader = self.setup_dataloaders(train_loader, move_to_device=False) + print("LEN TRAIN LOADER", len(train_loader)) + optimizer, scheduler = fetch_optimizer(args, model) + + total_steps = 0 + if self.global_rank == 0: + logger = Logger(model, scheduler, args.ckpt_path) + + folder_ckpts = [ + f + for f in os.listdir(args.ckpt_path) + if not os.path.isdir(f) and f.endswith(".pth") and not "final" in f + ] + if len(folder_ckpts) > 0: + ckpt_path = sorted(folder_ckpts)[-1] + ckpt = self.load(os.path.join(args.ckpt_path, ckpt_path)) + logging.info(f"Loading checkpoint {ckpt_path}") + if "model" in ckpt: + model.load_state_dict(ckpt["model"]) + else: + model.load_state_dict(ckpt) + if "optimizer" in ckpt: + logging.info("Load optimizer") + optimizer.load_state_dict(ckpt["optimizer"]) + if "scheduler" in ckpt: + logging.info("Load scheduler") + scheduler.load_state_dict(ckpt["scheduler"]) + if "total_steps" in ckpt: + total_steps = ckpt["total_steps"] + logging.info(f"Load total_steps {total_steps}") + + elif args.restore_ckpt is not None: + assert args.restore_ckpt.endswith(".pth") or args.restore_ckpt.endswith( + ".pt" + ) + logging.info("Loading checkpoint...") + + state_dict = self.load(args.restore_ckpt) + if "model" in state_dict: + state_dict = state_dict["model"] + + if list(state_dict.keys())[0].startswith("module."): + state_dict = { + k.replace("module.", ""): v for k, v in state_dict.items() + } + model.load_state_dict(state_dict, strict=True) + + logging.info(f"Done loading checkpoint") + model, optimizer = self.setup(model, optimizer, move_to_device=False) + # model.cuda() + model.train() + + save_freq = args.save_freq + scaler = GradScaler(enabled=False) + + should_keep_training = True + global_batch_num = 0 + epoch = -1 + + if self.global_rank == 0 and args.validate_at_start: + run_test_eval( + evaluator, + model, + eval_dataloaders, + logger.writer, + total_steps, + ) + model.train() + torch.cuda.empty_cache() + + while should_keep_training: + epoch += 1 + for i_batch, batch in enumerate(tqdm(train_loader)): + batch, gotit = batch + if not all(gotit): + print("batch is None") + continue + dataclass_to_cuda_(batch) + + optimizer.zero_grad() + + assert model.training + + output = forward_batch( + batch, model, args, teacher_models=teacher_models + ) + + loss = 0 + for k, v in output.items(): + if "loss" in v: + loss += v["loss"] + + if self.global_rank == 0: + for k, v in output.items(): + if "loss" in v: + logger.writer.add_scalar( + f"live_{k}_loss", v["loss"].item(), total_steps + ) + if "metrics" in v: + logger.push(v["metrics"], k) + if total_steps % save_freq == save_freq - 1: + visualizer.visualize( + video=batch.video.clone(), + tracks=batch.trajectory.clone(), + visibility=batch.visibility.clone(), + filename="train_gt_traj", + query_frame=output["flow"]["query_frame"], + writer=logger.writer, + step=total_steps, + ) + + visualizer.visualize( + video=batch.video.clone(), + tracks=output["flow"]["predictions"][None], + visibility=output["visibility"]["predictions"][None] > 0.6, + filename="train_pred_traj", + query_frame=output["flow"]["query_frame"], + writer=logger.writer, + step=total_steps, + ) + + if len(output) > 1: + logger.writer.add_scalar( + f"live_total_loss", loss.item(), total_steps + ) + logger.writer.add_scalar( + f"learning_rate", optimizer.param_groups[0]["lr"], total_steps + ) + global_batch_num += 1 + + self.barrier() + + self.backward(scaler.scale(loss)) + + scaler.unscale_(optimizer) + torch.nn.utils.clip_grad_norm_(model.parameters(), 10.0) + + scaler.step(optimizer) + scheduler.step() + scaler.update() + total_steps += 1 + if self.global_rank == 0: + if i_batch >= len(train_loader) - 1: + if (epoch + 1) % args.save_every_n_epoch == 0: + ckpt_iter = "0" * (6 - len(str(total_steps))) + str( + total_steps + ) + save_path = Path( + f"{args.ckpt_path}/model_{args.model_name}_{ckpt_iter}.pth" + ) + + save_dict = { + "model": model.module.module.state_dict(), + "optimizer": optimizer.state_dict(), + "scheduler": scheduler.state_dict(), + "total_steps": total_steps, + } + + logging.info(f"Saving file {save_path}") + self.save(save_dict, save_path) + + if (epoch + 1) % args.evaluate_every_n_epoch == 0: + run_test_eval( + evaluator, + model, + eval_dataloaders, + logger.writer, + total_steps, + ) + model.train() + torch.cuda.empty_cache() + + self.barrier() + if total_steps > args.num_steps: + should_keep_training = False + break + if self.global_rank == 0: + print("FINISHED TRAINING") + + PATH = f"{args.ckpt_path}/{args.model_name}_final.pth" + torch.save(model.module.module.state_dict(), PATH) + run_test_eval( + evaluator, model, final_dataloaders, logger.writer, total_steps + ) + logger.close() + + +if __name__ == "__main__": + signal.signal(signal.SIGUSR1, sig_handler) + signal.signal(signal.SIGTERM, term_handler) + parser = argparse.ArgumentParser() + parser.add_argument("--model_name", default="cotracker_three", help="model name") + parser.add_argument("--restore_ckpt", help="path to restore a checkpoint") + parser.add_argument("--ckpt_path", help="path to save checkpoints") + parser.add_argument( + "--batch_size", type=int, default=4, help="batch size used during training." + ) + parser.add_argument("--num_nodes", type=int, default=1) + parser.add_argument( + "--num_workers", type=int, default=10, help="number of dataloader workers" + ) + + parser.add_argument( + "--mixed_precision", action="store_true", help="use mixed precision" + ) + parser.add_argument("--lr", type=float, default=0.0005, help="max learning rate.") + parser.add_argument( + "--wdecay", type=float, default=0.00001, help="Weight decay in optimizer." + ) + parser.add_argument( + "--num_steps", type=int, default=200000, help="length of training schedule." + ) + parser.add_argument( + "--evaluate_every_n_epoch", + type=int, + default=1, + help="evaluate during training after every n epochs, after every epoch by default", + ) + parser.add_argument( + "--save_every_n_epoch", + type=int, + default=1, + help="save checkpoints during training after every n epochs, after every epoch by default", + ) + parser.add_argument( + "--validate_at_start", + action="store_true", + help="whether to run evaluation before training starts", + ) + parser.add_argument( + "--save_freq", + type=int, + default=100, + help="frequency of trajectory visualization during training", + ) + parser.add_argument( + "--traj_per_sample", + type=int, + default=768, + help="the number of trajectories to sample for training", + ) + parser.add_argument( + "--dataset_root", type=str, help="path lo all the datasets (train and eval)" + ) + + parser.add_argument( + "--train_iters", + type=int, + default=4, + help="number of updates to the disparity field in each forward pass.", + ) + parser.add_argument( + "--sequence_len", type=int, default=8, help="train sequence length" + ) + parser.add_argument( + "--eval_datasets", + nargs="+", + default=["tapvid_davis_first"], + help="what datasets to use for evaluation", + ) + parser.add_argument( + "--num_virtual_tracks", + type=int, + default=None, + help="stride of the CoTracker feature network", + ) + parser.add_argument( + "--dont_use_augs", + action="store_true", + help="don't apply augmentations during training", + ) + parser.add_argument( + "--sample_vis_1st_frame", + action="store_true", + help="only sample trajectories with points visible on the first frame", + ) + parser.add_argument( + "--sliding_window_len", + type=int, + default=8, + help="length of the CoTracker sliding window", + ) + parser.add_argument( + "--model_stride", + type=int, + default=8, + help="stride of the CoTracker feature network", + ) + parser.add_argument( + "--crop_size", + type=int, + nargs="+", + default=[384, 512], + help="crop videos to this resolution during training", + ) + parser.add_argument( + "--eval_max_seq_len", + type=int, + default=1000, + help="maximum length of evaluation videos", + ) + parser.add_argument( + "--debug", + action="store_true", + help="saves launch time for faster debug", + ) + parser.add_argument( + "--random_frame_rate", + action="store_true", + help="random_frame_rate", + ) + parser.add_argument( + "--real_data_splits", + type=int, + nargs="+", + default=[0], + help="real data folders", + ) + parser.add_argument( + "--loss_only_for_visible_pts", + action="store_true", + help="compute sequence loss only for visible points", + ) + parser.add_argument( + "--real_data_filter_sift", + action="store_true", + help="select point to track based on SIFT features", + ) + parser.add_argument( + "--train_grid_size", + type=int, + default=5, + help="number of extra regular grid points that we sample at training. This number will be squared", + ) + parser.add_argument( + "--train_sift_size", + type=int, + default=0, + help="number of extra SIFT points that we sample at training.", + ) + parser.add_argument( + "--real_data_filter_superpoint", + action="store_true", + help="select point to track based on SuperPoint features", + ) + parser.add_argument( + "--train_only_visible_points", + action="store_true", + help="Loss only for visible points", + ) + parser.add_argument( + "--offline_model", + action="store_true", + help="training the offline model", + ) + parser.add_argument( + "--clean_kubric", + action="store_true", + help="filtering out bad tracks in Kubric", + ) + parser.add_argument( + "--random_number_traj", + action="store_true", + help="when training on Kubric, sampling a random number \ + of tracks between 1 and args.traj_per_sample", + ) + parser.add_argument( + "--random_seq_len", + action="store_true", + help="when training on Kubric, cropping the sequence \ + to have a length between 10 and args.sequence_len frames", + ) + parser.add_argument( + "--uniform_query_sampling_method", + action="store_true", + help="Whether to sample points uniformly across time. Kubric training only", + ) + parser.add_argument( + "--limit_samples", type=int, default=10000, help="limit samples on real data" + ) + + args = parser.parse_args() + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(levelname)-8s [%(filename)s:%(lineno)d] %(message)s", + ) + + Path(args.ckpt_path).mkdir(exist_ok=True, parents=True) + from pytorch_lightning.strategies import DDPStrategy + + Lite( + strategy=DDPStrategy(find_unused_parameters=False), + devices="auto", + accelerator="gpu", + precision="bf16" if args.mixed_precision else 32, + num_nodes=args.num_nodes, + # precision=32, + ).run(args) diff --git a/torch_hub/facebookresearch_dino_main/.github/CODE_OF_CONDUCT.md b/torch_hub/facebookresearch_dino_main/.github/CODE_OF_CONDUCT.md new file mode 100644 index 0000000000000000000000000000000000000000..0f7ad8bfc173eac554f0b6ef7c684861e8014bbe --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/.github/CODE_OF_CONDUCT.md @@ -0,0 +1,5 @@ +# Code of Conduct + +Facebook has adopted a Code of Conduct that we expect project participants to adhere to. +Please read the [full text](https://code.fb.com/codeofconduct/) +so that you can understand what actions will and will not be tolerated. diff --git a/torch_hub/facebookresearch_dino_main/.github/attention_maps.png b/torch_hub/facebookresearch_dino_main/.github/attention_maps.png new file mode 100644 index 0000000000000000000000000000000000000000..b8f5c64d5949eea913c37ac24e738817a2ab477d --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/.github/attention_maps.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2e5c9ee7de5f3b79bb726be1c38ec664c8ce4a4602e4887a873528fc193646e1 +size 4205906 diff --git a/torch_hub/facebookresearch_dino_main/LICENSE b/torch_hub/facebookresearch_dino_main/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..b09cd7856d58590578ee1a4f3ad45d1310a97f87 --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/LICENSE @@ -0,0 +1,201 @@ +Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/torch_hub/facebookresearch_dino_main/README.md b/torch_hub/facebookresearch_dino_main/README.md new file mode 100644 index 0000000000000000000000000000000000000000..affca658f19b5509920ec065797b6abb88149982 --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/README.md @@ -0,0 +1,411 @@ +:new: *Please check out our more recent [DINOv2](https://github.com/facebookresearch/dinov2) effort in the same line of work.* + +# Self-Supervised Vision Transformers with DINO + +PyTorch implementation and pretrained models for DINO. For details, see **Emerging Properties in Self-Supervised Vision Transformers**. +[[`blogpost`](https://ai.facebook.com/blog/dino-paws-computer-vision-with-self-supervised-transformers-and-10x-more-efficient-training)] [[`arXiv`](https://arxiv.org/abs/2104.14294)] [[`Yannic Kilcher's video`](https://www.youtube.com/watch?v=h3ij3F3cPIk)] + +
+ DINO illustration +
+ +## Pretrained models +You can choose to download only the weights of the pretrained backbone used for downstream tasks, or the full checkpoint which contains backbone and projection head weights for both student and teacher networks. We also provide the backbone in `onnx` format, as well as detailed arguments and training/evaluation logs. Note that `DeiT-S` and `ViT-S` names refer exactly to the same architecture. + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
archparamsk-nnlineardownload
ViT-S/1621M74.5%77.0%backbone onlyfull ckptonnxargslogseval logs
ViT-S/821M78.3%79.7%backbone onlyfull ckptonnxargslogseval logs
ViT-B/1685M76.1%78.2%backbone onlyfull ckptonnxargslogseval logs
ViT-B/885M77.4%80.1%backbone onlyfull ckptonnxargslogseval logs
ResNet-5023M67.5%75.3%backbone onlyfull ckptonnxargslogseval logs
+ +We also release XCiT models ([[`arXiv`](https://arxiv.org/abs/2106.09681)] [[`code`](https://github.com/facebookresearch/xcit)]) trained with DINO: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
archparamsk-nnlineardownload
xcit_small_12_p1626M76.0%77.8%backbone onlyfull ckptargslogseval
xcit_small_12_p826M77.1%79.2%backbone onlyfull ckptargslogseval
xcit_medium_24_p1684M76.4%78.8%backbone onlyfull ckptargslogseval
xcit_medium_24_p884M77.9%80.3%backbone onlyfull ckptargslogseval
+ +### Pretrained models on PyTorch Hub +```python +import torch +vits16 = torch.hub.load('facebookresearch/dino:main', 'dino_vits16') +vits8 = torch.hub.load('facebookresearch/dino:main', 'dino_vits8') +vitb16 = torch.hub.load('facebookresearch/dino:main', 'dino_vitb16') +vitb8 = torch.hub.load('facebookresearch/dino:main', 'dino_vitb8') +xcit_small_12_p16 = torch.hub.load('facebookresearch/dino:main', 'dino_xcit_small_12_p16') +xcit_small_12_p8 = torch.hub.load('facebookresearch/dino:main', 'dino_xcit_small_12_p8') +xcit_medium_24_p16 = torch.hub.load('facebookresearch/dino:main', 'dino_xcit_medium_24_p16') +xcit_medium_24_p8 = torch.hub.load('facebookresearch/dino:main', 'dino_xcit_medium_24_p8') +resnet50 = torch.hub.load('facebookresearch/dino:main', 'dino_resnet50') +``` + +## Training + +### Documentation +Please install [PyTorch](https://pytorch.org/) and download the [ImageNet](https://imagenet.stanford.edu/) dataset. This codebase has been developed with python version 3.6, PyTorch version 1.7.1, CUDA 11.0 and torchvision 0.8.2. The exact arguments to reproduce the models presented in our paper can be found in the `args` column of the [pretrained models section](https://github.com/facebookresearch/dino#pretrained-models). For a glimpse at the full documentation of DINO training please run: +``` +python main_dino.py --help +``` + +### Vanilla DINO training :sauropod: +Run DINO with ViT-small network on a single node with 8 GPUs for 100 epochs with the following command. Training time is 1.75 day and the resulting checkpoint should reach 69.3% on k-NN eval and 74.0% on linear eval. We provide [training](https://dl.fbaipublicfiles.com/dino/example_runs_logs/dino_vanilla_deitsmall16_log.txt) and [linear evaluation](https://dl.fbaipublicfiles.com/dino/example_runs_logs/dino_vanilla_deitsmall16_eval.txt) logs (with batch size 256 at evaluation time) for this run to help reproducibility. +``` +python -m torch.distributed.launch --nproc_per_node=8 main_dino.py --arch vit_small --data_path /path/to/imagenet/train --output_dir /path/to/saving_dir +``` + +### Multi-node training +We use Slurm and [submitit](https://github.com/facebookincubator/submitit) (`pip install submitit`). To train on 2 nodes with 8 GPUs each (total 16 GPUs): +``` +python run_with_submitit.py --nodes 2 --ngpus 8 --arch vit_small --data_path /path/to/imagenet/train --output_dir /path/to/saving_dir +``` + +
+ +DINO with ViT-base network. + + +``` +python run_with_submitit.py --nodes 2 --ngpus 8 --use_volta32 --arch vit_base --data_path /path/to/imagenet/train --output_dir /path/to/saving_dir +``` + +
+ +### Boosting DINO performance :t-rex: +You can improve the performance of the vanilla run by: +- training for more epochs: `--epochs 300`, +- increasing the teacher temperature: `--teacher_temp 0.07 --warmup_teacher_temp_epochs 30`. +- removing last layer normalization (only safe with `--arch vit_small`): `--norm_last_layer false`, + +
+ +Full command. + + +``` +python run_with_submitit.py --arch vit_small --epochs 300 --teacher_temp 0.07 --warmup_teacher_temp_epochs 30 --norm_last_layer false --data_path /path/to/imagenet/train --output_dir /path/to/saving_dir +``` + +
+ +The resulting pretrained model should reach 73.3% on k-NN eval and 76.0% on linear eval. Training time is 2.6 days with 16 GPUs. We provide [training](https://dl.fbaipublicfiles.com/dino/example_runs_logs/dino_boost_deitsmall16_log.txt) and [linear evaluation](https://dl.fbaipublicfiles.com/dino/example_runs_logs/dino_boost_deitsmall16_eval.txt) logs (with batch size 256 at evaluation time) for this run to help reproducibility. + +### ResNet-50 and other convnets trainings +This code also works for training DINO on convolutional networks, like ResNet-50 for example. We highly recommend to adapt some optimization arguments in this case. For example following is a command to train DINO on ResNet-50 on a single node with 8 GPUs for 100 epochs. We provide [training logs](https://dl.fbaipublicfiles.com/dino/example_runs_logs/dino_rn50_log.txt) and [final checkpoint](https://dl.fbaipublicfiles.com/dino/example_runs_logs/dino_rn50_checkpoint.pth) for this run. +``` +python -m torch.distributed.launch --nproc_per_node=8 main_dino.py --arch resnet50 --optimizer sgd --lr 0.03 --weight_decay 1e-4 --weight_decay_end 1e-4 --global_crops_scale 0.14 1 --local_crops_scale 0.05 0.14 --data_path /path/to/imagenet/train --output_dir /path/to/saving_dir +``` + +## Self-attention visualization +You can look at the self-attention of the [CLS] token on the different heads of the last layer by running: +``` +python visualize_attention.py +``` + +
+ Self-attention from a Vision Transformer with 8x8 patches trained with DINO +
+ +## Self-attention video generation +You can generate videos like the one on the blog post with `video_generation.py`. + +https://user-images.githubusercontent.com/46140458/116817761-47885e80-ab68-11eb-9975-d61d5a919e13.mp4 + +Extract frames from input video and generate attention video: +``` +python video_generation.py --pretrained_weights dino_deitsmall8_pretrain.pth \ + --input_path input/video.mp4 \ + --output_path output/ \ + --fps 25 +``` + +Use folder of frames already extracted and generate attention video: +``` +python video_generation.py --pretrained_weights dino_deitsmall8_pretrain.pth \ + --input_path output/frames/ \ + --output_path output/ \ + --resize 256 \ +``` + +Only generate video from folder of attention maps images: +``` +python video_generation.py --input_path output/attention \ + --output_path output/ \ + --video_only \ + --video_format avi +``` + + +## Evaluation: k-NN classification on ImageNet +To evaluate a simple k-NN classifier with a single GPU on a pre-trained model, run: +``` +python -m torch.distributed.launch --nproc_per_node=1 eval_knn.py --data_path /path/to/imagenet +``` +If you choose not to specify `--pretrained_weights`, then DINO reference weights are used by default. If you want instead to evaluate checkpoints from a run of your own, you can run for example: +``` +python -m torch.distributed.launch --nproc_per_node=1 eval_knn.py --pretrained_weights /path/to/checkpoint.pth --checkpoint_key teacher --data_path /path/to/imagenet +``` + +## Evaluation: Linear classification on ImageNet +To train a supervised linear classifier on frozen weights on a single node with 8 gpus, run: +``` +python -m torch.distributed.launch --nproc_per_node=8 eval_linear.py --data_path /path/to/imagenet +``` + +We release the logs and weights from evaluating the different models: + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
archtop-1 ImageNetlinear evaluation
ViT-S/1677.0%linear weightslogs
ViT-S/879.7%linear weightslogs
ViT-B/1678.2%linear weightslogs
ViT-B/880.1%linear weightslogs
xcit_small_12_p1677.8%linear weightslogs
xcit_small_12_p879.2%linear weightslogs
xcit_medium_24_p1678.8%linear weightslogs
xcit_medium_24_p880.3%linear weightslogs
ResNet-5075.3%linear weightslogs
+ +You can check the performance of the pretrained weights on ImageNet validation set by running the following command lines: +``` +python eval_linear.py --evaluate --arch vit_small --patch_size 16 --data_path /path/to/imagenet/train +``` + +``` +python eval_linear.py --evaluate --arch vit_small --patch_size 8 --data_path /path/to/imagenet/train +``` + +``` +python eval_linear.py --evaluate --arch vit_base --patch_size 16 --n_last_blocks 1 --avgpool_patchtokens true --data_path /path/to/imagenet/train +``` + +``` +python eval_linear.py --evaluate --arch vit_base --patch_size 8 --n_last_blocks 1 --avgpool_patchtokens true --data_path /path/to/imagenet/train +``` + +``` +python eval_linear.py --evaluate --arch resnet50 --data_path /path/to/imagenet/train +``` + +## Evaluation: DAVIS 2017 Video object segmentation +Please verify that you're using pytorch version 1.7.1 since we are not able to reproduce the results with most recent pytorch 1.8.1 at the moment. + +**Step 1: Prepare DAVIS 2017 data** +``` +cd $HOME +git clone https://github.com/davisvideochallenge/davis-2017 && cd davis-2017 +./data/get_davis.sh +``` + +**Step 2: Video object segmentation** +``` +python eval_video_segmentation.py --data_path $HOME/davis-2017/DAVIS/ --output_dir /path/to/saving_dir +``` + +**Step 3: Evaluate the obtained segmentation** +``` +git clone https://github.com/davisvideochallenge/davis2017-evaluation $HOME/davis2017-evaluation +python $HOME/davis2017-evaluation/evaluation_method.py --task semi-supervised --results_path /path/to/saving_dir --davis_path $HOME/davis-2017/DAVIS/ +``` + +## Evaluation: Image Retrieval on revisited Oxford and Paris +Step 1: Prepare revisited Oxford and Paris by following [this repo](https://github.com/filipradenovic/revisitop). + +Step 2: Image retrieval (if you do not specify weights with `--pretrained_weights` then by default [DINO weights pretrained on Google Landmark v2 dataset](https://dl.fbaipublicfiles.com/dino/dino_vitsmall16_googlelandmark_pretrain/dino_vitsmall16_googlelandmark_pretrain.pth) will be used). + +Paris: +``` +python -m torch.distributed.launch --use_env --nproc_per_node=1 eval_image_retrieval.py --imsize 512 --multiscale 1 --data_path /path/to/revisited_paris_oxford/ --dataset rparis6k +``` + +Oxford: +``` +python -m torch.distributed.launch --use_env --nproc_per_node=1 eval_image_retrieval.py --imsize 224 --multiscale 0 --data_path /path/to/revisited_paris_oxford/ --dataset roxford5k +``` + +## Evaluation: Copy detection on Copydays +Step 1: Prepare [Copydays dataset](https://lear.inrialpes.fr/~jegou/data.php#copydays). + +Step 2 (opt): Prepare a set of image distractors and a set of images on which to learn the whitening operator. +In our paper, we use 10k random images from YFCC100M as distractors and 20k random images from YFCC100M (different from the distractors) for computing the whitening operation. + +Step 3: Run copy detection: +``` +python -m torch.distributed.launch --use_env --nproc_per_node=1 eval_copy_detection.py --data_path /path/to/copydays/ --whitening_path /path/to/whitening_data/ --distractors_path /path/to/distractors/ +``` +We report result on the strong subset. For example in the stdout from the command above we get: `eval on strong mAP=0.858`. + +## License +This repository is released under the Apache 2.0 license as found in the [LICENSE](LICENSE) file. + +## Citation +If you find this repository useful, please consider giving a star :star: and citation :t-rex:: +``` +@inproceedings{caron2021emerging, + title={Emerging Properties in Self-Supervised Vision Transformers}, + author={Caron, Mathilde and Touvron, Hugo and Misra, Ishan and J\'egou, Herv\'e and Mairal, Julien and Bojanowski, Piotr and Joulin, Armand}, + booktitle={Proceedings of the International Conference on Computer Vision (ICCV)}, + year={2021} +} +``` diff --git a/torch_hub/facebookresearch_dino_main/__pycache__/utils.cpython-310.pyc b/torch_hub/facebookresearch_dino_main/__pycache__/utils.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..c0aa79f0eeaa5aaddc84ef8f4b47e757eb10b71c Binary files /dev/null and b/torch_hub/facebookresearch_dino_main/__pycache__/utils.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_dino_main/eval_copy_detection.py b/torch_hub/facebookresearch_dino_main/eval_copy_detection.py new file mode 100644 index 0000000000000000000000000000000000000000..73dcd507893f204a47a5036cc61bd65b30cf1ead --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/eval_copy_detection.py @@ -0,0 +1,301 @@ +# Copyright (c) Facebook, Inc. and its affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import os +import sys +import pickle +import argparse + +import torch +from torch import nn +import torch.distributed as dist +import torch.backends.cudnn as cudnn +from torchvision import models as torchvision_models +from torchvision import transforms as pth_transforms +from PIL import Image, ImageFile +import numpy as np + +import utils +import vision_transformer as vits +from eval_knn import extract_features + + +class CopydaysDataset(): + def __init__(self, basedir): + self.basedir = basedir + self.block_names = ( + ['original', 'strong'] + + ['jpegqual/%d' % i for i in + [3, 5, 8, 10, 15, 20, 30, 50, 75]] + + ['crops/%d' % i for i in + [10, 15, 20, 30, 40, 50, 60, 70, 80]]) + self.nblocks = len(self.block_names) + + self.query_blocks = range(self.nblocks) + self.q_block_sizes = np.ones(self.nblocks, dtype=int) * 157 + self.q_block_sizes[1] = 229 + # search only among originals + self.database_blocks = [0] + + def get_block(self, i): + dirname = self.basedir + '/' + self.block_names[i] + fnames = [dirname + '/' + fname + for fname in sorted(os.listdir(dirname)) + if fname.endswith('.jpg')] + return fnames + + def get_block_filenames(self, subdir_name): + dirname = self.basedir + '/' + subdir_name + return [fname + for fname in sorted(os.listdir(dirname)) + if fname.endswith('.jpg')] + + def eval_result(self, ids, distances): + j0 = 0 + for i in range(self.nblocks): + j1 = j0 + self.q_block_sizes[i] + block_name = self.block_names[i] + I = ids[j0:j1] # block size + sum_AP = 0 + if block_name != 'strong': + # 1:1 mapping of files to names + positives_per_query = [[i] for i in range(j1 - j0)] + else: + originals = self.get_block_filenames('original') + strongs = self.get_block_filenames('strong') + + # check if prefixes match + positives_per_query = [ + [j for j, bname in enumerate(originals) + if bname[:4] == qname[:4]] + for qname in strongs] + + for qno, Iline in enumerate(I): + positives = positives_per_query[qno] + ranks = [] + for rank, bno in enumerate(Iline): + if bno in positives: + ranks.append(rank) + sum_AP += score_ap_from_ranks_1(ranks, len(positives)) + + print("eval on %s mAP=%.3f" % ( + block_name, sum_AP / (j1 - j0))) + j0 = j1 + + +# from the Holidays evaluation package +def score_ap_from_ranks_1(ranks, nres): + """ Compute the average precision of one search. + ranks = ordered list of ranks of true positives + nres = total number of positives in dataset + """ + + # accumulate trapezoids in PR-plot + ap = 0.0 + + # All have an x-size of: + recall_step = 1.0 / nres + + for ntp, rank in enumerate(ranks): + + # y-size on left side of trapezoid: + # ntp = nb of true positives so far + # rank = nb of retrieved items so far + if rank == 0: + precision_0 = 1.0 + else: + precision_0 = ntp / float(rank) + + # y-size on right side of trapezoid: + # ntp and rank are increased by one + precision_1 = (ntp + 1) / float(rank + 1) + + ap += (precision_1 + precision_0) * recall_step / 2.0 + + return ap + + +class ImgListDataset(torch.utils.data.Dataset): + def __init__(self, img_list, transform=None): + self.samples = img_list + self.transform = transform + + def __getitem__(self, i): + with open(self.samples[i], 'rb') as f: + img = Image.open(f) + img = img.convert('RGB') + if self.transform is not None: + img = self.transform(img) + return img, i + + def __len__(self): + return len(self.samples) + + +def is_image_file(s): + ext = s.split(".")[-1] + if ext in ['jpg', 'jpeg', 'png', 'ppm', 'bmp', 'pgm', 'tif', 'tiff', 'webp']: + return True + return False + + +@torch.no_grad() +def extract_features(image_list, model, args): + transform = pth_transforms.Compose([ + pth_transforms.Resize((args.imsize, args.imsize), interpolation=3), + pth_transforms.ToTensor(), + pth_transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), + ]) + tempdataset = ImgListDataset(image_list, transform=transform) + data_loader = torch.utils.data.DataLoader(tempdataset, batch_size=args.batch_size_per_gpu, + num_workers=args.num_workers, drop_last=False, + sampler=torch.utils.data.DistributedSampler(tempdataset, shuffle=False)) + features = None + for samples, index in utils.MetricLogger(delimiter=" ").log_every(data_loader, 10): + samples, index = samples.cuda(non_blocking=True), index.cuda(non_blocking=True) + feats = model.get_intermediate_layers(samples, n=1)[0].clone() + + cls_output_token = feats[:, 0, :] # [CLS] token + # GeM with exponent 4 for output patch tokens + b, h, w, d = len(samples), int(samples.shape[-2] / model.patch_embed.patch_size), int(samples.shape[-1] / model.patch_embed.patch_size), feats.shape[-1] + feats = feats[:, 1:, :].reshape(b, h, w, d) + feats = feats.clamp(min=1e-6).permute(0, 3, 1, 2) + feats = nn.functional.avg_pool2d(feats.pow(4), (h, w)).pow(1. / 4).reshape(b, -1) + # concatenate [CLS] token and GeM pooled patch tokens + feats = torch.cat((cls_output_token, feats), dim=1) + + # init storage feature matrix + if dist.get_rank() == 0 and features is None: + features = torch.zeros(len(data_loader.dataset), feats.shape[-1]) + if args.use_cuda: + features = features.cuda(non_blocking=True) + + # get indexes from all processes + y_all = torch.empty(dist.get_world_size(), index.size(0), dtype=index.dtype, device=index.device) + y_l = list(y_all.unbind(0)) + y_all_reduce = torch.distributed.all_gather(y_l, index, async_op=True) + y_all_reduce.wait() + index_all = torch.cat(y_l) + + # share features between processes + feats_all = torch.empty(dist.get_world_size(), feats.size(0), feats.size(1), + dtype=feats.dtype, device=feats.device) + output_l = list(feats_all.unbind(0)) + output_all_reduce = torch.distributed.all_gather(output_l, feats, async_op=True) + output_all_reduce.wait() + + # update storage feature matrix + if dist.get_rank() == 0: + if args.use_cuda: + features.index_copy_(0, index_all, torch.cat(output_l)) + else: + features.index_copy_(0, index_all.cpu(), torch.cat(output_l).cpu()) + return features # features is still None for every rank which is not 0 (main) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser('Copy detection on Copydays') + parser.add_argument('--data_path', default='/path/to/copydays/', type=str, + help="See https://lear.inrialpes.fr/~jegou/data.php#copydays") + parser.add_argument('--whitening_path', default='/path/to/whitening_data/', type=str, + help="""Path to directory with images used for computing the whitening operator. + In our paper, we use 20k random images from YFCC100M.""") + parser.add_argument('--distractors_path', default='/path/to/distractors/', type=str, + help="Path to directory with distractors images. In our paper, we use 10k random images from YFCC100M.") + parser.add_argument('--imsize', default=320, type=int, help='Image size (square image)') + parser.add_argument('--batch_size_per_gpu', default=16, type=int, help='Per-GPU batch-size') + parser.add_argument('--pretrained_weights', default='', type=str, help="Path to pretrained weights to evaluate.") + parser.add_argument('--use_cuda', default=True, type=utils.bool_flag) + parser.add_argument('--arch', default='vit_base', type=str, help='Architecture') + parser.add_argument('--patch_size', default=8, type=int, help='Patch resolution of the model.') + parser.add_argument("--checkpoint_key", default="teacher", type=str, + help='Key to use in the checkpoint (example: "teacher")') + parser.add_argument('--num_workers', default=10, type=int, help='Number of data loading workers per GPU.') + parser.add_argument("--dist_url", default="env://", type=str, help="""url used to set up + distributed training; see https://pytorch.org/docs/stable/distributed.html""") + parser.add_argument("--local_rank", default=0, type=int, help="Please ignore and do not set this argument.") + args = parser.parse_args() + + utils.init_distributed_mode(args) + print("git:\n {}\n".format(utils.get_sha())) + print("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items()))) + cudnn.benchmark = True + + # ============ building network ... ============ + if "vit" in args.arch: + model = vits.__dict__[args.arch](patch_size=args.patch_size, num_classes=0) + print(f"Model {args.arch} {args.patch_size}x{args.patch_size} built.") + else: + print(f"Architecture {args.arch} non supported") + sys.exit(1) + if args.use_cuda: + model.cuda() + model.eval() + utils.load_pretrained_weights(model, args.pretrained_weights, args.checkpoint_key, args.arch, args.patch_size) + + dataset = CopydaysDataset(args.data_path) + + # ============ Extract features ... ============ + # extract features for queries + queries = [] + for q in dataset.query_blocks: + queries.append(extract_features(dataset.get_block(q), model, args)) + if utils.get_rank() == 0: + queries = torch.cat(queries) + print(f"Extraction of queries features done. Shape: {queries.shape}") + + # extract features for database + database = [] + for b in dataset.database_blocks: + database.append(extract_features(dataset.get_block(b), model, args)) + + # extract features for distractors + if os.path.isdir(args.distractors_path): + print("Using distractors...") + list_distractors = [os.path.join(args.distractors_path, s) for s in os.listdir(args.distractors_path) if is_image_file(s)] + database.append(extract_features(list_distractors, model, args)) + if utils.get_rank() == 0: + database = torch.cat(database) + print(f"Extraction of database and distractors features done. Shape: {database.shape}") + + # ============ Whitening ... ============ + if os.path.isdir(args.whitening_path): + print(f"Extracting features on images from {args.whitening_path} for learning the whitening operator.") + list_whit = [os.path.join(args.whitening_path, s) for s in os.listdir(args.whitening_path) if is_image_file(s)] + features_for_whitening = extract_features(list_whit, model, args) + if utils.get_rank() == 0: + # center + mean_feature = torch.mean(features_for_whitening, dim=0) + database -= mean_feature + queries -= mean_feature + pca = utils.PCA(dim=database.shape[-1], whit=0.5) + # compute covariance + cov = torch.mm(features_for_whitening.T, features_for_whitening) / features_for_whitening.shape[0] + pca.train_pca(cov.cpu().numpy()) + database = pca.apply(database) + queries = pca.apply(queries) + + # ============ Copy detection ... ============ + if utils.get_rank() == 0: + # l2 normalize the features + database = nn.functional.normalize(database, dim=1, p=2) + queries = nn.functional.normalize(queries, dim=1, p=2) + + # similarity + similarity = torch.mm(queries, database.T) + distances, indices = similarity.topk(20, largest=True, sorted=True) + + # evaluate + retrieved = dataset.eval_result(indices, distances) + dist.barrier() + diff --git a/torch_hub/facebookresearch_dino_main/eval_image_retrieval.py b/torch_hub/facebookresearch_dino_main/eval_image_retrieval.py new file mode 100644 index 0000000000000000000000000000000000000000..999f8c9009a9abcc28308c5995c286f65b1522ac --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/eval_image_retrieval.py @@ -0,0 +1,201 @@ +# Copyright (c) Facebook, Inc. and its affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import os +import sys +import pickle +import argparse + +import torch +from torch import nn +import torch.distributed as dist +import torch.backends.cudnn as cudnn +from torchvision import models as torchvision_models +from torchvision import transforms as pth_transforms +from PIL import Image, ImageFile +import numpy as np + +import utils +import vision_transformer as vits +from eval_knn import extract_features + + +class OxfordParisDataset(torch.utils.data.Dataset): + def __init__(self, dir_main, dataset, split, transform=None, imsize=None): + if dataset not in ['roxford5k', 'rparis6k']: + raise ValueError('Unknown dataset: {}!'.format(dataset)) + + # loading imlist, qimlist, and gnd, in cfg as a dict + gnd_fname = os.path.join(dir_main, dataset, 'gnd_{}.pkl'.format(dataset)) + with open(gnd_fname, 'rb') as f: + cfg = pickle.load(f) + cfg['gnd_fname'] = gnd_fname + cfg['ext'] = '.jpg' + cfg['qext'] = '.jpg' + cfg['dir_data'] = os.path.join(dir_main, dataset) + cfg['dir_images'] = os.path.join(cfg['dir_data'], 'jpg') + cfg['n'] = len(cfg['imlist']) + cfg['nq'] = len(cfg['qimlist']) + cfg['im_fname'] = config_imname + cfg['qim_fname'] = config_qimname + cfg['dataset'] = dataset + self.cfg = cfg + + self.samples = cfg["qimlist"] if split == "query" else cfg["imlist"] + self.transform = transform + self.imsize = imsize + + def __len__(self): + return len(self.samples) + + def __getitem__(self, index): + path = os.path.join(self.cfg["dir_images"], self.samples[index] + ".jpg") + ImageFile.LOAD_TRUNCATED_IMAGES = True + with open(path, 'rb') as f: + img = Image.open(f) + img = img.convert('RGB') + if self.imsize is not None: + img.thumbnail((self.imsize, self.imsize), Image.ANTIALIAS) + if self.transform is not None: + img = self.transform(img) + return img, index + + +def config_imname(cfg, i): + return os.path.join(cfg['dir_images'], cfg['imlist'][i] + cfg['ext']) + + +def config_qimname(cfg, i): + return os.path.join(cfg['dir_images'], cfg['qimlist'][i] + cfg['qext']) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser('Image Retrieval on revisited Paris and Oxford') + parser.add_argument('--data_path', default='/path/to/revisited_paris_oxford/', type=str) + parser.add_argument('--dataset', default='roxford5k', type=str, choices=['roxford5k', 'rparis6k']) + parser.add_argument('--multiscale', default=False, type=utils.bool_flag) + parser.add_argument('--imsize', default=224, type=int, help='Image size') + parser.add_argument('--pretrained_weights', default='', type=str, help="Path to pretrained weights to evaluate.") + parser.add_argument('--use_cuda', default=True, type=utils.bool_flag) + parser.add_argument('--arch', default='vit_small', type=str, help='Architecture') + parser.add_argument('--patch_size', default=16, type=int, help='Patch resolution of the model.') + parser.add_argument("--checkpoint_key", default="teacher", type=str, + help='Key to use in the checkpoint (example: "teacher")') + parser.add_argument('--num_workers', default=10, type=int, help='Number of data loading workers per GPU.') + parser.add_argument("--dist_url", default="env://", type=str, help="""url used to set up + distributed training; see https://pytorch.org/docs/stable/distributed.html""") + parser.add_argument("--local_rank", default=0, type=int, help="Please ignore and do not set this argument.") + args = parser.parse_args() + + utils.init_distributed_mode(args) + print("git:\n {}\n".format(utils.get_sha())) + print("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items()))) + cudnn.benchmark = True + + # ============ preparing data ... ============ + transform = pth_transforms.Compose([ + pth_transforms.ToTensor(), + pth_transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), + ]) + dataset_train = OxfordParisDataset(args.data_path, args.dataset, split="train", transform=transform, imsize=args.imsize) + dataset_query = OxfordParisDataset(args.data_path, args.dataset, split="query", transform=transform, imsize=args.imsize) + sampler = torch.utils.data.DistributedSampler(dataset_train, shuffle=False) + data_loader_train = torch.utils.data.DataLoader( + dataset_train, + sampler=sampler, + batch_size=1, + num_workers=args.num_workers, + pin_memory=True, + drop_last=False, + ) + data_loader_query = torch.utils.data.DataLoader( + dataset_query, + batch_size=1, + num_workers=args.num_workers, + pin_memory=True, + drop_last=False, + ) + print(f"train: {len(dataset_train)} imgs / query: {len(dataset_query)} imgs") + + # ============ building network ... ============ + if "vit" in args.arch: + model = vits.__dict__[args.arch](patch_size=args.patch_size, num_classes=0) + print(f"Model {args.arch} {args.patch_size}x{args.patch_size} built.") + elif "xcit" in args.arch: + model = torch.hub.load('facebookresearch/xcit:main', args.arch, num_classes=0) + elif args.arch in torchvision_models.__dict__.keys(): + model = torchvision_models.__dict__[args.arch](num_classes=0) + else: + print(f"Architecture {args.arch} non supported") + sys.exit(1) + if args.use_cuda: + model.cuda() + model.eval() + + # load pretrained weights + if os.path.isfile(args.pretrained_weights): + state_dict = torch.load(args.pretrained_weights, map_location="cpu") + if args.checkpoint_key is not None and args.checkpoint_key in state_dict: + print(f"Take key {args.checkpoint_key} in provided checkpoint dict") + state_dict = state_dict[args.checkpoint_key] + # remove `module.` prefix + state_dict = {k.replace("module.", ""): v for k, v in state_dict.items()} + # remove `backbone.` prefix induced by multicrop wrapper + state_dict = {k.replace("backbone.", ""): v for k, v in state_dict.items()} + msg = model.load_state_dict(state_dict, strict=False) + print('Pretrained weights found at {} and loaded with msg: {}'.format(args.pretrained_weights, msg)) + elif args.arch == "vit_small" and args.patch_size == 16: + print("Since no pretrained weights have been provided, we load pretrained DINO weights on Google Landmark v2.") + model.load_state_dict(torch.hub.load_state_dict_from_url(url="https://dl.fbaipublicfiles.com/dino/dino_vitsmall16_googlelandmark_pretrain/dino_vitsmall16_googlelandmark_pretrain.pth")) + else: + print("Warning: We use random weights.") + + ############################################################################ + # Step 1: extract features + train_features = extract_features(model, data_loader_train, args.use_cuda, multiscale=args.multiscale) + query_features = extract_features(model, data_loader_query, args.use_cuda, multiscale=args.multiscale) + + if utils.get_rank() == 0: # only rank 0 will work from now on + # normalize features + train_features = nn.functional.normalize(train_features, dim=1, p=2) + query_features = nn.functional.normalize(query_features, dim=1, p=2) + + ############################################################################ + # Step 2: similarity + sim = torch.mm(train_features, query_features.T) + ranks = torch.argsort(-sim, dim=0).cpu().numpy() + + ############################################################################ + # Step 3: evaluate + gnd = dataset_train.cfg['gnd'] + # evaluate ranks + ks = [1, 5, 10] + # search for easy & hard + gnd_t = [] + for i in range(len(gnd)): + g = {} + g['ok'] = np.concatenate([gnd[i]['easy'], gnd[i]['hard']]) + g['junk'] = np.concatenate([gnd[i]['junk']]) + gnd_t.append(g) + mapM, apsM, mprM, prsM = utils.compute_map(ranks, gnd_t, ks) + # search for hard + gnd_t = [] + for i in range(len(gnd)): + g = {} + g['ok'] = np.concatenate([gnd[i]['hard']]) + g['junk'] = np.concatenate([gnd[i]['junk'], gnd[i]['easy']]) + gnd_t.append(g) + mapH, apsH, mprH, prsH = utils.compute_map(ranks, gnd_t, ks) + print('>> {}: mAP M: {}, H: {}'.format(args.dataset, np.around(mapM*100, decimals=2), np.around(mapH*100, decimals=2))) + print('>> {}: mP@k{} M: {}, H: {}'.format(args.dataset, np.array(ks), np.around(mprM*100, decimals=2), np.around(mprH*100, decimals=2))) + dist.barrier() diff --git a/torch_hub/facebookresearch_dino_main/eval_linear.py b/torch_hub/facebookresearch_dino_main/eval_linear.py new file mode 100644 index 0000000000000000000000000000000000000000..cdef16b473d216889b493aa0c7a63e15f945092c --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/eval_linear.py @@ -0,0 +1,281 @@ +# Copyright (c) Facebook, Inc. and its affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import os +import argparse +import json +from pathlib import Path + +import torch +from torch import nn +import torch.distributed as dist +import torch.backends.cudnn as cudnn +from torchvision import datasets +from torchvision import transforms as pth_transforms +from torchvision import models as torchvision_models + +import utils +import vision_transformer as vits + + +def eval_linear(args): + utils.init_distributed_mode(args) + print("git:\n {}\n".format(utils.get_sha())) + print("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items()))) + cudnn.benchmark = True + + # ============ building network ... ============ + # if the network is a Vision Transformer (i.e. vit_tiny, vit_small, vit_base) + if args.arch in vits.__dict__.keys(): + model = vits.__dict__[args.arch](patch_size=args.patch_size, num_classes=0) + embed_dim = model.embed_dim * (args.n_last_blocks + int(args.avgpool_patchtokens)) + # if the network is a XCiT + elif "xcit" in args.arch: + model = torch.hub.load('facebookresearch/xcit:main', args.arch, num_classes=0) + embed_dim = model.embed_dim + # otherwise, we check if the architecture is in torchvision models + elif args.arch in torchvision_models.__dict__.keys(): + model = torchvision_models.__dict__[args.arch]() + embed_dim = model.fc.weight.shape[1] + model.fc = nn.Identity() + else: + print(f"Unknow architecture: {args.arch}") + sys.exit(1) + model.cuda() + model.eval() + # load weights to evaluate + utils.load_pretrained_weights(model, args.pretrained_weights, args.checkpoint_key, args.arch, args.patch_size) + print(f"Model {args.arch} built.") + + linear_classifier = LinearClassifier(embed_dim, num_labels=args.num_labels) + linear_classifier = linear_classifier.cuda() + linear_classifier = nn.parallel.DistributedDataParallel(linear_classifier, device_ids=[args.gpu]) + + # ============ preparing data ... ============ + val_transform = pth_transforms.Compose([ + pth_transforms.Resize(256, interpolation=3), + pth_transforms.CenterCrop(224), + pth_transforms.ToTensor(), + pth_transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), + ]) + dataset_val = datasets.ImageFolder(os.path.join(args.data_path, "val"), transform=val_transform) + val_loader = torch.utils.data.DataLoader( + dataset_val, + batch_size=args.batch_size_per_gpu, + num_workers=args.num_workers, + pin_memory=True, + ) + + if args.evaluate: + utils.load_pretrained_linear_weights(linear_classifier, args.arch, args.patch_size) + test_stats = validate_network(val_loader, model, linear_classifier, args.n_last_blocks, args.avgpool_patchtokens) + print(f"Accuracy of the network on the {len(dataset_val)} test images: {test_stats['acc1']:.1f}%") + return + + train_transform = pth_transforms.Compose([ + pth_transforms.RandomResizedCrop(224), + pth_transforms.RandomHorizontalFlip(), + pth_transforms.ToTensor(), + pth_transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), + ]) + dataset_train = datasets.ImageFolder(os.path.join(args.data_path, "train"), transform=train_transform) + sampler = torch.utils.data.distributed.DistributedSampler(dataset_train) + train_loader = torch.utils.data.DataLoader( + dataset_train, + sampler=sampler, + batch_size=args.batch_size_per_gpu, + num_workers=args.num_workers, + pin_memory=True, + ) + print(f"Data loaded with {len(dataset_train)} train and {len(dataset_val)} val imgs.") + + # set optimizer + optimizer = torch.optim.SGD( + linear_classifier.parameters(), + args.lr * (args.batch_size_per_gpu * utils.get_world_size()) / 256., # linear scaling rule + momentum=0.9, + weight_decay=0, # we do not apply weight decay + ) + scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, args.epochs, eta_min=0) + + # Optionally resume from a checkpoint + to_restore = {"epoch": 0, "best_acc": 0.} + utils.restart_from_checkpoint( + os.path.join(args.output_dir, "checkpoint.pth.tar"), + run_variables=to_restore, + state_dict=linear_classifier, + optimizer=optimizer, + scheduler=scheduler, + ) + start_epoch = to_restore["epoch"] + best_acc = to_restore["best_acc"] + + for epoch in range(start_epoch, args.epochs): + train_loader.sampler.set_epoch(epoch) + + train_stats = train(model, linear_classifier, optimizer, train_loader, epoch, args.n_last_blocks, args.avgpool_patchtokens) + scheduler.step() + + log_stats = {**{f'train_{k}': v for k, v in train_stats.items()}, + 'epoch': epoch} + if epoch % args.val_freq == 0 or epoch == args.epochs - 1: + test_stats = validate_network(val_loader, model, linear_classifier, args.n_last_blocks, args.avgpool_patchtokens) + print(f"Accuracy at epoch {epoch} of the network on the {len(dataset_val)} test images: {test_stats['acc1']:.1f}%") + best_acc = max(best_acc, test_stats["acc1"]) + print(f'Max accuracy so far: {best_acc:.2f}%') + log_stats = {**{k: v for k, v in log_stats.items()}, + **{f'test_{k}': v for k, v in test_stats.items()}} + if utils.is_main_process(): + with (Path(args.output_dir) / "log.txt").open("a") as f: + f.write(json.dumps(log_stats) + "\n") + save_dict = { + "epoch": epoch + 1, + "state_dict": linear_classifier.state_dict(), + "optimizer": optimizer.state_dict(), + "scheduler": scheduler.state_dict(), + "best_acc": best_acc, + } + torch.save(save_dict, os.path.join(args.output_dir, "checkpoint.pth.tar")) + print("Training of the supervised linear classifier on frozen features completed.\n" + "Top-1 test accuracy: {acc:.1f}".format(acc=best_acc)) + + +def train(model, linear_classifier, optimizer, loader, epoch, n, avgpool): + linear_classifier.train() + metric_logger = utils.MetricLogger(delimiter=" ") + metric_logger.add_meter('lr', utils.SmoothedValue(window_size=1, fmt='{value:.6f}')) + header = 'Epoch: [{}]'.format(epoch) + for (inp, target) in metric_logger.log_every(loader, 20, header): + # move to gpu + inp = inp.cuda(non_blocking=True) + target = target.cuda(non_blocking=True) + + # forward + with torch.no_grad(): + if "vit" in args.arch: + intermediate_output = model.get_intermediate_layers(inp, n) + output = torch.cat([x[:, 0] for x in intermediate_output], dim=-1) + if avgpool: + output = torch.cat((output.unsqueeze(-1), torch.mean(intermediate_output[-1][:, 1:], dim=1).unsqueeze(-1)), dim=-1) + output = output.reshape(output.shape[0], -1) + else: + output = model(inp) + output = linear_classifier(output) + + # compute cross entropy loss + loss = nn.CrossEntropyLoss()(output, target) + + # compute the gradients + optimizer.zero_grad() + loss.backward() + + # step + optimizer.step() + + # log + torch.cuda.synchronize() + metric_logger.update(loss=loss.item()) + metric_logger.update(lr=optimizer.param_groups[0]["lr"]) + # gather the stats from all processes + metric_logger.synchronize_between_processes() + print("Averaged stats:", metric_logger) + return {k: meter.global_avg for k, meter in metric_logger.meters.items()} + + +@torch.no_grad() +def validate_network(val_loader, model, linear_classifier, n, avgpool): + linear_classifier.eval() + metric_logger = utils.MetricLogger(delimiter=" ") + header = 'Test:' + for inp, target in metric_logger.log_every(val_loader, 20, header): + # move to gpu + inp = inp.cuda(non_blocking=True) + target = target.cuda(non_blocking=True) + + # forward + with torch.no_grad(): + if "vit" in args.arch: + intermediate_output = model.get_intermediate_layers(inp, n) + output = torch.cat([x[:, 0] for x in intermediate_output], dim=-1) + if avgpool: + output = torch.cat((output.unsqueeze(-1), torch.mean(intermediate_output[-1][:, 1:], dim=1).unsqueeze(-1)), dim=-1) + output = output.reshape(output.shape[0], -1) + else: + output = model(inp) + output = linear_classifier(output) + loss = nn.CrossEntropyLoss()(output, target) + + if linear_classifier.module.num_labels >= 5: + acc1, acc5 = utils.accuracy(output, target, topk=(1, 5)) + else: + acc1, = utils.accuracy(output, target, topk=(1,)) + + batch_size = inp.shape[0] + metric_logger.update(loss=loss.item()) + metric_logger.meters['acc1'].update(acc1.item(), n=batch_size) + if linear_classifier.module.num_labels >= 5: + metric_logger.meters['acc5'].update(acc5.item(), n=batch_size) + if linear_classifier.module.num_labels >= 5: + print('* Acc@1 {top1.global_avg:.3f} Acc@5 {top5.global_avg:.3f} loss {losses.global_avg:.3f}' + .format(top1=metric_logger.acc1, top5=metric_logger.acc5, losses=metric_logger.loss)) + else: + print('* Acc@1 {top1.global_avg:.3f} loss {losses.global_avg:.3f}' + .format(top1=metric_logger.acc1, losses=metric_logger.loss)) + return {k: meter.global_avg for k, meter in metric_logger.meters.items()} + + +class LinearClassifier(nn.Module): + """Linear layer to train on top of frozen features""" + def __init__(self, dim, num_labels=1000): + super(LinearClassifier, self).__init__() + self.num_labels = num_labels + self.linear = nn.Linear(dim, num_labels) + self.linear.weight.data.normal_(mean=0.0, std=0.01) + self.linear.bias.data.zero_() + + def forward(self, x): + # flatten + x = x.view(x.size(0), -1) + + # linear layer + return self.linear(x) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser('Evaluation with linear classification on ImageNet') + parser.add_argument('--n_last_blocks', default=4, type=int, help="""Concatenate [CLS] tokens + for the `n` last blocks. We use `n=4` when evaluating ViT-Small and `n=1` with ViT-Base.""") + parser.add_argument('--avgpool_patchtokens', default=False, type=utils.bool_flag, + help="""Whether ot not to concatenate the global average pooled features to the [CLS] token. + We typically set this to False for ViT-Small and to True with ViT-Base.""") + parser.add_argument('--arch', default='vit_small', type=str, help='Architecture') + parser.add_argument('--patch_size', default=16, type=int, help='Patch resolution of the model.') + parser.add_argument('--pretrained_weights', default='', type=str, help="Path to pretrained weights to evaluate.") + parser.add_argument("--checkpoint_key", default="teacher", type=str, help='Key to use in the checkpoint (example: "teacher")') + parser.add_argument('--epochs', default=100, type=int, help='Number of epochs of training.') + parser.add_argument("--lr", default=0.001, type=float, help="""Learning rate at the beginning of + training (highest LR used during training). The learning rate is linearly scaled + with the batch size, and specified here for a reference batch size of 256. + We recommend tweaking the LR depending on the checkpoint evaluated.""") + parser.add_argument('--batch_size_per_gpu', default=128, type=int, help='Per-GPU batch-size') + parser.add_argument("--dist_url", default="env://", type=str, help="""url used to set up + distributed training; see https://pytorch.org/docs/stable/distributed.html""") + parser.add_argument("--local_rank", default=0, type=int, help="Please ignore and do not set this argument.") + parser.add_argument('--data_path', default='/path/to/imagenet/', type=str) + parser.add_argument('--num_workers', default=10, type=int, help='Number of data loading workers per GPU.') + parser.add_argument('--val_freq', default=1, type=int, help="Epoch frequency for validation.") + parser.add_argument('--output_dir', default=".", help='Path to save logs and checkpoints') + parser.add_argument('--num_labels', default=1000, type=int, help='Number of labels for linear classifier') + parser.add_argument('--evaluate', dest='evaluate', action='store_true', help='evaluate model on validation set') + args = parser.parse_args() + eval_linear(args) diff --git a/torch_hub/facebookresearch_dino_main/main_dino.py b/torch_hub/facebookresearch_dino_main/main_dino.py new file mode 100644 index 0000000000000000000000000000000000000000..cade9873dcb1d1c6c69ddba61dbdcf3f01dd7540 --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/main_dino.py @@ -0,0 +1,471 @@ +# Copyright (c) Facebook, Inc. and its affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import argparse +import os +import sys +import datetime +import time +import math +import json +from pathlib import Path + +import numpy as np +from PIL import Image +import torch +import torch.nn as nn +import torch.distributed as dist +import torch.backends.cudnn as cudnn +import torch.nn.functional as F +from torchvision import datasets, transforms +from torchvision import models as torchvision_models + +import utils +import vision_transformer as vits +from vision_transformer import DINOHead + +torchvision_archs = sorted(name for name in torchvision_models.__dict__ + if name.islower() and not name.startswith("__") + and callable(torchvision_models.__dict__[name])) + +def get_args_parser(): + parser = argparse.ArgumentParser('DINO', add_help=False) + + # Model parameters + parser.add_argument('--arch', default='vit_small', type=str, + choices=['vit_tiny', 'vit_small', 'vit_base', 'xcit', 'deit_tiny', 'deit_small'] \ + + torchvision_archs + torch.hub.list("facebookresearch/xcit:main"), + help="""Name of architecture to train. For quick experiments with ViTs, + we recommend using vit_tiny or vit_small.""") + parser.add_argument('--patch_size', default=16, type=int, help="""Size in pixels + of input square patches - default 16 (for 16x16 patches). Using smaller + values leads to better performance but requires more memory. Applies only + for ViTs (vit_tiny, vit_small and vit_base). If <16, we recommend disabling + mixed precision training (--use_fp16 false) to avoid unstabilities.""") + parser.add_argument('--out_dim', default=65536, type=int, help="""Dimensionality of + the DINO head output. For complex and large datasets large values (like 65k) work well.""") + parser.add_argument('--norm_last_layer', default=True, type=utils.bool_flag, + help="""Whether or not to weight normalize the last layer of the DINO head. + Not normalizing leads to better performance but can make the training unstable. + In our experiments, we typically set this paramater to False with vit_small and True with vit_base.""") + parser.add_argument('--momentum_teacher', default=0.996, type=float, help="""Base EMA + parameter for teacher update. The value is increased to 1 during training with cosine schedule. + We recommend setting a higher value with small batches: for example use 0.9995 with batch size of 256.""") + parser.add_argument('--use_bn_in_head', default=False, type=utils.bool_flag, + help="Whether to use batch normalizations in projection head (Default: False)") + + # Temperature teacher parameters + parser.add_argument('--warmup_teacher_temp', default=0.04, type=float, + help="""Initial value for the teacher temperature: 0.04 works well in most cases. + Try decreasing it if the training loss does not decrease.""") + parser.add_argument('--teacher_temp', default=0.04, type=float, help="""Final value (after linear warmup) + of the teacher temperature. For most experiments, anything above 0.07 is unstable. We recommend + starting with the default value of 0.04 and increase this slightly if needed.""") + parser.add_argument('--warmup_teacher_temp_epochs', default=0, type=int, + help='Number of warmup epochs for the teacher temperature (Default: 30).') + + # Training/Optimization parameters + parser.add_argument('--use_fp16', type=utils.bool_flag, default=True, help="""Whether or not + to use half precision for training. Improves training time and memory requirements, + but can provoke instability and slight decay of performance. We recommend disabling + mixed precision if the loss is unstable, if reducing the patch size or if training with bigger ViTs.""") + parser.add_argument('--weight_decay', type=float, default=0.04, help="""Initial value of the + weight decay. With ViT, a smaller value at the beginning of training works well.""") + parser.add_argument('--weight_decay_end', type=float, default=0.4, help="""Final value of the + weight decay. We use a cosine schedule for WD and using a larger decay by + the end of training improves performance for ViTs.""") + parser.add_argument('--clip_grad', type=float, default=3.0, help="""Maximal parameter + gradient norm if using gradient clipping. Clipping with norm .3 ~ 1.0 can + help optimization for larger ViT architectures. 0 for disabling.""") + parser.add_argument('--batch_size_per_gpu', default=64, type=int, + help='Per-GPU batch-size : number of distinct images loaded on one GPU.') + parser.add_argument('--epochs', default=100, type=int, help='Number of epochs of training.') + parser.add_argument('--freeze_last_layer', default=1, type=int, help="""Number of epochs + during which we keep the output layer fixed. Typically doing so during + the first epoch helps training. Try increasing this value if the loss does not decrease.""") + parser.add_argument("--lr", default=0.0005, type=float, help="""Learning rate at the end of + linear warmup (highest LR used during training). The learning rate is linearly scaled + with the batch size, and specified here for a reference batch size of 256.""") + parser.add_argument("--warmup_epochs", default=10, type=int, + help="Number of epochs for the linear learning-rate warm up.") + parser.add_argument('--min_lr', type=float, default=1e-6, help="""Target LR at the + end of optimization. We use a cosine LR schedule with linear warmup.""") + parser.add_argument('--optimizer', default='adamw', type=str, + choices=['adamw', 'sgd', 'lars'], help="""Type of optimizer. We recommend using adamw with ViTs.""") + parser.add_argument('--drop_path_rate', type=float, default=0.1, help="stochastic depth rate") + + # Multi-crop parameters + parser.add_argument('--global_crops_scale', type=float, nargs='+', default=(0.4, 1.), + help="""Scale range of the cropped image before resizing, relatively to the origin image. + Used for large global view cropping. When disabling multi-crop (--local_crops_number 0), we + recommand using a wider range of scale ("--global_crops_scale 0.14 1." for example)""") + parser.add_argument('--local_crops_number', type=int, default=8, help="""Number of small + local views to generate. Set this parameter to 0 to disable multi-crop training. + When disabling multi-crop we recommend to use "--global_crops_scale 0.14 1." """) + parser.add_argument('--local_crops_scale', type=float, nargs='+', default=(0.05, 0.4), + help="""Scale range of the cropped image before resizing, relatively to the origin image. + Used for small local view cropping of multi-crop.""") + + # Misc + parser.add_argument('--data_path', default='/path/to/imagenet/train/', type=str, + help='Please specify path to the ImageNet training data.') + parser.add_argument('--output_dir', default=".", type=str, help='Path to save logs and checkpoints.') + parser.add_argument('--saveckp_freq', default=20, type=int, help='Save checkpoint every x epochs.') + parser.add_argument('--seed', default=0, type=int, help='Random seed.') + parser.add_argument('--num_workers', default=10, type=int, help='Number of data loading workers per GPU.') + parser.add_argument("--dist_url", default="env://", type=str, help="""url used to set up + distributed training; see https://pytorch.org/docs/stable/distributed.html""") + parser.add_argument("--local_rank", default=0, type=int, help="Please ignore and do not set this argument.") + return parser + + +def train_dino(args): + utils.init_distributed_mode(args) + utils.fix_random_seeds(args.seed) + print("git:\n {}\n".format(utils.get_sha())) + print("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items()))) + cudnn.benchmark = True + + # ============ preparing data ... ============ + transform = DataAugmentationDINO( + args.global_crops_scale, + args.local_crops_scale, + args.local_crops_number, + ) + dataset = datasets.ImageFolder(args.data_path, transform=transform) + sampler = torch.utils.data.DistributedSampler(dataset, shuffle=True) + data_loader = torch.utils.data.DataLoader( + dataset, + sampler=sampler, + batch_size=args.batch_size_per_gpu, + num_workers=args.num_workers, + pin_memory=True, + drop_last=True, + ) + print(f"Data loaded: there are {len(dataset)} images.") + + # ============ building student and teacher networks ... ============ + # we changed the name DeiT-S for ViT-S to avoid confusions + args.arch = args.arch.replace("deit", "vit") + # if the network is a Vision Transformer (i.e. vit_tiny, vit_small, vit_base) + if args.arch in vits.__dict__.keys(): + student = vits.__dict__[args.arch]( + patch_size=args.patch_size, + drop_path_rate=args.drop_path_rate, # stochastic depth + ) + teacher = vits.__dict__[args.arch](patch_size=args.patch_size) + embed_dim = student.embed_dim + # if the network is a XCiT + elif args.arch in torch.hub.list("facebookresearch/xcit:main"): + student = torch.hub.load('facebookresearch/xcit:main', args.arch, + pretrained=False, drop_path_rate=args.drop_path_rate) + teacher = torch.hub.load('facebookresearch/xcit:main', args.arch, pretrained=False) + embed_dim = student.embed_dim + # otherwise, we check if the architecture is in torchvision models + elif args.arch in torchvision_models.__dict__.keys(): + student = torchvision_models.__dict__[args.arch]() + teacher = torchvision_models.__dict__[args.arch]() + embed_dim = student.fc.weight.shape[1] + else: + print(f"Unknow architecture: {args.arch}") + + # multi-crop wrapper handles forward with inputs of different resolutions + student = utils.MultiCropWrapper(student, DINOHead( + embed_dim, + args.out_dim, + use_bn=args.use_bn_in_head, + norm_last_layer=args.norm_last_layer, + )) + teacher = utils.MultiCropWrapper( + teacher, + DINOHead(embed_dim, args.out_dim, args.use_bn_in_head), + ) + # move networks to gpu + student, teacher = student.cuda(), teacher.cuda() + # synchronize batch norms (if any) + if utils.has_batchnorms(student): + student = nn.SyncBatchNorm.convert_sync_batchnorm(student) + teacher = nn.SyncBatchNorm.convert_sync_batchnorm(teacher) + + # we need DDP wrapper to have synchro batch norms working... + teacher = nn.parallel.DistributedDataParallel(teacher, device_ids=[args.gpu]) + teacher_without_ddp = teacher.module + else: + # teacher_without_ddp and teacher are the same thing + teacher_without_ddp = teacher + student = nn.parallel.DistributedDataParallel(student, device_ids=[args.gpu]) + # teacher and student start with the same weights + teacher_without_ddp.load_state_dict(student.module.state_dict()) + # there is no backpropagation through the teacher, so no need for gradients + for p in teacher.parameters(): + p.requires_grad = False + print(f"Student and Teacher are built: they are both {args.arch} network.") + + # ============ preparing loss ... ============ + dino_loss = DINOLoss( + args.out_dim, + args.local_crops_number + 2, # total number of crops = 2 global crops + local_crops_number + args.warmup_teacher_temp, + args.teacher_temp, + args.warmup_teacher_temp_epochs, + args.epochs, + ).cuda() + + # ============ preparing optimizer ... ============ + params_groups = utils.get_params_groups(student) + if args.optimizer == "adamw": + optimizer = torch.optim.AdamW(params_groups) # to use with ViTs + elif args.optimizer == "sgd": + optimizer = torch.optim.SGD(params_groups, lr=0, momentum=0.9) # lr is set by scheduler + elif args.optimizer == "lars": + optimizer = utils.LARS(params_groups) # to use with convnet and large batches + # for mixed precision training + fp16_scaler = None + if args.use_fp16: + fp16_scaler = torch.cuda.amp.GradScaler() + + # ============ init schedulers ... ============ + lr_schedule = utils.cosine_scheduler( + args.lr * (args.batch_size_per_gpu * utils.get_world_size()) / 256., # linear scaling rule + args.min_lr, + args.epochs, len(data_loader), + warmup_epochs=args.warmup_epochs, + ) + wd_schedule = utils.cosine_scheduler( + args.weight_decay, + args.weight_decay_end, + args.epochs, len(data_loader), + ) + # momentum parameter is increased to 1. during training with a cosine schedule + momentum_schedule = utils.cosine_scheduler(args.momentum_teacher, 1, + args.epochs, len(data_loader)) + print(f"Loss, optimizer and schedulers ready.") + + # ============ optionally resume training ... ============ + to_restore = {"epoch": 0} + utils.restart_from_checkpoint( + os.path.join(args.output_dir, "checkpoint.pth"), + run_variables=to_restore, + student=student, + teacher=teacher, + optimizer=optimizer, + fp16_scaler=fp16_scaler, + dino_loss=dino_loss, + ) + start_epoch = to_restore["epoch"] + + start_time = time.time() + print("Starting DINO training !") + for epoch in range(start_epoch, args.epochs): + data_loader.sampler.set_epoch(epoch) + + # ============ training one epoch of DINO ... ============ + train_stats = train_one_epoch(student, teacher, teacher_without_ddp, dino_loss, + data_loader, optimizer, lr_schedule, wd_schedule, momentum_schedule, + epoch, fp16_scaler, args) + + # ============ writing logs ... ============ + save_dict = { + 'student': student.state_dict(), + 'teacher': teacher.state_dict(), + 'optimizer': optimizer.state_dict(), + 'epoch': epoch + 1, + 'args': args, + 'dino_loss': dino_loss.state_dict(), + } + if fp16_scaler is not None: + save_dict['fp16_scaler'] = fp16_scaler.state_dict() + utils.save_on_master(save_dict, os.path.join(args.output_dir, 'checkpoint.pth')) + if args.saveckp_freq and epoch % args.saveckp_freq == 0: + utils.save_on_master(save_dict, os.path.join(args.output_dir, f'checkpoint{epoch:04}.pth')) + log_stats = {**{f'train_{k}': v for k, v in train_stats.items()}, + 'epoch': epoch} + if utils.is_main_process(): + with (Path(args.output_dir) / "log.txt").open("a") as f: + f.write(json.dumps(log_stats) + "\n") + total_time = time.time() - start_time + total_time_str = str(datetime.timedelta(seconds=int(total_time))) + print('Training time {}'.format(total_time_str)) + + +def train_one_epoch(student, teacher, teacher_without_ddp, dino_loss, data_loader, + optimizer, lr_schedule, wd_schedule, momentum_schedule,epoch, + fp16_scaler, args): + metric_logger = utils.MetricLogger(delimiter=" ") + header = 'Epoch: [{}/{}]'.format(epoch, args.epochs) + for it, (images, _) in enumerate(metric_logger.log_every(data_loader, 10, header)): + # update weight decay and learning rate according to their schedule + it = len(data_loader) * epoch + it # global training iteration + for i, param_group in enumerate(optimizer.param_groups): + param_group["lr"] = lr_schedule[it] + if i == 0: # only the first group is regularized + param_group["weight_decay"] = wd_schedule[it] + + # move images to gpu + images = [im.cuda(non_blocking=True) for im in images] + # teacher and student forward passes + compute dino loss + with torch.cuda.amp.autocast(fp16_scaler is not None): + teacher_output = teacher(images[:2]) # only the 2 global views pass through the teacher + student_output = student(images) + loss = dino_loss(student_output, teacher_output, epoch) + + if not math.isfinite(loss.item()): + print("Loss is {}, stopping training".format(loss.item()), force=True) + sys.exit(1) + + # student update + optimizer.zero_grad() + param_norms = None + if fp16_scaler is None: + loss.backward() + if args.clip_grad: + param_norms = utils.clip_gradients(student, args.clip_grad) + utils.cancel_gradients_last_layer(epoch, student, + args.freeze_last_layer) + optimizer.step() + else: + fp16_scaler.scale(loss).backward() + if args.clip_grad: + fp16_scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place + param_norms = utils.clip_gradients(student, args.clip_grad) + utils.cancel_gradients_last_layer(epoch, student, + args.freeze_last_layer) + fp16_scaler.step(optimizer) + fp16_scaler.update() + + # EMA update for the teacher + with torch.no_grad(): + m = momentum_schedule[it] # momentum parameter + for param_q, param_k in zip(student.module.parameters(), teacher_without_ddp.parameters()): + param_k.data.mul_(m).add_((1 - m) * param_q.detach().data) + + # logging + torch.cuda.synchronize() + metric_logger.update(loss=loss.item()) + metric_logger.update(lr=optimizer.param_groups[0]["lr"]) + metric_logger.update(wd=optimizer.param_groups[0]["weight_decay"]) + # gather the stats from all processes + metric_logger.synchronize_between_processes() + print("Averaged stats:", metric_logger) + return {k: meter.global_avg for k, meter in metric_logger.meters.items()} + + +class DINOLoss(nn.Module): + def __init__(self, out_dim, ncrops, warmup_teacher_temp, teacher_temp, + warmup_teacher_temp_epochs, nepochs, student_temp=0.1, + center_momentum=0.9): + super().__init__() + self.student_temp = student_temp + self.center_momentum = center_momentum + self.ncrops = ncrops + self.register_buffer("center", torch.zeros(1, out_dim)) + # we apply a warm up for the teacher temperature because + # a too high temperature makes the training instable at the beginning + self.teacher_temp_schedule = np.concatenate(( + np.linspace(warmup_teacher_temp, + teacher_temp, warmup_teacher_temp_epochs), + np.ones(nepochs - warmup_teacher_temp_epochs) * teacher_temp + )) + + def forward(self, student_output, teacher_output, epoch): + """ + Cross-entropy between softmax outputs of the teacher and student networks. + """ + student_out = student_output / self.student_temp + student_out = student_out.chunk(self.ncrops) + + # teacher centering and sharpening + temp = self.teacher_temp_schedule[epoch] + teacher_out = F.softmax((teacher_output - self.center) / temp, dim=-1) + teacher_out = teacher_out.detach().chunk(2) + + total_loss = 0 + n_loss_terms = 0 + for iq, q in enumerate(teacher_out): + for v in range(len(student_out)): + if v == iq: + # we skip cases where student and teacher operate on the same view + continue + loss = torch.sum(-q * F.log_softmax(student_out[v], dim=-1), dim=-1) + total_loss += loss.mean() + n_loss_terms += 1 + total_loss /= n_loss_terms + self.update_center(teacher_output) + return total_loss + + @torch.no_grad() + def update_center(self, teacher_output): + """ + Update center used for teacher output. + """ + batch_center = torch.sum(teacher_output, dim=0, keepdim=True) + dist.all_reduce(batch_center) + batch_center = batch_center / (len(teacher_output) * dist.get_world_size()) + + # ema update + self.center = self.center * self.center_momentum + batch_center * (1 - self.center_momentum) + + +class DataAugmentationDINO(object): + def __init__(self, global_crops_scale, local_crops_scale, local_crops_number): + flip_and_color_jitter = transforms.Compose([ + transforms.RandomHorizontalFlip(p=0.5), + transforms.RandomApply( + [transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.2, hue=0.1)], + p=0.8 + ), + transforms.RandomGrayscale(p=0.2), + ]) + normalize = transforms.Compose([ + transforms.ToTensor(), + transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), + ]) + + # first global crop + self.global_transfo1 = transforms.Compose([ + transforms.RandomResizedCrop(224, scale=global_crops_scale, interpolation=Image.BICUBIC), + flip_and_color_jitter, + utils.GaussianBlur(1.0), + normalize, + ]) + # second global crop + self.global_transfo2 = transforms.Compose([ + transforms.RandomResizedCrop(224, scale=global_crops_scale, interpolation=Image.BICUBIC), + flip_and_color_jitter, + utils.GaussianBlur(0.1), + utils.Solarization(0.2), + normalize, + ]) + # transformation for the local small crops + self.local_crops_number = local_crops_number + self.local_transfo = transforms.Compose([ + transforms.RandomResizedCrop(96, scale=local_crops_scale, interpolation=Image.BICUBIC), + flip_and_color_jitter, + utils.GaussianBlur(p=0.5), + normalize, + ]) + + def __call__(self, image): + crops = [] + crops.append(self.global_transfo1(image)) + crops.append(self.global_transfo2(image)) + for _ in range(self.local_crops_number): + crops.append(self.local_transfo(image)) + return crops + + +if __name__ == '__main__': + parser = argparse.ArgumentParser('DINO', parents=[get_args_parser()]) + args = parser.parse_args() + Path(args.output_dir).mkdir(parents=True, exist_ok=True) + train_dino(args) diff --git a/torch_hub/facebookresearch_dino_main/run_with_submitit.py b/torch_hub/facebookresearch_dino_main/run_with_submitit.py new file mode 100644 index 0000000000000000000000000000000000000000..33d4116f2ff512b39d0cec5c936f999df1ac80fe --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/run_with_submitit.py @@ -0,0 +1,132 @@ +# Copyright (c) Facebook, Inc. and its affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +A script to run multinode training with submitit. +Almost copy-paste from https://github.com/facebookresearch/deit/blob/main/run_with_submitit.py +""" +import argparse +import os +import uuid +from pathlib import Path + +import main_dino +import submitit + + +def parse_args(): + parser = argparse.ArgumentParser("Submitit for DINO", parents=[main_dino.get_args_parser()]) + parser.add_argument("--ngpus", default=8, type=int, help="Number of gpus to request on each node") + parser.add_argument("--nodes", default=2, type=int, help="Number of nodes to request") + parser.add_argument("--timeout", default=2800, type=int, help="Duration of the job") + + parser.add_argument("--partition", default="learnfair", type=str, help="Partition where to submit") + parser.add_argument("--use_volta32", action='store_true', help="Big models? Use this") + parser.add_argument('--comment', default="", type=str, + help='Comment to pass to scheduler, e.g. priority message') + return parser.parse_args() + + +def get_shared_folder() -> Path: + user = os.getenv("USER") + if Path("/checkpoint/").is_dir(): + p = Path(f"/checkpoint/{user}/experiments") + p.mkdir(exist_ok=True) + return p + raise RuntimeError("No shared folder available") + + +def get_init_file(): + # Init file must not exist, but it's parent dir must exist. + os.makedirs(str(get_shared_folder()), exist_ok=True) + init_file = get_shared_folder() / f"{uuid.uuid4().hex}_init" + if init_file.exists(): + os.remove(str(init_file)) + return init_file + + +class Trainer(object): + def __init__(self, args): + self.args = args + + def __call__(self): + import main_dino + + self._setup_gpu_args() + main_dino.train_dino(self.args) + + def checkpoint(self): + import os + import submitit + + self.args.dist_url = get_init_file().as_uri() + print("Requeuing ", self.args) + empty_trainer = type(self)(self.args) + return submitit.helpers.DelayedSubmission(empty_trainer) + + def _setup_gpu_args(self): + import submitit + from pathlib import Path + + job_env = submitit.JobEnvironment() + self.args.output_dir = Path(str(self.args.output_dir).replace("%j", str(job_env.job_id))) + self.args.gpu = job_env.local_rank + self.args.rank = job_env.global_rank + self.args.world_size = job_env.num_tasks + print(f"Process group: {job_env.num_tasks} tasks, rank: {job_env.global_rank}") + + +def main(): + args = parse_args() + if args.output_dir == "": + args.output_dir = get_shared_folder() / "%j" + Path(args.output_dir).mkdir(parents=True, exist_ok=True) + executor = submitit.AutoExecutor(folder=args.output_dir, slurm_max_num_timeout=30) + + num_gpus_per_node = args.ngpus + nodes = args.nodes + timeout_min = args.timeout + + partition = args.partition + kwargs = {} + if args.use_volta32: + kwargs['slurm_constraint'] = 'volta32gb' + if args.comment: + kwargs['slurm_comment'] = args.comment + + executor.update_parameters( + mem_gb=40 * num_gpus_per_node, + gpus_per_node=num_gpus_per_node, + tasks_per_node=num_gpus_per_node, # one task per GPU + cpus_per_task=10, + nodes=nodes, + timeout_min=timeout_min, # max is 60 * 72 + # Below are cluster dependent parameters + slurm_partition=partition, + slurm_signal_delay_s=120, + **kwargs + ) + + executor.update_parameters(name="dino") + + args.dist_url = get_init_file().as_uri() + + trainer = Trainer(args) + job = executor.submit(trainer) + + print(f"Submitted job_id: {job.job_id}") + print(f"Logs and checkpoints will be saved at: {args.output_dir}") + + +if __name__ == "__main__": + main() diff --git a/torch_hub/facebookresearch_dino_main/video_generation.py b/torch_hub/facebookresearch_dino_main/video_generation.py new file mode 100644 index 0000000000000000000000000000000000000000..94da9836ad0e9bd8dccf0f989b93a93ed11cfd7e --- /dev/null +++ b/torch_hub/facebookresearch_dino_main/video_generation.py @@ -0,0 +1,378 @@ +# Copyright (c) Facebook, Inc. and its affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import os +import glob +import sys +import argparse +import cv2 + +from tqdm import tqdm +import matplotlib.pyplot as plt +import torch +import torch.nn as nn +import torchvision +from torchvision import transforms as pth_transforms +import numpy as np +from PIL import Image + +import utils +import vision_transformer as vits + + +FOURCC = { + "mp4": cv2.VideoWriter_fourcc(*"MP4V"), + "avi": cv2.VideoWriter_fourcc(*"XVID"), +} +DEVICE = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") + + +class VideoGenerator: + def __init__(self, args): + self.args = args + # self.model = None + # Don't need to load model if you only want a video + if not self.args.video_only: + self.model = self.__load_model() + + def run(self): + if self.args.input_path is None: + print(f"Provided input path {self.args.input_path} is non valid.") + sys.exit(1) + else: + if self.args.video_only: + self._generate_video_from_images( + self.args.input_path, self.args.output_path + ) + else: + # If input path exists + if os.path.exists(self.args.input_path): + # If input is a video file + if os.path.isfile(self.args.input_path): + frames_folder = os.path.join(self.args.output_path, "frames") + attention_folder = os.path.join( + self.args.output_path, "attention" + ) + + os.makedirs(frames_folder, exist_ok=True) + os.makedirs(attention_folder, exist_ok=True) + + self._extract_frames_from_video( + self.args.input_path, frames_folder + ) + + self._inference( + frames_folder, + attention_folder, + ) + + self._generate_video_from_images( + attention_folder, self.args.output_path + ) + + # If input is a folder of already extracted frames + if os.path.isdir(self.args.input_path): + attention_folder = os.path.join( + self.args.output_path, "attention" + ) + + os.makedirs(attention_folder, exist_ok=True) + + self._inference(self.args.input_path, attention_folder) + + self._generate_video_from_images( + attention_folder, self.args.output_path + ) + + # If input path doesn't exists + else: + print(f"Provided input path {self.args.input_path} doesn't exists.") + sys.exit(1) + + def _extract_frames_from_video(self, inp: str, out: str): + vidcap = cv2.VideoCapture(inp) + self.args.fps = vidcap.get(cv2.CAP_PROP_FPS) + + print(f"Video: {inp} ({self.args.fps} fps)") + print(f"Extracting frames to {out}") + + success, image = vidcap.read() + count = 0 + while success: + cv2.imwrite( + os.path.join(out, f"frame-{count:04}.jpg"), + image, + ) + success, image = vidcap.read() + count += 1 + + def _generate_video_from_images(self, inp: str, out: str): + img_array = [] + attention_images_list = sorted(glob.glob(os.path.join(inp, "attn-*.jpg"))) + + # Get size of the first image + with open(attention_images_list[0], "rb") as f: + img = Image.open(f) + img = img.convert("RGB") + size = (img.width, img.height) + img_array.append(cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR)) + + print(f"Generating video {size} to {out}") + + for filename in tqdm(attention_images_list[1:]): + with open(filename, "rb") as f: + img = Image.open(f) + img = img.convert("RGB") + img_array.append(cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR)) + + out = cv2.VideoWriter( + os.path.join(out, "video." + self.args.video_format), + FOURCC[self.args.video_format], + self.args.fps, + size, + ) + + for i in range(len(img_array)): + out.write(img_array[i]) + out.release() + print("Done") + + def _inference(self, inp: str, out: str): + print(f"Generating attention images to {out}") + + for img_path in tqdm(sorted(glob.glob(os.path.join(inp, "*.jpg")))): + with open(img_path, "rb") as f: + img = Image.open(f) + img = img.convert("RGB") + + if self.args.resize is not None: + transform = pth_transforms.Compose( + [ + pth_transforms.ToTensor(), + pth_transforms.Resize(self.args.resize), + pth_transforms.Normalize( + (0.485, 0.456, 0.406), (0.229, 0.224, 0.225) + ), + ] + ) + else: + transform = pth_transforms.Compose( + [ + pth_transforms.ToTensor(), + pth_transforms.Normalize( + (0.485, 0.456, 0.406), (0.229, 0.224, 0.225) + ), + ] + ) + + img = transform(img) + + # make the image divisible by the patch size + w, h = ( + img.shape[1] - img.shape[1] % self.args.patch_size, + img.shape[2] - img.shape[2] % self.args.patch_size, + ) + img = img[:, :w, :h].unsqueeze(0) + + w_featmap = img.shape[-2] // self.args.patch_size + h_featmap = img.shape[-1] // self.args.patch_size + + attentions = self.model.get_last_selfattention(img.to(DEVICE)) + + nh = attentions.shape[1] # number of head + + # we keep only the output patch attention + attentions = attentions[0, :, 0, 1:].reshape(nh, -1) + + # we keep only a certain percentage of the mass + val, idx = torch.sort(attentions) + val /= torch.sum(val, dim=1, keepdim=True) + cumval = torch.cumsum(val, dim=1) + th_attn = cumval > (1 - self.args.threshold) + idx2 = torch.argsort(idx) + for head in range(nh): + th_attn[head] = th_attn[head][idx2[head]] + th_attn = th_attn.reshape(nh, w_featmap, h_featmap).float() + # interpolate + th_attn = ( + nn.functional.interpolate( + th_attn.unsqueeze(0), + scale_factor=self.args.patch_size, + mode="nearest", + )[0] + .cpu() + .numpy() + ) + + attentions = attentions.reshape(nh, w_featmap, h_featmap) + attentions = ( + nn.functional.interpolate( + attentions.unsqueeze(0), + scale_factor=self.args.patch_size, + mode="nearest", + )[0] + .cpu() + .numpy() + ) + + # save attentions heatmaps + fname = os.path.join(out, "attn-" + os.path.basename(img_path)) + plt.imsave( + fname=fname, + arr=sum( + attentions[i] * 1 / attentions.shape[0] + for i in range(attentions.shape[0]) + ), + cmap="inferno", + format="jpg", + ) + + def __load_model(self): + # build model + model = vits.__dict__[self.args.arch]( + patch_size=self.args.patch_size, num_classes=0 + ) + for p in model.parameters(): + p.requires_grad = False + model.eval() + model.to(DEVICE) + + if os.path.isfile(self.args.pretrained_weights): + state_dict = torch.load(self.args.pretrained_weights, map_location="cpu") + if ( + self.args.checkpoint_key is not None + and self.args.checkpoint_key in state_dict + ): + print( + f"Take key {self.args.checkpoint_key} in provided checkpoint dict" + ) + state_dict = state_dict[self.args.checkpoint_key] + state_dict = {k.replace("module.", ""): v for k, v in state_dict.items()} + # remove `backbone.` prefix induced by multicrop wrapper + state_dict = {k.replace("backbone.", ""): v for k, v in state_dict.items()} + msg = model.load_state_dict(state_dict, strict=False) + print( + "Pretrained weights found at {} and loaded with msg: {}".format( + self.args.pretrained_weights, msg + ) + ) + else: + print( + "Please use the `--pretrained_weights` argument to indicate the path of the checkpoint to evaluate." + ) + url = None + if self.args.arch == "vit_small" and self.args.patch_size == 16: + url = "dino_deitsmall16_pretrain/dino_deitsmall16_pretrain.pth" + elif self.args.arch == "vit_small" and self.args.patch_size == 8: + url = "dino_deitsmall8_300ep_pretrain/dino_deitsmall8_300ep_pretrain.pth" # model used for visualizations in our paper + elif self.args.arch == "vit_base" and self.args.patch_size == 16: + url = "dino_vitbase16_pretrain/dino_vitbase16_pretrain.pth" + elif self.args.arch == "vit_base" and self.args.patch_size == 8: + url = "dino_vitbase8_pretrain/dino_vitbase8_pretrain.pth" + if url is not None: + print( + "Since no pretrained weights have been provided, we load the reference pretrained DINO weights." + ) + state_dict = torch.hub.load_state_dict_from_url( + url="https://dl.fbaipublicfiles.com/dino/" + url + ) + model.load_state_dict(state_dict, strict=True) + else: + print( + "There is no reference weights available for this model => We use random weights." + ) + return model + + +def parse_args(): + parser = argparse.ArgumentParser("Generation self-attention video") + parser.add_argument( + "--arch", + default="vit_small", + type=str, + choices=["vit_tiny", "vit_small", "vit_base"], + help="Architecture (support only ViT atm).", + ) + parser.add_argument( + "--patch_size", default=8, type=int, help="Patch resolution of the self.model." + ) + parser.add_argument( + "--pretrained_weights", + default="", + type=str, + help="Path to pretrained weights to load.", + ) + parser.add_argument( + "--checkpoint_key", + default="teacher", + type=str, + help='Key to use in the checkpoint (example: "teacher")', + ) + parser.add_argument( + "--input_path", + required=True, + type=str, + help="""Path to a video file if you want to extract frames + or to a folder of images already extracted by yourself. + or to a folder of attention images.""", + ) + parser.add_argument( + "--output_path", + default="./", + type=str, + help="""Path to store a folder of frames and / or a folder of attention images. + and / or a final video. Default to current directory.""", + ) + parser.add_argument( + "--threshold", + type=float, + default=0.6, + help="""We visualize masks + obtained by thresholding the self-attention maps to keep xx percent of the mass.""", + ) + parser.add_argument( + "--resize", + default=None, + type=int, + nargs="+", + help="""Apply a resize transformation to input image(s). Use if OOM error. + Usage (single or W H): --resize 512, --resize 720 1280""", + ) + parser.add_argument( + "--video_only", + action="store_true", + help="""Use this flag if you only want to generate a video and not all attention images. + If used, --input_path must be set to the folder of attention images. Ex: ./attention/""", + ) + parser.add_argument( + "--fps", + default=30.0, + type=float, + help="FPS of input / output video. Automatically set if you extract frames from a video.", + ) + parser.add_argument( + "--video_format", + default="mp4", + type=str, + choices=["mp4", "avi"], + help="Format of generated video (mp4 or avi).", + ) + + return parser.parse_args() + + +if __name__ == "__main__": + args = parse_args() + + vg = VideoGenerator(args) + vg.run() diff --git a/torch_hub/facebookresearch_dinov2_main/.gitignore b/torch_hub/facebookresearch_dinov2_main/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..f6893ca30f324f6ed3e18ae9c726af8377c57c69 --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/.gitignore @@ -0,0 +1,11 @@ +build/ +dist/ +*.egg-info/ +**/__pycache__/ + +**/.ipynb_checkpoints +**/.ipynb_checkpoints/** + +*.swp + +.vscode/ diff --git a/torch_hub/facebookresearch_dinov2_main/LICENSE_XRAY_DINO_MODEL b/torch_hub/facebookresearch_dinov2_main/LICENSE_XRAY_DINO_MODEL new file mode 100644 index 0000000000000000000000000000000000000000..0561abf417f16f579987c7d4a0340005b2fee0fe --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/LICENSE_XRAY_DINO_MODEL @@ -0,0 +1,39 @@ +X-Ray DINO Research License [2025/12/18] + +This X-Ray DINO Research License (“Agreement”) contains the terms and conditions that govern your access and use of the Materials (as defined below). You may not use the Materials if you do not accept this Agreement. By clicking “I Accept” to accept, or accessing, using, or distributing any portion or element of the Materials you hereby agree to be bound by the terms of this Agreement. If you are agreeing to be bound by the Agreement on behalf of your employer or other entity, you represent and warrant to Meta Platforms Ireland Limited (if you are located in or, if you are an entity, your principal place of business is in the EEA or Switzerland) and Meta Platforms, Inc. (if you are located outside of the EEA or Switzerland) (“Meta”) that you have full legal authority to bind your employer or such entity to this Agreement. If you do not have requisite authority, you may not accept the Agreement or access the Materials on behalf of your employer or other entity. + +This Agreement is effective upon the earlier of the date that you first access the Materials or accept this Agreement (“Effective Date”), and is entered into by and between Meta, and you, or if you are entering into this Agreement on behalf of your employer or other entity (if you are entering into this Agreement on such person or entity’s behalf), of the age required under applicable laws, rules, or regulations to provide legal consent and, your employer or other entity and that has legal authority to bind your employer or such other person or entity if you are entering in this Agreement on their behalf (“Licensee” or “You”). + +1. Definitions. +“Documentation” means the specifications, manuals and documentation accompanying this release distributed by Meta at [https://github.com/facebookresearch/dinov2/blob/main/README.md]. + +“Noncommercial Research Use” means noncommercial research use cases related to research, development, education, processing, or analysis and in each case, is not primarily intended for commercial advantage or monetary compensation to you or others. + +“Materials” means, collectively, Documentation and the models and software and algorithms, including machine-learning model code, trained model weights, inference-enabling code, training-enabling code, fine-tuning enabling code, demonstration materials and other elements of the foregoing distributed by Meta at [https://github.com/facebookresearch/dinov2/blob/main/README.md] and made available under this Agreement. + +“Trade Control Laws” means any applicable U.S. and non-U.S. export control and trade sanctions laws and regulations. + +2. License Rights and Redistribution. Subject to Your compliance with the terms and conditions of this Agreement, Meta hereby grants you the following: +Grant of Rights. You are hereby granted a non-exclusive, worldwide, non-transferable and royalty-free limited license under Meta’s intellectual property or other rights owned by Meta embodied in the Materials to use, reproduce, distribute, copy, create derivative works of, and make modifications to the Materials solely for Noncommercial Research Uses. +Redistribution and Use. +Distribution of Materials, and any derivative works thereof, are subject to the terms of this Agreement. If you distribute or make the Materials, or any derivative works thereof, available to a third party, you may only do so under the terms of this Agreement. You shall also provide a copy of this Agreement to such third party. +If you submit for publication the results of research you perform on, using, or otherwise in connection with Materials, you must acknowledge the use of Materials in your publication. +You must retain in all copies of the Materials that you distribute and include the following attribution notice within a “Notice” text file distributed as a part of such copies: “Materials are licensed under the X-Ray DINO Research License, Copyright © Meta Platforms, Inc. All Rights Reserved.” +Your use of the Materials must comply with applicable laws and regulations (including Trade Control Laws). +You agree to report any violation of this X-Ray DINO Research License. +3. Restrictions. You will not, and will not permit, assist or cause any third party to: +use the Materials or any outputs or results of the Materials in connection with any commercial uses or for any uses other than Noncommercial Research Uses; +use the Materials or any outputs or results of the Materials for provisioning medical advice or in connection with any clinical purpose or medical procedures, including health preventative, mitigatory, diagnostic, treatment, or curative applications or as a substitute or adjunct to professional medical judgment as the Materials have not been reviewed or approved by the Food and Drug Administration, and are for non-clinical, Noncommerical Research Use only; +disguise your or their location through IP proxying or other methods; +use or download Materials if you or they are: (a) located in a comprehensively sanctioned jurisdiction, (b) currently listed on any U.S. or non-U.S. restricted parties list, or (c) will use Materials for any purpose prohibited by Trade Control Laws; or +directly or indirectly export, re-export, provide, or otherwise transfer Materials: (a) to any individual, entity, or country prohibited by Trade Control Laws; (b) to anyone on U.S. or non-U.S. government restricted parties lists; or (c) for any purpose prohibited by Trade Control Laws, including nuclear, chemical or biological weapons, or missile technology applications. +4. User Support. Your Noncommercial Research Use of the Materials is done at your own discretion; Meta does not process any information nor provide any service in relation to such use. Meta is under no obligation to provide any support services for the Materials. Any support provided is “as is”, “with all faults”, and without warranty of any kind. +5. Disclaimer of Warranty. UNLESS REQUIRED BY APPLICABLE LAW, THE MATERIALS AND ANY OUTPUT AND RESULTS THEREFROM ARE PROVIDED ON AN “AS IS” BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING, WITHOUT LIMITATION, ANY WARRANTIES OF TITLE, NON-INFRINGEMENT, MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE, THE ABSENCE OF LATENT OR OTHER DEFECTS, ACCURACY, OR THE PRESENCE OR ABSENCE OF ERRORS, WHETHER OR NOT DISCOVERABLE. YOU ARE SOLELY RESPONSIBLE FOR DETERMINING THE APPROPRIATENESS OF USING OR REDISTRIBUTING THE MATERIALS AND ASSUME ANY RISKS ASSOCIATED WITH YOUR USE OF THE MATERIALS AND ANY OUTPUT AND RESULTS. +6. Limitation of Liability. IN NO EVENT WILL META OR ITS AFFILIATES BE LIABLE UNDER ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, TORT, NEGLIGENCE, PRODUCTS LIABILITY, OR OTHERWISE, ARISING OUT OF THIS AGREEMENT, FOR ANY LOST PROFITS OR ANY DIRECT OR INDIRECT, SPECIAL, CONSEQUENTIAL, INCIDENTAL, EXEMPLARY OR PUNITIVE DAMAGES, EVEN IF META OR ITS AFFILIATES HAVE BEEN ADVISED OF THE POSSIBILITY OF ANY OF THE FOREGOING. +7. Intellectual Property. +No trademark licenses are granted under this Agreement, and in connection with the Materials, neither Meta nor Licensee may use any name or mark owned by or associated with the other or any of its affiliates, except as required for reasonable and customary use in describing and redistributing the Materials. +Subject to Meta’s ownership of Materials and derivatives made by or for Meta, with respect to any derivative works and modifications of the Materials that are made by you, as between you and Meta, you are and will be the owner of such derivative works and modifications. +If you institute litigation or other proceedings against Meta or any entity (including a cross-claim or counterclaim in a lawsuit) alleging that the Materials or outputs or results, or any portion of any of the foregoing, constitutes infringement of intellectual property or other rights owned or licensable by you, then any licenses and rights granted to you under this Agreement shall terminate as of the date such litigation or claim is filed or instituted. You will indemnify and hold harmless Meta from and against any claim by any third party arising out of or related to your use or distribution of the Materials. +8. Term and Termination. The term of this Agreement will commence upon your acceptance of this Agreement or access to the Materials and will continue in full force and effect until terminated in accordance with the terms and conditions herein. Meta may terminate this Agreement if you are in breach of any term or condition of this Agreement. Upon termination of this Agreement, you shall delete and cease use of the Materials. Sections 3, 4, 5, 6, 7, 8 and 9 shall survive the termination of this Agreement. +9. Governing Law and Jurisdiction. This Agreement will be governed and construed under the laws of the State of California without regard to choice of law principles, and the UN Convention on Contracts for the International Sale of Goods does not apply to this Agreement. The courts of California shall have exclusive jurisdiction of any dispute arising out of this Agreement. +10. Modifications and Amendments. Meta may modify this Agreement from time to time by posting a revised version at [https://ai.meta.com/resources/models-and-libraries/raydino-license/]; provided that they are similar in spirit to the current version of the Agreement, but may differ in detail to address new problems or concerns. All such changes will be effective immediately. Your continued use of the Materials after any modification to this Agreement constitutes your agreement to such modification. Except as provided in this Agreement, no other modification or addition to any provision of this Agreement will be binding unless it is in writing and signed by an authorized representative of both you and Meta. diff --git a/torch_hub/facebookresearch_dinov2_main/MODEL_CARD.md b/torch_hub/facebookresearch_dinov2_main/MODEL_CARD.md new file mode 100644 index 0000000000000000000000000000000000000000..21b9bf295c8cab14e782e1a7a1d051be9e501088 --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/MODEL_CARD.md @@ -0,0 +1,272 @@ +# Model Card for DINOv2-S/B/L/g + +These are Vision Transformer models trained following the method described in the papers: +"DINOv2: Learning Robust Visual Features without Supervision" +and +"Vision Transformers Need Registers". + +We provide 8 models: +- 1 ViT-g trained from scratch with 3 ViT-S/B/L models distilled from the ViT-g, without registers. +- 1 ViT-g trained from scratch with 3 ViT-S/B/L models distilled from the ViT-g, with registers. + +## Model Details +The model takes an image as input and returns a class token and patch tokens, and optionally 4 register tokens. + +The embedding dimension is: +- 384 for ViT-S. +- 768 for ViT-B. +- 1024 for ViT-L. +- 1536 for ViT-g. + +The models follow a Transformer architecture, with a patch size of 14. In the case of registers, we add 4 register tokens, learned during training, to the input sequence after the patch embedding. + +For a 224x224 image, this results in 1 class token + 256 patch tokens, and optionally 4 register tokens. + +The models can accept larger images provided the image shapes are multiples of the patch size (14). +If this condition is not verified, the model will crop to the closest smaller multiple of the patch size. + +### Model Description + +- **Developed by:** Meta AI +- **Model type:** Vision Transformer +- **License:** Apache License 2.0 + +- **Repository:** https://github.com/facebookresearch/dinov2 +- **Paper:** https://arxiv.org/abs/2304.07193 +- **Demo:** https://dinov2.metademolab.com/ + +## Uses + +The models are vision backbones providing multi-purpose features for downstream tasks. + +### Direct Use + +The models can be used without fine-tuning, with downstream classifiers as simple as linear layers, to obtain competitive results: +- on depth estimation, semantic segmentation, using linear layers. +- on image classification, using k-NN classifiers on the class token. +- on image classification, with logistic regression classifiers applied on the class token. +- on image classification, with a linear layer applied on the class token and the average of the patch tokens. +- on image retrieval using nearest neighbors. + +### Downstream Use + +It is technically possible to perform fine-tuning on the models, for small gains (we measured +2% on ImageNet-1k classification). +We recommend keeping this as a very last step and only when necessary, as the features already provide good performance out-of-the-box. + +## Bias, Risks, and Limitations + +Despite improvements thanks to the training method not using annotations, we still observe significant biases in our models toward rich households from Western countries. + +### Recommendations + +We expect fine-tuning will increase the biases in the features produced by the model as they will be tuned to the fine-tuning labels. + +## How to Get Started with the Model + +Use the code below to get started with the model. + +```python +import torch + +# DINOv2 +dinov2_vits14 = torch.hub.load('facebookresearch/dinov2', 'dinov2_vits14') +dinov2_vitb14 = torch.hub.load('facebookresearch/dinov2', 'dinov2_vitb14') +dinov2_vitl14 = torch.hub.load('facebookresearch/dinov2', 'dinov2_vitl14') +dinov2_vitg14 = torch.hub.load('facebookresearch/dinov2', 'dinov2_vitg14') + +# DINOv2 with registers +dinov2_vits14_reg = torch.hub.load('facebookresearch/dinov2', 'dinov2_vits14_reg') +dinov2_vitb14_reg = torch.hub.load('facebookresearch/dinov2', 'dinov2_vitb14_reg') +dinov2_vitl14_reg = torch.hub.load('facebookresearch/dinov2', 'dinov2_vitl14_reg') +dinov2_vitg14_reg = torch.hub.load('facebookresearch/dinov2', 'dinov2_vitg14_reg') +``` + +## Training Details + +### Training Data + +- **Training data:** LVD-142M (see paper) +- **Training regime:** fp16 using PyTorch-FSDP mixed-precision. + +### Training Procedure + +- **Training objective:** + - DINO self-distillation loss with multi-crop + - iBOT masked-image modeling loss + - KoLeo regularization on [CLS] tokens +- **Architectures:** + - ViT-S (21M params): Patch size 14, embedding dimension 384, 6 heads, MLP FFN + - ViT-B (86M params): Patch size 14, embedding dimension 768, 12 heads, MLP FFN + - ViT-L (0.3B params): Patch size 14, embedding dimension 1024, 16 heads, MLP FFN + - ViT-g (1.1B params): Patch size 14, embedding dimension 1536, 24 heads, SwiGLU FFN +- **Distillation:** + - Distillation follows the standard DINOv2 pretraining procedure, except the teacher is a pretrained ViT-g, frozen. + +## Evaluation + +We refer users to the associated papers for the evaluation protocols. + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
ImageNet-1kNYU-Depth v2SUN-RGBDADE20kiNaturalist 2018Oxford-H
modelwith
registers
classif. (acc)classif. (acc)classif. V2 (acc)depth (RMSE)depth (RMSE)segm. (mAP)classif. (acc)retrieval (mAP)
k-NNlinearlinearlinear
4 layers
NYU-D transfermultiscalelinearnearest neighbor
ViT-S/14:x:79.0%81.1%70.8%0.4170.43147.269.5%43.2
ViT-S/14:white_check_mark:79.1%80.9%71.0%N/AN/AN/A67.6%39.5
ViT-B/14:x:82.1%84.5%74.9%0.3620.40051.376.3%49.5
ViT-B/14:white_check_mark:82.0%84.6%75.6%N/AN/AN/A73.8%51.0
ViT-L/14:x:83.5%86.3%77.6%0.3330.39653.179.8%54.0
ViT-L/14:white_check_mark:83.8%86.7%78.5%N/AN/AN/A80.9%55.7
ViT-g/14:x:83.5%86.5%78.4%0.2980.36253.081.6%52.3
ViT-g/14:white_check_mark:83.7%87.1%78.8%N/AN/AN/A81.5%58.2
+ +## Environmental Impact + +- **Hardware Type:** Nvidia A100 +- **Hours used:** 22,000 for ViT-g, 4,500 for ViT-S distillation, 5,300 for ViT-B distillation, 8,000 for ViT-L distillation +- **Cloud Provider:** Private infra +- **Compute Region:** USA +- **Carbon Emitted:** 7t CO2eq + +#### Hardware + +Nvidia A100 GPUs + +#### Software + +PyTorch 2.0, +xFormers 0.0.18 + +**BibTeX** + +``` +@misc{oquab2023dinov2, + title={DINOv2: Learning Robust Visual Features without Supervision}, + author={Oquab, Maxime and Darcet, Timothée and Moutakanni, Theo and Vo, Huy and Szafraniec, Marc and Khalidov, Vasil and Fernandez, Pierre and Haziza, Daniel and Massa, Francisco and El-Nouby, Alaaeldin and Howes, Russell and Huang, Po-Yao and Xu, Hu and Sharma, Vasu and Li, Shang-Wen and Galuba, Wojciech and Rabbat, Mike and Assran, Mido and Ballas, Nicolas and Synnaeve, Gabriel and Misra, Ishan and Jegou, Herve and Mairal, Julien and Labatut, Patrick and Joulin, Armand and Bojanowski, Piotr}, + journal={arXiv:2304.07193}, + year={2023} +} +@misc{darcet2023vitneedreg, + title={Vision Transformers Need Registers}, + author={Darcet, Timothée and Oquab, Maxime and Mairal, Julien and Bojanowski, Piotr}, + journal={arXiv:2309.16588}, + year={2023} +} +``` diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/__init__.py b/torch_hub/facebookresearch_dinov2_main/dinov2/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..ae847e46898077fe3d8701b8a181d7b4e3d41cd9 --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/dinov2/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +__version__ = "0.0.1" diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/__pycache__/__init__.cpython-310.pyc b/torch_hub/facebookresearch_dinov2_main/dinov2/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..c7976b52883132711af247524914420127ff6128 Binary files /dev/null and b/torch_hub/facebookresearch_dinov2_main/dinov2/__pycache__/__init__.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/configs/eval/vitb14_reg4_pretrain.yaml b/torch_hub/facebookresearch_dinov2_main/dinov2/configs/eval/vitb14_reg4_pretrain.yaml new file mode 100644 index 0000000000000000000000000000000000000000..d53edc04a0761b4b35c147d63e04d55c90092c8f --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/dinov2/configs/eval/vitb14_reg4_pretrain.yaml @@ -0,0 +1,9 @@ +student: + arch: vit_base + patch_size: 14 + num_register_tokens: 4 + interpolate_antialias: true + interpolate_offset: 0.0 +crops: + global_crops_size: 518 # this is to set up the position embeddings properly + local_crops_size: 98 \ No newline at end of file diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/configs/ssl_default_config.yaml b/torch_hub/facebookresearch_dinov2_main/dinov2/configs/ssl_default_config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..cdbea43d8e497baec8ab5172b1f83d037f29a4d9 --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/dinov2/configs/ssl_default_config.yaml @@ -0,0 +1,123 @@ +MODEL: + WEIGHTS: '' +compute_precision: + grad_scaler: true + teacher: + backbone: + sharding_strategy: SHARD_GRAD_OP + mixed_precision: + param_dtype: fp16 + reduce_dtype: fp16 + buffer_dtype: fp32 + dino_head: + sharding_strategy: SHARD_GRAD_OP + mixed_precision: + param_dtype: fp16 + reduce_dtype: fp16 + buffer_dtype: fp32 + ibot_head: + sharding_strategy: SHARD_GRAD_OP + mixed_precision: + param_dtype: fp16 + reduce_dtype: fp16 + buffer_dtype: fp32 + student: + backbone: + sharding_strategy: SHARD_GRAD_OP + mixed_precision: + param_dtype: fp16 + reduce_dtype: fp16 + buffer_dtype: fp32 + dino_head: + sharding_strategy: SHARD_GRAD_OP + mixed_precision: + param_dtype: fp16 + reduce_dtype: fp32 + buffer_dtype: fp32 + ibot_head: + sharding_strategy: SHARD_GRAD_OP + mixed_precision: + param_dtype: fp16 + reduce_dtype: fp32 + buffer_dtype: fp32 +dino: + loss_weight: 1.0 + head_n_prototypes: 65536 + head_bottleneck_dim: 256 + head_nlayers: 3 + head_hidden_dim: 2048 + koleo_loss_weight: 0.1 +ibot: + loss_weight: 1.0 + mask_sample_probability: 0.5 + mask_ratio_min_max: + - 0.1 + - 0.5 + separate_head: false + head_n_prototypes: 65536 + head_bottleneck_dim: 256 + head_nlayers: 3 + head_hidden_dim: 2048 +train: + batch_size_per_gpu: 64 + dataset_path: ImageNet:split=TRAIN + output_dir: . + saveckp_freq: 20 + seed: 0 + num_workers: 10 + OFFICIAL_EPOCH_LENGTH: 1250 + cache_dataset: true + centering: "centering" # or "sinkhorn_knopp" + cell_augmentation: false +student: + arch: vit_large + patch_size: 16 + drop_path_rate: 0.3 + layerscale: 1.0e-05 + drop_path_uniform: true + pretrained_weights: '' + ffn_layer: "mlp" + block_chunks: 0 + qkv_bias: true + proj_bias: true + ffn_bias: true + num_register_tokens: 0 + interpolate_antialias: false + interpolate_offset: 0.1 + in_chans: 3 + channel_adaptive: false +teacher: + momentum_teacher: 0.992 + final_momentum_teacher: 1 + warmup_teacher_temp: 0.04 + teacher_temp: 0.07 + warmup_teacher_temp_epochs: 30 + in_chans: 3 + channel_adaptive: false +optim: + epochs: 100 + weight_decay: 0.04 + weight_decay_end: 0.4 + base_lr: 0.004 # learning rate for a batch size of 1024 + lr: 0. # will be set after applying scaling rule + warmup_epochs: 10 + min_lr: 1.0e-06 + clip_grad: 3.0 + freeze_last_layer_epochs: 1 + scaling_rule: sqrt_wrt_1024 + patch_embed_lr_mult: 0.2 + layerwise_decay: 0.9 + adamw_beta1: 0.9 + adamw_beta2: 0.999 +crops: + global_crops_scale: + - 0.32 + - 1.0 + local_crops_number: 8 + local_crops_scale: + - 0.05 + - 0.32 + global_crops_size: 224 + local_crops_size: 96 +evaluation: + eval_period_iterations: 12500 diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/eval/segmentation_m2f/models/losses/cross_entropy_loss.py b/torch_hub/facebookresearch_dinov2_main/dinov2/eval/segmentation_m2f/models/losses/cross_entropy_loss.py new file mode 100644 index 0000000000000000000000000000000000000000..0a1f9dd4aa52ebe94cc527db36b1c7fa2f53813e --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/dinov2/eval/segmentation_m2f/models/losses/cross_entropy_loss.py @@ -0,0 +1,279 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# This source code is licensed under the Apache License, Version 2.0 +# found in the LICENSE file in the root directory of this source tree. + +import warnings + +import torch +import torch.nn as nn +import torch.nn.functional as F +from mmseg.models.builder import LOSSES +from mmseg.models.losses.utils import get_class_weight, weight_reduce_loss + + +def cross_entropy( + pred, + label, + weight=None, + class_weight=None, + reduction="mean", + avg_factor=None, + ignore_index=-100, + avg_non_ignore=False, +): + """cross_entropy. The wrapper function for :func:`F.cross_entropy` + + Args: + pred (torch.Tensor): The prediction with shape (N, 1). + label (torch.Tensor): The learning label of the prediction. + weight (torch.Tensor, optional): Sample-wise loss weight. + Default: None. + class_weight (list[float], optional): The weight for each class. + Default: None. + reduction (str, optional): The method used to reduce the loss. + Options are 'none', 'mean' and 'sum'. Default: 'mean'. + avg_factor (int, optional): Average factor that is used to average + the loss. Default: None. + ignore_index (int): Specifies a target value that is ignored and + does not contribute to the input gradients. When + ``avg_non_ignore `` is ``True``, and the ``reduction`` is + ``''mean''``, the loss is averaged over non-ignored targets. + Defaults: -100. + avg_non_ignore (bool): The flag decides to whether the loss is + only averaged over non-ignored targets. Default: False. + `New in version 0.23.0.` + """ + + # class_weight is a manual rescaling weight given to each class. + # If given, has to be a Tensor of size C element-wise losses + loss = F.cross_entropy(pred, label, weight=class_weight, reduction="none", ignore_index=ignore_index) + + # apply weights and do the reduction + # average loss over non-ignored elements + # pytorch's official cross_entropy average loss over non-ignored elements + # refer to https://github.com/pytorch/pytorch/blob/56b43f4fec1f76953f15a627694d4bba34588969/torch/nn/functional.py#L2660 # noqa + if (avg_factor is None) and avg_non_ignore and reduction == "mean": + avg_factor = label.numel() - (label == ignore_index).sum().item() + if weight is not None: + weight = weight.float() + loss = weight_reduce_loss(loss, weight=weight, reduction=reduction, avg_factor=avg_factor) + + return loss + + +def _expand_onehot_labels(labels, label_weights, target_shape, ignore_index): + """Expand onehot labels to match the size of prediction.""" + bin_labels = labels.new_zeros(target_shape) + valid_mask = (labels >= 0) & (labels != ignore_index) + inds = torch.nonzero(valid_mask, as_tuple=True) + + if inds[0].numel() > 0: + if labels.dim() == 3: + bin_labels[inds[0], labels[valid_mask], inds[1], inds[2]] = 1 + else: + bin_labels[inds[0], labels[valid_mask]] = 1 + + valid_mask = valid_mask.unsqueeze(1).expand(target_shape).float() + + if label_weights is None: + bin_label_weights = valid_mask + else: + bin_label_weights = label_weights.unsqueeze(1).expand(target_shape) + bin_label_weights = bin_label_weights * valid_mask + + return bin_labels, bin_label_weights, valid_mask + + +def binary_cross_entropy( + pred, + label, + weight=None, + reduction="mean", + avg_factor=None, + class_weight=None, + ignore_index=-100, + avg_non_ignore=False, + **kwargs, +): + """Calculate the binary CrossEntropy loss. + + Args: + pred (torch.Tensor): The prediction with shape (N, 1). + label (torch.Tensor): The learning label of the prediction. + Note: In bce loss, label < 0 is invalid. + weight (torch.Tensor, optional): Sample-wise loss weight. + reduction (str, optional): The method used to reduce the loss. + Options are "none", "mean" and "sum". + avg_factor (int, optional): Average factor that is used to average + the loss. Defaults to None. + class_weight (list[float], optional): The weight for each class. + ignore_index (int): The label index to be ignored. Default: -100. + avg_non_ignore (bool): The flag decides to whether the loss is + only averaged over non-ignored targets. Default: False. + `New in version 0.23.0.` + + Returns: + torch.Tensor: The calculated loss + """ + if pred.size(1) == 1: + # For binary class segmentation, the shape of pred is + # [N, 1, H, W] and that of label is [N, H, W]. + assert label.max() <= 1, "For pred with shape [N, 1, H, W], its label must have at " "most 2 classes" + pred = pred.squeeze() + if pred.dim() != label.dim(): + assert (pred.dim() == 2 and label.dim() == 1) or (pred.dim() == 4 and label.dim() == 3), ( + "Only pred shape [N, C], label shape [N] or pred shape [N, C, " "H, W], label shape [N, H, W] are supported" + ) + # `weight` returned from `_expand_onehot_labels` + # has been treated for valid (non-ignore) pixels + label, weight, valid_mask = _expand_onehot_labels(label, weight, pred.shape, ignore_index) + else: + # should mask out the ignored elements + valid_mask = ((label >= 0) & (label != ignore_index)).float() + if weight is not None: + weight = weight * valid_mask + else: + weight = valid_mask + # average loss over non-ignored and valid elements + if reduction == "mean" and avg_factor is None and avg_non_ignore: + avg_factor = valid_mask.sum().item() + + loss = F.binary_cross_entropy_with_logits(pred, label.float(), pos_weight=class_weight, reduction="none") + # do the reduction for the weighted loss + loss = weight_reduce_loss(loss, weight, reduction=reduction, avg_factor=avg_factor) + + return loss + + +def mask_cross_entropy( + pred, target, label, reduction="mean", avg_factor=None, class_weight=None, ignore_index=None, **kwargs +): + """Calculate the CrossEntropy loss for masks. + + Args: + pred (torch.Tensor): The prediction with shape (N, C), C is the number + of classes. + target (torch.Tensor): The learning label of the prediction. + label (torch.Tensor): ``label`` indicates the class label of the mask' + corresponding object. This will be used to select the mask in the + of the class which the object belongs to when the mask prediction + if not class-agnostic. + reduction (str, optional): The method used to reduce the loss. + Options are "none", "mean" and "sum". + avg_factor (int, optional): Average factor that is used to average + the loss. Defaults to None. + class_weight (list[float], optional): The weight for each class. + ignore_index (None): Placeholder, to be consistent with other loss. + Default: None. + + Returns: + torch.Tensor: The calculated loss + """ + assert ignore_index is None, "BCE loss does not support ignore_index" + assert reduction == "mean" and avg_factor is None + num_rois = pred.size()[0] + inds = torch.arange(0, num_rois, dtype=torch.long, device=pred.device) + pred_slice = pred[inds, label].squeeze(1) + return F.binary_cross_entropy_with_logits(pred_slice, target, weight=class_weight, reduction="mean")[None] + + +@LOSSES.register_module(force=True) +class CrossEntropyLoss(nn.Module): + """CrossEntropyLoss. + + Args: + use_sigmoid (bool, optional): Whether the prediction uses sigmoid + of softmax. Defaults to False. + use_mask (bool, optional): Whether to use mask cross entropy loss. + Defaults to False. + reduction (str, optional): . Defaults to 'mean'. + Options are "none", "mean" and "sum". + class_weight (list[float] | str, optional): Weight of each class. If in + str format, read them from a file. Defaults to None. + loss_weight (float, optional): Weight of the loss. Defaults to 1.0. + loss_name (str, optional): Name of the loss item. If you want this loss + item to be included into the backward graph, `loss_` must be the + prefix of the name. Defaults to 'loss_ce'. + avg_non_ignore (bool): The flag decides to whether the loss is + only averaged over non-ignored targets. Default: False. + `New in version 0.23.0.` + """ + + def __init__( + self, + use_sigmoid=False, + use_mask=False, + reduction="mean", + class_weight=None, + loss_weight=1.0, + loss_name="loss_ce", + avg_non_ignore=False, + ): + super(CrossEntropyLoss, self).__init__() + assert (use_sigmoid is False) or (use_mask is False) + self.use_sigmoid = use_sigmoid + self.use_mask = use_mask + self.reduction = reduction + self.loss_weight = loss_weight + self.class_weight = get_class_weight(class_weight) + self.avg_non_ignore = avg_non_ignore + if not self.avg_non_ignore and self.reduction == "mean": + warnings.warn( + "Default ``avg_non_ignore`` is False, if you would like to " + "ignore the certain label and average loss over non-ignore " + "labels, which is the same with PyTorch official " + "cross_entropy, set ``avg_non_ignore=True``." + ) + + if self.use_sigmoid: + self.cls_criterion = binary_cross_entropy + elif self.use_mask: + self.cls_criterion = mask_cross_entropy + else: + self.cls_criterion = cross_entropy + self._loss_name = loss_name + + def extra_repr(self): + """Extra repr.""" + s = f"avg_non_ignore={self.avg_non_ignore}" + return s + + def forward( + self, cls_score, label, weight=None, avg_factor=None, reduction_override=None, ignore_index=-100, **kwargs + ): + """Forward function.""" + assert reduction_override in (None, "none", "mean", "sum") + reduction = reduction_override if reduction_override else self.reduction + if self.class_weight is not None: + class_weight = cls_score.new_tensor(self.class_weight) + else: + class_weight = None + # Note: for BCE loss, label < 0 is invalid. + loss_cls = self.loss_weight * self.cls_criterion( + cls_score, + label, + weight, + class_weight=class_weight, + reduction=reduction, + avg_factor=avg_factor, + avg_non_ignore=self.avg_non_ignore, + ignore_index=ignore_index, + **kwargs, + ) + return loss_cls + + @property + def loss_name(self): + """Loss Name. + + This function must be implemented and will return the name of this + loss function. This name will be used to combine different loss items + by simple sum operation. In addition, if you want this loss item to be + included into the backward graph, `loss_` must be the prefix of the + name. + + Returns: + str: The name of this loss item. + """ + return self._loss_name diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/hub/__pycache__/depthers.cpython-310.pyc b/torch_hub/facebookresearch_dinov2_main/dinov2/hub/__pycache__/depthers.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..491cd270d2530f82408f2ddfd1b019e426b37b1e Binary files /dev/null and b/torch_hub/facebookresearch_dinov2_main/dinov2/hub/__pycache__/depthers.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/hub/text/dinotxt_model.py b/torch_hub/facebookresearch_dinov2_main/dinov2/hub/text/dinotxt_model.py new file mode 100644 index 0000000000000000000000000000000000000000..38d3ccc0ef887f05151202d7f42437601488d446 --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/dinov2/hub/text/dinotxt_model.py @@ -0,0 +1,130 @@ +import math +from dataclasses import dataclass +from typing import Optional, Tuple + +import torch +import torch.nn.functional as F +from torch import nn, Tensor + +from .vision_tower import VisionTower +from .text_tower import TextTower + + +@dataclass +class DinoTxtConfig: + embed_dim: int + vision_model_freeze_backbone: bool = True + vision_model_train_img_size: int = 224 + vision_model_use_class_token: bool = True + vision_model_use_patch_tokens: bool = False + vision_model_num_head_blocks: int = 0 + vision_model_head_blocks_drop_path: float = 0.3 + vision_model_use_linear_projection: bool = False + vision_model_patch_tokens_pooler_type: str = "mean" + vision_model_patch_token_layer: int = 1 # which layer to take patch tokens from + # 1 - last layer, 2 - second last layer, etc. + text_model_freeze_backbone: bool = False + text_model_num_head_blocks: int = 0 + text_model_head_blocks_is_causal: bool = False + text_model_head_blocks_drop_prob: float = 0.0 + text_model_tokens_pooler_type: str = "first" + text_model_use_linear_projection: bool = False + init_logit_scale: float = math.log(1 / 0.07) + init_logit_bias: Optional[float] = None + freeze_logit_scale: bool = False + + +class DinoTxt(nn.Module): + def __init__( + self, + model_config: DinoTxtConfig, + vision_backbone: nn.Module, + text_backbone: nn.Module, + ): + super().__init__() + self.model_config = model_config + self.visual_model = VisionTower( + vision_backbone, + model_config.vision_model_freeze_backbone, + model_config.embed_dim, + model_config.vision_model_num_head_blocks, + model_config.vision_model_head_blocks_drop_path, + model_config.vision_model_use_class_token, + model_config.vision_model_use_patch_tokens, + model_config.vision_model_patch_token_layer, + model_config.vision_model_patch_tokens_pooler_type, + model_config.vision_model_use_linear_projection, + ) + self.text_model = TextTower( + text_backbone, + model_config.text_model_freeze_backbone, + model_config.embed_dim, + model_config.text_model_num_head_blocks, + model_config.text_model_head_blocks_is_causal, + model_config.text_model_head_blocks_drop_prob, + model_config.text_model_tokens_pooler_type, + model_config.text_model_use_linear_projection, + ) + self.logit_scale = nn.Parameter(torch.ones(1) * model_config.init_logit_scale) + if model_config.freeze_logit_scale: + self.logit_scale.requires_grad = False + + def init_weights(self): + self.visual_model.init_weights() + self.text_model.init_weights() + + def get_visual_class_and_patch_tokens(self, image: Tensor) -> Tuple[Tensor, Tensor]: + return self.visual_model.get_class_and_patch_tokens(image) + + def encode_image( + self, + image: Tensor, + normalize: bool = False, + ) -> Tensor: + """ + Encode an image into a vector descriptor containing both global and local features. + + Args: + image (Tensor): Tensor of shape `(batch_size, rgb, height, width)`, normalized using ImageNet mean and std. + normalize (bool, optional): Whether to normalize the output vectors. Default is False. + Image features should always be normalized before comparing them with text features: + Returns: + Tensor: Tensor of shape `(batch_size, embed_dim)` containing the image features. + The first half of the vector corresponds to the global features (class token), + and the second half corresponds to the pooled patch features. + """ + features = self.visual_model(image) + return F.normalize(features, dim=-1) if normalize else features + + def encode_text(self, text: Tensor, normalize: bool = False) -> Tensor: + """ + Encode a text input into a vector descriptor. + + Args: + text (Tensor): Tensor of shape `(batch_size, seq_len)` containing token indices. + normalize (bool, optional): Whether to normalize the output vectors. Default is False. + Text features should be normalized before comparing them with image features: + Returns: + Tensor: Tensor of shape `(batch_size, embed_dim)` containing the text features. + As a consequence of the training procedure, assume that the first half of the tensor corresponds + to global image features and the second half to pooled patch features. + """ + features = self.text_model(text) + return F.normalize(features, dim=-1) if normalize else features + + def get_logits(self, image: Tensor, text: Tensor) -> Tuple[Tensor, Tensor]: + text_features = self.encode_text(text, normalize=True) + image_features = self.encode_image(image, normalize=True) + image_logits = self.logit_scale.exp() * image_features @ text_features.T + text_logits = image_logits.T + return image_logits, text_logits + + def forward( + self, + image: Tensor, + text: Tensor, + ) -> Tuple[Tensor, Tensor, Tensor]: + + text_features = self.encode_text(text, normalize=True) + image_features = self.encode_image(image, normalize=True) + return image_features, text_features, self.logit_scale.exp() diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/layers/__pycache__/attention.cpython-310.pyc b/torch_hub/facebookresearch_dinov2_main/dinov2/layers/__pycache__/attention.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a1c361266cf98ddf7f4df49e878295b7a3846574 Binary files /dev/null and b/torch_hub/facebookresearch_dinov2_main/dinov2/layers/__pycache__/attention.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_dinov2_main/dinov2/layers/__pycache__/mlp.cpython-310.pyc b/torch_hub/facebookresearch_dinov2_main/dinov2/layers/__pycache__/mlp.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..db6bc42f9904778a05af29aeca89f63ea94459b7 Binary files /dev/null and b/torch_hub/facebookresearch_dinov2_main/dinov2/layers/__pycache__/mlp.cpython-310.pyc differ diff --git a/torch_hub/facebookresearch_dinov2_main/docs/Cell-DINO.png b/torch_hub/facebookresearch_dinov2_main/docs/Cell-DINO.png new file mode 100644 index 0000000000000000000000000000000000000000..afc0e14657a78e2b3e40a70189b94aba2ed344e4 --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/docs/Cell-DINO.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:709d2fba153499262d190e22126b10337a4f76c137849d9f65cd0ade8f8d91aa +size 1101690 diff --git a/torch_hub/facebookresearch_dinov2_main/docs/ChannelAdaptiveDINO.png b/torch_hub/facebookresearch_dinov2_main/docs/ChannelAdaptiveDINO.png new file mode 100644 index 0000000000000000000000000000000000000000..0b44be0e4810bec317201bf5707ab7ccaf358fb1 --- /dev/null +++ b/torch_hub/facebookresearch_dinov2_main/docs/ChannelAdaptiveDINO.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29e92f5853abbd9126eab3956022b0e969b83aaaba7e4397d6679f5b279153d2 +size 665674