diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..0ba3e15bf3d96aa973cf9ad6c5540e8944be6aae --- /dev/null +++ b/.gitignore @@ -0,0 +1,8 @@ +models +test_results + +# Python bytecode/cache files +__pycache__/ +**/__pycache__/ +*.py[cod] +*$py.class \ No newline at end of file diff --git a/CITATION.cff b/CITATION.cff new file mode 100644 index 0000000000000000000000000000000000000000..903885201a0c96a34f6fd20dc7e0f468fac08601 --- /dev/null +++ b/CITATION.cff @@ -0,0 +1,27 @@ +cff-version: 1.2.0 +message: "If you use this work, please cite it as below." +title: "ScanDL 2.0: A Generative Model of Eye Movements in Reading Synthesizing Scanpaths and Fixation Durations" +authors: + - family-names: Bolliger + given-names: Lena S. + - family-names: Reich + given-names: David R. + - family-names: Jäger + given-names: Lena A. +date-released: 2025-05-01 +journal: "Proceedings of the ACM on Human-Computer Interaction" +publisher: "Association for Computing Machinery" +location: "New York, NY, USA" +volume: "9" +issue: "ETRA5" +article-number: "5" +numpages: "30" +doi: "10.1145/3725830" +url: "https://doi.org/10.1145/3725830" +abstract: "Eye movements in reading have become a vital tool for investigating the cognitive mechanisms involved in language processing. They are not only used within psycholinguistics but have also been leveraged within the field of NLP to improve the performance of language models on downstream tasks. However, the scarcity and limited generalizability of real eye-tracking data present challenges for data-driven approaches. In response, synthetic scanpaths have emerged as a promising alternative. Despite advances, however, existing machine learning-based methods, including the state-of-the-art ScanDL (Bolliger et al. 2023), fail to incorporate fixation durations into the generated scanpaths, which are crucial for a complete representation of reading behavior. We therefore propose a novel model, denoted ScanDL 2.0, which synthesizes both fixation locations and durations. It sets a new benchmark in generating human-like synthetic scanpaths, demonstrating superior performance across various evaluation settings. Furthermore, psycholinguistic analyses confirm its ability to emulate key phenomena in human reading. Our code as well as pre-trained model weights are available via https://github.com/DiLi-Lab/ScanDL-2.0." +keywords: + - neural networks + - scanpath generation + - eye movements + - reading + - diffusion models \ No newline at end of file diff --git a/CONSTANTS.py b/CONSTANTS.py new file mode 100644 index 0000000000000000000000000000000000000000..c8706ed1251b751c33b456416353c01d6f05af85 --- /dev/null +++ b/CONSTANTS.py @@ -0,0 +1,62 @@ +##### data paths ##### + +# TODO adapt your paths to folder that contains the celer and zuco folders +path_to_celer = "/data/lenbol/data/" # e.g., path_to_celer = 'data/' if 'data/celer/...' +path_to_zuco = "/data/lenbol/data/" # e.g., path_to_zuco = 'data/' if 'data/zuco/...' +path_to_emtec = "/data/lenbol/data/" # e.g., path_to_copco = 'data/' if 'data/copco/...' +path_to_bsc = "/data/lenbol/data/" # e.g., path_to_copco = 'data/' if 'data/BSC/...' + +PATH_TO_FIX = f"{path_to_celer}CELER/data_v2.0/sent_fix.tsv" +PATH_TO_IA = f"{path_to_celer}CELER/data_v2.0/sent_ia.tsv" +SUB_METADATA_PATH = f"{path_to_celer}/CELER/participant_metadata/metadata.tsv" +PATH_TO_EMTEC_FIX = f"{path_to_emtec}/EMTeC/fixations_corrected.csv" +PATH_TO_EMTEC_STIM = f"{path_to_emtec}/EMTeC/stimuli.csv" +PATH_TO_BSC_WORD = f"{path_to_bsc}/BSC/BSC.Word.Info.v2.xlsx" +PATH_TO_BSC_FIX = f"{path_to_bsc}/BSC/BSC.EMD.txt" + + +##### model paths ##### + +# TODO adapt your paths + +# training of original ScanDL for modular use with seq2seq fixdur module +SCANDL_MODULE_TRAIN_PATH = "" +SCANDL_MODULE_INF_PATH = "" + +# training of the fixation duration module +FIXDUR_MODULE_TRAIN_PATH = "" +FIXDUR_MODULE_INF_PATH = "" + +# training and inference of the diffusion-only architecture +DIFFUSION_ONLY_TRAIN_PATH = "" +DIFFUSION_ONLY_INF_PATH = "" + + +# names for EMTeC +# training of original ScanDL for modular use with seq2seq fixdur module +SCANDL_MODULE_TRAIN_PATH_EMTEC = "" +SCANDL_MODULE_INF_PATH_EMTEC = "" + +# training of the fixation duration module +FIXDUR_MODULE_TRAIN_PATH_EMTEC = "" +FIXDUR_MODULE_INF_PATH_EMTEC = "" + + +# names for BSC +# training of original ScanDL for modular use with seq2seq fixdur module +SCANDL_MODULE_TRAIN_PATH_BSC = "" +SCANDL_MODULE_INF_PATH_BSC = "" + +# training of the fixation duration module +FIXDUR_MODULE_TRAIN_PATH_BSC = "" +FIXDUR_MODULE_INF_PATH_BSC = "" + + +# training of ScanDL 2.0 on all EMTeC data for paragraph-level ScanDL 2.0 +COMPLETE_SCANDL_MODULE_TRAIN_PATH_EMTEC = "" +COMPLETE_FIXDUR_MODULE_TRAIN_PATH_EMTEC = "" + + +# training of ScanDL 2.0 on all CELER data for sentence-level ScanDL 2.0 +COMPLETE_SCANDL_MODULE_TRAIN_PATH_CELER = "" +COMPLETE_FIXDUR_MODULE_TRAIN_PATH_CELER = "" diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..354f1e04f1247b3ddcbfc83b7519594d0f1ba261 --- /dev/null +++ b/LICENSE @@ -0,0 +1,121 @@ +Creative Commons Legal Code + +CC0 1.0 Universal + + CREATIVE COMMONS CORPORATION IS NOT A LAW FIRM AND DOES NOT PROVIDE + LEGAL SERVICES. DISTRIBUTION OF THIS DOCUMENT DOES NOT CREATE AN + ATTORNEY-CLIENT RELATIONSHIP. CREATIVE COMMONS PROVIDES THIS + INFORMATION ON AN "AS-IS" BASIS. CREATIVE COMMONS MAKES NO WARRANTIES + REGARDING THE USE OF THIS DOCUMENT OR THE INFORMATION OR WORKS + PROVIDED HEREUNDER, AND DISCLAIMS LIABILITY FOR DAMAGES RESULTING FROM + THE USE OF THIS DOCUMENT OR THE INFORMATION OR WORKS PROVIDED + HEREUNDER. + +Statement of Purpose + +The laws of most jurisdictions throughout the world automatically confer +exclusive Copyright and Related Rights (defined below) upon the creator +and subsequent owner(s) (each and all, an "owner") of an original work of +authorship and/or a database (each, a "Work"). + +Certain owners wish to permanently relinquish those rights to a Work for +the purpose of contributing to a commons of creative, cultural and +scientific works ("Commons") that the public can reliably and without fear +of later claims of infringement build upon, modify, incorporate in other +works, reuse and redistribute as freely as possible in any form whatsoever +and for any purposes, including without limitation commercial purposes. +These owners may contribute to the Commons to promote the ideal of a free +culture and the further production of creative, cultural and scientific +works, or to gain reputation or greater distribution for their Work in +part through the use and efforts of others. + +For these and/or other purposes and motivations, and without any +expectation of additional consideration or compensation, the person +associating CC0 with a Work (the "Affirmer"), to the extent that he or she +is an owner of Copyright and Related Rights in the Work, voluntarily +elects to apply CC0 to the Work and publicly distribute the Work under its +terms, with knowledge of his or her Copyright and Related Rights in the +Work and the meaning and intended legal effect of CC0 on those rights. + +1. Copyright and Related Rights. A Work made available under CC0 may be +protected by copyright and related or neighboring rights ("Copyright and +Related Rights"). Copyright and Related Rights include, but are not +limited to, the following: + + i. the right to reproduce, adapt, distribute, perform, display, + communicate, and translate a Work; + ii. moral rights retained by the original author(s) and/or performer(s); +iii. publicity and privacy rights pertaining to a person's image or + likeness depicted in a Work; + iv. rights protecting against unfair competition in regards to a Work, + subject to the limitations in paragraph 4(a), below; + v. rights protecting the extraction, dissemination, use and reuse of data + in a Work; + vi. database rights (such as those arising under Directive 96/9/EC of the + European Parliament and of the Council of 11 March 1996 on the legal + protection of databases, and under any national implementation + thereof, including any amended or successor version of such + directive); and +vii. other similar, equivalent or corresponding rights throughout the + world based on applicable law or treaty, and any national + implementations thereof. + +2. Waiver. To the greatest extent permitted by, but not in contravention +of, applicable law, Affirmer hereby overtly, fully, permanently, +irrevocably and unconditionally waives, abandons, and surrenders all of +Affirmer's Copyright and Related Rights and associated claims and causes +of action, whether now known or unknown (including existing as well as +future claims and causes of action), in the Work (i) in all territories +worldwide, (ii) for the maximum duration provided by applicable law or +treaty (including future time extensions), (iii) in any current or future +medium and for any number of copies, and (iv) for any purpose whatsoever, +including without limitation commercial, advertising or promotional +purposes (the "Waiver"). Affirmer makes the Waiver for the benefit of each +member of the public at large and to the detriment of Affirmer's heirs and +successors, fully intending that such Waiver shall not be subject to +revocation, rescission, cancellation, termination, or any other legal or +equitable action to disrupt the quiet enjoyment of the Work by the public +as contemplated by Affirmer's express Statement of Purpose. + +3. Public License Fallback. Should any part of the Waiver for any reason +be judged legally invalid or ineffective under applicable law, then the +Waiver shall be preserved to the maximum extent permitted taking into +account Affirmer's express Statement of Purpose. In addition, to the +extent the Waiver is so judged Affirmer hereby grants to each affected +person a royalty-free, non transferable, non sublicensable, non exclusive, +irrevocable and unconditional license to exercise Affirmer's Copyright and +Related Rights in the Work (i) in all territories worldwide, (ii) for the +maximum duration provided by applicable law or treaty (including future +time extensions), (iii) in any current or future medium and for any number +of copies, and (iv) for any purpose whatsoever, including without +limitation commercial, advertising or promotional purposes (the +"License"). The License shall be deemed effective as of the date CC0 was +applied by Affirmer to the Work. Should any part of the License for any +reason be judged legally invalid or ineffective under applicable law, such +partial invalidity or ineffectiveness shall not invalidate the remainder +of the License, and in such case Affirmer hereby affirms that he or she +will not (i) exercise any of his or her remaining Copyright and Related +Rights in the Work or (ii) assert any associated claims and causes of +action with respect to the Work, in either case contrary to Affirmer's +express Statement of Purpose. + +4. Limitations and Disclaimers. + + a. No trademark or patent rights held by Affirmer are waived, abandoned, + surrendered, licensed or otherwise affected by this document. + b. Affirmer offers the Work as-is and makes no representations or + warranties of any kind concerning the Work, express, implied, + statutory or otherwise, including without limitation warranties of + title, merchantability, fitness for a particular purpose, non + infringement, or the absence of latent or other defects, accuracy, or + the present or absence of errors, whether or not discoverable, all to + the greatest extent permissible under applicable law. + c. Affirmer disclaims responsibility for clearing rights of other persons + that may apply to the Work or any use thereof, including without + limitation any person's Copyright and Related Rights in the Work. + Further, Affirmer disclaims responsibility for obtaining any necessary + consents, permissions or other rights required for any use of the + Work. + d. Affirmer understands and acknowledges that Creative Commons is not a + party to this document and has no duty or obligation with respect to + this CC0 or use of the Work. diff --git a/PATHS.py b/PATHS.py new file mode 100644 index 0000000000000000000000000000000000000000..4d24437034fbe52b62f522d389de3bdb73a56d59 --- /dev/null +++ b/PATHS.py @@ -0,0 +1,4 @@ +SENT_SCANDL_MODULE = "ScanDL2/models/sentence/scandl-module/" +SENT_FIXDUR_MODULE = "ScanDL2/models/sentence/fixdur-module/" +PAR_SCANDL_MODULE = "ScanDL2/models/paragraph/scandl-module/" +PAR_FIXDUR_MODULE = "ScanDL2/models/paragraph/fixdur-module/" diff --git a/README.md b/README.md new file mode 100644 index 0000000000000000000000000000000000000000..d1fd2481367672445598490d04b08de9309f80a6 --- /dev/null +++ b/README.md @@ -0,0 +1,212 @@ +# ScanDL 2.0: A Generative Model of Eye Movements in Reading Synthesizing Scanpaths and Fixation Durations + +This repository contains ScanDL 2.0, described in [ScanDL 2.0: A Generative Model of Eye Movements in Reading Synthesizing Scanpaths and Fixation Durations](https://doi.org/10.1145/3725830), together with pretrained weights for paragraph-level and sentence-level scanpath generation. + +The model, pretrained inference API, and research functionality originate from the [authors' implementation](https://github.com/DiLi-Lab/ScanDL-2.0). In this project, the repository has been reorganized under the `ScanDL2` package, and `handler.py` and a Gradio interface in `app.py` have been added. The reorganization does not introduce a new model architecture. + +## Setup + +Run the commands and Python examples below from the project root—the directory containing `ScanDL2/`. This ensures package imports and the relative model paths resolve correctly. + +### Install requirements + +The code uses PyTorch and Hugging Face libraries. + +```bash +python -m pip install -r ScanDL2/requirements.txt +``` + +Install a PyTorch build appropriate for your platform, and install Gradio to use the web interface; neither is included in this requirements file. + +```bash +python -m pip install torch gradio +``` + +Use a separate environment from Eyettention because their dependency versions differ. The requirements contain older version pins and require a compatible Python version. For GPU inference, install CUDA-enabled PyTorch. The local implementation selects CUDA when available and otherwise CPU, where diffusion inference can be slow. Its distributed initialization requires a local socket. BERT and GPT-2 assets must be available from Hugging Face or the local cache. + +## Using pre-trained ScanDL 2.0 + +The authors pretrained ScanDL 2.0 on EMTeC for paragraph-level generation and CELER for sentence-level generation. Both versions generate fixation locations and fixation durations without training from scratch. + +To obtain missing weights, download `models.zip` from the [upstream releases](https://github.com/DiLi-Lab/ScanDL-2.0/releases). Place the model directories under `ScanDL2/models/` and ensure that the paths in [PATHS.py](PATHS.py) match their locations: + +```text +ScanDL2/models/ +├── sentence/ +│ ├── scandl-module/ +│ │ ├── ema_0.9999_080000.pt +│ │ └── training_args.json +│ └── fixdur-module/ +│ ├── seq2seq_fixdur.pt +│ ├── hyperparameters.json +│ └── min_max_scaler.pkl +└── paragraph/ + ├── scandl-module/ # same filenames as above + └── fixdur-module/ # same filenames as above +``` + +### Python example + +```python +import torch +from ScanDL2 import ScanDL2 + +model = ScanDL2( + text_type="sentence", + bsz=2, + save=None, + filename=None, +) +model.eval() +with torch.no_grad(): + output = model(texts=["The quick brown fox jumps over the lazy dog."]) + +print(output) +``` + +Set `text_type="paragraph"` to use the paragraph model. + +### Parameters + +| Parameter | Default | Description | +| --- | --- | --- | +| `text_type` | `"sentence"` | Either `"sentence"` or `"paragraph"`; selects the corresponding pretrained modules | +| `bsz` | `2` | Inference batch size; adjust to available memory | +| `save` | `None` | Optional directory in which to save the output as JSON | +| `filename` | `None` | Optional output filename; defaults to `scandl2_outputs.json` when `save` is set | + +For example, setting `save="outputs"` and `filename="example.json"` saves results to `outputs/example.json`. Repeated calls with the same output name overwrite that file. + +### Input and output + +The model accepts a single string or a list of strings. Use nonempty sentences or paragraphs appropriate for the selected checkpoint. Input and scanpath lengths are bounded by the model configuration; split longer documents before inference. + +The returned dictionary contains: + +| Key | Contents | +| --- | --- | +| `predicted_sp_words` | Predicted scanpaths as lists of words in fixation order | +| `predicted_sp_ids` | Corresponding word-position indices | +| `original_sn` | Each original input as a list of words | +| `predicted_fix_durs` | Predicted fixation durations in milliseconds | +| `unique_idx` | An identifier for each input within the current call | + +Word indices originate from the model's sequence including special tokens; do not assume they are zero-based offsets into `original_sn`. Use the returned word, index, and duration lists together. Generation is stochastic. + +## Gradio interface + +The added [app.py](app.py) provides a web interface to the existing model API. + +```bash +python -m ScanDL2.app +``` + +Open `http://localhost:7860`, enter text, select **sentence** or **paragraph**, and click **Run**. Each nonempty line is processed as a separate input in either mode, so keep each paragraph on one line. Results include a fixation table and raw JSON output. The model is loaded and cached when first requested. + +The app uses port 7860 by default, configurable through `PORT`, and binds to `0.0.0.0`. + +## Endpoint handler + +The added [handler.py](handler.py) defines an `EndpointHandler` adapter intended to accept requests with text in `inputs` and inference settings in `parameters`: + +```json +{ + "inputs": ["The quick brown fox jumps over the lazy dog."], + "parameters": { + "text_type": "sentence", + "bsz": 2 + } +} +``` + +The adapter is intended to select the sentence or paragraph model and return the model's output dictionary. It does not itself start an HTTP server. + +If omitted, `text_type` defaults to `"sentence"` and `bsz` to `2`. The batch size must be a positive integer and is applied to both model components. + +## Training, inference, and evaluation + +The original research workflow trains the ScanDL module for fixation locations and the fixation-duration module for durations, then evaluates their combined predictions. The following describes the corresponding data and code in the reorganized repository. + +### Download the data + +- **CELER:** follow the instructions in the [dataset repository](https://github.com/berzak/celer). +- **ZuCo:** download from the [OSF repository](https://osf.io/q3zws/). The dataset requires substantial storage. +- **Beijing Sentence Corpus (BSC):** download from the [OSF repository](https://osf.io/vr3k8/). +- **EMTeC:** download from the [OSF repository](https://osf.io/ajqze/) or use the authors' [Python download utility](https://github.com/DiLi-Lab/EMTeC/blob/main/get_et_data.py). + +Adapt dataset and output paths in [CONSTANTS.py](CONSTANTS.py). Check the expected filenames and directory spelling in [sp_load_celer_zuco.py](scandl_module/scripts/sp_load_celer_zuco.py). + +### Preprocess the training and test data + +Preprocessing takes time, so save processed datasets and reuse the same split definitions across comparable experiments. The data-loading and processing utilities are in `scandl_module/scripts/sp_load_celer_zuco.py`; [create_data.py](create_data.py) contains preprocessing for training the sentence and paragraph models on CELER and EMTeC. + +Some internal paths still need alignment with the rearranged repository: `create_data.py` reads configuration from an old directory and writes beneath `scandl2_pkg`. Its BSC branch is not implemented. Review those paths before using it for a new training run. + +### ScanDL module + +The fixation-location training code is in [sp_train.py](scandl_module/scripts/sp_train.py), with its launcher in [sp_run_train.py](scandl_module/scripts/sp_run_train.py). Configure the appropriate `SCANDL_MODULE_TRAIN_PATH*` and `SCANDL_MODULE_INF_PATH*` values in `CONSTANTS.py` for the chosen dataset and experiment. + +The launcher retains working-directory and distributed-environment assumptions from the original layout. Review those alongside your processed-data and GPU settings before training. Pretrained location prediction is also exposed by `ScanDLModule` in [model.py](model.py). + +### Fixation-duration module + +The duration-training code is in [train_seq2seq.py](fix_dur_module/train_seq2seq.py). Configure the corresponding `FIXDUR_MODULE_TRAIN_PATH*` and `FIXDUR_MODULE_INF_PATH*` values in `CONSTANTS.py`. Keep the fitted duration scaler, model weights, and hyperparameters together. Pretrained duration prediction is exposed by `FixdurModule` in `model.py` and is applied automatically by `ScanDL2`. + +### Sentence-level and paragraph-level training + +The sentence version uses CELER and the paragraph version uses EMTeC. Their full-data training output locations are configured through the `COMPLETE_SCANDL_MODULE_TRAIN_PATH_*` and `COMPLETE_FIXDUR_MODULE_TRAIN_PATH_*` constants. Train the location and duration components with matching data and configurations, then point `PATHS.py` at their output directories for inference. + +### Evaluation and the diffusion-only ablation + +The paper includes evaluation against human scanpaths and a diffusion-only duration ablation, ScanDL diff-dur. Refer to the [original repository](https://github.com/DiLi-Lab/ScanDL-2.0) for those experiment definitions and evaluation tools; their original evaluation modules are not present under those names in this checkout. + +Local checks can be run from the project root: + +```bash +python -m unittest discover -s ScanDL2/tests +``` + +These include real inference and require dependencies, model assets, and local socket access. They check execution and output structure rather than reproducing the paper's reported metrics. + +## Citation + +```bibtex +@article{bolliger2025scandl2, + author = {Bolliger, Lena S. and Reich, David R. and J\"{a}ger, Lena A.}, + title = {ScanDL 2.0: A Generative Model of Eye Movements in Reading Synthesizing Scanpaths and Fixation Durations}, + year = {2025}, + issue_date = {May 2025}, + publisher = {Association for Computing Machinery}, + address = {New York, NY, USA}, + volume = {9}, + number = {ETRA5}, + url = {https://doi.org/10.1145/3725830}, + doi = {10.1145/3725830}, + abstract = {Eye movements in reading have become a vital tool for investigating the cognitive mechanisms involved in language processing. They are not only used within psycholinguistics but have also been leveraged within the field of NLP to improve the performance of language models on downstream tasks. However, the scarcity and limited generalizability of real eye-tracking data present challenges for data-driven approaches. In response, synthetic scanpaths have emerged as a promising alternative. Despite advances, however, existing machine learning-based methods, including the state-of-the-art ScanDL (Bolliger et al. 2023), fail to incorporate fixation durations into the generated scanpaths, which are crucial for a complete representation of reading behavior. We therefore propose a novel model, denoted ScanDL 2.0, which synthesizes both fixation locations and durations. It sets a new benchmark in generating human-like synthetic scanpaths, demonstrating superior performance across various evaluation settings. Furthermore, psycholinguistic analyses confirm its ability to emulate key phenomena in human reading. Our code as well as pre-trained model weights are available via https://github.com/DiLi-Lab/ScanDL-2.0.}, + journal = {Proceedings of the ACM on Human-Computer Interaction}, + month = may, + articleno = {5}, + numpages = {30}, + keywords = {neural networks, scanpath generation, eye movements, reading, diffusion models} +} +``` + +## Related paper + +The fixation-location module builds on **ScanDL: A Diffusion Model for Generating Synthetic Scanpaths on Texts** (Bolliger et al., EMNLP 2023). [Paper and citation metadata](https://aclanthology.org/2023.emnlp-main.960/). + +```bibtex +@inproceedings{bolliger2023scandl, + title={ScanDL: A Diffusion Model for Generating Synthetic Scanpaths on Texts}, + author={Bolliger, Lena S. and Reich, David R. and Haller, Patrick and Jakobi, Deborah N. and Prasse, Paul and Jäger, Lena A.}, + booktitle={Proceedings of the 2023 Conference on Empirical Methods in Natural Language Processing}, + year={2023}, + pages={15513--15538}, + doi={10.18653/v1/2023.emnlp-main.960}, + url={https://aclanthology.org/2023.emnlp-main.960/} +} +``` + +## License + +The included [LICENSE](LICENSE) is **CC0 1.0 Universal**, a public-domain dedication with a fallback license. The full file contains the applicable terms and disclaimers. Preserve provenance and cite the research when using it in scientific work. External datasets, pretrained language-model assets, and dependencies retain their own terms; this README does not assign them a new license. diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..97744a373eec34859526c56c612d74fdc69b609e --- /dev/null +++ b/__init__.py @@ -0,0 +1,27 @@ +from ScanDL2.model import ScanDL2 +from ScanDL2.model import ScanDLModule +from ScanDL2.model import FixdurModule +from ScanDL2.scandl_module.original_scandl.sp_transformer_model import TransformerNetModel +from ScanDL2.fix_dur_module.model_seq2seq import Seq2SeqModel + +from ScanDL2.scandl_module.original_scandl.sp_gaussian_diffusion import GaussianDiffusion +from ScanDL2.scandl_module.original_scandl.sp_gaussian_diffusion import SpacedDiffusion +from ScanDL2.fix_dur_module.model_seq2seq import Pooler +from ScanDL2.scandl_module.original_scandl.sp_rounding import denoised_fn_round + +from ScanDL2 import utils +from ScanDL2 import training_utils + +__all__ = [ + "ScanDL2", + "ScanDL", + "FixdurModule", + "TransformerNetModel", + "ScanDLModule", + "GaussianDiffusion", + "SpacedDiffusion", + "Pooler", + "denoised_fn_round", + "utils", + "training_utils", +] diff --git a/app.py b/app.py new file mode 100644 index 0000000000000000000000000000000000000000..1dc4a303c9f6129ba83f15128b0d731d7ddb7ce3 --- /dev/null +++ b/app.py @@ -0,0 +1,185 @@ +import os +import sys +import json +import tempfile +from typing import Union, List, Dict, Any + +import gradio as gr +import torch + + +from ScanDL2 import ScanDL2 + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + + +_MODELS: Dict[str, ScanDL2] = {} + + +def get_model(text_type: str) -> ScanDL2: + """Load and cache a ScanDL2 model for the given text type.""" + if text_type not in _MODELS: + # `save=None` -> we handle saving ourselves in the Gradio app + _MODELS[text_type] = ScanDL2(text_type=text_type, save=None) + return _MODELS[text_type] + + +def predict( + text: str, + text_type: str, + progress=gr.Progress(track_tqdm=True), +) -> Dict[str, Any]: + """ + Run ScanDL2 on the input text and return a structured result. + """ + if text is None or text.strip() == "": + raise gr.Error("Please provide some input text.") + + lines = [ln.strip() for ln in text.strip().split("\n") if ln.strip()] + if len(lines) == 0: + raise gr.Error("Input text is empty after cleaning.") + + progress(0.05, desc="Loading ScanDL2 model...") + model = get_model(text_type) + + progress(0.15, desc="Running ScanDL + FixDur modules...") + with torch.no_grad(): + output = model(texts=lines) + + return output + + +def format_output(output: Dict[str, Any]) -> str: + """Pretty-print the ScanDL2 output for display in the UI.""" + if output is None: + return "" + + n = len(output.get("original_sn", [])) + lines: List[str] = [] + + for i in range(n): + original_sn = output["original_sn"][i] + sp_words = output["predicted_sp_words"][i] + sp_ids = output["predicted_sp_ids"][i] + fix_durs = output["predicted_fix_durs"][i] + + lines.append(f"### Example {i + 1}") + lines.append("") + lines.append(f"**Original text:** {' '.join(original_sn)}") + lines.append("") + + # Build a readable scanpath table + lines.append("| # | Word | Word index | Fixation duration (ms) |") + lines.append("|---|------|-----------|------------------------|") + for step, (w, wid, dur) in enumerate(zip(sp_words, sp_ids, fix_durs), start=1): + lines.append(f"| {step} | {w} | {wid} | {dur} |") + lines.append("") + + return "\n".join(lines) + + +def format_json(output: Dict[str, Any]) -> str: + """Return the raw JSON string of the output.""" + if output is None: + return "{}" + return json.dumps(output, indent=2, ensure_ascii=False) + + +DESCRIPTION = """ +# ScanDL 2.0 + +**ScanDL 2.0** predicts human-like **eye-movement scanpaths** (which words are fixated, in what order) +and their **fixation durations** (in milliseconds) directly from text. + +This Space wraps two jointly-trained modules: + +1. **ScanDL module** — a discrete diffusion model that generates fixation *locations* (a scanpath) over the input text. +2. **FixDur module** — a sequence-to-sequence model that predicts the *duration* of each fixation. + +### How to use +1. Paste your text in the box below. + - In **sentence** mode, each line is treated as a separate sentence. + - In **paragraph** mode, each line is treated as a separate paragraph. +2. Choose the text type (must match the model checkpoint you want to use). +3. Click **Run**. + +### Output +- A human-readable scanpath table with predicted fixation durations per word. +- The raw JSON output (fixated words, word indices, and durations). +""" + +EXAMPLES = [ + [ + "The quick brown fox jumps over the lazy dog.", + "sentence", + ], + [ + "Researchers have long been interested in how humans process written language.\n" + "Eye-tracking studies reveal where and for how long readers fixate on words.", + "paragraph", + ], +] + + +def build_demo() -> gr.Blocks: + with gr.Blocks( + title="ScanDL 2.0 — Eye-Movement Scanpath Prediction", + theme=gr.themes.Soft(), + ) as demo: + gr.Markdown(DESCRIPTION) + + with gr.Row(): + with gr.Column(scale=3): + text_in = gr.Textbox( + label="Input text", + placeholder="Paste a sentence or paragraph here...", + lines=8, + ) + text_type_in = gr.Radio( + choices=["sentence", "paragraph"], + value="sentence", + label="Text type", + info="Must match an available ScanDL2 checkpoint.", + ) + with gr.Row(): + run_btn = gr.Button("Run", variant="primary") + clear_btn = gr.Button("Clear") + + with gr.Column(scale=4): + table_out = gr.Markdown( + label="Predicted scanpath", + value="_Results will appear here._", + ) + json_out = gr.Code( + label="Raw JSON output", + language="json", + value="{}", + ) + + gr.Examples(examples=EXAMPLES, inputs=[text_in, text_type_in]) + + def _run(text, text_type): + output = predict(text, text_type) + return format_output(output), format_json(output) + + run_btn.click( + fn=_run, + inputs=[text_in, text_type_in], + outputs=[table_out, json_out], + ) + clear_btn.click( + fn=lambda: ("", "sentence", "_Results will appear here._", "{}"), + inputs=None, + outputs=[text_in, text_type_in, table_out, json_out], + ) + + return demo + + +if __name__ == "__main__": + demo = build_demo() + demo.queue(max_size=16).launch( + server_name="0.0.0.0", + server_port=int(os.environ.get("PORT", 7860)), + show_error=True, + ) diff --git a/config.json b/config.json new file mode 100644 index 0000000000000000000000000000000000000000..db9b7571b1c2574d85be334a2a64a5bc935f70e3 --- /dev/null +++ b/config.json @@ -0,0 +1,53 @@ +{ + "lr": 0.0001, + "batch_size": 128, + "microbatch": 64, + "learning_steps": 80000, + "log_interval": 50, + "save_interval": 5000, + "eval_interval": 500, + "ema_rate": "0.9999", + "resume_checkpoint": "none", + "schedule_sampler": "lossaware", + "diffusion_steps": 2000, + "noise_schedule": "sqrt", + "timestep_respacing": "", + "vocab": "bert", + "use_plm_init": "no", + "vocab_size": 0, + "config_name": "bert-base-cased", + "gpt_config_name": "gpt2", + "notes": "folder-notes", + "data_dir": "processed_data", + "dataset": "dataset-name", + "checkpoint_path": "checkpoint-path/test-run", + "seq_len": 128, + "hidden_t_dim": 128, + "hidden_dim": 256, + "dropout": 0.1, + "use_fp16": false, + "fp16_scale_growth": 0.001, + "seed": 102, + "gradient_clipping": -1.0, + "weight_decay": 0.0, + "learn_sigma": false, + "use_kl": false, + "predict_xstart": true, + "rescale_timesteps": true, + "rescale_learned_sigmas": false, + "sigma_small": false, + "emb_scale_factor": 1.0, + "num_transformer_layers": 12, + "num_transformer_heads": 8, + "one_noise_step": true, + "mask_padding": false, + "celer_only_L1": true, + "data_split_criterion": "scanpath", + "corpus": "celer", + "inference": "none", + "n_folds": 5, + "ablation_type": "none", + "nll_in_loss": false, + "load_from_checkpoint": false, + "load_train_data": "-" +} \ No newline at end of file diff --git a/config_bsc.json b/config_bsc.json new file mode 100644 index 0000000000000000000000000000000000000000..2882ca299dc8406a34caf5d6b3ebc07182f493be --- /dev/null +++ b/config_bsc.json @@ -0,0 +1,54 @@ +{ + "lr": 0.0001, + "batch_size": 128, + "microbatch": 64, + "learning_steps": 80000, + "log_interval": 50, + "save_interval": 5000, + "eval_interval": 500, + "ema_rate": "0.9999", + "resume_checkpoint": "none", + "schedule_sampler": "lossaware", + "diffusion_steps": 2000, + "noise_schedule": "sqrt", + "timestep_respacing": "", + "vocab": "bert", + "use_plm_init": "no", + "vocab_size": 0, + "config_name": "bert-base-chinese", + "gpt_config_name": "benjamin/gpt2-wechsel-chinese", + "notes": "folder-notes", + "data_dir": "processed_data_bsc", + "dataset": "dataset-name", + "checkpoint_path": "checkpoint-path/test-run", + "seq_len": 68, + "hidden_t_dim": 68, + "hidden_dim": 256, + "dropout": 0.1, + "use_fp16": false, + "fp16_scale_growth": 0.001, + "seed": 102, + "gradient_clipping": -1.0, + "weight_decay": 0.0, + "learn_sigma": false, + "use_kl": false, + "predict_xstart": true, + "rescale_timesteps": true, + "rescale_learned_sigmas": false, + "sigma_small": false, + "emb_scale_factor": 1.0, + "num_transformer_layers": 12, + "num_transformer_heads": 8, + "one_noise_step": true, + "mask_padding": false, + "celer_only_L1": true, + "data_split_criterion": "scanpath", + "corpus": "bsc", + "inference": "none", + "n_folds": 5, + "ablation_type": "none", + "nll_in_loss": false, + "load_from_checkpoint": false, + "load_train_data": "-" + } + \ No newline at end of file diff --git a/config_emtec.json b/config_emtec.json new file mode 100644 index 0000000000000000000000000000000000000000..7c3ae13a86b5d939395de24e86acb8f7a336599d --- /dev/null +++ b/config_emtec.json @@ -0,0 +1,54 @@ +{ + "lr": 0.0001, + "batch_size": 128, + "microbatch": 64, + "learning_steps": 80000, + "log_interval": 50, + "save_interval": 5000, + "eval_interval": 500, + "ema_rate": "0.9999", + "resume_checkpoint": "none", + "schedule_sampler": "lossaware", + "diffusion_steps": 2000, + "noise_schedule": "sqrt", + "timestep_respacing": "", + "vocab": "bert", + "use_plm_init": "no", + "vocab_size": 0, + "config_name": "bert-base-cased", + "gpt_config_name": "gpt2", + "notes": "folder-notes", + "data_dir": "processed_data_emtec", + "dataset": "dataset-name", + "checkpoint_path": "checkpoint-path/test-run", + "seq_len": 352, + "hidden_t_dim": 352, + "hidden_dim": 256, + "dropout": 0.1, + "use_fp16": false, + "fp16_scale_growth": 0.001, + "seed": 102, + "gradient_clipping": -1.0, + "weight_decay": 0.0, + "learn_sigma": false, + "use_kl": false, + "predict_xstart": true, + "rescale_timesteps": true, + "rescale_learned_sigmas": false, + "sigma_small": false, + "emb_scale_factor": 1.0, + "num_transformer_layers": 12, + "num_transformer_heads": 8, + "one_noise_step": true, + "mask_padding": false, + "celer_only_L1": true, + "data_split_criterion": "scanpath", + "corpus": "emtec", + "inference": "none", + "n_folds": 5, + "ablation_type": "none", + "nll_in_loss": false, + "load_from_checkpoint": false, + "load_train_data": "-" + } + \ No newline at end of file diff --git a/create_data.py b/create_data.py new file mode 100644 index 0000000000000000000000000000000000000000..e506ed722e7116ceaced9d4b82fe02f2a443d210 --- /dev/null +++ b/create_data.py @@ -0,0 +1,167 @@ +""" +Create the data for training ScanDL on all data. +""" + +import argparse +import os +import json +import numpy as np +import pandas as pd +import sys + + +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import ( + load_celer, + load_celer_speakers, + process_celer, +) +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import ( + load_zuco, + process_zuco, + get_kfold, + get_kfold_indices_combined, +) +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_emtec, process_emtec +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_bsc, process_bsc +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import flatten_data, unflatten_data +from transformers import set_seed, BertTokenizerFast + +sys.path.append("./") +sys.path.append("../") + + +def create_argparser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser() + parser.add_argument( + "--folder-name", + type=str, + default="processed_data_all", + help="Name of the folder to save the processed data in.", + ) + parser.add_argument( + "--max-fix-dur", + type=int, + help="max fixatino duration value. greater fixation durations are replaced with this value.", + default=999, + ) + parser.add_argument( + "--data", + type=str, + choices=["celer", "emtec", "bsc"], + required=True, + ) + defaults = dict() + defaults.update(load_defaults_config(parser.parse_args())) + + add_dict_to_argparser(parser, defaults) + return parser + + +def load_defaults_config(args): + """ + Load defaults for training args. + """ + if args.data == "emtec": + config_name = "config_emtec.json" + elif args.data == "bsc": + config_name = "config_bsc.json" + else: + config_name = "config.json" + with open(f"diffusion_only/scandl_diff_dur/{config_name}", "r") as f: + return json.load(f) + + +def add_dict_to_argparser(parser, default_dict): + for k, v in default_dict.items(): + v_type = type(v) + if v is None: + v_type = str + elif isinstance(v, bool): + v_type = str2bool + parser.add_argument(f"--{k}", default=v, type=v_type) + + +def str2bool(v): + """ + https://stackoverflow.com/questions/15008758/parsing-boolean-values-with-argparse + """ + if isinstance(v, bool): + return v + if v.lower() in ("yes", "true", "t", "y", "1"): + return True + elif v.lower() in ("no", "false", "f", "n", "0"): + return False + else: + raise argparse.ArgumentTypeError("boolean value expected") + + +def main(): + + base_folder_name = "scandl2_pkg" + + print("Loading argument parser...") + args = create_argparser().parse_args() + set_seed(args.seed) + + if args.data == "celer": + + tokenizer = BertTokenizerFast.from_pretrained(args.config_name) + data_path = args.folder_name + "_celer" + if not os.path.exists(os.path.join(base_folder_name, data_path)): + os.makedirs(os.path.join(base_folder_name, data_path)) + + # load Celer data + word_info_df, eyemovement_df = load_celer() + reader_list = load_celer_speakers(only_native_speakers=args.celer_only_L1) + sn_list = np.unique( + word_info_df[word_info_df["list"].isin(reader_list)].sentenceid.values + ).tolist() + + data, splitting_IDs_dict = process_celer( + sn_list=sn_list, + reader_list=reader_list, + word_info_df=word_info_df, + eyemovement_df=eyemovement_df, + tokenizer=tokenizer, + args=args, + inference="cv", + max_fix_dur=args.max_fix_dur, + ) + flattened_data = flatten_data(data) + flattened_data = np.array(flattened_data, dtype=object).tolist() + train_data = unflatten_data(flattened_data=flattened_data, split="train") + train_data.save_to_disk(os.path.join(base_folder_name, data_path)) + + elif args.data == "bsc": + + raise NotImplementedError("BSC data not implemented yet.") + + elif args.data == "emtec": + + tokenizer = BertTokenizerFast.from_pretrained(args.config_name) + data_path = args.folder_name + "_emtec" + if not os.path.exists(os.path.join(base_folder_name, data_path)): + os.makedirs(os.path.join(base_folder_name, data_path)) + + # load EMTeC data + print("Loading EMTeC data...") + fixations_df, stimuli_df = load_emtec() + data, splitting_IDs_dict = process_emtec( + fixations_df=fixations_df, + stimuli_df=stimuli_df, + tokenizer=tokenizer, + args=args, + inference="cv", + max_fix_dur=args.max_fix_dur, + ) + flattened_data = flatten_data(data) + flattened_data = np.array(flattened_data, dtype=object).tolist() + train_data = unflatten_data(flattened_data=flattened_data, split="train") + train_data.save_to_disk(os.path.join(base_folder_name, data_path)) + + else: + raise NotImplementedError("Data not implemented yet.") + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/fix_dur_module/__init__.py b/fix_dur_module/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/fix_dur_module/__pycache__/__init__.cpython-313.pyc b/fix_dur_module/__pycache__/__init__.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..380f16a9033846c57c1860634a6d093ec2ef23f6 Binary files /dev/null and b/fix_dur_module/__pycache__/__init__.cpython-313.pyc differ diff --git a/fix_dur_module/__pycache__/model_seq2seq.cpython-313.pyc b/fix_dur_module/__pycache__/model_seq2seq.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..daa7a862dcf2095f5cecf5020ce070217413061e Binary files /dev/null and b/fix_dur_module/__pycache__/model_seq2seq.cpython-313.pyc differ diff --git a/fix_dur_module/__pycache__/scasim.cpython-313.pyc b/fix_dur_module/__pycache__/scasim.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..caae5ae6a089fc3fcf7da6d03555fe760a7613c3 Binary files /dev/null and b/fix_dur_module/__pycache__/scasim.cpython-313.pyc differ diff --git a/fix_dur_module/__pycache__/utils_data.cpython-313.pyc b/fix_dur_module/__pycache__/utils_data.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..1cef1f18bd4b889b1d9ebead5347bb7cf681bc29 Binary files /dev/null and b/fix_dur_module/__pycache__/utils_data.cpython-313.pyc differ diff --git a/fix_dur_module/__pycache__/utils_train.cpython-313.pyc b/fix_dur_module/__pycache__/utils_train.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..3c634ecd1a53f208c5ef8de6aee2d410823cb993 Binary files /dev/null and b/fix_dur_module/__pycache__/utils_train.cpython-313.pyc differ diff --git a/fix_dur_module/model_seq2seq.py b/fix_dur_module/model_seq2seq.py new file mode 100644 index 0000000000000000000000000000000000000000..e518c579108141d2dfaa5e297b30bdd8292e8614 --- /dev/null +++ b/fix_dur_module/model_seq2seq.py @@ -0,0 +1,89 @@ +import torch +import torch.nn as nn +from transformers.models.bert.modeling_bert import BertEncoder +from typing import Optional + + +class Seq2SeqModel(nn.Module): + + def __init__( + self, + config, + output_dim, + num_linear, + dropout, + ): + super().__init__() + self.config = config + self.output_dim = output_dim + self.encoder = BertEncoder(config) + self.pooler = Pooler(config) + + layers_list = list() + for i in range(num_linear): + layers_list.append(nn.Linear(config.hidden_size, config.hidden_size)) + layers_list.append(nn.ReLU()) + layers_list.append(nn.Dropout(dropout)) + self.ff = nn.Sequential(*layers_list) + # self.ff = nn.Linear(config.hidden_size, config.hidden_size) + self.ff_out = nn.Linear(config.hidden_size, output_dim) + + def _invert_attention_mask(self, attention_mask): + if attention_mask.dim() == 3: + extended_attention_mask = attention_mask[:, None, :, :] + elif attention_mask.dim() == 2: + extended_attention_mask = attention_mask[:, None, None, :] + extended_attention_mask = (1.0 - extended_attention_mask) * torch.finfo(torch.float32).min + return extended_attention_mask + + def forward( + self, + sp_embeddings, + attention_mask: Optional[torch.Tensor] = None, + output_attentions: Optional[bool] = None, + ): + # get the extended attention mask + # zeros and ones are inverted such that what is not maked is 0 and what is masked is -inf + if attention_mask is not None: + attention_mask = self._invert_attention_mask(attention_mask) + encoder_outputs = self.encoder( + sp_embeddings, + attention_mask=attention_mask, + output_attentions=output_attentions, + ) + else: + encoder_outputs = self.encoder( + sp_embeddings, + output_attentions=output_attentions, + ) + + last_hidden_state = encoder_outputs.last_hidden_state + + # pool the encoder output: the hidden state of the CLS token is passed through another linear layer + pooled_output = self.pooler(last_hidden_state) + + # map to the output dimension + out = self.ff(pooled_output) + out = self.ff_out(out) + + if output_attentions: + attentions = encoder_outputs.attentions + return out, attentions + + else: + return out + + +class Pooler(nn.Module): + def __init__(self, config): + super().__init__() + self.dense = nn.Linear(config.hidden_size, config.hidden_size) + self.activation = nn.Tanh() + + def forward(self, hidden_states): + # pool the output by taking the hidden state of the first token (the CLS token) + # and pass it through another linear layer wtih tanh activation + cls_out = hidden_states[:, 0] + pooled_output = self.dense(cls_out) + pooled_output = self.activation(pooled_output) + return pooled_output diff --git a/fix_dur_module/scasim.py b/fix_dur_module/scasim.py new file mode 100644 index 0000000000000000000000000000000000000000..5468c0e09ef41a828b38ef2297f6c529b523e7c0 --- /dev/null +++ b/fix_dur_module/scasim.py @@ -0,0 +1,185 @@ +""" +Script that implements the scanpath similarity metric ScaSim by +Von der Malsburg, Titus, and Shravan Vasishth. +"What is the scanpath signature of syntactic reanalysis?." +Journal of Memory and Language 65.2 (2011): 109-127. +""" + +from __future__ import annotations +from math import pi, sin, cos, acos +import numpy as np +from typing import List, Tuple, Optional, Any, Union + + +# only need 0, 2 of s/t due to word index instead of x/y location +def scasim( + s: List[Tuple[int, int, Union[int, float]]], + t: List[Tuple[int, int, Union[int, float]]], + modulator: Optional[float] = 0.83, + normalize: Optional[str] = None, # fixations, durations, None +) -> float: + """ + Calculate the similarity between two scanpaths s and t. + :param s: scanpath s, consisting of fixation locations (word indices) and fixation durations + :param t: scanpath t, consisting of fixation locations (word indices) and fixation durations + :param modulator: modulator specifies how spatial distances between fixations are assessed. When set to 0, any spatial divergence of two + compared scanpaths is penalized independently of its degree. When set to 1, the scanpaths are compared only with respect to their + temporal patterns. The default value approximates the sensitivity to spatial distance found in the human visual system. + :param normalize: if 'fixations', the similarity score is normalized by the number of fixations in the two scanpaths. If 'durations', + the similarity score is normalized by the sum of fixation durations in the two scanpaths. If None, no normalization is applied. + + :return: similarity between scanpaths s and t + """ + m, n = len(s), len(t) + d = [list(map(lambda i: 0, range(n + 1))) for _ in range(m + 1)] + + # sum of fixation durations of the two scanpaths + s_fixdur_sum = sum([fix[2] for fix in s]) + t_fixdur_sum = sum([fix[2] for fix in t]) + # number of fixations in the two scanpaths + s_nfix = len(s) + t_nfix = len(t) + + acc = 0 + # sequence alignment + # loop over fixations in scanpath s: + for fix_i in range(1, m + 1): + acc += s[fix_i - 1][2] + d[fix_i][0] = acc + + # loop over fixations in scanpath t: + acc = 0 + for fix_j in range(1, n + 1): + acc += t[fix_j - 1][2] + d[0][fix_j] = acc + + # Compute similarity: + for fix_i in range(n): + for fix_j in range(m): + # calculating angle between fixation targets: + slon = s[fix_j][0] / (180 / pi) # longitude (x-axis) + tlon = t[fix_i][0] / (180 / pi) + slat = s[fix_j][1] / (180 / pi) # latitude (y-axis) + tlat = t[fix_i][1] / (180 / pi) + + angle = acos(sin(slat) * sin(tlat) + cos(slat) * cos(tlat) * cos(slon - tlon)) * ( + 180 / pi + ) + + # approximation of cortical magnification: + mixer = modulator**angle + + # cost for substitution: + cost = abs(t[fix_i][2] - s[fix_j][2]) * mixer + (t[fix_i][2] + s[fix_j][2]) * ( + 1.0 - mixer + ) + + # select optimal edit operation + ops = ( + d[fix_j][fix_i + 1] + s[fix_j][2], + d[fix_j + 1][fix_i] + t[fix_i][2], + d[fix_j][fix_i] + cost, + ) + + # mi = which_min(*ops) + mi = np.argmin(ops) + + d[fix_j + 1][fix_i + 1] = ops[mi] + + result = d[-1][-1] + if normalize == "fixations": + result /= s_nfix + t_nfix + elif normalize == "durations": + result /= s_fixdur_sum + t_fixdur_sum + + return result + + +def main(): + + predicted_sp_ids = [ + [0, 1, 2, 3, 4, 5, 6, 8, 10, 10, 10, 11], + [0, 1, 1, 2, 4, 5, 7, 4, 3, 7, 1, 8], + [0, 1, 2, 4, 4, 5, 7, 7, 8, 9], + ] + original_sp_ids = [ + [0, 1, 2, 4, 2, 3, 5, 6, 8, 9, 10, 4, 11], + [0, 1, 2, 4, 3, 8], + [0, 1, 6, 8, 9], + ] + predicted_fix_durs = [ + [69, 52, 374, 374, 374, 374, 256, 423, 423, 423, 423, 188], + [69, 52, 374, 384, 374, 374, 423, 423, 423, 423, 52, 423], + [69, 52, 374, 502, 502, 374, 423, 423, 374, 374], + ] + original_fix_durs = [ + [0, 208, 232, 197, 314, 151, 219, 308, 195, 280, 260, 102], + [0, 192, 182, 297, 134], + [0, 195, 130, 101], + ] + + # remove last element in each sublist of list for predicted_sp_ids, original_sp_ids, and predicted_fix_durs + # these are the pad tokens and they are not contained in original_fix_durs + predicted_sp_ids = [sublist[:-1] for sublist in predicted_sp_ids] + original_sp_ids = [sublist[:-1] for sublist in original_sp_ids] + predicted_fix_durs = [sublist[:-1] for sublist in predicted_fix_durs] + + # create dummy y values for original_sp_ids and predicted_sp_ids + dummy_y_original_sp_ids = [[1] * len(sublist) for sublist in original_sp_ids] + dummy_y_predicted_sp_ids = [[1] * len(sublist) for sublist in predicted_sp_ids] + + # zip together the predicted_sp_ids and predicted_fix_durs lists as list of list of tuples + predicted_sp = list( + map( + lambda x, y, z: list(zip(x, y, z)), + predicted_sp_ids, + dummy_y_predicted_sp_ids, + predicted_fix_durs, + ) + ) + # zip together the original_sp_ids and original_fix_durs lists as list of list of tuples + original_sp = list( + map( + lambda x, y, z: list(zip(x, y, z)), + original_sp_ids, + dummy_y_original_sp_ids, + original_fix_durs, + ) + ) + + s1 = predicted_sp[0] + t1 = original_sp[0] + sim1 = scasim(s=s1, t=t1) + + s2 = predicted_sp[1] + t2 = original_sp[1] + sim2 = scasim(s=s2, t=t2) + + s3 = predicted_sp[2] + t3 = original_sp[2] + sim3 = scasim(s=s3, t=t3) + + # normalize by fixations + sim10 = scasim(s=s1, t=t1, normalize="fixations") + sim11 = scasim(s=s2, t=t2, normalize="fixations") + sim12 = scasim(s=s3, t=t3, normalize="fixations") + + # normalize by durations + sim13 = scasim(s=s1, t=t1, normalize="durations") + sim14 = scasim(s=s2, t=t2, normalize="durations") + sim15 = scasim(s=s3, t=t3, normalize="durations") + + print("normalize by fixations") + print("sim10:", sim10) + print("sim11:", sim11) + print("sim12:", sim12) + print("normalize by durations") + print("sim13:", sim13) + print("sim14:", sim14) + print("sim15:", sim15) + + breakpoint() + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/fix_dur_module/train_seq2seq.py b/fix_dur_module/train_seq2seq.py new file mode 100644 index 0000000000000000000000000000000000000000..6d329c8d1e0c0ab279d13862e60ec4ecfe661a43 --- /dev/null +++ b/fix_dur_module/train_seq2seq.py @@ -0,0 +1,283 @@ +""" +The training script for training the fixation duration module. +""" + +import joblib +import sys +import json +import os +from typing import Dict +from argparse import ArgumentParser + +from transformers import GPT2TokenizerFast, GPT2LMHeadModel, GPT2Model, AutoConfig, BertModel +from transformers.models.bert.modeling_bert import BertEncoder, BertPooler +from transformers import AdamW, get_linear_schedule_with_warmup + +import numpy as np +import torch +import torch.nn as nn +from torch.utils.data import Dataset, DataLoader +from datasets import load_from_disk, DatasetDict +from sklearn.preprocessing import MinMaxScaler + +from ScanDL2.fix_dur_module.utils_data import ( + prepare_seq2seq_data, + get_embeddings_seq2seq, + Seq2SeqDataset, + split_train_val_data, +) +from ScanDL2.fix_dur_module.model_seq2seq import Seq2SeqModel +from ScanDL2.fix_dur_module.utils_train import EarlyStopping, train + +sys.path.append("./") +sys.path.append("../") +sys.path.append("../../") + +from ScanDL2.CONSTANTS import ( + COMPLETE_FIXDUR_MODULE_TRAIN_PATH_BSC, + COMPLETE_FIXDUR_MODULE_TRAIN_PATH_CELER, + COMPLETE_FIXDUR_MODULE_TRAIN_PATH_EMTEC, +) + + +def get_parser() -> ArgumentParser: + parser = ArgumentParser() + parser.add_argument( + "--max-length", + type=int, + default=128, + help="The maximum sequence length.", + ) + parser.add_argument( + "--num-heads", + type=int, + default=12, + help="The number of attention heads in the Transformer encoder.", + ) + parser.add_argument( + "--num-layers", + type=int, + default=12, + help="The number of layers in the Transformer encoder.", + ) + parser.add_argument( + "--num-linear", + type=int, + default=8, + help="The number of linear layers.", + ) + parser.add_argument( + "--bsz", + type=int, + default=128, + help="The batch size.", + ) + parser.add_argument( + "--dropout", + type=float, + default=0.5, + help="The dropout rate.", + ) + parser.add_argument( + "--num-epochs", + type=int, + default=400, + ) + parser.add_argument( + "--sp-pad-token", + type=int, + default=127, + help="the padding token appended to the sp, usually seq_len-1", + ) + parser.add_argument( + "--use-attention-mask", + action="store_true", + help="Whether to use the attention mask in the Transformer encoder.", + ) + parser.add_argument( + "--data", + type=str, + required=True, + choices=["emtec", "bsc", "celer"], + help="The dataset to train on.", + ) + return parser + + +def main(): + + args = get_parser().parse_args() + + max_length = args.max_length + output_attentions = False + learning_rate = 1e-4 + num_epochs = args.num_epochs + patience = 25 + normalize = True + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + if args.data == "emtec": + path_save_model = COMPLETE_FIXDUR_MODULE_TRAIN_PATH_EMTEC + path_to_data = "processed_data_all_emtec" + elif args.data == "bsc": + raise NotImplementedError("Training on BSC data is not yet implemented.") + path_save_model = COMPLETE_FIXDUR_MODULE_TRAIN_PATH_BSC + path_to_data = "processed_data_all_bsc" + elif args.data == "celer": + path_save_model = COMPLETE_FIXDUR_MODULE_TRAIN_PATH_CELER + path_to_data = "processed_data_all_celer" + else: + raise ValueError("Unknown dataset.") + + if not os.path.exists(path_save_model): + os.makedirs(path_save_model) + model_name = "seq2seq_fixdur.pt" + + hypeparameters = { + "num_heads": args.num_heads, + "num_layers": args.num_layers, + "num_linear": args.num_linear, + "bsz": args.bsz, + "dropout": args.dropout, + "use_attention_mask": args.use_attention_mask, + } + with open(os.path.join(path_save_model, "hyperparameters.json"), "w") as f: + json.dump(hypeparameters, f) + + # load GPT-2 and GPT-2 tokenizer to get the contextualized embeddings + if args.data == "bsc": + raise NotImplementedError("Training on BSC data is not yet implemented.") + gpt_config_name = "benjamin/gpt2-wechsel-chinese" + else: + gpt_config_name = "gpt2" + + tokenizer = GPT2TokenizerFast.from_pretrained(gpt_config_name, add_prefix_space=True) + gpt2_model = GPT2Model.from_pretrained(gpt_config_name) + tokenizer.pad_token = tokenizer.eos_token + # freeze parameters + for param in gpt2_model.parameters(): + param.requires_grad = False + + # load BERT config (for model architecture) and BERT model (for embeddings of CLS and PAD tokens) + + if args.data == "bsc": + raise NotImplementedError("Training on BSC data is not yet implemented.") + bert_config_name = "bert-base-chinese" + else: + bert_config_name = "bert-base-cased" + config = AutoConfig.from_pretrained(bert_config_name) + bert_embeddings = BertModel.from_pretrained(bert_config_name).embeddings.word_embeddings + # freeze parameters + for param in bert_embeddings.parameters(): + param.requires_grad = False + + # change the parameters in the config + config.num_attention_heads = args.num_heads + config.num_hidden_layers = args.num_layers + + # training + print("--- load and prepare data ...") + train_data = load_from_disk(os.path.join("scandl2_pkg", path_to_data, "train")) + new_data = DatasetDict() + new_data["train"] = train_data + + # prepare the data for training + data = prepare_seq2seq_data( + data=new_data, + tokenizer=tokenizer, + gpt2_model=gpt2_model, + bert_embeddings=bert_embeddings, + aggregate="mean", + max_length=max_length, + sp_pad_token=args.sp_pad_token, + ) + + fix_dur_colname = "fix_durs" + + if normalize: + min_max_scaler = MinMaxScaler() + fix_durs = [t.cpu().detach().numpy() for t in data["fix_durs"]] + flattened = np.concatenate(fix_durs).reshape(-1, 1) + # fit the scaler on the training data + min_max_scaler.fit(flattened) + # normalize the fixation durations + flattened_normalized = min_max_scaler.transform(flattened) + # reshape + split_indices = [len(t) for t in fix_durs] + normalized_data = np.split(flattened_normalized.flatten(), np.cumsum(split_indices)[:-1]) + # convert back to tensors + normalized_tensors = [torch.tensor(t) for t in normalized_data] + data["fix_durs_normalized"] = normalized_tensors + # save the scaler (needed for inference) + joblib.dump(min_max_scaler, os.path.join(path_save_model, "min_max_scaler.pkl")) + fix_dur_colname = "fix_durs_normalized" + + # split data into train and val data (val data for early stopping) + train_data, val_data = split_train_val_data( + data=data, + val_size=0.1, + ) + + # create dataset and dataloader + train_dataset = Seq2SeqDataset( + data=train_data, + normalize=normalize, + ) + val_dataset = Seq2SeqDataset( + data=val_data, + normalize=normalize, + ) + train_loader = DataLoader( + train_dataset, + batch_size=args.bsz, + shuffle=True, + ) + val_loader = DataLoader( + val_dataset, + batch_size=args.bsz, + shuffle=False, + ) + + # model, loss, optimizer, scheduler, early stopping + + model = Seq2SeqModel( + config=config, + output_dim=max_length, + num_linear=args.num_linear, + dropout=args.dropout, + ) + model.to(device) + criterion = nn.MSELoss(reduction="mean") + optimizer = AdamW(model.parameters(), lr=learning_rate) + early_stopping = EarlyStopping( + patience=patience, + path=os.path.join(path_save_model, model_name), + ) + + num_training_steps = len(train_loader) * num_epochs + num_warmup_steps = int(0.05 * num_training_steps) + scheduler = get_linear_schedule_with_warmup( + optimizer, + num_warmup_steps=num_warmup_steps, + num_training_steps=num_training_steps, + ) + + # training + train( + model=model, + num_epochs=num_epochs, + train_loader=train_loader, + val_loader=val_loader, + criterion=criterion, + optimizer=optimizer, + early_stopping=early_stopping, + scheduler=scheduler, + device=device, + fix_dur_colname=fix_dur_colname, + output_attentions=output_attentions, + use_attention_mask=args.use_attention_mask, + ) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/fix_dur_module/utils_data.py b/fix_dur_module/utils_data.py new file mode 100644 index 0000000000000000000000000000000000000000..ef125e8d9970356b454212d0524f6acde9b4bd77 --- /dev/null +++ b/fix_dur_module/utils_data.py @@ -0,0 +1,530 @@ +""" +Utils for the fixation duration module. +""" + +import torch +from datasets import load_from_disk +from typing import Dict, Any, List, Optional, Union +import transformers +from torch.utils.data import Dataset +import datasets +from tqdm import tqdm +import random + +# def get_input_embeddings( +# data_instance: Dict[str, Any], +# tokenizer: transformers.GPT2TokenizerFast, +# gpt2_model: transformers.GPT2Model, +# aggregate: str = 'mean', # 'mean', 'sum' +# #max_length: int = 128, +# ): +# # dummy code for now +# sn_repr_len = data_instance['sn_repr_len'] + +# # the sentence +# sn_words = data_instance['words_for_mapping'].split() +# while sn_words[-1] == '[PAD]': +# sn_words.pop() + +# # the scanpath +# # remove the CLS token (already have a SEP token at the end of the sentence and the two will beconcatenated) +# # the SEP token can stay +# sp_ids = data_instance['sn_sp_repr'][sn_repr_len:][1:] +# # cut off the trialing pad tokens +# while sp_ids[-1] == 127: +# sp_ids.pop() +# # get the scanpath as fixated words +# sp_words = list() +# for sp_id in sp_ids: +# sp_words.append(sn_words[sp_id]) + +# # TODO doesn't make sense to have CLS sn SEP sp SEP for auto-regressive model. BOS token +# # join sentence and scanpath into strings and concatenate them +# sn = ' '.join(sn_words) +# sp = ' '.join(sp_words) +# sn_sp = sn + ' ' + sp + +# encoded = tokenizer.encode_plus( +# sn_sp, +# add_special_tokens=False, +# return_tensors='pt', +# return_attention_mask=True, +# ) +# word_ids = torch.Tensor(encoded.word_ids()) + +# last_hidden = gpt2_model(encoded.input_ids).last_hidden_state + +# # aggregate the embeddings to word-level +# embeddings = aggregate_input_embeddings( +# embeddings=last_hidden, +# word_ids=word_ids, +# aggregate=aggregate, +# ) + +# return embeddings, word_ids, sn_sp + + +def get_embeddings_seq2seq( + data_instance: Dict[str, Any], + tokenizer: transformers.GPT2TokenizerFast, + gpt2_model: transformers.GPT2Model, + bert_embeddings: torch.nn.Embedding, + instance_idx: int, + aggregate: str = "mean", # 'mean', 'sum' + max_length: int = 128, + sp_pad_token: int = 127, +): + """ + Get the embeddings of the scanpath (fixated words) from the encoder. + :param data_instance: the data instance from the dataset. + :param tokenizer: the tokenizer. + :param gpt2_model: the GPT2 model. + :param bert_embeddings: the BERT embeddings. + :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings. + :return: the embeddings of the scanpath, the padded fixation durations, and the attention mask. + """ + + sn_repr_len = data_instance["sn_repr_len"] + + # the sentence + sn_words = data_instance["words_for_mapping"].split() + while sn_words[-1] == "[PAD]": + sn_words.pop() + # remove the CLS and SEP tokens + sn_words = sn_words[1:-1] + # chinese characters are one string in a list + if sp_pad_token == 67: # chinese pad token + sn_words = list(sn_words[0]) + + # the scanpath + # remove the CLS token + sp_ids = data_instance["sn_sp_repr"][sn_repr_len:][1:] + # cut off the trailing pad tokens + while sp_ids[-1] == sp_pad_token: + sp_ids.pop() + # remove the SEP token + sp_ids = sp_ids[:-1] + + # make the scanpath ids start from 0 for re-ordering of the embeddings + sp_ids = [sp_id - 1 for sp_id in sp_ids] + + # get the scanpath as fixated words + sp_words = list() + try: + for sp_id in sp_ids: + sp_words.append(sn_words[sp_id]) + except: + print(f"Error at index {instance_idx}") + # breakpoint() + return None, None, None + + # get the fixation durations + fix_durs = data_instance["sn_sp_fix_dur"][sn_repr_len + 1 :] + while fix_durs[-1] == 0: + fix_durs.pop() + # convert to tensor + fix_durs = torch.Tensor(fix_durs) + + # get the sentence encoding + sn_enc = tokenizer.encode_plus( + sn_words, + add_special_tokens=False, + return_tensors="pt", + is_split_into_words=True, + ) + sn_word_ids = torch.Tensor(sn_enc.word_ids()) + + # get the embeddings + with torch.no_grad(): + last_hidden = gpt2_model(sn_enc.input_ids).last_hidden_state + + # aggregate the embeddings to word-level + sn_embeddings = aggregate_input_embeddings( + embeddings=last_hidden, + word_ids=sn_word_ids, + aggregate=aggregate, + ) + + # convert sp_ids to tensor + sp_ids = torch.Tensor(sp_ids).long() + + # re-order the embeddings as scanpath + sp_embeddings = sn_embeddings[:, sp_ids, :] + + # pad the embeddings and fixation durations to max input length + # and get the attention mask + sp_embeddings_padded, fix_durs_padded, attention_mask = padding_and_mask_seq2seq( + sp_embeddings=sp_embeddings, + fix_durs=fix_durs, + bert_embeddings=bert_embeddings, + max_length=max_length, + ) + + return sp_embeddings_padded.squeeze(0), fix_durs_padded, attention_mask.squeeze(0) + + +def padding_and_mask_seq2seq( + sp_embeddings: torch.Tensor, + bert_embeddings: torch.nn.Embedding, + max_length: int, + fix_durs: Optional[torch.Tensor] = None, + inference: Optional[bool] = None, +): + """ + Add the BERT CLS token to the beginning of the scanpath embedding (needed for pooler output). + Pad the scanpath embeddings and fixation durations to max input lenght. + Use the PAD token embedding for padding. + """ + # get the embedding for the pad token + pad_emb = bert_embeddings(torch.Tensor([0]).long()) + cls_emb = bert_embeddings(torch.Tensor([101]).long()) + + # prepend the cls emb to the sp_embeddings + sp_embeddings = torch.cat((cls_emb.unsqueeze(0), sp_embeddings), dim=1) + + # pad the embeddings + current_length = sp_embeddings.size(1) + padding_needed = max_length - current_length + pad_tensor = pad_emb.unsqueeze(0).expand(1, padding_needed, -1) + sp_embeddings_padded = torch.cat((sp_embeddings, pad_tensor), dim=1) + + # create attention mask + sp_mask = torch.ones((1, current_length), dtype=torch.long) + pad_mask = torch.zeros((1, padding_needed), dtype=torch.long) + attention_mask = torch.cat((sp_mask, pad_mask), dim=1) + + if inference: + return sp_embeddings_padded, attention_mask + + # prepend 0 to the fixation durations because the first word is the CLS token + fix_durs = torch.cat((torch.Tensor([0]), fix_durs), dim=0) + + # pad the fixation durations + fix_dur_pad = torch.zeros(padding_needed) + fix_durs_padded = torch.cat((fix_durs, fix_dur_pad), dim=0) + + return sp_embeddings_padded, fix_durs_padded, attention_mask + + +def aggregate_input_embeddings( + embeddings: torch.Tensor, + word_ids: torch.Tensor, + aggregate: str = "mean", # 'mean', 'sum' +): + """ + Aggregate the embeddings that are input to the fixation module to word-level. + :param embeddings: the last hidden state (contextualised embeddings) of the sentence-scanpath concatenation + when passed through the GPT2 model. + :param word_ids: the word ids of the sentence-scanpath concatenation. + :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings. + :return: the aggregated word embeddings. + """ + # get the unique indices and inverse + unique_indices, inverse_indices = torch.unique(word_ids, return_inverse=True) + + # sum the tensor along the dimension 1 (sequence length) for the same word ids + summed_tensor = torch.zeros((1, unique_indices.size(0), embeddings.size(2))) + summed_tensor = summed_tensor.scatter_add( + 1, inverse_indices.unsqueeze(0).unsqueeze(-1).expand_as(embeddings), embeddings + ) + + if aggregate == "sum": + return summed_tensor + + elif aggregate == "mean": + + # count the occurrences of each word id (how many sub-words per word) + counts = torch.zeros(unique_indices.size(0)).scatter_add( + 0, inverse_indices, torch.ones_like(inverse_indices, dtype=torch.float) + ) + + # average the summed tensor + averaged_tensor = summed_tensor / counts.view(1, -1, 1) + return averaged_tensor + + +class Seq2SeqDataset(Dataset): + def __init__( + self, + data: Dict[str, torch.Tensor], + normalize: Optional[bool] = None, + inference: Optional[bool] = None, + ): + super().__init__() + self.data = data + self.normalize = normalize + self.inference = inference + + def __len__(self): + return len(self.data["sp_embeddings"]) + + def __getitem__(self, idx): + if self.inference: + sample = { + "sp_embeddings": self.data["sp_embeddings"][idx], + "attention_masks": self.data["attention_masks"][idx], + } + return sample + else: + sample = { + "sp_embeddings": self.data["sp_embeddings"][idx], + "attention_masks": self.data["attention_masks"][idx], + "fix_durs": self.data["fix_durs"][idx], + } + if self.normalize: + sample["fix_durs_normalized"] = self.data["fix_durs_normalized"][idx] + return sample + + +def prepare_seq2seq_data( + data: datasets.DatasetDict, + tokenizer: transformers.GPT2TokenizerFast, + gpt2_model: transformers.GPT2Model, + bert_embeddings: torch.nn.Embedding, + aggregate: str = "mean", + max_length: int = 128, + sp_pad_token: int = 127, +): + """ + Prepare the data for training the fixation duration module. + :param data: the dataset. + :param tokenizer: the tokenizer. + :param gpt2_model: the GPT2 model. + :param bert_embeddings: the BERT embeddings. + :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings. + :param max_length: the maximum input length. + :return: the data for training the fixation duration module. + """ + data_dict = { + "sp_embeddings": [], + "attention_masks": [], + "fix_durs": [], + } + + for idx, instance in tqdm(enumerate(data["train"])): + + sp_embeddings, fix_durs, attention_mask = get_embeddings_seq2seq( + data_instance=instance, + tokenizer=tokenizer, + gpt2_model=gpt2_model, + bert_embeddings=bert_embeddings, + instance_idx=idx, + aggregate=aggregate, + max_length=max_length, + sp_pad_token=sp_pad_token, + ) + if sp_embeddings is None: + continue + + data_dict["sp_embeddings"].append(sp_embeddings) + data_dict["attention_masks"].append(attention_mask) + data_dict["fix_durs"].append(fix_durs) + + return data_dict + + +def split_train_val_data( + data: Dict[str, List[torch.Tensor]], + val_size: float = 0.1, +): + """ + Split the train data into train and validation data. + :param data: the data. + :param val_size: the size of the validation data. + :return: the train and validation data. + """ + num_samples = len(next(iter(data.values()))) + # shuffle the indices + indices = list(range(num_samples)) + random.shuffle(indices) + + # compute the split point + split_point = int(num_samples * val_size) + train_indices = indices[split_point:] + val_indices = indices[:split_point] + + train_data = {key: [value[i] for i in train_indices] for key, value in data.items()} + val_data = {key: [value[i] for i in val_indices] for key, value in data.items()} + + return train_data, val_data + + +def get_embeddings_seq2seq_hp( + sn_repr_len: int, + sn_words: List[str], + sp_ids: List[int], + tokenizer: transformers.GPT2TokenizerFast, + gpt2_model: transformers.GPT2Model, + bert_embeddings: torch.nn.Embedding, + aggregate: str = "mean", + max_length: int = 128, + sp_pad_token: int = 127, +): + """ + Get the embeddings of the scanpath (fixated words) from the encoder. + :param sn_repr_len: the length of the sentence representation. + :param sn_words: the words of the sentence. + :param sp_ids: the scanpath ids. + :param tokenizer: the tokenizer. + :param gpt2_model: the GPT2 model. + :param bert_embeddings: the BERT embeddings. + :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings. + :return: the embeddings of the scanpath, the padded fixation durations, and the attention mask. + """ + + pad_idx = [i for i, word in enumerate(sn_words) if word == "[PAD]"] + sep_idx = [sn_words.index("[SEP]")] + all_remove_idx = [0] # for CLS + all_remove_idx += sep_idx + all_remove_idx += pad_idx + + # get rid of trailing pad tokens in sentence + while sn_words[-1] == "[PAD]": + sn_words.pop() + # get rid of the CLS and SEP tokens + sn_words = sn_words[1:-1] + + # the scanpath + # get rid of predicted CLS, SEP and wrongly predicted PAD tokens (will throw error) + sp_ids = [sp_id for sp_id in sp_ids if sp_id not in all_remove_idx] + + # make the scanpath ids start from 0 for re-ordering of the embeddings + sp_ids = [sp_id - 1 for sp_id in sp_ids] + + # get the scanpath as fixated words + sp_words = list() + for sp_id in sp_ids: + sp_words.append(sn_words[sp_id]) + + # get the sentence encoding + sn_enc = tokenizer.encode_plus( + sn_words, + add_special_tokens=False, + return_tensors="pt", + is_split_into_words=True, + ) + sn_word_ids = torch.Tensor(sn_enc.word_ids()) + + # get the embeddings + with torch.no_grad(): + last_hidden = gpt2_model(sn_enc.input_ids).last_hidden_state + + # aggregate the embeddings to word-level + sn_embeddings = aggregate_input_embeddings( + embeddings=last_hidden, + word_ids=sn_word_ids, + aggregate=aggregate, + ) + + # convert sp_ids to tensor + sp_ids = torch.Tensor(sp_ids).long() + + # re-order the embeddings as scanpath + sp_embeddings = sn_embeddings[:, sp_ids, :] + + # pad the embeddings to max input length and get the attention mask + sp_embeddings_padded, attention_mask = padding_and_mask_seq2seq( + sp_embeddings=sp_embeddings, + bert_embeddings=bert_embeddings, + max_length=max_length, + inference=True, + ) + + return sp_embeddings_padded.squeeze(0), attention_mask.squeeze(0) + + +def prepare_seq2seq_data_hp( + scandl_output: Dict[str, Any], + tokenizer: transformers.GPT2TokenizerFast, + gpt2_model: transformers.GPT2Model, + bert_embeddings: torch.nn.Embedding, + aggregate: str = "mean", + max_length: int = 128, + sp_pad_token: int = 127, +): + """ + Prepare the scandl output for inference of the hyper-parameter search of the Seq2Seq fixation duration model. + :param scandl_output: the ScanDL output. + :return: the data for inference. + """ + data_dict = { + "sp_embeddings": [], + "attention_masks": [], + "original_fix_durs": [], + "predicted_sp_ids": [], + "reader_ids": [], + "sn_ids": [], + } + + for idx in tqdm(range(len(scandl_output["predicted_sp_ids"]))): + + sn_repr_len = scandl_output["sn_repr_len"][idx] + if sp_pad_token == 67: + # for Chinese: make sure the words are split correctly (chinese characters have no whitespace) + sn_words = scandl_output["words_for_mapping"][idx].split() + sn_words = [sn_words[0]] + list(sn_words[1]) + sn_words[2:] + else: + sn_words = scandl_output["words_for_mapping"][idx].split() + sp_ids = scandl_output["predicted_sp_ids"][idx] + + try: + sp_embeddings, attention_mask = get_embeddings_seq2seq_hp( + sn_repr_len=sn_repr_len, + sn_words=sn_words, + sp_ids=sp_ids, + tokenizer=tokenizer, + gpt2_model=gpt2_model, + bert_embeddings=bert_embeddings, + aggregate=aggregate, + max_length=max_length, + sp_pad_token=sp_pad_token, + ) + + # get the original fixation durations + fix_durs = scandl_output["sn_sp_fix_dur"][idx][sn_repr_len:] + while fix_durs[-1] == 0: + fix_durs.pop() + fix_durs.append(0) + + data_dict["sp_embeddings"].append(sp_embeddings) + data_dict["attention_masks"].append(attention_mask) + data_dict["original_fix_durs"].append(str(fix_durs)) + data_dict["predicted_sp_ids"].append(str(sp_ids)) + data_dict["reader_ids"].append(scandl_output["reader_ids"][idx]) + data_dict["sn_ids"].append(scandl_output["sn_ids"][idx]) + except: + print(f"Error at index {idx}") + continue + + return data_dict + + +class Seq2SeqDatasetHP(Dataset): + def __init__( + self, + data: Dict[str, Union[torch.Tensor, Any]], + ): + super().__init__() + self.data = data + + def __len__(self): + return len(self.data["sp_embeddings"]) + + def __getitem__(self, idx): + sample = { + "sp_embeddings": self.data["sp_embeddings"][idx], + "attention_masks": self.data["attention_masks"][idx], + #'predicted_sp_words': self.data['predicted_sp_words'][idx], + #'original_sp_words': self.data['original_sp_words'][idx], + "predicted_sp_ids": self.data["predicted_sp_ids"][idx], + # 'original_sp_ids': self.data['original_sp_ids'][idx], + # 'original_sn': self.data['original_sn'][idx], + "sn_ids": self.data["sn_ids"][idx], + "reader_ids": self.data["reader_ids"][idx], + # 'sn_repr_len': self.data['sn_repr_len'][idx], + # 'words_for_mapping': self.data['words_for_mapping'][idx], + # 'sn_sp_repr': self.data['sn_sp_repr'][idx], + # 'sn_sp_fix_dur': self.data['sn_sp_fix_dur'][idx], + "original_fix_durs": self.data["original_fix_durs"][idx], + } + return sample diff --git a/fix_dur_module/utils_train.py b/fix_dur_module/utils_train.py new file mode 100644 index 0000000000000000000000000000000000000000..6ebeee23034db7bdbd7db0b2d25c644c589ff356 --- /dev/null +++ b/fix_dur_module/utils_train.py @@ -0,0 +1,195 @@ +import torch +import torch.nn as nn +import numpy as np +import transformers + +from typing import Optional + + +class EarlyStopping: + def __init__( + self, + patience: int, + path: str, + delta: Optional[int] = 0, + ): + self.patience = patience + self.delta = delta + self.best_score = None + self.early_stop = False + self.counter = 0 + self.best_loss = np.inf + self.path = path + + def __call__( + self, + val_loss, + model, + ): + score = -val_loss + + if self.best_score is None: + self.best_score = score + self.save_checkpoint(val_loss, model) + elif score < self.best_score + self.delta: + self.counter += 1 + print(f"EarlyStopping counter: {self.counter} out of {self.patience}") + if self.counter >= self.patience: + self.early_stop = True + else: + self.best_score = score + self.save_checkpoint(val_loss, model) + self.counter = 0 + + def save_checkpoint( + self, + val_loss, + model, + ): + """Saves model when validation loss decreases.""" + print( + f"Validation loss decreased ({self.best_loss:.6f} --> {val_loss:.6f}). Saving model..." + ) + torch.save(model.state_dict(), self.path) + self.best_loss = val_loss + + +def train( + model, + num_epochs: int, + train_loader: torch.utils.data.DataLoader, + val_loader: torch.utils.data.DataLoader, + criterion: nn.MSELoss, + optimizer: transformers.AdamW, + early_stopping: EarlyStopping, + scheduler: transformers.get_linear_schedule_with_warmup, + device: torch.device, + fix_dur_colname: str, + output_attentions: Optional[bool] = None, + use_attention_mask: Optional[bool] = None, +): + """ + Train loop to train the Seq2Seq model. + :param model: the model to train + :param num_epochs: number of epochs to train + :param train_loader: the training data loader + :param val_loader: the validation data loader + :param criterion: the loss function (MSE Loss) + :param optimizer: the optimizer (AdamW) + :param early_stopping: the early stopping object + :param scheduler: the learning rate scheduler + :param device: the device to train on + :param fix_dur_colname: the name of the column containing the fixations durations + """ + for epoch in range(num_epochs): + + model.train() + + for batch_idx, train_batch in enumerate(train_loader): + + optimizer.zero_grad() + + sp_embeddings = train_batch["sp_embeddings"].to(device) + attention_mask = train_batch["attention_masks"].to(device) + fix_durs = train_batch[fix_dur_colname].to(device) + + # forward pass + if use_attention_mask: + + if output_attentions: + + out, _ = model( + sp_embeddings=sp_embeddings, + attention_mask=attention_mask, + output_attentions=output_attentions, + ) + else: + + out = model( + sp_embeddings=sp_embeddings, + attention_mask=attention_mask, + output_attentions=output_attentions, + ) + else: + + if output_attentions: + + out, _ = model( + sp_embeddings=sp_embeddings, + output_attentions=output_attentions, + ) + else: + out = model( + sp_embeddings=sp_embeddings, + output_attentions=output_attentions, + ) + + # train_loss = criterion(out, fix_durs) + # mask the padding in the loss computation + loss_mask = (fix_durs != 0).float() + # train_loss = criterion(out * loss_mask, fix_durs * loss_mask) + train_loss = criterion(out, fix_durs) + + train_loss.backward() + optimizer.step() + scheduler.step() + + print(f"\t epoch {epoch+1}, batch {batch_idx+1}, loss: {train_loss.item():.4f}") + + # validation + + model.eval() + val_loss = 0.0 + + with torch.no_grad(): + + for val_batch in val_loader: + + sp_embeddings = val_batch["sp_embeddings"].to(device) + attention_mask = val_batch["attention_masks"].to(device) + fix_durs = val_batch["fix_durs"].to(device) + + if use_attention_mask: + + if output_attentions: + out, attentions = model( + sp_embeddings=sp_embeddings, + attention_mask=attention_mask, + output_attentions=output_attentions, + ) + else: + out = model( + sp_embeddings=sp_embeddings, + attention_mask=attention_mask, + output_attentions=output_attentions, + ) + else: + + # forward pass + if output_attentions: + out, attentions = model( + sp_embeddings=sp_embeddings, + output_attentions=output_attentions, + ) + + else: + out = model( + sp_embeddings=sp_embeddings, + output_attentions=output_attentions, + ) + + val_loss_mask = (fix_durs != 0).float() + # val_loss += criterion(out * val_loss_mask, fix_durs * val_loss_mask).item() + val_loss += criterion(out, fix_durs).item() + + # average the losses + val_loss /= len(val_loader) + train_loss /= len(train_loader) + + print(f"Epoch {epoch+1}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}") + + # check for early stopping + early_stopping(val_loss, model) + if early_stopping.early_stop: + print("Early stopping") + break diff --git a/handler.py b/handler.py new file mode 100644 index 0000000000000000000000000000000000000000..be6a5a108e1e254601b6bb6a6795b4f026fb4dd7 --- /dev/null +++ b/handler.py @@ -0,0 +1,53 @@ +import torch + +from ScanDL2 import ScanDL2 + + +class EndpointHandler: + def __init__(self, path: str = ""): + + self.models = { + "sentence": ScanDL2( + text_type="sentence", + bsz=2, + save=None, + filename=None, + ), + "paragraph": ScanDL2( + text_type="paragraph", + bsz=2, + save=None, + filename=None, + ), + } + + for m in self.models.values(): + # m.to(self.device) + m.eval() + + def __call__(self, data): + + inputs = data.get("inputs", data) + + parameters = data.get("parameters", {}) + + text_type = parameters.get("text_type", "sentence") + model = self.models[text_type] + bsz = parameters.get("bsz", 2) + + if model.scandl_module.args.batch_size != bsz: + model.scandl_module.args.batch_size = bsz + model.fixdur_module.bsz = bsz + model.fixdur_module.args["bsz"] = bsz + + if isinstance(inputs, str): + texts = [inputs] + elif isinstance(inputs, list): + texts = inputs + else: + raise ValueError("'inputs' must be a string or list of strings.") + + with torch.no_grad(): + output = model(texts=texts) + + return output diff --git a/model.py b/model.py new file mode 100644 index 0000000000000000000000000000000000000000..07c1e3501711202b06b6c15a74d7ba85343dcc85 --- /dev/null +++ b/model.py @@ -0,0 +1,701 @@ +import os +import sys +import time +import json +import joblib +import argparse + +from functools import partial + +from typing import Union, List, Dict, Optional, Any + +import torch +import torch.nn as nn +import torch.distributed as dist +from torch.utils.data import DataLoader + +import numpy as np +import pandas as pd + +from tqdm import tqdm + +from transformers import ( + set_seed, + BertTokenizerFast, + GPT2TokenizerFast, + GPT2LMHeadModel, + GPT2Model, + AutoConfig, + BertModel, +) +from transformers.models.bert.modeling_bert import BertEncoder +from datasets import DatasetDict +from datasets import Dataset as Dataset2 + +from ScanDL2.scandl_module.original_scandl.sp_rounding import denoised_fn_round +from ScanDL2.scandl_module.original_scandl.utils import dist_util, logger +from ScanDL2.scandl_module.original_scandl.utils.nn import * + +from ScanDL2.scandl2_utils import text_dataset_loader, FixdurDataset + +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import _collate_batch_helper +from ScanDL2.scandl_module.scripts.sp_basic_utils import ( + load_defaults_config, + create_model_and_diffusion, + add_dict_to_argparser, + args_to_dict, +) + +from ScanDL2.fix_dur_module.model_seq2seq import Seq2SeqModel +from ScanDL2.fix_dur_module.utils_data import aggregate_input_embeddings, padding_and_mask_seq2seq + +from ScanDL2.PATHS import ( + SENT_SCANDL_MODULE, + SENT_FIXDUR_MODULE, + PAR_SCANDL_MODULE, + PAR_FIXDUR_MODULE, +) + + +class ScanDL2(nn.Module): + def __init__( + self, + text_type: str = "sentence", # sentence, paragraph + bsz: Optional[int] = 2, + save: Optional[str] = None, + filename: Optional[str] = None, + ): + super(ScanDL2, self).__init__() + + self.save = save + self.filename = filename + + # initialize the ScanDL module and the Fixdur Module + self.scandl_module = ScanDLModule( + text_type=text_type, + bsz=bsz, + ) + self.fixdur_module = FixdurModule( + text_type=text_type, + bsz=bsz, + ) + + def forward( + self, + texts: Union[str, List[str]], + ): + # check if the input is in the correct format + self._validate_inputs(texts=texts) + + # get the fixation location predictions from the ScanDL module + scandl_module_output = self.scandl_module(texts=texts) + + # get the fixation duration predictions from the Fixdur module + fixdur_module_output = self.fixdur_module(scandl_module_output=scandl_module_output) + + if self.save is not None: + filename = self.filename if self.filename is not None else f"scandl2_outputs.json" + filename = f"{filename}.json" if not filename.endswith(".json") else filename + self._save_results(results=fixdur_module_output, filename=filename) + + return fixdur_module_output + + def _save_results(self, results, filename): + if not os.path.exists(self.save): + os.makedirs(self.save) + with open(os.path.join(self.save, filename), "w") as f: + json.dump(results, f) + print(f"--- ScanDL 2.0 outputs saved to {os.path.join(self.save, filename)}.") + + def _validate_inputs(self, texts: Union[str, List[str]]) -> None: + if not isinstance(texts, (str, list)) or ( + isinstance(texts, list) and not all(isinstance(t, str) for t in texts) + ): + raise TypeError("Invalid input: 'texts' must be of type 'str' or 'List[str]'.") + + +class ScanDLModule(nn.Module): + + def __init__( + self, + text_type: str, # sentence, paragraph + bsz: int, + ): + super(ScanDLModule, self).__init__() + + base_path = os.path.dirname(__file__) + if text_type == "paragraph": + self.path_to_config = os.path.join(base_path, "config_emtec.json") + self.path_to_scandl_module = PAR_SCANDL_MODULE + elif text_type == "sentence": + self.path_to_config = os.path.join(base_path, "config.json") + self.path_to_scandl_module = SENT_SCANDL_MODULE + else: + raise NotImplementedError(f"Text type {text_type} not implemented.") + + # get the args + self.args = self._get_args() + self.args.batch_size = bsz + + # seting up the environment + dist_util.setup_dist() + logger.configure() + self.world_size = dist.get_world_size() or 1 + self.rank = dist.get_rank() or 0 + # set_seed(self.args.seed2) + + # load the tokenizer + self.tokenizer = self._load_tokenizer() + + # load the ScanDL module and the Diffusion + self.scandl_module, self.diffusion = self._load_scandl_module( + path_to_scandl_module=self.path_to_scandl_module + ) + self.sn_sp_repr_embedding = self._get_sn_sp_repr_emb() + + def forward( + self, + texts: Union[str, List[str]], + ) -> Dict[str, Union[List[List[str]], List[List[int]], List[str]]]: + + data_loader = self._preprocess_text(texts=texts) + + predicted_sp_words, predicted_sp_ids = [], [] + original_sn = [] + + print("\t\t### ScanDL Module generates fixation locations ...") + + unique_idx = list() + idx_ctr = 0 + + for batch_idx, batch in tqdm(enumerate(data_loader)): + + mask = batch["mask"].to(dist_util.dev()) + sn_sp_repr = batch["sn_sp_repr"].to(dist_util.dev()) + sn_input_ids = batch["sn_input_ids"].to(dist_util.dev()) + indices_pos_enc = batch["indices_pos_enc"].to(dist_util.dev()) + sn_repr_len = batch["sn_repr_len"].to(dist_util.dev()) + words_for_mapping = batch["words_for_mapping"] + + sn_sp_emb, pos_enc, sn_input_ids_emb = self.scandl_module.get_embeds( + sn_sp_repr=sn_sp_repr, + sn_input_ids=sn_input_ids, + indices_pos_enc=indices_pos_enc, + ) + + x_start = sn_sp_emb + noise = torch.randn_like(x_start) + mask = torch.broadcast_to(mask.unsqueeze(dim=-1), x_start.shape).to(dist_util.dev()) + x_noised = torch.where(mask == 0, x_start, noise) + + self.args.use_ddim = False + step_gap = 1 + + sample_fn = ( + self.diffusion.p_sample_loop + if not self.args.use_ddim + else self.diffusion.ddim_sample_loop + ) + + sample_shape = (x_start.shape[0], self.args.seq_len, self.args.hidden_dim) + subwords = [self.tokenizer.convert_ids_to_tokens(i) for i in sn_input_ids] + + samples = sample_fn( + model=self.scandl_module, + shape=sample_shape, + noise=x_noised, + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + mask_sn_padding=None, + mask_transformer_att=None, + clip_denoised=self.args.clip_denoised, + denoised_fn=partial(denoised_fn_round, self.args, self.sn_sp_repr_embedding), + model_kwargs=None, + top_p=self.args.top_p, + clamp_step=self.args.clamp_step, + clamp_first=self.args.clamp_first_bool, + mask=mask, + x_start=x_start, + gap=step_gap, + ) + sample = samples[-1] + + logits = self.scandl_module.get_logits(sample) + cands = torch.topk(logits, k=1, dim=-1) + + for instance_idx, (pred_seq, orig_words, sn_len) in enumerate( + zip(cands.indices, words_for_mapping, sn_repr_len) + ): + pred_seq_sp = pred_seq[sn_len:] + words_split = orig_words.split() + predicted_sp = [words_split[i] for i in pred_seq_sp] + pred_sp_ids = [e.item() for e in pred_seq_sp] + + # cut off trailing pad tokens + while len(predicted_sp) > 1 and predicted_sp[-1] == "[PAD]": + predicted_sp.pop() + while len(pred_sp_ids) > 1 and pred_sp_ids[-1] == self.args.seq_len - 1: + pred_sp_ids.pop() + while len(words_split) > 1 and words_split[-1] == "[PAD]": + words_split.pop() + + # remove CLS and SEP tokens from predictions + if predicted_sp[0] == "[CLS]": + predicted_sp = predicted_sp[1:] + pred_sp_ids = pred_sp_ids[1:] + if predicted_sp[-1] == "[SEP]": + predicted_sp = predicted_sp[:-1] + pred_sp_ids = pred_sp_ids[:-1] + words_split = words_split[1:-1] + + # filter out erroneously predicted PAD tokens (they will raise an error in the fixdur module) + pred_sp_ids, predicted_sp = self._remove_special_tokens( + predicted_sp_ids=pred_sp_ids, + predicted_sp_words=predicted_sp, + token="[PAD]", + ) + pred_sp_ids, predicted_sp = self._remove_special_tokens( + predicted_sp_ids=pred_sp_ids, + predicted_sp_words=predicted_sp, + token="[CLS]", + ) + pred_sp_ids, predicted_sp = self._remove_special_tokens( + predicted_sp_ids=pred_sp_ids, + predicted_sp_words=predicted_sp, + token="[SEP]", + ) + + predicted_sp_words.append(predicted_sp) + predicted_sp_ids.append(pred_sp_ids) + original_sn.append(words_split) + + idx_ctr += 1 + unique_idx.append(idx_ctr) + + predictions = { + "predicted_sp_words": predicted_sp_words, + "predicted_sp_ids": predicted_sp_ids, + "original_sn": original_sn, + "unique_idx": unique_idx, + } + return predictions + + def _remove_special_tokens( + self, + predicted_sp_ids: List[int], + predicted_sp_words: List[str], + token: str, # '[CLS]' or '[SEP]' or '[PAD]' + ): + filtered_sp_ids, filtered_sp_words = [], [] + for sp_word, sp_id in zip(predicted_sp_words, predicted_sp_ids): + if sp_word != token: + filtered_sp_ids.append(sp_id) + filtered_sp_words.append(sp_word) + return filtered_sp_ids, filtered_sp_words + + def _preprocess_text( + self, + texts: Union[str, List[str]], + ): + data = { + "mask": [], + "sn_sp_repr": [], + "sn_input_ids": [], + "indices_pos_enc": [], + "words_for_mapping": [], + "sn_repr_len": [], + } + + if isinstance(texts, str): + texts = [texts] + + for sn_idx, sn in enumerate(texts): + + if sn.startswith("[CLS]") and sn.endswith("[SEP]"): + sn = sn + elif sn.startswith("[CLS]"): + sn = sn + " [SEP]" + elif sn.endswith("[SEP]"): + sn = "[CLS] " + sn + else: + sn = "[CLS] " + sn + " [SEP]" + + encoded_sn = self.tokenizer.encode_plus( + sn.split(), + add_special_tokens=False, + padding=False, + return_attention_mask=False, + is_split_into_words=True, + truncation=False, + ) + + if len(encoded_sn) > self.args.seq_len / 2: + print(f"Sentence {sn} is too long. Continue.") + + sn_word_ids = encoded_sn.word_ids() + sn_input_ids = encoded_sn["input_ids"] + + sn_sp_repr = sn_word_ids + + mask = [0] * len(sn_word_ids) + indices_pos_enc = list(range(0, len(sn_word_ids))) + list( + range(0, self.args.seq_len - len(sn_word_ids)) + ) + words_for_mapping = sn.split() + (self.args.seq_len - len(sn.split())) * ["[PAD]"] + + data["mask"].append(mask) + data["sn_sp_repr"].append(sn_sp_repr) + data["sn_input_ids"].append(sn_input_ids) + data["indices_pos_enc"].append(indices_pos_enc) + data["words_for_mapping"].append(" ".join(words_for_mapping)) + data["sn_repr_len"].append(len(sn_word_ids)) + + # padding + data["mask"] = _collate_batch_helper( + examples=data["mask"], + pad_token_id=1, + max_length=self.args.seq_len, + ) + data["sn_sp_repr"] = _collate_batch_helper( + examples=data["sn_sp_repr"], + pad_token_id=self.args.seq_len - 1, + max_length=self.args.seq_len, + ) + data["sn_input_ids"] = _collate_batch_helper( + examples=data["sn_input_ids"], + pad_token_id=self.tokenizer.pad_token_id, + max_length=self.args.seq_len, + ) + + split = "inference" + dataset = Dataset2.from_dict(data) + dataset_dict = DatasetDict() + dataset_dict[split] = dataset + data_loader = text_dataset_loader( + data=dataset_dict, + data_args=self.args, + split=split, + deterministic=True, + ) + return data_loader + + def _load_scandl_module( + self, + path_to_scandl_module: str, + ): + logger.log("### Loading ScanDL Diffusion Module ...") + scandl_module, diffusion = create_model_and_diffusion( + **args_to_dict(self.args, load_defaults_config(config_path=self.path_to_config).keys()) + ) + # TODO Name scandl module, not model + scandl_module.load_state_dict( + dist_util.load_state_dict( + os.path.join(self.path_to_scandl_module, "ema_0.9999_080000.pt"), map_location="cpu" + ) + ) + pytorch_total_params = sum(p.numel() for p in scandl_module.parameters()) + logger.log(f"### Total number of parameters: {pytorch_total_params}") + scandl_module.eval().requires_grad_(False).to(dist_util.dev()) + return scandl_module, diffusion + + def _get_sn_sp_repr_emb(self): + sn_sp_repr_embedding = nn.Embedding( + num_embeddings=self.args.hidden_t_dim, + embedding_dim=self.args.hidden_dim, + _weight=self.scandl_module.sn_sp_repr_embedding.weight.clone().cpu(), + ) + return sn_sp_repr_embedding + + def _get_args(self): + args = self._get_parser().parse_args() + # load the training arguments + with open(os.path.join(self.path_to_scandl_module, "training_args.json")) as f: + training_args = json.load(f) + training_args["batch_size"] = args.batch_size + args.__dict__.update(training_args) + if args.clamp_first == "yes": + args.clamp_first_bool = True + else: + args.clamp_first_bool = False + # TODO self.args.clamp_first_bool as argument + # set mask_padding to False + args.mask_padding = False + return args + + def _load_tokenizer(self): + tokenizer = BertTokenizerFast.from_pretrained(self.args.config_name) + self.args.vocab_size = tokenizer.vocab_size + return tokenizer + + def _get_parser(self) -> argparse.ArgumentParser: + defaults = dict( + model_path="", + step=0, + out_dir="", + top_p=0, + clamp_first="yes", + test_set_sns="mixed", + atten_vis=False, + notes="-", + tsne_vis=False, + sp_vis=False, + no_inst=0, + atten_vis_sp=False, + load_ids="-", + load_test_data="-", + setting="-", + fold=0, + ) + decode_defaults = dict( + split="valid", + clamp_step=0, + seed2=105, + clip_denoised=False, + ) + + defaults.update(load_defaults_config(config_path=self.path_to_config)) + defaults.update(decode_defaults) + parser = argparse.ArgumentParser() + add_dict_to_argparser(parser, defaults) + + return parser + + +class FixdurModule(nn.Module): + + def __init__( + self, + text_type: str, # sentence, paragraph + bsz: Optional[int] = 2, + ): + super(FixdurModule, self).__init__() + + base_path = os.path.dirname(__file__) + if text_type == "paragraph": + self.path_to_config = os.path.join(base_path, "config_emtec.json") + self.path_to_fixdur_module = PAR_FIXDUR_MODULE + elif text_type == "sentence": + self.path_to_config = os.path.join(base_path, "config.json") + self.path_to_fixdur_module = SENT_FIXDUR_MODULE + else: + raise NotImplementedError(f"Text type {text_type} not implemented.") + + self.bsz = bsz + self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + self.config = load_defaults_config(config_path=self.path_to_config) + self.args = self._get_args(config=self.config) + self.hyperparameters = self._get_hyperparams( + path_to_fixdur_module=self.path_to_fixdur_module + ) + + # load GPT-2 model and tokenizer, and BERT embeddings + self.gpt2_model, self.tokenizer, self.bert_embeddings = self._load_gpt_and_bert( + config=self.config + ) + + # load the Fixdur module and the MinMax Scaler + self.fixdur_module = self._load_fixdur_module() + self.scaler = self._load_scaler() + + def _get_args(self, config: Dict[str, Any]) -> Dict[str, Any]: + args = { + "max_length": config["seq_len"], + "normalize": True, + "output_attentions": False, + "bsz": self.bsz, + "corpus": config["corpus"], + "sp_pad_token": config["seq_len"] - 1, + } + return args + + def _get_hyperparams(self, path_to_fixdur_module: str) -> Dict[str, Any]: + with open(os.path.join(path_to_fixdur_module, "hyperparameters.json")) as f: + return json.load(f) + + def _load_gpt_and_bert(self, config: Dict[str, Any]): + """ + Load GPT-2 and GPT-2 tokenizer to get the contextualized embeddings. + Load BERT model (for embeddings of CLS and PAD tokens) + """ + # GPT-2 + gpt_config_name = config["gpt_config_name"] + tokenizer = GPT2TokenizerFast.from_pretrained(gpt_config_name, add_prefix_space=True) + gpt2_model = GPT2Model.from_pretrained(gpt_config_name) + tokenizer.pad_token = tokenizer.eos_token + # freeze parameters + for param in gpt2_model.parameters(): + param.requires_grad = False + + # BERT + bert_config_name = config["config_name"] + bert_embeddings = BertModel.from_pretrained(bert_config_name).embeddings.word_embeddings + # freeze parameters + for param in bert_embeddings.parameters(): + param.requires_grad = False + + return gpt2_model, tokenizer, bert_embeddings + + def _load_fixdur_module(self): + fixdur_module_config = AutoConfig.from_pretrained("bert-base-cased") + fixdur_module_config.num_attention_heads = self.hyperparameters["num_heads"] + fixdur_module_config.num_hidden_layers = self.hyperparameters["num_layers"] + fixdur_module = Seq2SeqModel( + config=fixdur_module_config, + output_dim=self.args["max_length"], + num_linear=self.hyperparameters["num_linear"], + dropout=self.hyperparameters["dropout"], + ) + fixdur_module.load_state_dict( + torch.load( + os.path.join(self.path_to_fixdur_module, "seq2seq_fixdur.pt"), + map_location=self.device, + ) + ) + fixdur_module.eval() + fixdur_module.to(self.device) + return fixdur_module + + def _load_scaler(self): + scaler = joblib.load(os.path.join(self.path_to_fixdur_module, "min_max_scaler.pkl")) + return scaler + + def _prepare_data( + self, + scandl_module_output: Dict[str, Union[List[List[str]], List[List[int]], List[str]]], + ): + data_dict = { + "sp_embeddings": [], + "attention_mask": [], + "unique_idx": [], + } + for idx in range(len(scandl_module_output["predicted_sp_words"])): + + sn_words = scandl_module_output["original_sn"][idx] + sp_ids = scandl_module_output["predicted_sp_ids"][idx] + unique_id = scandl_module_output["unique_idx"][idx] + + # make the scanpath ids start at 0 for re-ordering of the embeddings + sp_ids = [i - 1 for i in sp_ids] + + sp_words = scandl_module_output["predicted_sp_words"][idx] + + # get the sentence encoding + sn_enc = self.tokenizer( + sn_words, + add_special_tokens=False, + return_tensors="pt", + is_split_into_words=True, + ) + sn_word_ids = torch.Tensor(sn_enc.word_ids()) + + # get the embeddings + with torch.no_grad(): + last_hidden = self.gpt2_model(sn_enc.input_ids).last_hidden_state + + # aggregate the embeddings to word level + sn_embeddings = aggregate_input_embeddings( + embeddings=last_hidden, + word_ids=sn_word_ids, + aggregate="mean", + ) + + # convert sp_ids to tensor + sp_ids = torch.Tensor(sp_ids).long() + + # re-order the embeddings as scanpath + try: + sp_embeddings = sn_embeddings[:, sp_ids, :] + except: + breakpoint() + + # pad the embeddings to max input length and get the attentino mask + sp_embeddings_padded, attention_mask = padding_and_mask_seq2seq( + sp_embeddings=sp_embeddings, + bert_embeddings=self.bert_embeddings, + max_length=self.args["max_length"], + inference=True, + ) + + data_dict["sp_embeddings"].append(sp_embeddings_padded) + data_dict["attention_mask"].append(attention_mask) + data_dict["unique_idx"].append(unique_id) + + return data_dict + + def forward( + self, + scandl_module_output: Dict[str, Union[List[List[str]], List[List[int]], List[str]]], + ) -> Dict[str, Union[List[List[str]], List[List[int]], List[str], List[List[float]]]]: + + output_dict = { + "predicted_sp_words": [], + "predicted_sp_ids": [], + "original_sn": [], + "predicted_fix_durs": [], + "unique_idx": [], + } + + data_df = pd.DataFrame(scandl_module_output) + + data_dict = self._prepare_data( + scandl_module_output=scandl_module_output, + ) + dataset = FixdurDataset(data=data_dict) + data_loader = DataLoader( + dataset, + batch_size=self.bsz, + shuffle=False, + ) + + print("\t\t### FixDur Module generates fixation durations ...") + for batch_idx, batch in tqdm(enumerate(data_loader)): + + sp_embeddings = batch["sp_embeddings"].squeeze(1).to(self.device) + attention_mask = batch["attention_mask"].squeeze(1).to(self.device) + unique_indices = batch["unique_idx"] + + out = self.fixdur_module( + sp_embeddings=sp_embeddings, + attention_mask=attention_mask, + output_attentions=self.args["output_attentions"], + ) + + # scale the output back to the original range + out_transformed = self.scaler.inverse_transform(out.detach().cpu().numpy()) + out_transformed_rounded = np.round(out_transformed, 2) + + # iterate over the individual predictions + for out_idx, out_instance in enumerate(out_transformed_rounded): + + predicted_fix_durs = out_instance + unique_idx = unique_indices[out_idx].item() + + # find predicted_sp_words, predicted_sp_ids, and original_sn in data_df conditioned on unique_idx + predicted_sp_words = data_df.loc[ + data_df["unique_idx"] == unique_idx, "predicted_sp_words" + ].values[0] + predicted_sp_ids = data_df.loc[ + data_df["unique_idx"] == unique_idx, "predicted_sp_ids" + ].values[0] + original_sn = data_df.loc[ + data_df["unique_idx"] == unique_idx, "original_sn" + ].values[0] + + sp_len = len(predicted_sp_ids) + + # cut off the predicted_fix_durs to the length of the scanpath + # the predicted fixation durations still contain predictions for the CLS and SEP token as well + pred_fix_durs = predicted_fix_durs[: sp_len + 2].tolist()[1:-1] + pred_fix_durs = [round(d, 2) for d in pred_fix_durs] + + # add to output_dict + output_dict["predicted_sp_words"].append(predicted_sp_words) + output_dict["predicted_sp_ids"].append(predicted_sp_ids) + output_dict["original_sn"].append(original_sn) + output_dict["predicted_fix_durs"].append(pred_fix_durs) + output_dict["unique_idx"].append(unique_idx) + + print(f"fixdur original sn: {original_sn}") + + return output_dict diff --git a/models/paragraph/fixdur-module/hyperparameters.json b/models/paragraph/fixdur-module/hyperparameters.json new file mode 100644 index 0000000000000000000000000000000000000000..16b572a4a7a5660b256d95754ccaf4e76dc2605a --- /dev/null +++ b/models/paragraph/fixdur-module/hyperparameters.json @@ -0,0 +1 @@ +{"num_heads": 12, "num_layers": 12, "num_linear": 8, "bsz": 48, "dropout": 0.5, "use_attention_mask": true} \ No newline at end of file diff --git a/models/paragraph/fixdur-module/min_max_scaler.pkl b/models/paragraph/fixdur-module/min_max_scaler.pkl new file mode 100644 index 0000000000000000000000000000000000000000..754d7bbc3870dd752b0f199699d17bd9b84e370c --- /dev/null +++ b/models/paragraph/fixdur-module/min_max_scaler.pkl @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0b234e10b46971a6d87694bdaa150483394cb952ebf984ab1c93801caaec927c +size 667 diff --git a/models/paragraph/fixdur-module/seq2seq_fixdur.pt b/models/paragraph/fixdur-module/seq2seq_fixdur.pt new file mode 100644 index 0000000000000000000000000000000000000000..c127e568615d6a186f921a6e894f0fc82b0a2f41 --- /dev/null +++ b/models/paragraph/fixdur-module/seq2seq_fixdur.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f4cef30a1665cb4d09c5622b2bd91622f8a12ab9d5cdc1c1f6da347ffd4d8885 +size 362646514 diff --git a/models/paragraph/scandl-module/ema_0.9999_080000.pt b/models/paragraph/scandl-module/ema_0.9999_080000.pt new file mode 100644 index 0000000000000000000000000000000000000000..b14240bbd2b93be6db5188aebda7d5b149d10370 --- /dev/null +++ b/models/paragraph/scandl-module/ema_0.9999_080000.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eea605cedc021bc437f965febf00ab01fba073ca9b0d955ba282af99876ac236 +size 616139530 diff --git a/models/paragraph/scandl-module/training_args.json b/models/paragraph/scandl-module/training_args.json new file mode 100644 index 0000000000000000000000000000000000000000..749774f30e968d8828ba88ceda9716d6dce43bf5 --- /dev/null +++ b/models/paragraph/scandl-module/training_args.json @@ -0,0 +1,52 @@ +{ + "checkpoint_path": "/data/lenbol/projects/ScanDL-fix-dur/complete/EMTeC/scandl-module/checkpoint-path", + "vocab": "bert", + "use_plm_init": "no", + "lr": 0.0001, + "batch_size": 64, + "microbatch": 64, + "diffusion_steps": 2000, + "noise_schedule": "sqrt", + "schedule_sampler": "lossaware", + "seq_len": 352, + "resume_checkpoint": "none", + "hidden_t_dim": 352, + "seed": 101, + "hidden_dim": 256, + "learning_steps": 80000, + "save_interval": 5000, + "notes": "-", + "data_split_criterion": "reader", + "num_transformer_layers": 12, + "num_transformer_heads": 8, + "corpus": "emtec", + "inference": "cv", + "load_train_data": "processed_data_all_emtec", + "log_interval": 50, + "eval_interval": 500, + "ema_rate": "0.9999", + "timestep_respacing": "", + "vocab_size": 28996, + "config_name": "bert-base-cased", + "data_dir": "processed_data", + "dataset": "dataset-name", + "dropout": 0.1, + "use_fp16": false, + "fp16_scale_growth": 0.001, + "gradient_clipping": -1.0, + "weight_decay": 0.0, + "learn_sigma": false, + "use_kl": false, + "predict_xstart": true, + "rescale_timesteps": true, + "rescale_learned_sigmas": false, + "sigma_small": false, + "emb_scale_factor": 1.0, + "one_noise_step": true, + "mask_padding": false, + "celer_only_L1": true, + "n_folds": 5, + "ablation_type": "none", + "nll_in_loss": false, + "load_from_checkpoint": false +} \ No newline at end of file diff --git a/models/sentence/fixdur-module/hyperparameters.json b/models/sentence/fixdur-module/hyperparameters.json new file mode 100644 index 0000000000000000000000000000000000000000..04057c5b24824af610b7618532f6eb1516630250 --- /dev/null +++ b/models/sentence/fixdur-module/hyperparameters.json @@ -0,0 +1 @@ +{"num_heads": 12, "num_layers": 12, "num_linear": 8, "bsz": 128, "dropout": 0.5, "use_attention_mask": true} \ No newline at end of file diff --git a/models/sentence/fixdur-module/min_max_scaler.pkl b/models/sentence/fixdur-module/min_max_scaler.pkl new file mode 100644 index 0000000000000000000000000000000000000000..f7529a5b08c8ceb6d6ac06efee121753ab2f329c --- /dev/null +++ b/models/sentence/fixdur-module/min_max_scaler.pkl @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9cea38a7b6b840f0ea990707db171ae219ecf0796d1824e57be935cd66b16dd8 +size 667 diff --git a/models/sentence/fixdur-module/seq2seq_fixdur.pt b/models/sentence/fixdur-module/seq2seq_fixdur.pt new file mode 100644 index 0000000000000000000000000000000000000000..31ce38cc35b590fac9f16355aafcf3a912c2a9dc --- /dev/null +++ b/models/sentence/fixdur-module/seq2seq_fixdur.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:75e490252c13d83cc099e1a2bc1c63bd75373e17fd29229c58fa10a9e9cc00e5 +size 361957490 diff --git a/models/sentence/scandl-module/ema_0.9999_080000.pt b/models/sentence/scandl-module/ema_0.9999_080000.pt new file mode 100644 index 0000000000000000000000000000000000000000..bb9ea7c67c8224a92d0c94e53d44d7c8a2a53be5 --- /dev/null +++ b/models/sentence/scandl-module/ema_0.9999_080000.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1109935957d7a54efcdd0979dcdeb2a4cde0ae1e50d6274a5ed4ce16e301b8a9 +size 612809098 diff --git a/models/sentence/scandl-module/training_args.json b/models/sentence/scandl-module/training_args.json new file mode 100644 index 0000000000000000000000000000000000000000..e3530cec2c26083f8a34e77da4b4e2a2c0fa5a50 --- /dev/null +++ b/models/sentence/scandl-module/training_args.json @@ -0,0 +1,52 @@ +{ + "checkpoint_path": "/data/lenbol/projects/ScanDL-fix-dur/complete/CELER/scandl-module/checkpoint-path", + "vocab": "bert", + "use_plm_init": "no", + "lr": 0.0001, + "batch_size": 64, + "microbatch": 64, + "diffusion_steps": 2000, + "noise_schedule": "sqrt", + "schedule_sampler": "lossaware", + "seq_len": 128, + "resume_checkpoint": "none", + "hidden_t_dim": 128, + "seed": 101, + "hidden_dim": 256, + "learning_steps": 80000, + "save_interval": 5000, + "notes": "-", + "data_split_criterion": "reader", + "num_transformer_layers": 12, + "num_transformer_heads": 8, + "corpus": "celer", + "inference": "cv", + "load_train_data": "processed_data_all_celer", + "log_interval": 50, + "eval_interval": 500, + "ema_rate": "0.9999", + "timestep_respacing": "", + "vocab_size": 28996, + "config_name": "bert-base-cased", + "data_dir": "processed_data", + "dataset": "dataset-name", + "dropout": 0.1, + "use_fp16": false, + "fp16_scale_growth": 0.001, + "gradient_clipping": -1.0, + "weight_decay": 0.0, + "learn_sigma": false, + "use_kl": false, + "predict_xstart": true, + "rescale_timesteps": true, + "rescale_learned_sigmas": false, + "sigma_small": false, + "emb_scale_factor": 1.0, + "one_noise_step": true, + "mask_padding": false, + "celer_only_L1": true, + "n_folds": 5, + "ablation_type": "none", + "nll_in_loss": false, + "load_from_checkpoint": false +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..ae12fbd77b865bf5ade267f448fe16e7a6ce7768 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,15 @@ +blobfile==2.0.1 +datasets==2.14.7 +huggingface-hub==0.17.3 +joblib +matplotlib>=3.7,<3.9 +numpy==1.23.5 +openpyxl==3.0.10 +pandas==1.5.3 +scikit-learn==1.6.1 +seaborn==0.12.2 +textdistance +tqdm==4.66.4 +transformers==4.34.1 +wandb==0.14.0 +setuptools==68.2.2 #for compatibility with wandb 0.14.0 diff --git a/scandl2_utils.py b/scandl2_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..e5a51d58ce842931216456868f07f9f9632438a2 --- /dev/null +++ b/scandl2_utils.py @@ -0,0 +1,77 @@ +import pandas as pd +import numpy as np +from tqdm import tqdm +import os +import random +import torch +from torch.utils.data import Dataset, DataLoader +from typing import Dict, Union, Any, Optional, List + + +class TextDataset(Dataset): + + def __init__( + self, + dataset, + data_args, + split, # 'train', 'test', 'val' + ): + super().__init__() + self.dataset = dataset + self.length = len(self.dataset[split]) + self.data_args = data_args + self.split = split + + def __len__(self): + return self.length + + def __getitem__(self, idx): + sample = { + "mask": np.array(self.dataset[self.split][idx]["mask"]), + "sn_sp_repr": np.array(self.dataset[self.split][idx]["sn_sp_repr"]), + "sn_input_ids": np.array(self.dataset[self.split][idx]["sn_input_ids"]), + "indices_pos_enc": np.array(self.dataset[self.split][idx]["indices_pos_enc"]), + "sn_repr_len": np.array(self.dataset[self.split][idx]["sn_repr_len"]), + "words_for_mapping": self.dataset[self.split][idx]["words_for_mapping"], + } + return sample + + +def text_dataset_loader( + data, + data_args, + split: str, + deterministic: bool = False, +): + dataset = TextDataset( + dataset=data, + data_args=data_args, + split=split, + ) + data_loader = DataLoader( + dataset, + batch_size=data_args.batch_size, + shuffle=not deterministic, + num_workers=0, + ) + return iter(data_loader) + + +class FixdurDataset(Dataset): + def __init__( + self, + data: Dict[str, Union[torch.Tensor, Any]], + ): + super().__init__() + self.data = data + + def __len__(self): + return len(self.data["sp_embeddings"]) + + def __getitem__(self, idx): + sample = { + "sp_embeddings": self.data["sp_embeddings"][idx], + "attention_mask": self.data["attention_mask"][idx], + "unique_idx": self.data["unique_idx"][idx], + } + return sample diff --git a/scandl_module/.DS_Store b/scandl_module/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..d530b51f4205a90ff76fc5bad73f46db4a10f6de Binary files /dev/null and b/scandl_module/.DS_Store differ diff --git a/scandl_module/__init__.py b/scandl_module/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..df01aedfc21a649cdf459db5561ebebd9629bd9a --- /dev/null +++ b/scandl_module/__init__.py @@ -0,0 +1,4 @@ +from .original_scandl import utils as utils +from . import original_scandl as original_scandl + +__all__ = ["utils", "original_scandl"] diff --git a/scandl_module/__pycache__/__init__.cpython-313.pyc b/scandl_module/__pycache__/__init__.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..e30424b5aa4b1df3249a7552e1240936a5cb51a1 Binary files /dev/null and b/scandl_module/__pycache__/__init__.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/__init__.py b/scandl_module/original_scandl/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/scandl_module/original_scandl/__pycache__/__init__.cpython-313.pyc b/scandl_module/original_scandl/__pycache__/__init__.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..b588b9fdc3223a9b354ac281073952a57b616576 Binary files /dev/null and b/scandl_module/original_scandl/__pycache__/__init__.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/__pycache__/sp_gaussian_diffusion.cpython-313.pyc b/scandl_module/original_scandl/__pycache__/sp_gaussian_diffusion.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..0751d74a1a909d4e29e697a6e8761d06f2e2a08b Binary files /dev/null and b/scandl_module/original_scandl/__pycache__/sp_gaussian_diffusion.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/__pycache__/sp_rounding.cpython-313.pyc b/scandl_module/original_scandl/__pycache__/sp_rounding.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..7d223a71194417b6fc077efa45616a41981bc1de Binary files /dev/null and b/scandl_module/original_scandl/__pycache__/sp_rounding.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/__pycache__/sp_transformer_model.cpython-313.pyc b/scandl_module/original_scandl/__pycache__/sp_transformer_model.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a99c1c815396998b3b6525d3968467c015214584 Binary files /dev/null and b/scandl_module/original_scandl/__pycache__/sp_transformer_model.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/__pycache__/step_sample.cpython-313.pyc b/scandl_module/original_scandl/__pycache__/step_sample.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..ae83cd7231e6ebc0733e19428345d6ace4cc104a Binary files /dev/null and b/scandl_module/original_scandl/__pycache__/step_sample.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/config.json b/scandl_module/original_scandl/config.json new file mode 100644 index 0000000000000000000000000000000000000000..2fec1ee2e14fcddd3b4cf6e0e5e784cbab0bfd45 --- /dev/null +++ b/scandl_module/original_scandl/config.json @@ -0,0 +1,52 @@ +{ + "lr": 0.0001, + "batch_size": 128, + "microbatch": 64, + "learning_steps": 80000, + "log_interval": 50, + "save_interval": 5000, + "eval_interval": 500, + "ema_rate": "0.9999", + "resume_checkpoint": "none", + "schedule_sampler": "lossaware", + "diffusion_steps": 2000, + "noise_schedule": "sqrt", + "timestep_respacing": "", + "vocab": "bert", + "use_plm_init": "no", + "vocab_size": 0, + "config_name": "bert-base-cased", + "notes": "folder-notes", + "data_dir": "processed_data", + "dataset": "dataset-name", + "checkpoint_path": "checkpoint-path/test-run", + "seq_len": 128, + "hidden_t_dim": 128, + "hidden_dim": 256, + "dropout": 0.1, + "use_fp16": false, + "fp16_scale_growth": 0.001, + "seed": 102, + "gradient_clipping": -1.0, + "weight_decay": 0.0, + "learn_sigma": false, + "use_kl": false, + "predict_xstart": true, + "rescale_timesteps": true, + "rescale_learned_sigmas": false, + "sigma_small": false, + "emb_scale_factor": 1.0, + "num_transformer_layers": 12, + "num_transformer_heads": 8, + "one_noise_step": true, + "mask_padding": false, + "celer_only_L1": true, + "data_split_criterion": "scanpath", + "corpus": "celer", + "inference": "none", + "n_folds": 5, + "ablation_type": "none", + "nll_in_loss": false, + "load_from_checkpoint": false, + "load_train_data": "-" +} diff --git a/scandl_module/original_scandl/sp_gaussian_diffusion.py b/scandl_module/original_scandl/sp_gaussian_diffusion.py new file mode 100644 index 0000000000000000000000000000000000000000..f7e285caec0070a0679a18726b00fffbd876d8b3 --- /dev/null +++ b/scandl_module/original_scandl/sp_gaussian_diffusion.py @@ -0,0 +1,1183 @@ +""" +This code is adapted from Gong et al.'s 2023 DiffuSeq Model: https://github.com/Shark-NLP/DiffuSeq +""" + +import math +import numpy as np +import torch as th +import sys +import os +import torch.nn + +from .utils.nn import mean_flat + +sys.path.append(".") + + +def get_named_beta_schedule(schedule_name, num_diffusion_timesteps): + """ + Get a pre-defined beta schedule for the given name. + + The beta schedule library consists of beta schedules which remain similar + in the limit of num_diffusion_timesteps. + Beta schedules may be added, but should not be removed or changed once + they are committed to maintain backwards compatibility. + """ + if schedule_name == "linear": + # Linear schedule from Ho et al, extended to work for any number of + # diffusion steps. + scale = 1000 / num_diffusion_timesteps + beta_start = scale * 0.0001 + beta_end = scale * 0.02 + return np.linspace(beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64) + elif schedule_name == "cosine": + return betas_for_alpha_bar( + num_diffusion_timesteps, + lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2, + ) + elif schedule_name == "sqrt": + return betas_for_alpha_bar( + num_diffusion_timesteps, + lambda t: 1 - np.sqrt(t + 0.0001), + ) + elif schedule_name == "trunc_cos": + return betas_for_alpha_bar_left( + num_diffusion_timesteps, + lambda t: np.cos((t + 0.1) / 1.1 * np.pi / 2) ** 2, + ) + elif schedule_name == "trunc_lin": + scale = 1000 / num_diffusion_timesteps + beta_start = scale * 0.0001 + 0.01 + beta_end = scale * 0.02 + 0.01 + return np.linspace(beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64) + elif schedule_name == "pw_lin": + scale = 1000 / num_diffusion_timesteps + beta_start = scale * 0.0001 + 0.01 + beta_mid = scale * 0.0001 # scale * 0.02 + beta_end = scale * 0.02 + first_part = np.linspace(beta_start, beta_mid, 10, dtype=np.float64) + second_part = np.linspace( + beta_mid, beta_end, num_diffusion_timesteps - 10, dtype=np.float64 + ) + return np.concatenate([first_part, second_part]) + else: + raise NotImplementedError(f"unknown beta schedule: {schedule_name}") + + +def betas_for_alpha_bar_left(num_diffusion_timesteps, alpha_bar, max_beta=0.999): + """ + Create a beta schedule that discretizes the given alpha_t_bar function, but shifts towards left interval starting from 0 + which defines the cumulative product of (1-beta) over time from t = [0,1]. + + :param num_diffusion_timesteps: the number of betas to produce. + :param alpha_bar: a lambda that takes an argument t from 0 to 1 and + produces the cumulative product of (1-beta) up to that + part of the diffusion process. + :param max_beta: the maximum beta to use; use values lower than 1 to + prevent singularities. + """ + betas = [] + betas.append(min(1 - alpha_bar(0), max_beta)) + for i in range(num_diffusion_timesteps - 1): + t1 = i / num_diffusion_timesteps + t2 = (i + 1) / num_diffusion_timesteps + betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta)) + return np.array(betas) + + +def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999): + """ + Create a beta schedule that discretizes the given alpha_t_bar function, + which defines the cumulative product of (1-beta) over time from t = [0,1]. + + :param num_diffusion_timesteps: the number of betas to produce. + :param alpha_bar: a lambda that takes an argument t from 0 to 1 and + produces the cumulative product of (1-beta) up to that + part of the diffusion process. + :param max_beta: the maximum beta to use; use values lower than 1 to + prevent singularities. + """ + betas = [] + for i in range(num_diffusion_timesteps): + t1 = i / num_diffusion_timesteps + t2 = (i + 1) / num_diffusion_timesteps + betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta)) + return np.array(betas) + + +class GaussianDiffusion: + """ + Utilities for training and sampling diffusion models. + """ + + def __init__( + self, + *, + betas, + predict_xstart, + rescale_learned_sigmas, + learn_sigmas, + sigma_small, + use_kl, + one_noise_step, + nll_in_loss, + mask_padding, + rescale_timesteps=False, + ): + self.rescale_timesteps = rescale_timesteps + self.predict_xstart = predict_xstart + self.rescale_learned_sigmas = rescale_learned_sigmas + self.learn_sigmas = learn_sigmas + self.sigma_small = sigma_small + self.use_kl = use_kl + self.one_noise_step = one_noise_step + self.nll_in_loss = nll_in_loss + self.mask_padding = mask_padding + + # Use float64 for accuracy. + betas = np.array(betas, dtype=np.float64) # shape [diffusion_steps] + self.betas = betas + assert len(betas.shape) == 1, "betas must be 1-D" + assert (betas > 0).all() and (betas <= 1).all() + + self.num_timesteps = int(betas.shape[0]) + + alphas = 1.0 - betas + + self.alphas_cumprod = np.cumprod(alphas, axis=0) # will approximate 0 + self.alphas_cumprod_prev = np.append( + 1.0, self.alphas_cumprod[:-1] + ) # shifted one to the right + self.alphas_cumprod_next = np.append( + self.alphas_cumprod[1:], 0.0 + ) # shifted one to the left + assert self.alphas_cumprod_prev.shape == (self.num_timesteps,) + + # calculations for diffusion q(x_t | x_{t-1}) and others + self.sqrt_alphas_cumprod = np.sqrt(self.alphas_cumprod) + self.sqrt_one_minus_alphas_cumprod = np.sqrt(1.0 - self.alphas_cumprod) + self.log_one_minus_alphas_cumprod = np.log(1.0 - self.alphas_cumprod) + self.sqrt_recip_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod) + self.sqrt_recipm1_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod - 1) + + # calculations for posterior q(x_{t-1} | x_t, x_0) + self.posterior_variance = ( + betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) + ) + # log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain. + self.posterior_log_variance_clipped = np.log( + np.append(self.posterior_variance[1], self.posterior_variance[1:]) + ) + self.posterior_mean_coef1 = ( + betas * np.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) + ) + self.posterior_mean_coef2 = ( + (1.0 - self.alphas_cumprod_prev) * np.sqrt(alphas) / (1.0 - self.alphas_cumprod) + ) + + self.mapping_func = None # implement in train main() + self.add_mask_noise = False # TODO + + def training_losses(self, model, *args, **kwargs): + self.model = model + return self.training_losses_seq2seq(model, *args, **kwargs) + + def _predict_xstart_from_eps(self, x_t, t, eps): + assert x_t.shape == eps.shape + return ( + _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t + - _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps + ) + + def _predict_eps_from_xstart(self, x_t, t, pred_xstart): + return ( + _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - pred_xstart + ) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) + + def _scale_timesteps(self, t): + if self.rescale_timesteps: + return t.float() * (1000.0 / self.num_timesteps) + return t + + def q_mean_variance(self, x_start, t): + """ + Get the distribution q(x_t | x_0). + + :param x_start: the [N x C x ...] tensor of noiseless inputs. + :param t: the number of diffusion steps (minus 1). Here, 0 means one step. + :return: A tuple (mean, variance, log_variance), all of x_start's shape. + """ + mean = _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start + variance = _extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape) + log_variance = _extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape) + return mean, variance, log_variance + + def q_sample(self, x_start, t, noise=None, mask=None): + """ + Diffuse the data for a given number of diffusion steps. + + In other words, sample from q(x_t | x_0). + + :param x_start: the initial data batch. + :param t: the number of diffusion steps (minus 1). Here, 0 means one step. + :param noise: if specified, the split-out normal noise. + :param mask: anchoring masked position + :return: A noisy version of x_start. + """ + if noise is None: + noise = th.randn_like(x_start) + + assert noise.shape == x_start.shape + x_t = ( + _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start # mu * x_0 + + _extract_into_tensor( + self.sqrt_one_minus_alphas_cumprod, t, x_start.shape + ) # sd * noise + * noise + ) + + if mask == None: + return x_t + else: + mask = th.broadcast_to(mask.unsqueeze(dim=-1), x_start.shape) + return th.where(mask == 0, x_start, x_t) + + def q_posterior_mean_variance(self, x_start, x_t, t): + """ + Compute the mean and variance of the diffusion posterior: + q(x_{t-1} | x_t, x_0) + + """ + assert x_start.shape == x_t.shape + posterior_mean = ( + _extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start + + _extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t + ) + posterior_variance = _extract_into_tensor(self.posterior_variance, t, x_t.shape) + posterior_log_variance_clipped = _extract_into_tensor( + self.posterior_log_variance_clipped, t, x_t.shape + ) + + assert ( + posterior_mean.shape[0] + == posterior_variance.shape[0] + == posterior_log_variance_clipped.shape[0] + == x_start.shape[0] + ) + return posterior_mean, posterior_variance, posterior_log_variance_clipped + + def p_mean_variance( + self, + model, + x, + sn_input_ids_emb, + pos_enc, + mask_sn_padding, + mask_transformer_att, + t, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + subwords_list=None, + atten_vis=None, + atten_vis_fn=None, + atten_vis_path=None, + batch_idx=None, + rank=None, + atten_vis_sp=None, + ): + """ + Apply the model to get p(x_{t-1} | x_t), as well as a prediction of + the initial x, x_0. + + :param model: the model, which takes a signal and a batch of timesteps + as input. + :param x: the [N x C x ...] tensor at time t; the noised input where the embedded condition/sn was not noised + and the target/sp was replaced with Gaussian noise; at time step t + :param sn_input_ids_emb: the BERT embeddings + :param pos_enc: the positional embeddings + :param mask_sn_padding: + :param mask_transformer_att: the attention mask for the transformer + :param t: a 1-D Tensor of timesteps. + :param clip_denoised: if True, clip the denoised signal into [-1, 1]. + :param denoised_fn: if not None, a function which applies to the x_start prediction before it is used to sample. + Applies before clip_denoised. + :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning. + :param subwords_list: list of list containing the subwortokens of each instance (for attention visualization) + :param atten_vis: bool: if True, attention is visualized (heatmaps) + :param atten_vis_fn: the attention visualization function + :param atten_vis_path: the path where to save the heatmaps to + :param batch_idx: index of the current batch + :paran rank: if parallel processing, the GPU index + :param atten_vis_sp: bool: save attention scores of the last timestep + :return: a dict with the following keys: + - 'mean': the model mean output. + - 'variance': the model variance output. + - 'log_variance': the log of 'variance'. + - 'pred_xstart': the prediction for x_0. + """ + if model_kwargs is None: + model_kwargs = {} + + B, C = x.size(0), x.size(-1) + assert t.shape == (B,) + + if not atten_vis and not atten_vis_sp: + model_output = model( + x=x, + ts=self._scale_timesteps(t), + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + attention_mask=mask_transformer_att, + **model_kwargs, + ) + else: + # visualise the attention (heatmaps) + if atten_vis: + model_output, attention_scores = model( + x=x, + ts=self._scale_timesteps(t), + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + attention_mask=mask_transformer_att, + atten_vis=atten_vis, + **model_kwargs, + ) + # visualise attention for last denoising step + if t[0] % 200 == 0: + + atten_vis_fn( + attention_scores=attention_scores, + subwords_list=subwords_list, + batch_idx=batch_idx, + path_to_dir=atten_vis_path, + denoising_step=t[0].item(), + aggregate=True, + rank=rank, + ) + if atten_vis_sp: + + model_output, attention_scores = model( + x=x, + ts=self._scale_timesteps(t), + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + attention_mask=mask_transformer_att, + atten_vis=True, + **model_kwargs, + ) + if t[0] == 0: + + out_path_heatmaps_sp = os.path.join(atten_vis_path, "heatmaps_sps") + if not os.path.exists(out_path_heatmaps_sp): + os.makedirs(out_path_heatmaps_sp) + filename = f"att_scores_rank{rank}_batch{batch_idx}.pt" + path_to_file = os.path.join(out_path_heatmaps_sp, filename) + torch.save(attention_scores, path_to_file) + + model_variance = np.append(self.posterior_variance[1], self.betas[1:]) + model_log_variance = np.log(np.append(self.posterior_variance[1], self.betas[1:])) + + model_variance = _extract_into_tensor(model_variance, t, x.shape) + model_log_variance = _extract_into_tensor(model_log_variance, t, x.shape) + + # The denoised_fn is applied to x_start (the model output) before it is used for sampling + def process_xstart(x): + """here x is the model output""" + if denoised_fn is not None: + # print(denoised_fn) + x = denoised_fn(x, t) + if clip_denoised: + return x.clamp(-1, 1) + return x + + if self.predict_xstart: + # the denoised fn is applied to the model output + pred_xstart = process_xstart(model_output) + else: + ### model is used to predict eps + pred_xstart = process_xstart( + self._predict_xstart_from_eps(x_t=x, t=t, eps=model_output) + ) + + # this is the mean of the posterior distribution q(x_{t-1} | x_t, x_0), estimated from x_t, which is the noised + # input, and pred_xstart, which is what the model predicted to be x_0 from the noised input x_noised/x_t + model_mean, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t) + + assert model_mean.shape == model_log_variance.shape == pred_xstart.shape == x.shape + return { + "mean": model_mean, + "variance": model_variance, + "log_variance": model_log_variance, + "pred_xstart": pred_xstart, + } + + def p_sample( + self, + model, + x, + sn_input_ids_emb, + pos_enc, + mask_sn_padding, + mask_transformer_att, + t, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + top_p=None, + mask=None, + x_start=None, + subwords_list=None, + atten_vis=None, + atten_vis_fn=None, + atten_vis_path=None, + batch_idx=None, + rank=None, + atten_vis_sp=None, + ): + """ + Sample x_{t-1} from the model at the given timestep. + + :param model: the model to sample from; the transformer model that learned the denoising + :param x: the current tensor at x_{t-1}. + :param sn_input_ids_emb: the BERT embeddings + :param pos_enc: the positional embeddings + :param mask_sn_padding: + :param mask_transformer_att: the attention mask for the transformer + :param t: the value of t, starting at 0 for the first diffusion step. + :param clip_denoised: if True, clip the x_start prediction to [-1, 1]. + :param denoised_fn: if not None, a function which applies to the x_start prediction before it is used to sample. + :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning. + :param top_p: + :param mask: anchoring masked position to x_start + :param x_start: + :param subwords_list: list of list containing the subwortokens of each instance (for attention visualization) + :param atten_vis: bool: if True, attention is visualized (heatmaps) + :param atten_vis_fn: the attention visualization function + :param atten_vis_path: the path where to save the heatmaps to + :param batch_idx: index of the current batch + :paran rank: if parallel processing, the GPU index + :param atten_vis_sp: bool: save attention scores of the last timestep + :return: a dict containing the following keys: + - 'sample': a random sample from the model. + - 'pred_xstart': a prediction of x_0. + """ + out = self.p_mean_variance( + model=model, + x=x, + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + mask_sn_padding=mask_sn_padding, + mask_transformer_att=mask_transformer_att, + t=t, + clip_denoised=clip_denoised, + denoised_fn=denoised_fn, + model_kwargs=model_kwargs, + subwords_list=subwords_list, + atten_vis=atten_vis, + atten_vis_fn=atten_vis_fn, + atten_vis_path=atten_vis_path, + batch_idx=batch_idx, + rank=rank, + atten_vis_sp=atten_vis_sp, + ) + + if top_p is not None and top_p > 0: + # print('top_p sampling') + noise = th.randn_like(x) + replace_mask = th.abs(noise) > top_p + while replace_mask.any(): + noise[replace_mask] = th.randn_like(noise[replace_mask]) + replace_mask = th.abs(noise) > top_p + assert (th.abs(noise) <= top_p).all() + + else: + noise = th.randn_like(x) + + nonzero_mask = ( + (t != 0).float().view(-1, *([1] * (len(x.shape) - 1))) + ) # no noise when t == 0 + + sample = out["mean"] + nonzero_mask * th.exp(0.5 * out["log_variance"]) * noise + + if mask == None: + pass + else: + # the original embedding for the sn, and the predicted sample for the sp + sample = th.where(mask == 0, x_start, sample) + + return { + "sample": sample, + "pred_xstart": out["pred_xstart"], + "greedy_mean": out["mean"], + "out": out, + } + + def p_sample_loop( + self, + model, + shape, + noise=None, + sn_input_ids_emb=None, + pos_enc=None, + mask_sn_padding=None, + mask_transformer_att=None, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + device=None, + progress=False, + top_p=None, + clamp_step=None, + clamp_first=None, + mask=None, + x_start=None, + subwords_list=None, + atten_vis=None, + atten_vis_fn=None, + atten_vis_path=None, + batch_idx=None, + gap=1, + rank=None, + atten_vis_sp=None, + ): + """ + Generate samples from the model. + + :param model: the transformer model that was trained to learn the denoising + :param shape: the shape of the samples, (N, C, H, W). + :param noise: the Gaussian noise that should be denoised at inference (the replaced word ID emb) + :param sn_input_ids_emb: the BERT embeddings + :param pos_enc: the positional embeddings + :param mask_sn_padding: + :param mask_transformer_att: the attention mask for the transformer + :param clip_denoised: if True, clip x_start predictions to [-1, 1]. + :param denoised_fn: if not None, a function which applies to the x_start prediction before it is used to sample. + :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning. + :param device: if specified, the device to create the samples on. If not specified, use a model parameter's device. + :param progress: if True, show a tqdm progress bar. + :param top_p: + :param clamp_step: in clamp_first mode, choose end clamp step, otherwise starting clamp step + :param clamp_first: bool, clamp_first mode + :param mask: anchoring masked position to x_start + :param x_start: the word ID embedding before replaced by noise + :param subwords_list: list of list containing the subwortokens of each instance (for attention visualization) + :param atten_vis: bool: if True, attention is visualized (heatmaps) + :param atten_vis_fn: the attention visualization function + :param atten_vis_path: the path where to save the heatmaps to + :param batch_idx: index of the current batch + :param gap: + :paran rank: if parallel processing, the GPU index + :param atten_vis_sp: bool: save attention scores of the last timestep + :return: a non-differentiable batch of samples. + """ + final = [] + for sample in self.p_sample_loop_progressive( + model, + shape, + noise=noise, + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + mask_sn_padding=mask_sn_padding, + mask_transformer_att=mask_transformer_att, + clip_denoised=clip_denoised, + denoised_fn=denoised_fn, + model_kwargs=model_kwargs, + device=device, + progress=progress, + top_p=top_p, + clamp_step=clamp_step, + clamp_first=clamp_first, + mask=mask, + x_start=x_start, + subwords_list=subwords_list, + atten_vis=atten_vis, + atten_vis_fn=atten_vis_fn, + atten_vis_path=atten_vis_path, + batch_idx=batch_idx, + rank=rank, + atten_vis_sp=atten_vis_sp, + ): + final.append(sample["sample"]) + return final + + def p_sample_loop_progressive( + self, + model, + shape, + noise=None, + sn_input_ids_emb=None, + pos_enc=None, + mask_sn_padding=None, + mask_transformer_att=None, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + device=None, + progress=False, + top_p=None, + clamp_step=None, + clamp_first=None, + mask=None, + x_start=None, + subwords_list=None, + atten_vis=None, + atten_vis_fn=None, + atten_vis_path=None, + batch_idx=None, + rank=None, + atten_vis_sp=None, + ): + """ + Generate samples from the model and yield intermediate samples from + each timestep of diffusion. + + Arguments are the same as p_sample_loop(). + Returns a generator over dicts, where each dict is the return value of + p_sample(). + """ + if device is None: + device = next(model.parameters()).device + assert isinstance(shape, (tuple, list)) + + # noise/sample_x is the input that was noised: the concatenated sn-sp embedding where the sp was completely + # replaced with Gaussian noise from the standard normal distribution + if noise is not None: + sample_x = noise + else: + sample_x = th.randn(*shape, device=device) + + # the number of diffusion steps in reverse order + indices = list(range(self.num_timesteps))[::-1] + + if progress: + # Lazy import so that we don't depend on tqdm. + from tqdm.auto import tqdm + + indices = tqdm(indices) + + # denoising from the number of diffusion steps T to t=0 + for i in indices: # from T to 0 + + t = th.tensor([i] * shape[0], device=device) + if not clamp_first: + if i > clamp_step: + denoised_fn_cur = None + else: + denoised_fn_cur = denoised_fn + else: + if i >= clamp_step: + denoised_fn_cur = denoised_fn + else: + denoised_fn_cur = None + + with th.no_grad(): + out = self.p_sample( + model=model, + x=sample_x, + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + mask_sn_padding=mask_sn_padding, + mask_transformer_att=mask_transformer_att, + t=t, + clip_denoised=clip_denoised, + denoised_fn=denoised_fn_cur, + model_kwargs=model_kwargs, + top_p=top_p, + mask=mask, + subwords_list=subwords_list, + x_start=x_start, + atten_vis=atten_vis, + atten_vis_fn=atten_vis_fn, + atten_vis_path=atten_vis_path, + batch_idx=batch_idx, + rank=rank, + atten_vis_sp=atten_vis_sp, + ) + yield out + sample_x = out["sample"] + + def _get_x_start(self, x_start_mean, std): + """ + Word embedding projection from {Emb(w)} to {x_0} + :param x_start_mean: word embedding + :return: x_0 + """ + noise = th.randn_like(x_start_mean) + assert noise.shape == x_start_mean.shape + # print(x_start_mean.device, noise.device) + return x_start_mean + std * noise + + def _token_discrete_loss(self, x_t, get_logits, input_ids, mask=None, truncate=False, t=None): + """ + the loss of -log p(w|z_0) + :param x_start_mean: word embedding + :return: x_0 + """ + reshaped_x_t = x_t + logits = get_logits(reshaped_x_t) # shape [microbatch size, seq_len, vocabulary] + # print(logits.shape) + loss_fct = th.nn.CrossEntropyLoss(reduction="none") + decoder_nll = loss_fct(logits.view(-1, logits.size(-1)), input_ids.view(-1)).view( + input_ids.shape + ) + if mask != None: + decoder_nll *= mask + # print(decoder_nll.shape) + if mask != None: + decoder_nll = decoder_nll.sum(dim=-1) / mask.sum(dim=-1) + else: + decoder_nll = decoder_nll.mean(dim=-1) + + return decoder_nll + + def _x0_helper(self, model_output, x, t): + + if self.predict_xstart: + pred_xstart = model_output + pred_prev, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t) + + else: # predict eps + pred_xstart = self._predict_xstart_from_eps(x_t=x, t=t, eps=model_output) + + pred_prev, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t) + + return {"pred_xprev": pred_prev, "pred_xstart": pred_xstart} + + def training_losses_seq2seq( + self, + model, # the transformer model + t, # the number of noise adding steps for each instance in the microbatch + sn_sp_repr, + mask, + sn_input_ids, + indices_pos_enc, + mask_sn_padding, + mask_transformer_att, + noise=None, + ): + """ + Compute training losses for a single timestep. + + :param model: the transformer model + :param t: a batch of timestep indices. + :param sn_sp_repr: the word IDs of sn and sp + :param mask: masking the sn + :param sn_input_ids: the tokenizer input IDs + :param indices_pos_enc: the indices for pos enc + :param_mask_sn_padding: + :param mask_transformer_att: the transformer att + :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning. + :param noise: if specified, the specific Gaussian noise to try to remove. + :return: a dict with the key "loss" containing a tensor of shape [N]. + Some mean or variance settings may also have other keys. + """ + + microbatch_size, seq_len = sn_sp_repr.shape + + # get the word ID embedding, BERT embedding, positional embedding + sn_sp_emb, pos_enc, sn_input_ids_emb = model.model.module.get_embeds( + sn_sp_repr=sn_sp_repr, + sn_input_ids=sn_input_ids, + indices_pos_enc=indices_pos_enc, + ) + + # get the standard deviation, shape [microbatch, args.seq_len, hidden_size=768] + std = _extract_into_tensor( + self.sqrt_one_minus_alphas_cumprod, th.tensor([0]).to(sn_sp_emb.device), sn_sp_emb.shape + ) + + # map sn_sp_emb to x_start, which is a one-step noised sn_sp_emb (in paper it's z_0) + if ( + self.one_noise_step + ): # this should always be true actually (without it performance is bad) + x_start = self._get_x_start(sn_sp_emb, std) + else: + x_start = sn_sp_emb + + # sample noise in the same shape as our input + if noise is None: + noise = th.randn_like(x_start) + + # get the noised sample x_t, which is still of shape [microbatch, args.seq_len, hidden_size=768] + # the condition/sn is not noised (hence the input mask) + # each instance in the microbatch receives a different amount of noise (t noising steps, as given in vector t) + x_t = self.q_sample( + x_start=x_start, + t=t, + noise=noise, + mask=mask, + ) + + terms = {} + + target = x_start + + # model_output is of shape [microbatch, args.seq_len, emb_dim=768] + model_output = model( + x=x_t, + ts=self._scale_timesteps(t), + sn_input_ids_emb=sn_input_ids_emb, + pos_enc=pos_enc, + attention_mask=mask_transformer_att, + ) + assert model_output.shape == target.shape == x_start.shape + + # Loss 1: Mean Squared Error (MSE) (L_{VLB}) + terms["mse"] = mean_flat((target - model_output) ** 2) + model_out_x_start = self._x0_helper(model_output, x_t, t)["pred_xstart"] + t0_mask = ( + t == 0 + ) # mask that says true for every instance where no noise was received, i.e. t=0 + # MSE between the model output and the embedded input before the one noise step + t0_loss = mean_flat((sn_sp_emb - model_out_x_start) ** 2) + # update the MSE between the model output and the one-step noised input embeddings with the MSE between the + # model output and the embeddings before the one noise step wherever there was no noise received in the noising + # process (i.e., wherever t was 0) + terms["mse"] = th.where(t0_mask, t0_loss, terms["mse"]) + + # Loss 2: L_{round} + out_mean, _, _ = self.q_mean_variance( + x_start, th.LongTensor([self.num_timesteps - 1]).to(x_start.device) + ) + tT_loss = mean_flat(out_mean**2) + + # for the NLL losses, we need to convert the model output into logits + get_logits = model.model.module.get_logits + + # Loss 3: L_{EMB} + # compute the NLL between the one-noised embeddings and the initial representation (word IDs) + # embedding regularisation + decoder_nll = self._token_discrete_loss(x_start, get_logits, sn_sp_repr) + + # unused Loss + terms["nll"] = self._token_discrete_loss(model_output, get_logits, sn_sp_repr, mask=mask) + + # combined loss + if self.nll_in_loss: # should be False; model performance drops if nll included + terms["loss"] = terms["mse"] + tT_loss + decoder_nll + terms["nll"] + else: + terms["loss"] = terms["mse"] + tT_loss + decoder_nll + + return terms + + def ddim_sample( + self, + model, + x, + t, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + eta=0.0, + langevin_fn=None, + mask=None, + x_start=None, + ): + """ + Sample x_{t-1} from the model using DDIM. + + Same usage as p_sample(). + """ + out = self.p_mean_variance( + model, + x, + t, + clip_denoised=clip_denoised, + denoised_fn=denoised_fn, + model_kwargs=model_kwargs, + ) + # Usually our model outputs epsilon, but we re-derive it + # in case we used x_start or x_prev prediction. + eps = self._predict_eps_from_xstart(x, t, out["pred_xstart"]) + alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape) + alpha_bar_prev = _extract_into_tensor(self.alphas_cumprod_prev, t, x.shape) + sigma = ( + eta + * th.sqrt((1 - alpha_bar_prev) / (1 - alpha_bar)) + * th.sqrt(1 - alpha_bar / alpha_bar_prev) + ) + # Equation 12. + noise = th.randn_like(x) + mean_pred = ( + out["pred_xstart"] * th.sqrt(alpha_bar_prev) + + th.sqrt(1 - alpha_bar_prev - sigma**2) * eps + ) + nonzero_mask = ( + (t != 0).float().view(-1, *([1] * (len(x.shape) - 1))) + ) # no noise when t == 0 + # print(sigma.mean()) + sample = mean_pred + nonzero_mask * sigma * noise + if langevin_fn: + print(t.shape) + sample = langevin_fn(sample, mean_pred, sigma, self.alphas_cumprod_prev[t[0]], t, x) + + if mask == None: + pass + else: + sample = th.where(mask == 0, x_start, sample) + + return {"sample": sample, "pred_xstart": out["pred_xstart"]} + + def ddim_reverse_sample( + self, + model, + x, + t, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + eta=0.0, + ): + """ + Sample x_{t+1} from the model using DDIM reverse ODE. + """ + assert eta == 0.0, "Reverse ODE only for deterministic path" + out = self.p_mean_variance( + model, + x, + t, + clip_denoised=clip_denoised, + denoised_fn=denoised_fn, + model_kwargs=model_kwargs, + ) + # Usually our model outputs epsilon, but we re-derive it + # in case we used x_start or x_prev prediction. + eps = ( + _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x.shape) * x + - out["pred_xstart"] + ) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x.shape) + alpha_bar_next = _extract_into_tensor(self.alphas_cumprod_next, t, x.shape) + + # Equation 12. reversed + mean_pred = out["pred_xstart"] * th.sqrt(alpha_bar_next) + th.sqrt(1 - alpha_bar_next) * eps + + return {"sample": mean_pred, "pred_xstart": out["pred_xstart"]} + + def ddim_sample_loop( + self, + model, + shape, + noise=None, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + device=None, + progress=False, + top_p=None, + clamp_step=None, + clamp_first=None, + mask=None, + x_start=None, + gap=1, + ): + """ + Generate samples from the model using DDIM. + :param gap: compute ddim sampling for each {gap} step + + Same usage as p_sample_loop(). + """ + final = [] + for sample in self.ddim_sample_loop_progressive( + model, + shape, + noise=noise, + clip_denoised=clip_denoised, + denoised_fn=denoised_fn, + model_kwargs=model_kwargs, + device=device, + progress=progress, + mask=mask, + x_start=x_start, + gap=gap, + ): + final.append(sample["sample"]) + return final + + def ddim_sample_loop_progressive( + self, + model, + shape, + noise=None, + clip_denoised=True, + denoised_fn=None, + model_kwargs=None, + device=None, + progress=False, + eta=0.0, + langevin_fn=None, + mask=None, + x_start=None, + gap=1, + ): + """ + Use DDIM to sample from the model and yield intermediate samples from + each timestep of DDIM. + + Same usage as p_sample_loop_progressive(). + """ + if device is None: + device = next(model.parameters()).device + assert isinstance(shape, (tuple, list)) + if noise is not None: + sample_x = noise + else: + sample_x = th.randn(*shape, device=device) + indices = list(range(self.num_timesteps))[::-1][::gap] + + if progress: + # Lazy import so that we don't depend on tqdm. + from tqdm.auto import tqdm + + indices = tqdm(indices) + + for i in indices: + t = th.tensor([i] * shape[0], device=device) + with th.no_grad(): + out = self.ddim_sample( + model, + sample_x, + t, + clip_denoised=clip_denoised, + denoised_fn=denoised_fn, + model_kwargs=model_kwargs, + mask=mask, + x_start=x_start, + ) + yield out + sample_x = out["sample"] + + +def _extract_into_tensor(arr, timesteps, broadcast_shape): + """ + Extract values from a 1-D numpy array for a batch of indices. + + :param arr: the 1-D numpy array. + :param timesteps: a tensor of indices into the array to extract. + :param broadcast_shape: a larger shape of K dimensions with the batch + dimension equal to the length of timesteps. + :return: a tensor of shape [batch_size, 1, ...] where the shape has K dims. + """ + res = th.from_numpy(arr).to(device=timesteps.device)[timesteps].float() + while len(res.shape) < len(broadcast_shape): + res = res[..., None] + return res.expand(broadcast_shape) + + +def space_timesteps(num_timesteps, section_counts): + """ + Create a list of timesteps to use from an original diffusion process, + given the number of timesteps we want to take from equally-sized portions + of the original process. + + For example, if there's 300 timesteps and the section counts are [10,15,20] + then the first 100 timesteps are strided to be 10 timesteps, the second 100 + are strided to be 15 timesteps, and the final 100 are strided to be 20. + + If the stride is a string starting with "ddim", then the fixed striding + from the DDIM paper is used, and only one section is allowed. + + :param num_timesteps: the number of diffusion steps in the original + process to divide up. + :param section_counts: either a list of numbers, or a string containing + comma-separated numbers, indicating the step count + per section. As a special case, use "ddimN" where N + is a number of steps to use the striding from the + DDIM paper. + :return: a set of diffusion steps from the original process to use. + """ + if isinstance(section_counts, str): + if section_counts.startswith("ddim"): + desired_count = int(section_counts[len("ddim") :]) + for i in range(1, num_timesteps): + if len(range(0, num_timesteps, i)) == desired_count: + return set(range(0, num_timesteps, i)) + raise ValueError(f"cannot create exactly {num_timesteps} steps with an integer stride") + section_counts = [int(x) for x in section_counts.split(",")] + size_per = num_timesteps // len(section_counts) + extra = num_timesteps % len(section_counts) + start_idx = 0 + all_steps = [] + for i, section_count in enumerate(section_counts): + size = size_per + (1 if i < extra else 0) + if size < section_count: + raise ValueError(f"cannot divide section of {size} steps into {section_count}") + if section_count <= 1: + frac_stride = 1 + else: + frac_stride = (size - 1) / (section_count - 1) + cur_idx = 0.0 + taken_steps = [] + for _ in range(section_count): + taken_steps.append(start_idx + round(cur_idx)) + cur_idx += frac_stride + all_steps += taken_steps + start_idx += size + return set(all_steps) + + +class SpacedDiffusion(GaussianDiffusion): + """ + A diffusion process which can skip steps in a base diffusion process. + + :param use_timesteps: a collection (sequence or set) of timesteps from the + original diffusion process to retain. + :param kwargs: the kwargs to create the base diffusion process. + """ + + def __init__(self, use_timesteps, **kwargs): + self.use_timesteps = set(use_timesteps) + self.timestep_map = [] + self.original_num_steps = len(kwargs["betas"]) + + # print(kwargs.keys()) + base_diffusion = GaussianDiffusion(**kwargs) # pylint: disable=missing-kwoa + last_alpha_cumprod = 1.0 + new_betas = [] + for i, alpha_cumprod in enumerate(base_diffusion.alphas_cumprod): + if i in self.use_timesteps: + new_betas.append(1 - alpha_cumprod / last_alpha_cumprod) + last_alpha_cumprod = alpha_cumprod + self.timestep_map.append(i) + kwargs["betas"] = np.array(new_betas) + super().__init__(**kwargs) + + def p_mean_variance(self, model, *args, **kwargs): # pylint: disable=signature-differs + # print('called p_mean_var') + return super().p_mean_variance(self._wrap_model(model), *args, **kwargs) + + def training_losses(self, model, *args, **kwargs): # pylint: disable=signature-differs + # print('called training_losses') + return super().training_losses(self._wrap_model(model), *args, **kwargs) + + def _wrap_model(self, model): + if isinstance(model, _WrappedModel): + return model + return _WrappedModel( + model, self.timestep_map, self.rescale_timesteps, self.original_num_steps + ) + + def _scale_timesteps(self, t): + # Scaling is done by the wrapped model. + return t + + +class _WrappedModel: + def __init__(self, model, timestep_map, rescale_timesteps, original_num_steps): + self.model = model + self.timestep_map = timestep_map + self.rescale_timesteps = rescale_timesteps + self.original_num_steps = original_num_steps + + def __call__(self, x, ts, **kwargs): + # print(ts) + map_tensor = th.tensor(self.timestep_map, device=ts.device, dtype=ts.dtype) + new_ts = map_tensor[ts] + # print(new_ts) + if self.rescale_timesteps: + new_ts = new_ts.float() * (1000.0 / self.original_num_steps) + # temp = self.model(x, new_ts, **kwargs) + # print(temp.shape) + # return temp + # print(new_ts) + return self.model(x, new_ts, **kwargs) diff --git a/scandl_module/original_scandl/sp_rounding.py b/scandl_module/original_scandl/sp_rounding.py new file mode 100644 index 0000000000000000000000000000000000000000..13888a1dc45b441e4b741ae2c08e4c9739199f44 --- /dev/null +++ b/scandl_module/original_scandl/sp_rounding.py @@ -0,0 +1,58 @@ +import numpy as np +import torch + + +def get_knn(model_emb, text_emb, dist="cos"): + if dist == "cos": + adjacency = model_emb @ text_emb.transpose(1, 0).to(model_emb.device) + elif dist == "l2": + adjacency = model_emb.unsqueeze(1).expand(-1, text_emb.size(0), -1) - text_emb.unsqueeze( + 0 + ).expand(model_emb.size(0), -1, -1) + adjacency = -torch.norm(adjacency, dim=-1) + topk_out = torch.topk(adjacency, k=6, dim=0) + return topk_out.values, topk_out.indices + + +def get_efficient_knn(sn_sp_repr_embedding_weight, text_emb): + """ + :param sn_sp_repr_embedding_weight: + :param text_emb: + """ + emb_norm = (sn_sp_repr_embedding_weight**2).sum(-1).view(-1, 1) + text_emb_t = torch.transpose(text_emb.view(-1, text_emb.size(-1)), 0, 1) + arr_norm = (text_emb**2).sum(-1).view(-1, 1) + dist = ( + emb_norm + + arr_norm.cpu().transpose(0, 1) + - 2.0 * torch.mm(sn_sp_repr_embedding_weight, text_emb_t.cpu()) + ) # (vocab, d) x (d, bsz*seqlen) + dist = torch.clamp(dist, 0.0, np.inf) + topk_out = torch.topk(-dist, k=1, dim=0) + return topk_out.values, topk_out.indices + + +def denoised_fn_round(args, sn_sp_repr_embedding, text_emb, t): + """ + :param sn_sp_repr_embedding: the weights/parameter of the embedding layer that embeds the concatenated word IDs + :param text_emb: the model output at denoising step t; the transformer received the noise as input; this is the pred. + shape [batch size, args.seq_len, hidden_dim=768] + :param t: the current time step, shape [batch size] (same t for each instance in the batch) + """ + sn_sp_repr_embedding_weight = sn_sp_repr_embedding.weight + old_shape = text_emb.shape + old_device = text_emb.device + + if len(text_emb.shape) > 2: + text_emb = text_emb.reshape(-1, text_emb.size(-1)) + else: + text_emb = text_emb + + text_emb.to(sn_sp_repr_embedding_weight.device) + + val, indices = get_efficient_knn( + sn_sp_repr_embedding_weight=sn_sp_repr_embedding_weight, text_emb=text_emb + ) + rounded_tokens = indices[0] + new_embeds = sn_sp_repr_embedding(rounded_tokens).view(old_shape).to(old_device) + return new_embeds diff --git a/scandl_module/original_scandl/sp_transformer_model.py b/scandl_module/original_scandl/sp_transformer_model.py new file mode 100644 index 0000000000000000000000000000000000000000..7383783b2470aef9ebb456cc80e1262e37eb6040 --- /dev/null +++ b/scandl_module/original_scandl/sp_transformer_model.py @@ -0,0 +1,167 @@ +from transformers import AutoConfig +from transformers.models.bert.modeling_bert import BertEncoder, BertModel +import torch +import torch as th +import torch.nn as nn +from typing import Optional + +from .utils.nn import ( + SiLU, + linear, + timestep_embedding, +) + + +class TransformerNetModel(nn.Module): + """ + The ScanDL transformer. + """ + + def __init__( + self, + input_dims, + output_dims, + hidden_t_dim, + num_transformer_layers, + num_transformer_heads, + one_noise_step, + mask_padding, + dropout=0, + config=None, + config_name="bert-base-uncased", + vocab_size=None, + init_pretrained="no", + logits_mode=1, + ): + super().__init__() + + if config is None: + config = AutoConfig.from_pretrained(config_name) + config.hidden_dropout_prob = dropout + config.num_hidden_layers = num_transformer_layers + config.num_attention_heads = num_transformer_heads + config.hidden_size = input_dims + + self.input_dims = input_dims + self.hidden_t_dim = hidden_t_dim + self.output_dims = output_dims + self.dropout = dropout + self.logits_mode = logits_mode + + self.mask_padding = mask_padding + self.one_noise_step = one_noise_step + + self.bert_for_embedding = BertModel.from_pretrained(config_name) + # freeze BERT parameters (so that embeddings are freezed) + for param in self.bert_for_embedding.parameters(): + param.requires_grad = False + + self.sn_sp_repr_embedding = nn.Embedding(self.hidden_t_dim, self.input_dims) + + self.positional_encoding = nn.Embedding(self.hidden_t_dim, self.input_dims) + self.sn_input_ids_embedding = self.bert_for_embedding.embeddings.word_embeddings + # additional linear layer needed if hidden is not 768 to map from pretrained BERT embeddings to other dim + if self.input_dims != 768: + self.proj_bert_emb = nn.Linear(768, self.input_dims) + + self.lm_head = nn.Linear(self.input_dims, self.hidden_t_dim) + with torch.no_grad(): + self.lm_head.weight = self.sn_sp_repr_embedding.weight + + time_embed_dim = hidden_t_dim * 4 + self.time_embed = nn.Sequential( + linear(hidden_t_dim, time_embed_dim), + SiLU(), + linear(time_embed_dim, config.hidden_size), + ) + + self.input_transformers = BertEncoder(config) + self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) + self.dropout = nn.Dropout(config.hidden_dropout_prob) + + def get_embeds( + self, + sn_sp_repr, + sn_input_ids, + indices_pos_enc, + ): + + sn_sp_emb = self.sn_sp_repr_embedding(sn_sp_repr) + pos_enc = self.positional_encoding(indices_pos_enc) + if self.input_dims == 768: + sn_input_ids_emb = self.sn_input_ids_embedding(sn_input_ids) + else: + sn_input_ids_emb_bert_embs = self.sn_input_ids_embedding(sn_input_ids) + sn_input_ids_emb = self.proj_bert_emb(sn_input_ids_emb_bert_embs) + return sn_sp_emb, pos_enc, sn_input_ids_emb + + def get_logits(self, model_output): + if self.logits_mode == 1: + return self.lm_head(model_output) + elif self.logits_mode == 2: # standard cosine similarity + raise NotImplementedError( + "standard cosine similarity not yet implemented for sp model output." + ) + else: + raise NotImplementedError + + def forward( + self, + x, # x_t + ts, + sn_input_ids_emb, + pos_enc, + attention_mask: Optional[torch.tensor] = None, + atten_vis: Optional[bool] = False, + ): + """ + Apply the model to an input batch. + + :param x: the noised input ID embeddings + :param ts: a 1-D batch of timesteps. + :param sn_input_ids_emb: the BERT embeddings + :param pos_enc: the positional embeddings + :param attention_mask: the attention mask (only given during training, not during inference) + :atten_vis: visualise attention + """ + # timestep embedding + emb_t = self.time_embed(timestep_embedding(ts, self.hidden_t_dim)) + + # add the input x_t, the positional encoding pos_enc, the word ID embedding word_id_emb, and the timestep emb + emb_inputs = x + pos_enc + sn_input_ids_emb + emb_t.unsqueeze(1).expand(-1, x.size(1), -1) + + # pipe through dropout and layer normalisation + emb_inputs = self.dropout(self.LayerNorm(emb_inputs)) + + if self.mask_padding: + if attention_mask == None: + raise ValueError("padding should be masked, but no attention mask given.") + + extended_attention_mask = attention_mask[:, None, None, :] + + if atten_vis: + model_out = self.input_transformers( + emb_inputs, attention_mask=extended_attention_mask, output_attentions=True + ) + input_trans_hidden_states = model_out.last_hidden_state + attention_scores = model_out.attentions + + else: + input_trans_hidden_states = self.input_transformers( + emb_inputs, attention_mask=extended_attention_mask + ).last_hidden_state + + else: + if atten_vis: + model_out = self.input_transformers(emb_inputs, output_attentions=True) + input_trans_hidden_states = model_out.last_hidden_state + attention_scores = model_out.attentions + else: + input_trans_hidden_states = self.input_transformers(emb_inputs).last_hidden_state + + h = input_trans_hidden_states + h = h.type(x.dtype) + if atten_vis: + return h, attention_scores + else: + return h diff --git a/scandl_module/original_scandl/step_sample.py b/scandl_module/original_scandl/step_sample.py new file mode 100644 index 0000000000000000000000000000000000000000..0e4f1ed543ebc1c24b8f1fd16393612122d0fb18 --- /dev/null +++ b/scandl_module/original_scandl/step_sample.py @@ -0,0 +1,170 @@ +from abc import ABC, abstractmethod + +import numpy as np +import torch as th +import torch.distributed as dist + + +def create_named_schedule_sampler(name, diffusion): + """ + Create a ScheduleSampler from a library of pre-defined samplers. + + :param name: the name of the sampler. + :param diffusion: the diffusion object to sample for. + """ + if name == "uniform": + return UniformSampler(diffusion) + elif name == "lossaware": + return LossSecondMomentResampler(diffusion) + elif name == "fixstep": + return FixSampler(diffusion) + else: + raise NotImplementedError(f"unknown schedule sampler: {name}") + + +class ScheduleSampler(ABC): + """ + A distribution over timesteps in the diffusion process, intended to reduce + variance of the objective. + + By default, samplers perform unbiased importance sampling, in which the + objective's mean is unchanged. + However, subclasses may override sample() to change how the resampled + terms are reweighted, allowing for actual changes in the objective. + """ + + @abstractmethod + def weights(self): + """ + Get a numpy array of weights, one per diffusion step. + + The weights needn't be normalized, but must be positive. + """ + + def sample(self, batch_size, device): + """ + Importance-sample timesteps for a batch. + + :param batch_size: the number of timesteps. + :param device: the torch device to save to. + :return: a tuple (timesteps, weights): + - timesteps: a tensor of timestep indices. + - weights: a tensor of weights to scale the resulting losses. + """ + w = self.weights() + p = w / np.sum(w) + indices_np = np.random.choice(len(p), size=(batch_size,), p=p) + indices = th.from_numpy(indices_np).long().to(device) + weights_np = 1 / (len(p) * p[indices_np]) + weights = th.from_numpy(weights_np).float().to(device) + return indices, weights + + +class UniformSampler(ScheduleSampler): + def __init__(self, diffusion): + self.diffusion = diffusion + self._weights = np.ones([diffusion.num_timesteps]) + + def weights(self): + return self._weights + + +class FixSampler(ScheduleSampler): + def __init__(self, diffusion): + self.diffusion = diffusion + + ############################################################### + ### You can custome your own sampling weight of steps here. ### + ############################################################### + self._weights = np.concatenate( + [ + np.ones([diffusion.num_timesteps // 2]), + np.zeros([diffusion.num_timesteps // 2]) + 0.5, + ] + ) + + def weights(self): + return self._weights + + +class LossAwareSampler(ScheduleSampler): + def update_with_local_losses(self, local_ts, local_losses): + """ + Update the reweighting using losses from a model. + + Call this method from each rank with a batch of timesteps and the + corresponding losses for each of those timesteps. + This method will perform synchronization to make sure all of the ranks + maintain the exact same reweighting. + + :param local_ts: an integer Tensor of timesteps. + :param local_losses: a 1D Tensor of losses. + """ + batch_sizes = [ + th.tensor([0], dtype=th.int32, device=local_ts.device) + for _ in range(dist.get_world_size()) + ] + dist.all_gather( + batch_sizes, + th.tensor([len(local_ts)], dtype=th.int32, device=local_ts.device), + ) + + # Pad all_gather batches to be the maximum batch size. + batch_sizes = [x.item() for x in batch_sizes] + max_bs = max(batch_sizes) + + timestep_batches = [th.zeros(max_bs).to(local_ts) for bs in batch_sizes] + loss_batches = [th.zeros(max_bs).to(local_losses) for bs in batch_sizes] + dist.all_gather(timestep_batches, local_ts) + dist.all_gather(loss_batches, local_losses) + timesteps = [x.item() for y, bs in zip(timestep_batches, batch_sizes) for x in y[:bs]] + losses = [x.item() for y, bs in zip(loss_batches, batch_sizes) for x in y[:bs]] + self.update_with_all_losses(timesteps, losses) + + @abstractmethod + def update_with_all_losses(self, ts, losses): + """ + Update the reweighting using losses from a model. + + Sub-classes should override this method to update the reweighting + using losses from the model. + + This method directly updates the reweighting without synchronizing + between workers. It is called by update_with_local_losses from all + ranks with identical arguments. Thus, it should have deterministic + behavior to maintain state across workers. + + :param ts: a list of int timesteps. + :param losses: a list of float losses, one per timestep. + """ + + +class LossSecondMomentResampler(LossAwareSampler): + def __init__(self, diffusion, history_per_term=10, uniform_prob=0.001): + self.diffusion = diffusion + self.history_per_term = history_per_term + self.uniform_prob = uniform_prob + self._loss_history = np.zeros([diffusion.num_timesteps, history_per_term], dtype=np.float64) + self._loss_counts = np.zeros([diffusion.num_timesteps], dtype=int) + + def weights(self): + if not self._warmed_up(): + return np.ones([self.diffusion.num_timesteps], dtype=np.float64) + weights = np.sqrt(np.mean(self._loss_history**2, axis=-1)) + weights /= np.sum(weights) + weights *= 1 - self.uniform_prob + weights += self.uniform_prob / len(weights) + return weights + + def update_with_all_losses(self, ts, losses): + for t, loss in zip(ts, losses): + if self._loss_counts[t] == self.history_per_term: + # Shift out the oldest loss term. + self._loss_history[t, :-1] = self._loss_history[t, 1:] + self._loss_history[t, -1] = loss + else: + self._loss_history[t, self._loss_counts[t]] = loss + self._loss_counts[t] += 1 + + def _warmed_up(self): + return (self._loss_counts == self.history_per_term).all() diff --git a/scandl_module/original_scandl/utils/__init__.py b/scandl_module/original_scandl/utils/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/scandl_module/original_scandl/utils/__pycache__/__init__.cpython-313.pyc b/scandl_module/original_scandl/utils/__pycache__/__init__.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..d765ae7cc5e00d0da02545d6c06cdd1893655fd6 Binary files /dev/null and b/scandl_module/original_scandl/utils/__pycache__/__init__.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/utils/__pycache__/dist_util.cpython-313.pyc b/scandl_module/original_scandl/utils/__pycache__/dist_util.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..581389d841491116c3a59034dd69a6dfb2bea933 Binary files /dev/null and b/scandl_module/original_scandl/utils/__pycache__/dist_util.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/utils/__pycache__/fp16_util.cpython-313.pyc b/scandl_module/original_scandl/utils/__pycache__/fp16_util.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..740f928708fd5b4d5e2e1e0564a53c01ae9bb0a8 Binary files /dev/null and b/scandl_module/original_scandl/utils/__pycache__/fp16_util.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/utils/__pycache__/logger.cpython-313.pyc b/scandl_module/original_scandl/utils/__pycache__/logger.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..83cfe3717e8db80427a7c7003c4cbcb1ae559837 Binary files /dev/null and b/scandl_module/original_scandl/utils/__pycache__/logger.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/utils/__pycache__/nn.cpython-313.pyc b/scandl_module/original_scandl/utils/__pycache__/nn.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..09ceac248afde65e5e72a5960ec81fdb8aca6bdb Binary files /dev/null and b/scandl_module/original_scandl/utils/__pycache__/nn.cpython-313.pyc differ diff --git a/scandl_module/original_scandl/utils/dist_util.py b/scandl_module/original_scandl/utils/dist_util.py new file mode 100644 index 0000000000000000000000000000000000000000..c46a4ef5b7ee273df1f3702fb3a6f87c742e17dc --- /dev/null +++ b/scandl_module/original_scandl/utils/dist_util.py @@ -0,0 +1,85 @@ +""" +Helpers for distributed training. +""" + +import io +import os +import socket + +import blobfile as bf + +import torch as th +import torch.distributed as dist + +# Change this to reflect your cluster layout. + + +def setup_dist(): + """ + Setup a distributed process group. + """ + if dist.is_initialized(): + return + + # UNCOMMENT IF ON LINUX/MAC + + # backend = "gloo" if not th.cuda.is_available() else "nccl" + + backend = "gloo" + + if backend == "gloo": + hostname = "localhost" + else: + hostname = socket.gethostbyname(socket.getfqdn()) + + if os.environ.get("LOCAL_RANK") is None: + os.environ["MASTER_ADDR"] = hostname + os.environ["RANK"] = str(0) + os.environ["WORLD_SIZE"] = str(1) + port = _find_free_port() + os.environ["MASTER_PORT"] = str(port) + os.environ["LOCAL_RANK"] = str(0) + + dist.init_process_group(backend=backend, init_method="env://") + + if th.cuda.is_available(): # This clears remaining caches in GPU 0 + th.cuda.set_device(dev()) + th.cuda.empty_cache() + + +def dev(): + """ + Get the device to use for torch.distributed. + """ + if th.cuda.is_available(): + return th.device(f"cuda:{os.environ['LOCAL_RANK']}") + return th.device("cpu") + + +def load_state_dict(path, **kwargs): + """ + Load a PyTorch file. + """ + # if int(os.environ['LOCAL_RANK']) == 0: + with bf.BlobFile(path, "rb") as f: + data = f.read() + return th.load(io.BytesIO(data), **kwargs) + + +def sync_params(params): + """ + Synchronize a sequence of Tensors across ranks from rank 0. + """ + for p in params: + with th.no_grad(): + dist.broadcast(p, 0) + + +def _find_free_port(): + try: + s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + s.bind(("", 0)) + s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + return s.getsockname()[1] + finally: + s.close() diff --git a/scandl_module/original_scandl/utils/fp16_util.py b/scandl_module/original_scandl/utils/fp16_util.py new file mode 100644 index 0000000000000000000000000000000000000000..57adf5049b9f7dda1b20f455e763163a7d160a5d --- /dev/null +++ b/scandl_module/original_scandl/utils/fp16_util.py @@ -0,0 +1,74 @@ +""" +Helpers to train with 16-bit precision. +""" + +import torch.nn as nn +from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors + + +def convert_module_to_f16(l): + """ + Convert primitive modules to float16. + """ + if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Conv3d)): + l.weight.data = l.weight.data.half() + l.bias.data = l.bias.data.half() + + +def convert_module_to_f32(l): + """ + Convert primitive modules to float32, undoing convert_module_to_f16(). + """ + if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Conv3d)): + l.weight.data = l.weight.data.float() + l.bias.data = l.bias.data.float() + + +def make_master_params(model_params): + """ + Copy model parameters into a (differently-shaped) list of full-precision + parameters. + """ + master_params = _flatten_dense_tensors([param.detach().float() for param in model_params]) + master_params = nn.Parameter(master_params) + master_params.requires_grad = True + return [master_params] + + +def model_grads_to_master_grads(model_params, master_params): + """ + Copy the gradients from the model parameters into the master parameters + from make_master_params(). + """ + master_params[0].grad = _flatten_dense_tensors( + [param.grad.data.detach().float() for param in model_params] + ) + + +def master_params_to_model_params(model_params, master_params): + """ + Copy the master parameter data back into the model parameters. + """ + # Without copying to a list, if a generator is passed, this will + # silently not copy any parameters. + model_params = list(model_params) + + for param, master_param in zip( + model_params, unflatten_master_params(model_params, master_params) + ): + param.detach().copy_(master_param) + + +def unflatten_master_params(model_params, master_params): + """ + Unflatten the master parameters to look like model_params. + """ + return _unflatten_dense_tensors(master_params[0].detach(), model_params) + + +def zero_grad(model_params): + for param in model_params: + # Taken from https://pytorch.org/docs/stable/_modules/torch/optim/optimizer.html#Optimizer.add_param_group + if param.grad is not None: + param.grad.detach_() + param.grad.zero_() diff --git a/scandl_module/original_scandl/utils/logger.py b/scandl_module/original_scandl/utils/logger.py new file mode 100644 index 0000000000000000000000000000000000000000..c22df648b5ba45325b7f85d8a8dd021cf76613e0 --- /dev/null +++ b/scandl_module/original_scandl/utils/logger.py @@ -0,0 +1,492 @@ +""" +Logger copied from OpenAI baselines to avoid extra RL-based dependencies: +https://github.com/openai/baselines/blob/ea25b9e8b234e6ee1bca43083f8f3cf974143998/baselines/logger.py +""" + +import os +import sys +import shutil +import os.path as osp +import json +import time +import datetime +import tempfile +import warnings +from collections import defaultdict +from contextlib import contextmanager + +# import wandb + +DEBUG = 10 +INFO = 20 +WARN = 30 +ERROR = 40 + +DISABLED = 50 + + +class KVWriter(object): + def writekvs(self, kvs): + raise NotImplementedError + + +class SeqWriter(object): + def writeseq(self, seq): + raise NotImplementedError + + +class HumanOutputFormat(KVWriter, SeqWriter): + def __init__(self, filename_or_file): + if isinstance(filename_or_file, str): + self.file = open(filename_or_file, "wt") + self.own_file = True + else: + assert hasattr(filename_or_file, "read"), ( + "expected file or str, got %s" % filename_or_file + ) + self.file = filename_or_file + self.own_file = False + + def writekvs(self, kvs): + # Create strings for printing + key2str = {} + for key, val in sorted(kvs.items()): + if hasattr(val, "__float__"): + valstr = "%-8.3g" % val + else: + valstr = str(val) + key2str[self._truncate(key)] = self._truncate(valstr) + + # Find max widths + if len(key2str) == 0: + print("WARNING: tried to write empty key-value dict") + return + else: + keywidth = max(map(len, key2str.keys())) + valwidth = max(map(len, key2str.values())) + + # Write out the data + dashes = "-" * (keywidth + valwidth + 7) + lines = [dashes] + for key, val in sorted(key2str.items(), key=lambda kv: kv[0].lower()): + lines.append( + "| %s%s | %s%s |" + % (key, " " * (keywidth - len(key)), val, " " * (valwidth - len(val))) + ) + lines.append(dashes) + self.file.write("\n".join(lines) + "\n") + + # Flush the output to the file + self.file.flush() + + def _truncate(self, s): + maxlen = 30 + return s[: maxlen - 3] + "..." if len(s) > maxlen else s + + def writeseq(self, seq): + seq = list(seq) + for i, elem in enumerate(seq): + self.file.write(elem) + if i < len(seq) - 1: # add space unless this is the last one + self.file.write(" ") + self.file.write("\n") + self.file.flush() + + def close(self): + if self.own_file: + self.file.close() + + +class JSONOutputFormat(KVWriter): + def __init__(self, filename): + self.file = open(filename, "wt") + + def writekvs(self, kvs): + for k, v in sorted(kvs.items()): + if hasattr(v, "dtype"): + kvs[k] = float(v) + self.file.write(json.dumps(kvs) + "\n") + self.file.flush() + + def close(self): + self.file.close() + + +class CSVOutputFormat(KVWriter): + def __init__(self, filename): + self.file = open(filename, "w+t") + self.keys = [] + self.sep = "," + + def writekvs(self, kvs): + # Add our current row to the history + extra_keys = list(kvs.keys() - self.keys) + extra_keys.sort() + if extra_keys: + self.keys.extend(extra_keys) + self.file.seek(0) + lines = self.file.readlines() + self.file.seek(0) + for i, k in enumerate(self.keys): + if i > 0: + self.file.write(",") + self.file.write(k) + self.file.write("\n") + for line in lines[1:]: + self.file.write(line[:-1]) + self.file.write(self.sep * len(extra_keys)) + self.file.write("\n") + for i, k in enumerate(self.keys): + if i > 0: + self.file.write(",") + v = kvs.get(k) + if v is not None: + self.file.write(str(v)) + self.file.write("\n") + self.file.flush() + + def close(self): + self.file.close() + + +class TensorBoardOutputFormat(KVWriter): + """ + Dumps key/value pairs into TensorBoard's numeric format. + """ + + def __init__(self, dir): + os.makedirs(dir, exist_ok=True) + self.dir = dir + self.step = 1 + prefix = "events" + path = osp.join(osp.abspath(dir), prefix) + import tensorflow as tf + from tensorflow.python import pywrap_tensorflow + from tensorflow.core.util import event_pb2 + from tensorflow.python.util import compat + + self.tf = tf + self.event_pb2 = event_pb2 + self.pywrap_tensorflow = pywrap_tensorflow + self.writer = pywrap_tensorflow.EventsWriter(compat.as_bytes(path)) + + def writekvs(self, kvs): + def summary_val(k, v): + kwargs = {"tag": k, "simple_value": float(v)} + return self.tf.Summary.Value(**kwargs) + + summary = self.tf.Summary(value=[summary_val(k, v) for k, v in kvs.items()]) + event = self.event_pb2.Event(wall_time=time.time(), summary=summary) + event.step = self.step # is there any reason why you'd want to specify the step? + self.writer.WriteEvent(event) + self.writer.Flush() + self.step += 1 + + def close(self): + if self.writer: + self.writer.Close() + self.writer = None + + +def make_output_format(format, ev_dir, log_suffix=""): + os.makedirs(ev_dir, exist_ok=True) + if format == "stdout": + return HumanOutputFormat(sys.stdout) + elif format == "log": + return HumanOutputFormat(osp.join(ev_dir, "log%s.txt" % log_suffix)) + elif format == "json": + return JSONOutputFormat(osp.join(ev_dir, "progress%s.json" % log_suffix)) + elif format == "csv": + return CSVOutputFormat(osp.join(ev_dir, "progress%s.csv" % log_suffix)) + elif format == "tensorboard": + return TensorBoardOutputFormat(osp.join(ev_dir, "tb%s" % log_suffix)) + else: + raise ValueError("Unknown format specified: %s" % (format,)) + + +# ================================================================ +# API +# ================================================================ + + +def logkv(key, val): + """ + Log a value of some diagnostic + Call this once for each diagnostic quantity, each iteration + If called many times, last value will be used. + """ + get_current().logkv(key, val) + + +def logkv_mean(key, val): + """ + The same as logkv(), but if called many times, values averaged. + """ + get_current().logkv_mean(key, val) + + +def logkvs(d): + """ + Log a dictionary of key-value pairs + """ + for k, v in d.items(): + logkv(k, v) + + +def dumpkvs(): + """ + Write all of the diagnostics from the current iteration + """ + return get_current().dumpkvs() + + +def getkvs(): + return get_current().name2val + + +def log(*args, level=INFO): + """ + Write the sequence of args, with no separators, to the console and output files (if you've configured an output file). + """ + get_current().log(*args, level=level) + + +def debug(*args): + log(*args, level=DEBUG) + + +def info(*args): + log(*args, level=INFO) + + +def warn(*args): + log(*args, level=WARN) + + +def error(*args): + log(*args, level=ERROR) + + +def set_level(level): + """ + Set logging threshold on current logger. + """ + get_current().set_level(level) + + +def set_comm(comm): + get_current().set_comm(comm) + + +def get_dir(): + """ + Get directory that log files are being written to. + will be None if there is no output directory (i.e., if you didn't call start) + """ + return get_current().get_dir() + + +record_tabular = logkv +dump_tabular = dumpkvs + + +@contextmanager +def profile_kv(scopename): + logkey = "wait_" + scopename + tstart = time.time() + try: + yield + finally: + get_current().name2val[logkey] += time.time() - tstart + + +def profile(n): + """ + Usage: + @profile("my_func") + def my_func(): code + """ + + def decorator_with_name(func): + def func_wrapper(*args, **kwargs): + with profile_kv(n): + return func(*args, **kwargs) + + return func_wrapper + + return decorator_with_name + + +# ================================================================ +# Backend +# ================================================================ + + +def get_current(): + if Logger.CURRENT is None: + _configure_default_logger() + + return Logger.CURRENT + + +class Logger(object): + DEFAULT = None # A logger with no output files. (See right below class definition) + # So that you can still log to the terminal without setting up any output files + CURRENT = None # Current logger being used by the free functions above + + def __init__(self, dir, output_formats, comm=None): + self.name2val = defaultdict(float) # values this iteration + self.name2cnt = defaultdict(int) + self.level = INFO + self.dir = dir + self.output_formats = output_formats + self.comm = comm + + # Logging API, forwarded + # ---------------------------------------- + def logkv(self, key, val): + self.name2val[key] = val + + def logkv_mean(self, key, val): + oldval, cnt = self.name2val[key], self.name2cnt[key] + self.name2val[key] = oldval * cnt / (cnt + 1) + val / (cnt + 1) + self.name2cnt[key] = cnt + 1 + + def dumpkvs(self, prefix=None): + if self.comm is None: + d = self.name2val + else: + d = mpi_weighted_mean( + self.comm, + {name: (val, self.name2cnt.get(name, 1)) for (name, val) in self.name2val.items()}, + ) + if self.comm.rank != 0: + d["dummy"] = 1 # so we don't get a warning about empty dict + # LISA + out = d.copy() # Return the dict for unit testing purposes + if int(os.environ["LOCAL_RANK"]) == 0: + # wandb.log({**d}) + for fmt in self.output_formats: + if isinstance(fmt, KVWriter): + fmt.writekvs(d) + self.name2val.clear() + self.name2cnt.clear() + return out + + def log(self, *args, level=INFO): + if self.level <= level: + self._do_log(args) + + # Configuration + # ---------------------------------------- + def set_level(self, level): + self.level = level + + def set_comm(self, comm): + self.comm = comm + + def get_dir(self): + return self.dir + + def close(self): + for fmt in self.output_formats: + fmt.close() + + # Misc + # ---------------------------------------- + def _do_log(self, args): + for fmt in self.output_formats: + if isinstance(fmt, SeqWriter): + fmt.writeseq(map(str, args)) + + +def get_rank_without_mpi_import(): + # check environment variables here instead of importing mpi4py + # to avoid calling MPI_Init() when this module is imported + for varname in ["PMI_RANK", "OMPI_COMM_WORLD_RANK"]: + if varname in os.environ: + return int(os.environ[varname]) + return 0 + + +def mpi_weighted_mean(comm, local_name2valcount): + """ + Copied from: https://github.com/openai/baselines/blob/ea25b9e8b234e6ee1bca43083f8f3cf974143998/baselines/common/mpi_util.py#L110 + Perform a weighted average over dicts that are each on a different node + Input: local_name2valcount: dict mapping key -> (value, count) + Returns: key -> mean + """ + all_name2valcount = comm.gather(local_name2valcount) + if comm.rank == 0: + name2sum = defaultdict(float) + name2count = defaultdict(float) + for n2vc in all_name2valcount: + for name, (val, count) in n2vc.items(): + try: + val = float(val) + except ValueError: + if comm.rank == 0: + warnings.warn( + "WARNING: tried to compute mean on non-float {}={}".format(name, val) + ) + else: + name2sum[name] += val * count + name2count[name] += count + return {name: name2sum[name] / name2count[name] for name in name2sum} + else: + return {} + + +def configure(dir=None, format_strs=None, comm=None, log_suffix=""): + """ + If comm is provided, average all numerical stats across that comm + """ + if dir is None: + dir = os.getenv("OPENAI_LOGDIR") + if dir is None: + dir = osp.join( + tempfile.gettempdir(), + datetime.datetime.now().strftime("openai-%Y-%m-%d-%H-%M-%S-%f"), + ) + assert isinstance(dir, str) + dir = os.path.expanduser(dir) + os.makedirs(os.path.expanduser(dir), exist_ok=True) + + rank = get_rank_without_mpi_import() + if rank > 0: + log_suffix = log_suffix + "-rank%03i" % rank + + if format_strs is None: + if rank == 0: + format_strs = os.getenv("OPENAI_LOG_FORMAT", "stdout,log,csv").split(",") + else: + format_strs = os.getenv("OPENAI_LOG_FORMAT_MPI", "log").split(",") + format_strs = filter(None, format_strs) + output_formats = [make_output_format(f, dir, log_suffix) for f in format_strs] + + Logger.CURRENT = Logger(dir=dir, output_formats=output_formats, comm=comm) + if output_formats: + log("Logging to %s" % dir) + + +def _configure_default_logger(): + configure() + Logger.DEFAULT = Logger.CURRENT + + +def reset(): + if Logger.CURRENT is not Logger.DEFAULT: + Logger.CURRENT.close() + Logger.CURRENT = Logger.DEFAULT + log("Reset logger") + + +@contextmanager +def scoped_configure(dir=None, format_strs=None, comm=None): + prevlogger = Logger.CURRENT + configure(dir=dir, format_strs=format_strs, comm=comm) + try: + yield + finally: + Logger.CURRENT.close() + Logger.CURRENT = prevlogger diff --git a/scandl_module/original_scandl/utils/losses.py b/scandl_module/original_scandl/utils/losses.py new file mode 100644 index 0000000000000000000000000000000000000000..e54b4492d4b666a288080d11e264edddbc505de2 --- /dev/null +++ b/scandl_module/original_scandl/utils/losses.py @@ -0,0 +1,120 @@ +""" +Helpers for various likelihood-based losses. These are ported from the original +Ho et al. diffusion models codebase: +https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/utils.py +""" + +import numpy as np + +import torch as th + + +def normal_kl(mean1, logvar1, mean2, logvar2): + """ + Compute the KL divergence between two gaussians. + + Shapes are automatically broadcasted, so batches can be compared to + scalars, among other use cases. + """ + tensor = None + for obj in (mean1, logvar1, mean2, logvar2): + if isinstance(obj, th.Tensor): + tensor = obj + break + assert tensor is not None, "at least one argument must be a Tensor" + + # Force variances to be Tensors. Broadcasting helps convert scalars to + # Tensors, but it does not work for th.exp(). + logvar1, logvar2 = [ + x if isinstance(x, th.Tensor) else th.tensor(x).to(tensor) for x in (logvar1, logvar2) + ] + + # print(logvar2.shape) + # temp1 = 0.5 * (-1.0 + logvar2 - logvar1 + th.exp(logvar1 - logvar2)) + # print(f'const = {temp1.mean()}, coef={(th.exp(-logvar2) * 0.5).mean()}, mse={((mean1 - mean2) ** 2).mean().item()}') + + return 0.5 * ( + -1.0 + + logvar2 + - logvar1 + + th.exp(logvar1 - logvar2) + + ((mean1 - mean2) ** 2) * th.exp(-logvar2) + ) + + +def approx_standard_normal_cdf(x): + """ + A fast approximation of the cumulative distribution function of the + standard normal. + """ + return 0.5 * (1.0 + th.tanh(np.sqrt(2.0 / np.pi) * (x + 0.044715 * th.pow(x, 3)))) + + +def discretized_gaussian_log_likelihood(x, *, means, log_scales): + """ + Compute the log-likelihood of a Gaussian distribution discretizing to a + given image. + + :param x: the target images. It is assumed that this was uint8 values, + rescaled to the range [-1, 1]. + :param means: the Gaussian mean Tensor. + :param log_scales: the Gaussian log stddev Tensor. + :return: a tensor like x of log probabilities (in nats). + """ + assert x.shape == means.shape == log_scales.shape + centered_x = x - means + inv_stdv = th.exp(-log_scales) + plus_in = inv_stdv * (centered_x + 1.0 / 255.0) + cdf_plus = approx_standard_normal_cdf(plus_in) + min_in = inv_stdv * (centered_x - 1.0 / 255.0) + cdf_min = approx_standard_normal_cdf(min_in) + log_cdf_plus = th.log(cdf_plus.clamp(min=1e-12)) + log_one_minus_cdf_min = th.log((1.0 - cdf_min).clamp(min=1e-12)) + cdf_delta = cdf_plus - cdf_min + log_probs = th.where( + x < -0.999, + log_cdf_plus, + th.where(x > 0.999, log_one_minus_cdf_min, th.log(cdf_delta.clamp(min=1e-12))), + ) + assert log_probs.shape == x.shape + return log_probs + + +def gaussian_density(x, *, means, log_scales): + from torch.distributions import Normal + + normal_dist = Normal(means, log_scales.exp()) + logp = normal_dist.log_prob(x) + return logp + + +def discretized_text_log_likelihood(x, *, means, log_scales): + """ + Compute the log-likelihood of a Gaussian distribution discretizing to a + given image. + + :param x: the target images. It is assumed that this was uint8 values, + rescaled to the range [-1, 1]. + :param means: the Gaussian mean Tensor. + :param log_scales: the Gaussian log stddev Tensor. + :return: a tensor like x of log probabilities (in nats). + """ + print(x.shape, means.shape) + # assert x.shape == means.shape == log_scales.shape + print(x, means) + centered_x = x - means + inv_stdv = th.exp(-log_scales) + plus_in = inv_stdv * (centered_x + 1.0 / 255.0) + cdf_plus = approx_standard_normal_cdf(plus_in) + min_in = inv_stdv * (centered_x - 1.0 / 255.0) + cdf_min = approx_standard_normal_cdf(min_in) + log_cdf_plus = th.log(cdf_plus.clamp(min=1e-12)) + log_one_minus_cdf_min = th.log((1.0 - cdf_min).clamp(min=1e-12)) + cdf_delta = cdf_plus - cdf_min + log_probs = th.where( + x < -0.999, + log_cdf_plus, + th.where(x > 0.999, log_one_minus_cdf_min, th.log(cdf_delta.clamp(min=1e-12))), + ) + assert log_probs.shape == x.shape + return log_probs diff --git a/scandl_module/original_scandl/utils/nn.py b/scandl_module/original_scandl/utils/nn.py new file mode 100644 index 0000000000000000000000000000000000000000..9a46a9566a992128a7a5052e1ba3e95fb2e7d5bb --- /dev/null +++ b/scandl_module/original_scandl/utils/nn.py @@ -0,0 +1,174 @@ +""" +Various utilities for neural networks. +""" + +import math + +import torch +import torch as th +import torch.nn as nn + + +# PyTorch 1.7 has SiLU, but we support PyTorch 1.5. +class SiLU(nn.Module): + def forward(self, x): + return x * th.sigmoid(x) + + +class GroupNorm32(nn.GroupNorm): + def forward(self, x): + return super().forward(x.float()).type(x.dtype) + + +def linear(*args, **kwargs): + """ + Create a linear module. + """ + return nn.Linear(*args, **kwargs) + + +def avg_pool_nd(dims, *args, **kwargs): + """ + Create a 1D, 2D, or 3D average pooling module. + """ + if dims == 1: + return nn.AvgPool1d(*args, **kwargs) + elif dims == 2: + return nn.AvgPool2d(*args, **kwargs) + elif dims == 3: + return nn.AvgPool3d(*args, **kwargs) + raise ValueError(f"unsupported dimensions: {dims}") + + +def update_ema(target_params, source_params, rate=0.99): + """ + Update target parameters to be closer to those of source parameters using + an exponential moving average. + + :param target_params: the target parameter sequence. + :param source_params: the source parameter sequence. + :param rate: the EMA rate (closer to 1 means slower). + """ + for targ, src in zip(target_params, source_params): + targ.detach().mul_(rate).add_(src, alpha=1 - rate) + + +def zero_module(module): + """ + Zero out the parameters of a module and return it. + """ + for p in module.parameters(): + p.detach().zero_() + return module + + +def scale_module(module, scale): + """ + Scale the parameters of a module and return it. + """ + for p in module.parameters(): + p.detach().mul_(scale) + return module + + +def mean_flat(tensor): + """ + Take the mean over all non-batch dimensions. + """ + return tensor.mean(dim=list(range(1, len(tensor.shape)))) + + +def normalization(channels): + """ + Make a standard normalization layer. + + :param channels: number of input channels. + :return: an nn.Module for normalization. + """ + return GroupNorm32(32, channels) + + +def timestep_embedding(timesteps, dim, max_period=10000): + """ + Create sinusoidal timestep embeddings. + + :param timesteps: a 1-D Tensor of N indices, one per batch element. + These may be fractional. + :param dim: the dimension of the output. + :param max_period: controls the minimum frequency of the embeddings. + :return: an [N x dim] Tensor of positional embeddings. + """ + half = dim // 2 + freqs = th.exp( + -math.log(max_period) * th.arange(start=0, end=half, dtype=th.float32) / half + ).to(device=timesteps.device) + args = timesteps[:, None].float() * freqs[None] + embedding = th.cat([th.cos(args), th.sin(args)], dim=-1) + if dim % 2: + embedding = th.cat([embedding, th.zeros_like(embedding[:, :1])], dim=-1) + + return embedding + + +def concatenate_sn_sp(sn_repr_emb, sp_repr_emb, sn_repr_len): + # Create an empty tensor to store the concatenated embeddings + concat_emb = torch.empty_like(sn_repr_emb) + + # Iterate over the batch size + for i in range(sn_repr_emb.size(0)): + # Get the true sequence length for the sn embedding + seq_len = sn_repr_len[i] + # get the true sequence length for the sp embedding + sp_idx = sp_repr_emb.size(1) - seq_len + # Cut off the extra dimensions in sn_repr_emb and sp_repr_emb + sn_emb = sn_repr_emb[i, :seq_len] + sp_emb = sp_repr_emb[i, :sp_idx] + + # Concatenate the embeddings along the sequence length dimension + concat_emb[i] = torch.cat([sn_emb, sp_emb], dim=0) + + return concat_emb + + +def split_into_sn_and_sp(model_output, sn_repr_len): + # create an empty tensor to store the split sn_repr + model_output_sn = torch.empty_like(model_output) + model_output_sp = torch.empty_like(model_output) + + microbatch_size, seq_len, emb_dim = model_output.shape + + model_output_sn_mask = torch.empty((microbatch_size, seq_len)) + model_output_sp_mask = torch.empty((microbatch_size, seq_len)) + + # iterate over the batc size + for i in range(microbatch_size): + # the start index of the sp output is the length of the sn + orig_sn_len = sn_repr_len[i] + + # the SP + # split the sp_repr of the current instance off from the model output + sp_repr_out = model_output[i, orig_sn_len:] + # TODO change this to more sensible padding + # pad the sp representation model output until it is again args.seq_len long + sp_padding = sp_repr_out[-1].repeat(orig_sn_len, 1) + # concatenate the embeddings along the + model_output_sp[i] = torch.cat([sp_repr_out, sp_padding], dim=0) + + # the SN + # split the sn_repr of the current instance off from the model output + sn_repr_out = model_output[i, :orig_sn_len] + sn_padding = sn_repr_out[-1].repeat(seq_len - orig_sn_len, 1) + model_output_sn[i] = torch.cat([sn_repr_out, sn_padding], dim=0) + + # the masks: they only mask the additional padding added now, not the original padding added after the sp + sp_mask = torch.ones(seq_len).to(model_output.device) + sp_mask[-orig_sn_len:] = 0 + model_output_sp_mask[i] = sp_mask + sn_mask = torch.ones(seq_len).to(model_output.device) + sn_mask[orig_sn_len:] = 0 + model_output_sn_mask[i] = sn_mask + + model_output_sn_mask = model_output_sn_mask.to(model_output.device) + model_output_sp_mask = model_output_sp_mask.to(model_output.device) + + return model_output_sn, model_output_sp, model_output_sn_mask, model_output_sp_mask diff --git a/scandl_module/scripts/__pycache__/sp_basic_utils.cpython-313.pyc b/scandl_module/scripts/__pycache__/sp_basic_utils.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..66a61673977c8e11d63879ab0fef1f520b0aad94 Binary files /dev/null and b/scandl_module/scripts/__pycache__/sp_basic_utils.cpython-313.pyc differ diff --git a/scandl_module/scripts/__pycache__/sp_load_celer_zuco.cpython-313.pyc b/scandl_module/scripts/__pycache__/sp_load_celer_zuco.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..09d7fcbb536a6ddb22de0908041ac2c5094e4734 Binary files /dev/null and b/scandl_module/scripts/__pycache__/sp_load_celer_zuco.cpython-313.pyc differ diff --git a/scandl_module/scripts/__pycache__/sp_train_util.cpython-313.pyc b/scandl_module/scripts/__pycache__/sp_train_util.cpython-313.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a8f308e4a4de3d6a990d4851963ef6e95101a184 Binary files /dev/null and b/scandl_module/scripts/__pycache__/sp_train_util.cpython-313.pyc differ diff --git a/scandl_module/scripts/sp_basic_utils.py b/scandl_module/scripts/sp_basic_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..f68926341bb031e24289262724a510129f46d1e9 --- /dev/null +++ b/scandl_module/scripts/sp_basic_utils.py @@ -0,0 +1,110 @@ +import argparse +import json +import sys + + +from ScanDL2.scandl_module.original_scandl import sp_gaussian_diffusion as gd +from ScanDL2.scandl_module.original_scandl.sp_gaussian_diffusion import ( + SpacedDiffusion, + space_timesteps, +) +from ScanDL2.scandl_module.original_scandl.sp_transformer_model import TransformerNetModel + +sys.path.append("./") +sys.path.append("../") + + +def load_defaults_config(config_path: str): + """ + Load defaults for training args. + """ + with open(config_path, "r") as f: + return json.load(f) + + +def create_model_and_diffusion( + hidden_t_dim, + hidden_dim, + vocab_size, + config_name, + use_plm_init, + dropout, + num_transformer_layers, + num_transformer_heads, + mask_padding, + diffusion_steps, + noise_schedule, + learn_sigma, + timestep_respacing, + predict_xstart, + rescale_timesteps, + sigma_small, + rescale_learned_sigmas, + use_kl, + one_noise_step, + nll_in_loss, + notes, + **kwargs, +): + model = TransformerNetModel( + input_dims=hidden_dim, + output_dims=(hidden_dim if not learn_sigma else hidden_dim * 2), + hidden_t_dim=hidden_t_dim, + num_transformer_layers=num_transformer_layers, + num_transformer_heads=num_transformer_heads, + one_noise_step=one_noise_step, + mask_padding=mask_padding, + dropout=dropout, + config_name=config_name, + vocab_size=vocab_size, + init_pretrained=use_plm_init, + ) + + betas = gd.get_named_beta_schedule(noise_schedule, diffusion_steps) + + if not timestep_respacing: + timestep_respacing = [diffusion_steps] + + diffusion = SpacedDiffusion( + use_timesteps=space_timesteps(diffusion_steps, timestep_respacing), + betas=betas, + rescale_timesteps=rescale_timesteps, + predict_xstart=predict_xstart, + learn_sigmas=learn_sigma, + sigma_small=sigma_small, + use_kl=use_kl, + one_noise_step=one_noise_step, + nll_in_loss=nll_in_loss, + mask_padding=mask_padding, + rescale_learned_sigmas=rescale_learned_sigmas, + ) + + return model, diffusion + + +def add_dict_to_argparser(parser, default_dict): + for k, v in default_dict.items(): + v_type = type(v) + if v is None: + v_type = str + elif isinstance(v, bool): + v_type = str2bool + parser.add_argument(f"--{k}", default=v, type=v_type) + + +def args_to_dict(args, keys): + return {k: getattr(args, k) for k in keys} + + +def str2bool(v): + """ + https://stackoverflow.com/questions/15008758/parsing-boolean-values-with-argparse + """ + if isinstance(v, bool): + return v + if v.lower() in ("yes", "true", "t", "y", "1"): + return True + elif v.lower() in ("no", "false", "f", "n", "0"): + return False + else: + raise argparse.ArgumentTypeError("boolean value expected") diff --git a/scandl_module/scripts/sp_load_celer_zuco.py b/scandl_module/scripts/sp_load_celer_zuco.py new file mode 100644 index 0000000000000000000000000000000000000000..9f2bbb3c7488fd15dafe9a25916c9c1518964497 --- /dev/null +++ b/scandl_module/scripts/sp_load_celer_zuco.py @@ -0,0 +1,1441 @@ +import pandas as pd +import numpy as np +from tqdm import tqdm +import os +import random +import torch +from torch.utils.data import Dataset, DataLoader +from datasets import Dataset as Dataset2 +import datasets +from sklearn.model_selection import ( + train_test_split, + GroupShuffleSplit, + KFold, + GroupKFold, + StratifiedKFold, +) +from typing import Optional, List, Tuple, Union, Any, Dict +import sys + +sys.path.append("./") +sys.path.append("../") +sys.path.append("../../") + +from ScanDL2.CONSTANTS import PATH_TO_IA, PATH_TO_FIX, SUB_METADATA_PATH, path_to_zuco + + +def load_celer(): + path_to_fix = PATH_TO_FIX + path_to_ia = PATH_TO_IA + eyemovement_df = pd.read_csv(path_to_fix, delimiter="\t", low_memory=False) + eyemovement_df["CURRENT_FIX_INTEREST_AREA_LABEL"] = ( + eyemovement_df.CURRENT_FIX_INTEREST_AREA_LABEL.replace("\t(.*)", "", regex=True) + ) + word_info_df = pd.read_csv(path_to_ia, delimiter="\t") + word_info_df["IA_LABEL"] = word_info_df.IA_LABEL.replace("\t(.*)", "", regex=True) + return word_info_df, eyemovement_df + + +def load_celer_speakers(only_native_speakers: bool = True): + sub_metadata_path = SUB_METADATA_PATH + sub_info = pd.read_csv(sub_metadata_path, delimiter="\t") + if only_native_speakers: + readers_list = sub_info[sub_info.L1 == "English"].List.values + else: + readers_list = sub_info.List.values + return readers_list.tolist() + + +def compute_word_length(arr): + # length of a punctuation is 0, plus an epsilon to avoid division output inf + arr = arr.astype("float64") + arr[arr == 0] = 1 / (0 + 0.5) + arr[arr != 0] = 1 / (arr[arr != 0]) + return arr + + +def compute_word_frequency(arr): + arr[arr == np.inf] = np.nan + arr[arr != np.inf] = np.log10(arr[arr != np.inf]) + return arr + + +def _collate_instance_helper( + instance, + pad_token_id, + padding_steps, # how many steps/dims to pad, +): + padding_list = [pad_token_id] * padding_steps + result = instance + padding_list + return result + + +def _collate_batch_helper( + examples, # List of Lists of input IDs + pad_token_id, + max_length, + return_mask=False, +): + result = torch.full([len(examples), max_length], pad_token_id, dtype=torch.int64).tolist() + mask_ = torch.full([len(examples), max_length], pad_token_id, dtype=torch.int64).tolist() + for i, example in enumerate(examples): + curr_len = min(len(example), max_length) + result[i][:curr_len] = example[:curr_len] + mask_[i][:curr_len] = [1] * curr_len + if return_mask: + return result, mask_ + return result + + +def _dummy_pad_words( + examples, + pad_token_id, + max_length, +): + padded_examples = list() + for instance in examples: + padded_examples.append(instance + (max_length - len(instance)) * [pad_token_id]) + return padded_examples + + +def infinite_loader(data_loader): + while True: + yield from data_loader + + +def celer_zuco_dataset_and_loader( + data, + data_args, + split: str, + deterministic=False, + loop=True, +): + + dataset = CelerZucoDataset( + dataset=data, + data_args=data_args, + split=split, + ) + data_loader = DataLoader( + dataset, + batch_size=data_args.batch_size, + shuffle=not deterministic, # ?? + num_workers=0, + ) + if loop: + return infinite_loader(data_loader) + else: + return iter(data_loader) + + +def combined_split(data, reader_IDs, sn_IDs, test_size): + """Splits the data so that the test data contains both unseen readers and unseen sentences.""" + random.seed(77) + unique_reader_IDs = set(reader_IDs) + unique_sn_IDs = set(sn_IDs) + + # sample the sentence and reader IDs that go into the test set + unique_reader_IDs_test = random.sample( + unique_reader_IDs, int(test_size * len(unique_reader_IDs)) + ) + unique_sn_IDs_test = random.sample(unique_sn_IDs, int(test_size * len(unique_sn_IDs))) + + train_data, test_data = [], [] + train_reader_IDs, test_reader_IDs = [], [] + train_sn_IDs, test_sn_IDs = [], [] + + for i in range(len(data)): + if reader_IDs[i] in unique_reader_IDs_test and sn_IDs[i] in unique_sn_IDs_test: + test_data.append(data[i]) + test_reader_IDs.append(reader_IDs[i]) + test_sn_IDs.append(sn_IDs[i]) + elif reader_IDs[i] not in unique_reader_IDs_test and sn_IDs[i] not in unique_sn_IDs_test: + train_data.append(data[i]) + train_reader_IDs.append(reader_IDs[i]) + train_sn_IDs.append(sn_IDs[i]) + else: + continue + + return train_data, test_data, train_reader_IDs, test_reader_IDs, train_sn_IDs, test_sn_IDs + + +def load_zuco(task: str = None): # 'zuco11', 'zuco12' + dir = path_to_zuco + if task.startswith("zuco1"): + dir = dir + "zuco/" + elif task == "zuco21": + dir = dir + "zuco2/" + dir = os.path.join(dir, f"task{task[-1]}", "Matlab_files") + word_info_path = dir + "/Word_Infor.csv" + word_info_df = pd.read_csv(word_info_path, sep="\t") + scanpath_path = dir + "/scanpath.csv" + eyemovement_df = pd.read_csv(scanpath_path, sep="\t") + return word_info_df, eyemovement_df + + +def get_kfold_indices_scanpath( + splitting_IDs_dict: Dict[str, Union[int, str]], + n_splits: int = 5, +): + """Function to implement the 'random'/'scanpath' split, in which the test set contains both sentences and + readers that were seen during training. Unfortunately, it is not possible to control for the ratios of both + readers and sentences (as they are co-dependent), so I will make sure that the data is shuffled and at least + the readers are stratified.""" + + reader_IDs = splitting_IDs_dict["reader"] + sn_IDs = splitting_IDs_dict["sentence"] + + tuple_ids = [ + (idx, reader_ID, sn_ID) for idx, (reader_ID, sn_ID) in enumerate(zip(reader_IDs, sn_IDs)) + ] + + # list of indices of the sn-reader pairs of unique sentences + unique_sns_ids = [idx for (idx, reader_ID, sn_ID) in tuple_ids if not sn_ID.startswith("en")] + # unique_sns_indices = {sn_ID: idx for (idx, reader_ID, sn_ID) in tuple_ids if not sn_ID.startswith('en')} + universal_sn_ids = [idx for (idx, reader_ID, sn_ID) in tuple_ids if sn_ID.startswith("en")] + + # get the reader IDs for the unique and universal sns + unique_reader_ids = [ + reader_ID for (idx, reader_ID, sn_ID) in tuple_ids if not sn_ID.startswith("en") + ] + universal_reader_ids = [ + reader_ID for (idx, reader_ID, sn_ID) in tuple_ids if sn_ID.startswith("en") + ] + + # stratify both the readers for the universal sentences and the ones for the unique sentences + kfold_unique = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=77) + kfold_universal = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=77) + train_idx_unique, test_idx_unique = list(), list() + train_idx_universal, test_idx_universal = list(), list() + kfold_unique.get_n_splits(unique_sns_ids, groups=unique_reader_ids) + kfold_universal.get_n_splits(universal_sn_ids, groups=universal_reader_ids) + for train_idx, test_idx in kfold_unique.split( + X=unique_sns_ids, y=unique_reader_ids, groups=unique_reader_ids + ): + train_idx_unique.append(train_idx) + test_idx_unique.append(test_idx) + for train_idx, test_idx in kfold_universal.split( + X=universal_sn_ids, y=universal_reader_ids, groups=universal_reader_ids + ): + train_idx_universal.append(train_idx) + test_idx_universal.append(test_idx) + + # concatenate the respective train and test idx together + all_train_idx, all_test_idx = list(), list() + + for train_idx_univ, train_idx_uniq in zip(train_idx_universal, train_idx_unique): + train_idx = np.concatenate([train_idx_univ, train_idx_uniq]) + all_train_idx.append(train_idx) + for test_idx_univ, test_idx_uniq in zip(test_idx_universal, test_idx_unique): + test_idx = np.concatenate([test_idx_univ, test_idx_uniq]) + all_test_idx.append(test_idx) + + all_idx = [(train_idx, test_idx) for train_idx, test_idx in zip(all_train_idx, all_test_idx)] + return all_idx + + +def get_kfold( + data: List[Tuple[Any]], + splitting_IDs_dict: Dict[str, Union[int, str]], + splitting_criterion: str = "scanpath", # 'scanpath', 'reader', 'sentence', 'combined', + n_splits: int = 5, +): + if splitting_criterion == "scanpath": + # kfold = KFold(n_splits=n_splits, random_state=77, shuffle=True) + # kfold.get_n_splits(data) + # return kfold.split(data) + return get_kfold_indices_scanpath( + splitting_IDs_dict=splitting_IDs_dict, + n_splits=n_splits, + ) + elif splitting_criterion in ["reader", "sentence"]: + splitting_group = splitting_IDs_dict[splitting_criterion] + kfold = GroupKFold(n_splits=n_splits) + kfold.get_n_splits(data, groups=splitting_group) + return kfold.split(data, groups=splitting_group) + else: # combined split + raise NotImplementedError + + +def get_kfold_indices_combined( + data: List[Tuple[Any]], + splitting_IDs_dict: Dict[str, Union[int, str]], + n_splits: int = 5, +): + kfold_reader = GroupKFold(n_splits=n_splits) + kfold_sentence = GroupKFold(n_splits=n_splits) + kfold_reader.get_n_splits(data, groups=splitting_IDs_dict["reader"]) + kfold_sentence.get_n_splits(data, groups=splitting_IDs_dict["sentence"]) + reader_indices, sentence_indices = list(), list() + for train_idx, test_idx in kfold_reader.split(data, groups=splitting_IDs_dict["reader"]): + reader_indices.append((train_idx, test_idx)) + for train_idx, test_idx in kfold_sentence.split(data, groups=splitting_IDs_dict["sentence"]): + sentence_indices.append((train_idx, test_idx)) + return reader_indices, sentence_indices + + +def flatten_data(data: Dict[str, List[Any]]): + flattened_data = list() + for i in range(len(data["sn_sp_repr"])): + flattened_data.append( + ( + data["mask"][i], + data["sn_sp_repr"][i], + data["sn_input_ids"][i], + data["indices_pos_enc"][i], + data["sn_sp_fix_dur"][i], + data["sn_repr_len"][i], + data["words_for_mapping"][i], + data["mask_sn_padding"][i], + data["mask_transformer_att"][i], + data["sn_ids"][i], + data["reader_ids"][i], + ) + ) + return flattened_data + + +def unflatten_data(flattened_data: List[Tuple[Any]], split: str): + dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in flattened_data], + "sn_sp_repr": [instance[1] for instance in flattened_data], + "sn_input_ids": [instance[2] for instance in flattened_data], + "indices_pos_enc": [instance[3] for instance in flattened_data], + "sn_sp_fix_dur": [instance[4] for instance in flattened_data], + "sn_repr_len": [instance[5] for instance in flattened_data], + "words_for_mapping": [instance[6] for instance in flattened_data], + "mask_sn_padding": [instance[7] for instance in flattened_data], + "mask_transformer_att": [instance[8] for instance in flattened_data], + "sn_ids": [instance[9] for instance in flattened_data], + "reader_ids": [instance[10] for instance in flattened_data], + } + ) + data_dict = datasets.DatasetDict() + data_dict[split] = dataset + return data_dict + + +def process_celer( + sn_list, + reader_list, + word_info_df, + eyemovement_df, + tokenizer, + args, + split: Optional[str] = "train", + subset_size: Optional[int] = None, + split_sizes: Optional[Dict[str, float]] = None, + splitting_criterion: Optional[str] = "scanpath", # 'reader', 'sentence', 'combined' + inference: Optional[str] = None, # cv, zuco +): + """ + Process the Celer corpus so that it can be used as input to the Diffusion model, where the original sentence (sn) + is the condition and the scan path (sp) is the target that will be noised. + :param sn_list: list of unique sentence IDs in celer + :param reader_list: list of reader IDs in celer + :param word_info_df: pd Dataframe with sentence info + :param eyemovement_df: pd Dataframe with fixation info + :param tokenizer: BertTokenizer + :param args: + :param split: 'train', 'train-test', 'train-test-val', 'train-val' + :param subset_size: for test runs: to not load whole dataset but specified no. of instances + :param split_sizes: proportion of data going into train, test and val + :param splitting_criterion: how the data should be split for testing and validation. + 'reader' = New Reader setting + 'sentence' = New Sentence setting + 'combined' = New Reader/New Sentence setting + 'scanpath' = data is split at random + :param inference: if inference is cross-validation, the data is returned before splitting into train test val + """ + SP_ordinal_pos = [] + SP_landing_pos = [] + SP_fix_dur = [] + + data = { + "mask": list(), # 0 for sn, 1 for sp + "sn_sp_repr": list(), # word IDs of sn and corresponding word IDs of sp (fixated words, interest area IDs) padded with args.seq_len -1 + "sn_input_ids": list(), # input IDs of tokenized sentence, padded with pad token ID + "indices_pos_enc": list(), # indices from 1 ... len(sn input ids) 1 ... (seq_len - len(sn input ids)) + "sn_repr_len": list(), # length of sentence in subword tokens + "words_for_mapping": list(), # original words of sentence, padded with PAD + "mask_sn_padding": list(), # masks both the sentence and the padding, for the loss computations + "mask_transformer_att": list(), # masks only the padding, for the transformer attention, + "sn_ids": list(), # the sentence IDs + "reader_ids": list(), # the reader IDs + } + + max_len = 0 + + reader_IDs, sn_IDs = list(), list() + + for sn_id_idx, sn_id in tqdm(enumerate(sn_list), total=len(sn_list)): # for text/sentence ID + + if subset_size is not None: + if sn_id_idx == subset_size + 1: + break + + # subset the fixations report DF to a DF containing only the current sentence/text ID (each sentence appears multiple times) + sn_df = eyemovement_df[eyemovement_df.sentenceid == sn_id] + # notice: Each sentence is recorded multiple times in file |word_info_df|. + # subset the interest area report DF to a DF containing only the current sentence/text ID + sn = word_info_df[word_info_df.sentenceid == sn_id] + # sn is a dataframe containing only one sentence (the sentence with the current sentence ID) + sn = sn[ + sn["list"] == sn.list.values.tolist()[0] + ] # list = experimental list number (unique to each participant). + # compute word length and frequency features for each word + sn_str = sn.sentence.iloc[-1] # the whole sentence as string + if ( + sn_id == "1987/w7_019/w7_019.295-3" + or sn_id == "1987/w7_036/w7_036.147-43" + or sn_id == "1987/w7_091/w7_091.360-6" + ): + # extra inverted commas at the end of the sentence + sn_str = sn_str[:-3] + sn_str[-1:] + if sn_id == "1987/w7_085/w7_085.200-18": + sn_str = sn_str[:43] + sn_str[44:] + + # skip nan values bc they are of type float (np.isnan raises an error) + if isinstance(sn_str, float): + continue + + sn_len = len(sn_str.split()) + + # add CLS and SEP 'manually' to the sentence so that they receive the word IDs 0 and len(sn)+1 + sn_str = "[CLS] " + sn_str + " [SEP]" + + tokenizer.padding_side = "right" + + for sub_id_idx, sub_id in enumerate(reader_list): + + if sub_id_idx == 5: + continue + + sub_df = sn_df[sn_df.list == sub_id] + # remove fixations on non-words + sub_df = sub_df.loc[ + sub_df.CURRENT_FIX_INTEREST_AREA_LABEL != "." + ] # Label for the interest area to which the currentixation is assigned + if len(sub_df) == 0: + # no scanpath data found for the subject + continue + + # prepare decoder input and output + sp_word_pos, sp_fix_loc, sp_fix_dur = ( + sub_df.CURRENT_FIX_INTEREST_AREA_ID.values, + sub_df.CURRENT_FIX_NEAREST_INTEREST_AREA_DISTANCE.values, + sub_df.CURRENT_FIX_DURATION.values, + ) + + # check if recorded fixation duration are within reasonable limits + # Less than 15ms attempt to merge with neighbouring fixation if fixate is on the same word, otherwise delete + outlier_indx = np.where(sp_fix_dur < 50)[ + 0 + ] # gives indices of the fixations in the fixations list that were shorter than 50ms + + if outlier_indx.size > 0: + for out_idx in range(len(outlier_indx)): + outlier_i = outlier_indx[out_idx] + merge_flag = False + + # outliers are commonly found in the fixation of the last record and the first record, and are removed directly + if outlier_i == len(sp_fix_dur) - 1 or outlier_i == 0: + merge_flag = True + + else: + if outlier_i - 1 >= 0 and not merge_flag: + # try to merge with the left fixation if they landed both on the same interest area + if ( + sub_df.iloc[outlier_i].CURRENT_FIX_INTEREST_AREA_LABEL + == sub_df.iloc[outlier_i - 1].CURRENT_FIX_INTEREST_AREA_LABEL + ): + sp_fix_dur[outlier_i - 1] = ( + sp_fix_dur[outlier_i - 1] + sp_fix_dur[outlier_i] + ) + merge_flag = True + + if outlier_i + 1 < len(sp_fix_dur) and not merge_flag: + # try to merge with the right fixation + if ( + sub_df.iloc[outlier_i].CURRENT_FIX_INTEREST_AREA_LABEL + == sub_df.iloc[outlier_i + 1].CURRENT_FIX_INTEREST_AREA_LABEL + ): + sp_fix_dur[outlier_i + 1] = ( + sp_fix_dur[outlier_i + 1] + sp_fix_dur[outlier_i] + ) + merge_flag = True + + # delete the position (interest area ID), the fixation location and the fixation duration from the respective arrays + sp_word_pos = np.delete(sp_word_pos, outlier_i) + sp_fix_loc = np.delete(sp_fix_loc, outlier_i) + sp_fix_dur = np.delete(sp_fix_dur, outlier_i) + sub_df.drop(sub_df.index[outlier_i], axis=0, inplace=True) + outlier_indx = outlier_indx - 1 + + # sanity check + # scanpath too long, remove outliers, speed up the inference; more than 50 fixations on sentence + if len(sp_word_pos) > 50: # 72/10684 + continue + # scanpath too short for a normal length sentence + if len(sp_word_pos) <= 1 and sn_len > 10: + continue + + sp_ordinal_pos = sp_word_pos.astype( + int + ) # interest area index, i.e., word IDs in fixation report + SP_ordinal_pos.append(sp_ordinal_pos) + SP_fix_dur.append(sp_fix_dur) + # preprocess landing position feature + # assign missing value to 'nan' + sp_fix_loc = np.where(sp_fix_loc == ".", np.nan, sp_fix_loc) + # convert string of number of float type + sp_fix_loc = [float(i) for i in sp_fix_loc] + # Outliers in calculated landing positions due to lack of valid AOI data, assign to 'nan' + if ( + np.nanmax(sp_fix_loc) > 35 + ): # returns fixation outliers (coordinates very off); np.nanmax returns the max value while igonoring nans + missing_idx = np.where(np.array(sp_fix_loc) > 5)[ + 0 + ] # array with indices where fix loc greater than 5 + for miss in missing_idx: + if sub_df.iloc[miss].CURRENT_FIX_INTEREST_AREA_LEFT in [ + "NONE", + "BEFORE", + "AFTER", + "BOTH", + ]: + sp_fix_loc[miss] = np.nan + else: + print( + "Landing position calculation error. Unknown cause, needs to be checked" + ) + SP_landing_pos.append(sp_fix_loc) + + encoded_sn = tokenizer.encode_plus( + sn_str.split(), + add_special_tokens=False, + padding=False, + return_attention_mask=False, + is_split_into_words=True, + truncation=False, + ) + + sn_word_ids = encoded_sn.word_ids() + sp_word_ids = [0] + sp_ordinal_pos.tolist() + [max(encoded_sn.word_ids())] + + sn_input_ids = encoded_sn["input_ids"] + assert len(sn_word_ids) == len(sn_input_ids) + + max_len = max(max_len, len(sn_word_ids) + len(sp_word_ids)) + + # truncating + sep_token_sn_word_ids = sn_word_ids[-1] + sep_token_sp_word_ids = sp_word_ids[-1] + sep_token_sn_input_ids = sn_input_ids[-1] + + sn_word_ids = sn_word_ids[:-1] + sp_word_ids = sp_word_ids[:-1] + sn_input_ids = sn_input_ids[:-1] + + while len(sn_word_ids) + len(sp_word_ids) > args.seq_len - 3: + if len(sn_word_ids) > len(sp_word_ids): + sn_word_ids.pop() + sn_input_ids.pop() + elif len(sp_word_ids) > len(sn_word_ids): + sp_word_ids.pop() + else: + sn_word_ids.pop() + sn_input_ids.pop() + sp_word_ids.pop() + + # add the SEP token word ID and input ID again + sn_word_ids.append(sep_token_sn_word_ids) + sp_word_ids.append(sep_token_sp_word_ids) + sn_input_ids.append(sep_token_sn_input_ids) + + sn_sp_repr = sn_word_ids + sp_word_ids + + mask = [0] * len(sn_word_ids) + mask_sn_padding = ( + [0] * len(sn_word_ids) + + [1] * len(sp_word_ids) + + [0] * (args.seq_len - len(sn_word_ids) - len(sp_word_ids)) + ) + mask_transformer_att = ( + [1] * len(sn_word_ids) + + [1] * len(sp_word_ids) + + [0] * (args.seq_len - len(sn_word_ids) - len(sp_word_ids)) + ) + + indices_pos_enc = list(range(0, len(sn_word_ids))) + list( + range(0, args.seq_len - len(sn_word_ids)) + ) + + sn_repr_len = len(sn_word_ids) + words_for_mapping = sn_str.split() + (args.seq_len - len(sn_str.split())) * ["[PAD]"] + + data["mask"].append(mask) + data["sn_sp_repr"].append(sn_sp_repr) + data["sn_input_ids"].append(sn_input_ids) + data["indices_pos_enc"].append(indices_pos_enc) + data["sn_repr_len"].append(sn_repr_len) + data["words_for_mapping"].append(" ".join(words_for_mapping)) + data["mask_sn_padding"].append(mask_sn_padding) + data["mask_transformer_att"].append(mask_transformer_att) + data["sn_ids"].append(sn_id) + data["reader_ids"].append(sub_id) + + reader_IDs.append(sub_id) + sn_IDs.append(sn_id) + + # padding + data["mask"] = _collate_batch_helper( + examples=data["mask"], + pad_token_id=1, + max_length=args.seq_len, + ) + data["sn_sp_repr"] = _collate_batch_helper( + examples=data["sn_sp_repr"], + pad_token_id=args.seq_len - 1, + max_length=args.seq_len, + ) + data["sn_input_ids"] = _collate_batch_helper( + examples=data["sn_input_ids"], + pad_token_id=tokenizer.pad_token_id, + max_length=args.seq_len, + ) + + splitting_IDs_dict = { + "reader": reader_IDs, + "sentence": sn_IDs, + } + + if inference == "cv": + return data, splitting_IDs_dict + + if split == "train": + + dataset = Dataset2.from_dict(data) + train_dataset = datasets.DatasetDict() + train_dataset["train"] = dataset + return train_dataset, splitting_IDs_dict + + else: + + # flatten the data + flattened_data = list() + for i in range(len(data["sn_sp_repr"])): + flattened_data.append( + ( + data["mask"][i], + data["sn_sp_repr"][i], + data["sn_input_ids"][i], + data["indices_pos_enc"][i], + data["sn_repr_len"][i], + data["words_for_mapping"][i], + data["mask_sn_padding"][i], + data["mask_transformer_att"][i], + data["sn_ids"][i], + data["reader_ids"][i], + ) + ) + + if split == "train-test": + + if split_sizes: + test_size = split_sizes["test_size"] + else: + test_size = 0.25 + + if splitting_criterion != "scanpath": + + if splitting_criterion == "combined": + + ( + train_data, + test_data, + train_reader_IDs, + test_reader_IDs, + train_sn_IDs, + test_sn_IDs, + ) = combined_split( + data=flattened_data, + reader_IDs=splitting_IDs_dict["reader"], + sn_IDs=splitting_IDs_dict["sentence"], + test_size=test_size, + ) + + else: + + splitting_IDs = splitting_IDs_dict[splitting_criterion] + gss = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=77) + for train_index, test_index in gss.split(flattened_data, groups=splitting_IDs): + train_data = np.array(flattened_data)[train_index].tolist() + test_data = np.array(flattened_data)[test_index].tolist() + + else: + train_data, test_data = train_test_split( + flattened_data, test_size=test_size, shuffle=True, random_state=77 + ) + + # unflatten the data + train_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in train_data], + "sn_sp_repr": [instance[1] for instance in train_data], + "sn_input_ids": [instance[2] for instance in train_data], + "indices_pos_enc": [instance[3] for instance in train_data], + "sn_repr_len": [instance[4] for instance in train_data], + "words_for_mapping": [instance[5] for instance in train_data], + "mask_sn_padding": [instance[6] for instance in train_data], + "mask_transformer_att": [instance[7] for instance in train_data], + "sn_ids": [instance[8] for instance in train_data], + "reader_ids": [instance[9] for instance in train_data], + } + ) + test_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in test_data], + "sn_sp_repr": [instance[1] for instance in test_data], + "sn_input_ids": [instance[2] for instance in test_data], + "indices_pos_enc": [instance[3] for instance in test_data], + "sn_repr_len": [instance[4] for instance in test_data], + "words_for_mapping": [instance[5] for instance in test_data], + "mask_sn_padding": [instance[6] for instance in test_data], + "mask_transformer_att": [instance[7] for instance in test_data], + "sn_ids": [instance[8] for instance in test_data], + "reader_ids": [instance[9] for instance in test_data], + } + ) + train_data_dict = datasets.DatasetDict() + test_data_dict = datasets.DatasetDict() + train_data_dict["train"] = train_dataset + test_data_dict["test"] = test_dataset + return train_data_dict, test_data_dict + + elif split == "train-val": + + if split_sizes: + val_size = split_sizes["val_size"] + else: + val_size = 0.1 + + if splitting_criterion != "scanpath": + + if splitting_criterion == "combined": + ( + train_data, + val_data, + train_reader_IDs, + val_reader_IDs, + train_sn_IDs, + val_sn_IDs, + ) = combined_split( + data=flattened_data, + reader_IDs=splitting_IDs_dict["reader"], + sn_IDs=splitting_IDs_dict["sentence"], + test_size=val_size, + ) + else: + splitting_IDs = splitting_IDs_dict[splitting_criterion] + gss = GroupShuffleSplit(n_splits=1, test_size=val_size, random_state=77) + for train_index, val_index in gss.split(flattened_data, groups=splitting_IDs): + train_data = np.array(flattened_data)[train_index].tolist() + val_data = np.array(flattened_data)[val_index].tolist() + + else: + train_data, val_data = train_test_split( + flattened_data, test_size=val_size, shuffle=True, random_state=77 + ) + + # unflatten the data + train_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in train_data], + "sn_sp_repr": [instance[1] for instance in train_data], + "sn_input_ids": [instance[2] for instance in train_data], + "indices_pos_enc": [instance[3] for instance in train_data], + "sn_repr_len": [instance[4] for instance in train_data], + "words_for_mapping": [instance[5] for instance in train_data], + "mask_sn_padding": [instance[6] for instance in train_data], + "mask_transformer_att": [instance[7] for instance in train_data], + "sn_ids": [instance[8] for instance in train_data], + "reader_ids": [instance[9] for instance in train_data], + } + ) + # unflatten the data + val_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in val_data], + "sn_sp_repr": [instance[1] for instance in val_data], + "sn_input_ids": [instance[2] for instance in val_data], + "indices_pos_enc": [instance[3] for instance in val_data], + "sn_repr_len": [instance[4] for instance in val_data], + "words_for_mapping": [instance[5] for instance in val_data], + "mask_sn_padding": [instance[6] for instance in val_data], + "mask_transformer_att": [instance[7] for instance in val_data], + "sn_ids": [instance[8] for instance in val_data], + "reader_ids": [instance[9] for instance in val_data], + } + ) + train_data_dict = datasets.DatasetDict() + val_data_dict = datasets.DatasetDict() + train_data_dict["train"] = train_dataset + val_data_dict["val"] = val_dataset + return train_data_dict, val_data_dict + + elif split == "train-val-test": + + if split_sizes: + val_size = split_sizes["val_size"] + test_size = split_sizes["test_size"] + else: + val_size = 0.1 + test_size = 0.25 + + if splitting_criterion != "scanpath": + + if splitting_criterion == "combined": + # split train and test data so that unseen readers and sentences are in the test data + ( + train_data, + test_data, + train_reader_IDs, + test_reader_IDs, + train_sn_IDs, + test_sn_IDs, + ) = combined_split( + data=flattened_data, + reader_IDs=splitting_IDs_dict["reader"], + sn_IDs=splitting_IDs_dict["sentence"], + test_size=test_size, + ) + # randomly split train data into train and validation + train_data, val_data = train_test_split( + train_data, test_size=val_size, random_state=77, shuffle=True + ) + + else: + splitting_IDs = splitting_IDs_dict[splitting_criterion] + + # split into train and test + gss = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=77) + for train_index, test_index in gss.split(flattened_data, groups=splitting_IDs): + train_data = np.array(flattened_data)[train_index].tolist() + test_data = np.array(flattened_data)[test_index].tolist() + train_ids = np.array(splitting_IDs)[train_index].tolist() + + # split into train and val + gss = GroupShuffleSplit(n_splits=1, test_size=val_size, random_state=77) + for train_index, val_index in gss.split(train_data, groups=train_ids): + val_data = np.array(train_data)[val_index].tolist() + train_data = np.array(train_data)[train_index].tolist() + + else: + train_data, test_data = train_test_split( + flattened_data, test_size=test_size, shuffle=True, random_state=77 + ) + train_data, val_data = train_test_split( + train_data, test_size=val_size, shuffle=True, random_state=77 + ) + + # unflatten the data + train_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in train_data], + "sn_sp_repr": [instance[1] for instance in train_data], + "sn_input_ids": [instance[2] for instance in train_data], + "indices_pos_enc": [instance[3] for instance in train_data], + "sn_repr_len": [instance[4] for instance in train_data], + "words_for_mapping": [instance[5] for instance in train_data], + "mask_sn_padding": [instance[6] for instance in train_data], + "mask_transformer_att": [instance[7] for instance in train_data], + "sn_ids": [instance[8] for instance in train_data], + "reader_ids": [instance[9] for instance in train_data], + } + ) + test_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in test_data], + "sn_sp_repr": [instance[1] for instance in test_data], + "sn_input_ids": [instance[2] for instance in test_data], + "indices_pos_enc": [instance[3] for instance in test_data], + "sn_repr_len": [instance[4] for instance in test_data], + "words_for_mapping": [instance[5] for instance in test_data], + "mask_sn_padding": [instance[6] for instance in test_data], + "mask_transformer_att": [instance[7] for instance in test_data], + "sn_ids": [instance[8] for instance in test_data], + "reader_ids": [instance[9] for instance in test_data], + } + ) + val_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in val_data], + "sn_sp_repr": [instance[1] for instance in val_data], + "sn_input_ids": [instance[2] for instance in val_data], + "indices_pos_enc": [instance[3] for instance in val_data], + "sn_repr_len": [instance[4] for instance in val_data], + "words_for_mapping": [instance[5] for instance in val_data], + "mask_sn_padding": [instance[6] for instance in val_data], + "mask_transformer_att": [instance[7] for instance in val_data], + "sn_ids": [instance[8] for instance in val_data], + "reader_ids": [instance[9] for instance in val_data], + } + ) + train_data_dict = datasets.DatasetDict() + test_data_dict = datasets.DatasetDict() + val_data_dict = datasets.DatasetDict() + train_data_dict["train"] = train_dataset + test_data_dict["test"] = test_dataset + val_data_dict["val"] = val_dataset + return train_data_dict, test_data_dict, val_data_dict + + +class CelerZucoDataset(Dataset): + + def __init__( + self, + dataset, + data_args, + split, # 'train', 'test', 'val' + ): + super().__init__() + self.dataset = dataset + self.length = len(self.dataset[split]) + self.data_args = data_args + self.split = split + + def __len__(self): + return self.length + + def __getitem__(self, idx): + sample = { + "mask": np.array(self.dataset[self.split][idx]["mask"]), + "sn_sp_repr": np.array(self.dataset[self.split][idx]["sn_sp_repr"]), + "sn_input_ids": np.array(self.dataset[self.split][idx]["sn_input_ids"]), + "indices_pos_enc": np.array(self.dataset[self.split][idx]["indices_pos_enc"]), + "sn_sp_fix_dur": np.array(self.dataset[self.split][idx]["sn_sp_fix_dur"]), + "sn_repr_len": np.array(self.dataset[self.split][idx]["sn_repr_len"]), + "words_for_mapping": self.dataset[self.split][idx]["words_for_mapping"], + "mask_sn_padding": np.array(self.dataset[self.split][idx]["mask_sn_padding"]), + "mask_transformer_att": np.array(self.dataset[self.split][idx]["mask_transformer_att"]), + "sn_ids": self.dataset[self.split][idx]["sn_ids"], + "reader_ids": self.dataset[self.split][idx]["reader_ids"], + } + return sample + + +def process_zuco( + sn_list, + reader_list, + word_info_df, + eyemovement_df, + tokenizer, + args, + split: Optional[str] = "train", + subset_size: Optional[int] = None, + split_sizes: Optional[Dict[str, float]] = None, + splitting_criterion: Optional[str] = "scanpath", # 'reader', 'sentence', 'combined' +): + """ + Process the ZuCo corpus so that it can be used as input to the Diffusion model, where the original sentence (sn) + is the condition and the scan path (sp) is the target that will be noised. + :param sn_list: list of unique sentence IDs in zuco + :param reader_list: list of reader IDs in zuco + :param word_info_df: pd Dataframe with sentence info + :param eyemovement_df: pd Dataframe with fixation info + :param tokenizer: BertTokenizer + :param args: + :param split: 'train', 'train-test', 'train-test-val', 'train-val' + :param subset_size: for test runs: to not load whole dataset but specified no. of instances + :param split_sizes: proportion of data going into train, test and val + :param splitting_criterion: how the data should be split for testing and validation. + 'reader' = New Reader setting + 'sentence' = New Sentence setting + 'combined' = New Reader/New Sentence setting + 'scanpath' = data is split at random + """ + SP_ordinal_pos = [] + SP_landing_pos = [] + SP_fix_dur = [] + + data = { + "mask": list(), # 0 for sn, 1 for sp + "sn_sp_repr": list(), # word IDs of sn and corresponding word IDs of sp (fixated words, interest area IDs), + # padded with args.seq_len -1 + "sn_input_ids": list(), # input IDs of tokenized sentence, padded with pad token ID + "indices_pos_enc": list(), # indices from 1 ... len(sn input ids) 1 ... (seq_len - len(sn input ids)) + "sn_repr_len": list(), # length of sentence in subword tokens + "words_for_mapping": list(), # original words of sentence, padded with PAD + "mask_sn_padding": list(), # masks both the sentence and the padding, for the loss computations + "mask_transformer_att": list(), # masks only the padding, for the transformer attention + "sn_ids": list(), # sentence IDs + "reader_ids": list(), # reader IDs + } + + max_len = 0 + all_lens = list() + + reader_IDs, sn_IDs = list(), list() + + for sn_id_idx, sn_id in tqdm(enumerate(sn_list), total=len(sn_list)): + + if subset_size is not None: + if sn_id_idx == subset_size + 1: + break + + sn_df = eyemovement_df[eyemovement_df.sn == sn_id] + sn = word_info_df[word_info_df.SN == sn_id] + sn_str = " ".join(sn.WORD.values) + sn_len = len(sn_str.split()) + + tokenizer.padding_side = "right" + sn_str = "[CLS] " + sn_str + " [SEP]" + + for sub_id_idx, sub_id in enumerate(reader_list): + + sub_df = sn_df[sn_df.id == sub_id] + # remove fixations on non-words + sub_df = sub_df.loc[sub_df.CURRENT_FIX_INTEREST_AREA_LABEL != ""] + if len(sub_df) == 0: + # no scanpath data found for the subject + continue + + sp_word_pos, sp_fix_loc, sp_fix_dur = ( + sub_df.wn.values, + sub_df.fl.values, + sub_df.dur.values, + ) + + # check if recorded fixation duration are within reasonable limits + # Less than 50ms attempt to merge with neighbouring fixation if fixate is on the same word, otherwise delete + outlier_indx = np.where(sp_fix_dur < 50)[0] + if outlier_indx.size > 0: + for out_idx in range(len(outlier_indx)): + outlier_i = outlier_indx[out_idx] + merge_flag = False + if outlier_i - 1 >= 0 and not merge_flag: + # try to merge with the left fixation + if ( + sub_df.iloc[outlier_i].CURRENT_FIX_INTEREST_AREA_LABEL + == sub_df.iloc[outlier_i - 1].CURRENT_FIX_INTEREST_AREA_LABEL + ): + sp_fix_dur[outlier_i - 1] = ( + sp_fix_dur[outlier_i - 1] + sp_fix_dur[outlier_i] + ) + merge_flag = True + + if outlier_i + 1 < len(sp_fix_dur) and not merge_flag: + # try to merge with the right fixation + if ( + sub_df.iloc[outlier_i].CURRENT_FIX_INTEREST_AREA_LABEL + == sub_df.iloc[outlier_i + 1].CURRENT_FIX_INTEREST_AREA_LABEL + ): + sp_fix_dur[outlier_i + 1] = ( + sp_fix_dur[outlier_i + 1] + sp_fix_dur[outlier_i] + ) + merge_flag = True + + sp_word_pos = np.delete(sp_word_pos, outlier_i) + sp_fix_loc = np.delete(sp_fix_loc, outlier_i) + sp_fix_dur = np.delete(sp_fix_dur, outlier_i) + sub_df.drop(sub_df.index[outlier_i], axis=0, inplace=True) + outlier_indx = outlier_indx - 1 + + # sanity check + # scanpath too short for a normal length sentence + if len(sp_word_pos) <= 1 and sn_len > 10: + continue + + sp_ordinal_pos = sp_word_pos.astype(int) + SP_ordinal_pos.append(sp_ordinal_pos) + SP_fix_dur.append(sp_fix_dur) + + # preprocess landing position feature + # assign missing value to 'nan' + # sp_fix_loc=np.where(sp_fix_loc=='.', np.nan, sp_fix_loc) + # convert string of number of float type + sp_fix_loc = [ + float(i) if isinstance(i, int) or isinstance(i, float) else np.nan + for i in sp_fix_loc + if isinstance(i, int) or isinstance(i, float) + ] + SP_landing_pos.append(sp_fix_loc) + + # encode the sentence + encoded_sn = tokenizer.encode_plus( + sn_str.split(), + add_special_tokens=False, + padding=False, + return_attention_mask=False, + is_split_into_words=True, + truncation=False, + ) + + sn_word_ids = encoded_sn.word_ids() + sp_word_ids = [0] + sp_ordinal_pos.tolist() + [max(encoded_sn.word_ids())] + + sn_input_ids = encoded_sn["input_ids"] + assert len(sn_word_ids) == len(sn_input_ids) + + max_len = max(max_len, len(sn_word_ids) + len(sp_word_ids)) + all_lens.append(len(sn_word_ids) + len(sp_word_ids)) + + # truncating + sep_token_sn_word_ids = sn_word_ids[-1] + sep_token_sp_word_ids = sp_word_ids[-1] + sep_token_sn_input_ids = sn_input_ids[-1] + + sn_word_ids = sn_word_ids[:-1] + sp_word_ids = sp_word_ids[:-1] + sn_input_ids = sn_input_ids[:-1] + + while len(sn_word_ids) + len(sp_word_ids) > args.seq_len - 3: + if len(sn_word_ids) > len(sp_word_ids): + sn_word_ids.pop() + sn_input_ids.pop() + elif len(sp_word_ids) > len(sn_word_ids): + sp_word_ids.pop() + else: + sn_word_ids.pop() + sn_input_ids.pop() + sp_word_ids.pop() + + # add the SEP token word ID and input ID again + sn_word_ids.append(sep_token_sn_word_ids) + sp_word_ids.append(sep_token_sp_word_ids) + sn_input_ids.append(sep_token_sn_input_ids) + + sn_sp_repr = sn_word_ids + sp_word_ids + + mask = [0] * len(sn_word_ids) + mask_sn_padding = ( + [0] * len(sn_word_ids) + + [1] * len(sp_word_ids) + + [0] * (args.seq_len - len(sn_word_ids) - len(sp_word_ids)) + ) + mask_transformer_att = ( + [1] * len(sn_word_ids) + + [1] * len(sp_word_ids) + + [0] * (args.seq_len - len(sn_word_ids) - len(sp_word_ids)) + ) + + indices_pos_enc = list(range(0, len(sn_word_ids))) + list( + range(0, args.seq_len - len(sn_word_ids)) + ) + + sn_repr_len = len(sn_word_ids) + words_for_mapping = sn_str.split() + (args.seq_len - len(sn_str.split())) * ["[PAD]"] + + data["mask"].append(mask) + data["sn_sp_repr"].append(sn_sp_repr) + data["sn_input_ids"].append(sn_input_ids) + data["indices_pos_enc"].append(indices_pos_enc) + data["sn_repr_len"].append(sn_repr_len) + data["words_for_mapping"].append(" ".join(words_for_mapping)) + data["mask_sn_padding"].append(mask_sn_padding) + data["mask_transformer_att"].append(mask_transformer_att) + data["sn_ids"].append(sn_id) + data["reader_ids"].append(sub_id) + + reader_IDs.append(sub_id) + sn_IDs.append(sn_id) + + # padding + data["mask"] = _collate_batch_helper( + examples=data["mask"], + pad_token_id=1, + max_length=args.seq_len, + ) + data["sn_sp_repr"] = _collate_batch_helper( + examples=data["sn_sp_repr"], + pad_token_id=args.seq_len - 1, + max_length=args.seq_len, + ) + data["sn_input_ids"] = _collate_batch_helper( + examples=data["sn_input_ids"], + pad_token_id=tokenizer.pad_token_id, + max_length=args.seq_len, + ) + + splitting_IDs_dict = { + "reader": reader_IDs, + "sentence": sn_IDs, + } + + if split == "train": + + dataset = Dataset2.from_dict(data) + train_dataset = datasets.DatasetDict() + train_dataset["train"] = dataset + return train_dataset + + else: + + # flatten the data + flattened_data = list() + for i in range(len(data["sn_sp_repr"])): + flattened_data.append( + ( + data["mask"][i], + data["sn_sp_repr"][i], + data["sn_input_ids"][i], + data["indices_pos_enc"][i], + data["sn_repr_len"][i], + data["words_for_mapping"][i], + data["mask_sn_padding"][i], + data["mask_transformer_att"][i], + data["sn_ids"][i], + data["reader_ids"][i], + ) + ) + + if split == "train-test": + + if split_sizes: + test_size = split_sizes["test_size"] + else: + test_size = 0.25 + + if splitting_criterion != "scanpath": + + if splitting_criterion == "combined": + + ( + train_data, + test_data, + train_reader_IDs, + test_reader_IDs, + train_sn_IDs, + test_sn_IDs, + ) = combined_split( + data=flattened_data, + reader_IDs=splitting_IDs_dict["reader"], + sn_IDs=splitting_IDs_dict["sentence"], + test_size=test_size, + ) + + else: + + splitting_IDs = splitting_IDs_dict[splitting_criterion] + gss = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=77) + for train_index, test_index in gss.split(flattened_data, groups=splitting_IDs): + train_data = np.array(flattened_data)[train_index].tolist() + test_data = np.array(flattened_data)[test_index].tolist() + + else: + train_data, test_data = train_test_split( + flattened_data, test_size=test_size, shuffle=True, random_state=77 + ) + + # unflatten the data + train_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in train_data], + "sn_sp_repr": [instance[1] for instance in train_data], + "sn_input_ids": [instance[2] for instance in train_data], + "indices_pos_enc": [instance[3] for instance in train_data], + "sn_repr_len": [instance[4] for instance in train_data], + "words_for_mapping": [instance[5] for instance in train_data], + "mask_sn_padding": [instance[6] for instance in train_data], + "mask_transformer_att": [instance[7] for instance in train_data], + "sn_ids": [instance[8] for instance in train_data], + "reader_ids": [instance[9] for instance in train_data], + } + ) + test_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in test_data], + "sn_sp_repr": [instance[1] for instance in test_data], + "sn_input_ids": [instance[2] for instance in test_data], + "indices_pos_enc": [instance[3] for instance in test_data], + "sn_repr_len": [instance[4] for instance in test_data], + "words_for_mapping": [instance[5] for instance in test_data], + "mask_sn_padding": [instance[6] for instance in test_data], + "mask_transformer_att": [instance[7] for instance in test_data], + "sn_ids": [instance[8] for instance in test_data], + "reader_ids": [instance[9] for instance in test_data], + } + ) + train_data_dict = datasets.DatasetDict() + test_data_dict = datasets.DatasetDict() + train_data_dict["train"] = train_dataset + test_data_dict["test"] = test_dataset + return train_data_dict, test_data_dict + + elif split == "train-val": + + if split_sizes: + val_size = split_sizes["val_size"] + else: + val_size = 0.1 + + if splitting_criterion != "scanpath": + + if splitting_criterion == "combined": + ( + train_data, + val_data, + train_reader_IDs, + val_reader_IDs, + train_sn_IDs, + val_sn_IDs, + ) = combined_split( + data=flattened_data, + reader_IDs=splitting_IDs_dict["reader"], + sn_IDs=splitting_IDs_dict["sentence"], + test_size=val_size, + ) + else: + splitting_IDs = splitting_IDs_dict[splitting_criterion] + gss = GroupShuffleSplit(n_splits=1, test_size=val_size, random_state=77) + for train_index, val_index in gss.split(flattened_data, groups=splitting_IDs): + train_data = np.array(flattened_data)[train_index].tolist() + val_data = np.array(flattened_data)[val_index].tolist() + + else: + train_data, val_data = train_test_split( + flattened_data, test_size=val_size, shuffle=True, random_state=77 + ) + + # unflatten the data + train_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in train_data], + "sn_sp_repr": [instance[1] for instance in train_data], + "sn_input_ids": [instance[2] for instance in train_data], + "indices_pos_enc": [instance[3] for instance in train_data], + "sn_repr_len": [instance[4] for instance in train_data], + "words_for_mapping": [instance[5] for instance in train_data], + "mask_sn_padding": [instance[6] for instance in train_data], + "mask_transformer_att": [instance[7] for instance in train_data], + "sn_ids": [instance[8] for instance in train_data], + "reader_ids": [instance[9] for instance in train_data], + } + ) + # unflatten the data + val_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in val_data], + "sn_sp_repr": [instance[1] for instance in val_data], + "sn_input_ids": [instance[2] for instance in val_data], + "indices_pos_enc": [instance[3] for instance in val_data], + "sn_repr_len": [instance[4] for instance in val_data], + "words_for_mapping": [instance[5] for instance in val_data], + "mask_sn_padding": [instance[6] for instance in val_data], + "mask_transformer_att": [instance[7] for instance in val_data], + "sn_ids": [instance[8] for instance in val_data], + "reader_ids": [instance[9] for instance in val_data], + } + ) + train_data_dict = datasets.DatasetDict() + val_data_dict = datasets.DatasetDict() + train_data_dict["train"] = train_dataset + val_data_dict["val"] = val_dataset + return train_data_dict, val_data_dict + + elif split == "train-val-test": + + if split_sizes: + val_size = split_sizes["val_size"] + test_size = split_sizes["test_size"] + else: + val_size = 0.1 + test_size = 0.25 + + if splitting_criterion != "scanpath": + + if splitting_criterion == "combined": + # split train and test data so that unseen readers and sentences are in the test data + ( + train_data, + test_data, + train_reader_IDs, + test_reader_IDs, + train_sn_IDs, + test_sn_IDs, + ) = combined_split( + data=flattened_data, + reader_IDs=splitting_IDs_dict["reader"], + sn_IDs=splitting_IDs_dict["sentence"], + test_size=test_size, + ) + # randomly split train data into train and validation + train_data, val_data = train_test_split( + train_data, test_size=val_size, random_state=77, shuffle=True + ) + + else: + splitting_IDs = splitting_IDs_dict[splitting_criterion] + + # split into train and test + gss = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=77) + for train_index, test_index in gss.split(flattened_data, groups=splitting_IDs): + train_data = np.array(flattened_data)[train_index].tolist() + test_data = np.array(flattened_data)[test_index].tolist() + train_ids = np.array(splitting_IDs)[train_index].tolist() + + # split into train and val + gss = GroupShuffleSplit(n_splits=1, test_size=val_size, random_state=77) + for train_index, val_index in gss.split(train_data, groups=train_ids): + val_data = np.array(train_data)[val_index].tolist() + train_data = np.array(train_data)[train_index].tolist() + + else: + train_data, test_data = train_test_split( + flattened_data, test_size=test_size, shuffle=True, random_state=77 + ) + train_data, val_data = train_test_split( + train_data, test_size=val_size, shuffle=True, random_state=77 + ) + + # unflatten the data + train_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in train_data], + "sn_sp_repr": [instance[1] for instance in train_data], + "sn_input_ids": [instance[2] for instance in train_data], + "indices_pos_enc": [instance[3] for instance in train_data], + "sn_repr_len": [instance[4] for instance in train_data], + "words_for_mapping": [instance[5] for instance in train_data], + "mask_sn_padding": [instance[6] for instance in train_data], + "mask_transformer_att": [instance[7] for instance in train_data], + "sn_ids": [instance[8] for instance in train_data], + "reader_ids": [instance[9] for instance in train_data], + } + ) + test_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in test_data], + "sn_sp_repr": [instance[1] for instance in test_data], + "sn_input_ids": [instance[2] for instance in test_data], + "indices_pos_enc": [instance[3] for instance in test_data], + "sn_repr_len": [instance[4] for instance in test_data], + "words_for_mapping": [instance[5] for instance in test_data], + "mask_sn_padding": [instance[6] for instance in test_data], + "mask_transformer_att": [instance[7] for instance in test_data], + "sn_ids": [instance[8] for instance in test_data], + "reader_ids": [instance[9] for instance in test_data], + } + ) + val_dataset = Dataset2.from_dict( + { + "mask": [instance[0] for instance in val_data], + "sn_sp_repr": [instance[1] for instance in val_data], + "sn_input_ids": [instance[2] for instance in val_data], + "indices_pos_enc": [instance[3] for instance in val_data], + "sn_repr_len": [instance[4] for instance in val_data], + "words_for_mapping": [instance[5] for instance in val_data], + "mask_sn_padding": [instance[6] for instance in val_data], + "mask_transformer_att": [instance[7] for instance in val_data], + "sn_ids": [instance[8] for instance in val_data], + } + ) + train_data_dict = datasets.DatasetDict() + test_data_dict = datasets.DatasetDict() + val_data_dict = datasets.DatasetDict() + train_data_dict["train"] = train_dataset + test_data_dict["test"] = test_dataset + val_data_dict["val"] = val_dataset + return train_data_dict, test_data_dict, val_data_dict diff --git a/scandl_module/scripts/sp_run_train.py b/scandl_module/scripts/sp_run_train.py new file mode 100644 index 0000000000000000000000000000000000000000..13dbeb9b9bd807271bfa776f4fa3932d0b1219c5 --- /dev/null +++ b/scandl_module/scripts/sp_run_train.py @@ -0,0 +1,196 @@ +import sys +import os +import argparse +import datetime +import time + +from ScanDL2.CONSTANTS import ( + COMPLETE_SCANDL_MODULE_TRAIN_PATH_BSC, + COMPLETE_SCANDL_MODULE_TRAIN_PATH_CELER, + COMPLETE_SCANDL_MODULE_TRAIN_PATH_EMTEC, +) + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="training args.") + parser.add_argument( + "--noise_schedule", + type=str, + default="sqrt", + choices=["linear", "cosine", "sqrt", "trunc_cos", "trunc_lin", "pw_lin"], + help="the distribution of noises", + ) + parser.add_argument("--diff_steps", type=int, default=2000, help="diffusion steps") + parser.add_argument( + "--schedule_sampler", + type=str, + default="lossaware", + choices=["uniform", "lossaware", "fixstep"], + help="schedule sampler of timesteps", + ) + + parser.add_argument("--seq_len", type=int, default=128, help="max len of input sequence") + parser.add_argument( + "--hidden_t_dim", type=int, default=128, help="hidden size of time embedding" + ) + parser.add_argument( + "--hidden_dim", + type=int, + default=768, + help="hidden size of word embedding and transformer hidden size", + ) + parser.add_argument("--learning_steps", type=int, default=60000, help="total steps of learning") + parser.add_argument("--save_interval", type=int, default=2000, help="save step") + parser.add_argument( + "--resume_checkpoint", + type=str, + default="none", + help="path to resume checkpoint, like xxx/xxx.pt", + ) + parser.add_argument("--lr", type=float, default=1e-04, help="learning rate") + parser.add_argument("--bsz", type=int, default=64, help="batch size") + parser.add_argument("--microbatch", type=int, default=64, help="microbatch size") + parser.add_argument("--seed", type=int, default=101, help="random seed") + + parser.add_argument( + "--config_name", type=str, default="bert-base-cased", help="config of pre-trained models" + ) + parser.add_argument( + "--vocab", + type=str, + default="bert", + help="use bert vocab or load external vocab dict if given as path", + ) + parser.add_argument( + "--use_plm_init", + type=str, + default="no", + choices=["no", "bert"], + help="load init parameter from the pre-trained lm", + ) + parser.add_argument("--log_interval", type=int, default=200, required=False) + parser.add_argument("--eval_interval", type=int, default=500, required=False) + + parser.add_argument( + "--notes", + type=str, + default="-", + help="as training notes or specifical args", + required=False, + ) + parser.add_argument("--app", type=str, default="", help="other input args") + + # further arguments + parser.add_argument( + "--data_split_criterion", + type=str, + help="how to split the data into train, val, test:" + " scanpath (random), reader, sentence, combined", + required=False, + default="reader", + ) + parser.add_argument( + "--num_transformer_layers", + type=int, + default=4, + required=False, + help="the number of encoder layers", + ) + parser.add_argument( + "--num_transformer_heads", + type=int, + default=8, + required=False, + help="the number of attention heads", + ) + parser.add_argument( + "--celer_only_L1", + required=False, + action="store_true", + help="if given, all celer speakers are used" "as opposed to only L1 speakers", + ) + parser.add_argument( + "--corpus", + type=str, + help="the eye-tracking corpus to use for training.", + required=False, + default="celer", + choices=["celer", "zuco", "emtec", "bsc"], + ) + parser.add_argument( + "--inference", + required=False, + default="cv", + choices=["cv", "zuco", "in-corpus"], + help="if zuco, inference is performed on zuco while trained on celer; if cv, inference is" + "done in k-fold Cross-Validation; if in-corpus, the training corpus is simply split into" + "train and test.", + ) + parser.add_argument( + "--mask_padding", + action="store_false", + required=False, + help="if given, padding will not be masked in transformer attention. if not given, mask_padding" + "is stored as True; padding will be masked.", + ) + parser.add_argument( + "--load_train_data", + type=str, + default="-", + help="if given, previously saved train data is loaded from the specified checkpoint path", + ) + + args = parser.parse_args() + + # set working dir to the upper folder + abspath = os.path.abspath(sys.argv[0]) + dname = os.path.dirname(abspath) + dname = os.path.dirname(dname) + os.chdir(dname) + + if args.corpus == "emtec": + model_file = COMPLETE_SCANDL_MODULE_TRAIN_PATH_EMTEC + elif args.corpus == "bsc": + model_file = COMPLETE_SCANDL_MODULE_TRAIN_PATH_BSC + elif args.corpus == "celer": + model_file = COMPLETE_SCANDL_MODULE_TRAIN_PATH_CELER + else: + raise NotImplementedError(f"Corpus {args.corpus} not implemented.") + + if int(os.environ["LOCAL_RANK"]) == 0: + if not os.path.exists(model_file): + os.makedirs(model_file) + + COMMANDLINE = ( + f"TOKENIZERS_PARALLELISM=FALSE " + f"python -m scripts.sp_train " + f"--checkpoint_path {model_file} " + f"--vocab {args.vocab} " + f"--use_plm_init {args.use_plm_init} " + f"--lr {args.lr} " + f"--batch_size {args.bsz} " + f"--microbatch {args.microbatch} " + f"--diffusion_steps {args.diff_steps} " + f"--noise_schedule {args.noise_schedule} " + f"--schedule_sampler {args.schedule_sampler} " + f"--seq_len {args.seq_len} " + f"--resume_checkpoint {args.resume_checkpoint} " + f"--hidden_t_dim {args.hidden_t_dim} " + f"--seed {args.seed} " + f"--hidden_dim {args.hidden_dim} " + f"--learning_steps {args.learning_steps} " + f"--save_interval {args.save_interval} " + f"--config_name {args.config_name} " + f"--notes {args.notes} " + f"--data_split_criterion {args.data_split_criterion} " + f"--num_transformer_layers {args.num_transformer_layers} " + f"--num_transformer_heads {args.num_transformer_heads} " + f"--corpus {args.corpus} " + f"--inference {args.inference} " + f"--load_train_data {args.load_train_data}" + ) + if int(os.environ["LOCAL_RANK"]) == 0: + with open(os.path.join(model_file, "saved_bash.sh"), "w") as f: + print(COMMANDLINE, file=f) + + print(COMMANDLINE) + os.system(COMMANDLINE) diff --git a/scandl_module/scripts/sp_train.py b/scandl_module/scripts/sp_train.py new file mode 100644 index 0000000000000000000000000000000000000000..24e3be16d3138da781bfe320a1974e5bd91963de --- /dev/null +++ b/scandl_module/scripts/sp_train.py @@ -0,0 +1,158 @@ +""" +Train ScanDL 2.0 on all data of the current dataset. +""" + +import argparse +import json +import os +import numpy as np + +# import wandb + +import torch +import torch.distributed as dist +from sklearn.model_selection import train_test_split +from transformers import set_seed, BertTokenizerFast +from datasets import load_from_disk, DatasetDict + +from ScanDL2.scandl_module.original_scandl.utils import dist_util, logger +from ScanDL2.scandl_module.original_scandl.step_sample import create_named_schedule_sampler +from ScanDL2.scandl_module.scripts.sp_basic_utils import ( + load_defaults_config, + create_model_and_diffusion, + args_to_dict, + add_dict_to_argparser, +) +from ScanDL2.scandl_module.scripts.sp_train_util import TrainLoop +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import ( + load_celer, + load_celer_speakers, + process_celer, + celer_zuco_dataset_and_loader, +) +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import get_kfold, get_kfold_indices_combined +from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import flatten_data, unflatten_data + + +def create_argparser(): + """Loads the config from the file scandl/config.json and adds all keys and values in the config dict + to the argument parser where config values are the argparse arguments' default values.""" + + defaults = dict( + checkpoint_path="", + vocab="bert", + use_plm_init="no", + lr=1e-4, + batch_size=64, + microbatch=64, + diffusion_steps=2000, + noise_schedule="sqrt", + schedule_sampler="lossaware", + seq_len=128, + resume_checkpoint="none", + hidden_t_dim=128, + seed=101, + hidden_dim=256, + learning_steps=80000, + save_interval=5000, + # config_name='bert-base-cased', + notes="-", + data_split_criterion="", + num_transformer_layers=12, + num_transformer_heads=8, + corpus="", + inference="", + load_train_data="-", + ) + defaults.update(load_defaults_config()) + parser = argparse.ArgumentParser() + add_dict_to_argparser(parser, defaults) # update latest args according to argparse + return parser + + +def main(): + args = create_argparser().parse_args() + set_seed(args.seed) + + assert args.seq_len == args.hidden_t_dim + + # set up distributed processing group + dist_util.setup_dist() + logger.configure() + logger.log("### Creating data loader...") + + rank = dist.get_rank() or 0 + + tokenizer = BertTokenizerFast.from_pretrained(args.config_name) + + args.vocab_size = tokenizer.vocab_size + + if rank == 0: + if not os.path.exists(args.checkpoint_path): + os.makedirs(args.checkpoint_path) + + # load train data + train = load_from_disk(os.path.join("..", args.load_train_data, "train")) + train_data = DatasetDict() + train_data["train"] = train + print("\t\t--- loaded train data ---") + + train_loader = celer_zuco_dataset_and_loader( + data=train_data, + data_args=args, + split="train", + ) + + logger.log("### Creating model and diffusion...") + if torch.cuda.is_available(): + print("#" * 30, "CUDA_VISIBLE_DEVICES", os.environ["CUDA_VISIBLE_DEVICES"]) + + model, diffusion = create_model_and_diffusion( + **args_to_dict(args, load_defaults_config().keys()) + ) + model.to(dist_util.dev()) + + pytorch_total_params = sum(p.numel() for p in model.parameters()) + + logger.log(f"### The parameter count is {pytorch_total_params}") + # args.schedule_sampler = lossaware + schedule_sampler = create_named_schedule_sampler(args.schedule_sampler, diffusion) + + logger.log(f"### Saving the hyperparameters to {args.checkpoint_path}/training_args.json") + with open(f"{args.checkpoint_path}/training_args.json", "w") as f: + json.dump(args.__dict__, f, indent=2) + + # if ('LOCAL_RANK' not in os.environ) or (int(os.environ['LOCAL_RANK']) == 0): + # wandb.init( + # project=os.getenv("WANDB_PROJECT", "ScanDL"), + # name=args.checkpoint_path, + # ) + # wandb.config.update(args.__dict__, allow_val_change=True) + + logger.log("### Training...") + + TrainLoop( + model=model, + diffusion=diffusion, + data=train_loader, + batch_size=args.batch_size, + microbatch=args.microbatch, + lr=args.lr, + ema_rate=args.ema_rate, + log_interval=args.log_interval, + save_interval=args.save_interval, + resume_checkpoint=args.resume_checkpoint, + use_fp16=args.use_fp16, + fp16_scale_growth=args.fp16_scale_growth, + schedule_sampler=schedule_sampler, + weight_decay=args.weight_decay, + learning_steps=args.learning_steps, + checkpoint_path=args.checkpoint_path, + gradient_clipping=args.gradient_clipping, + # eval_data=val_loader, + eval_interval=args.eval_interval, + ).run_loop() + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scandl_module/scripts/sp_train_util.py b/scandl_module/scripts/sp_train_util.py new file mode 100644 index 0000000000000000000000000000000000000000..a35d0a26fec82bd6f1c5240a609b59bfc5ac0cf0 --- /dev/null +++ b/scandl_module/scripts/sp_train_util.py @@ -0,0 +1,443 @@ +import copy +import functools +import os + +import blobfile as bf +import numpy as np +import torch as th +import torch.distributed as dist +from torch.nn.parallel.distributed import DistributedDataParallel as DDP +from torch.optim import AdamW + +from ScanDL2.scandl_module.original_scandl.utils import dist_util, logger +from ScanDL2.scandl_module.original_scandl.utils.fp16_util import ( + make_master_params, + master_params_to_model_params, + model_grads_to_master_grads, + unflatten_master_params, + zero_grad, +) +from ScanDL2.scandl_module.original_scandl.utils.nn import update_ema +from ScanDL2.scandl_module.original_scandl.step_sample import LossAwareSampler, UniformSampler + +INITIAL_LOG_LOSS_SCALE = 20.0 + + +class TrainLoop: + def __init__( + self, + *, + model, + diffusion, + data, + batch_size, + microbatch, + lr, + ema_rate, + log_interval, + save_interval, + resume_checkpoint, + use_fp16=False, + fp16_scale_growth=1e-3, + schedule_sampler=None, + weight_decay=0.0, + learning_steps=0, + checkpoint_path="", + gradient_clipping=-1.0, + eval_data=None, + eval_interval=-1, + ): + self.model = model + self.diffusion = diffusion + self.data = data + self.eval_data = eval_data + self.batch_size = batch_size + self.microbatch = microbatch if microbatch > 0 else batch_size + self.lr = lr + self.ema_rate = ( + [ema_rate] if isinstance(ema_rate, float) else [float(x) for x in ema_rate.split(",")] + ) + self.log_interval = log_interval + self.eval_interval = eval_interval + self.save_interval = save_interval + self.resume_checkpoint = resume_checkpoint + self.use_fp16 = use_fp16 + self.fp16_scale_growth = fp16_scale_growth + self.schedule_sampler = schedule_sampler or UniformSampler(diffusion) + self.weight_decay = weight_decay + self.learning_steps = learning_steps + self.gradient_clipping = gradient_clipping + + self.step = 0 + self.resume_step = 0 + self.global_batch = self.batch_size * dist.get_world_size() + + self.model_params = list(self.model.parameters()) + self.master_params = self.model_params + self.lg_loss_scale = INITIAL_LOG_LOSS_SCALE + self.sync_cuda = th.cuda.is_available() + + self.checkpoint_path = checkpoint_path # DEBUG ** + + self._load_and_sync_parameters() + if self.use_fp16: + self._setup_fp16() + + self.opt = AdamW(self.master_params, lr=self.lr, weight_decay=self.weight_decay) + if self.resume_step: + # self._load_optimizer_state() + frac_done = (self.step + self.resume_step) / self.learning_steps + lr = self.lr * (1 - frac_done) + self.opt = AdamW(self.master_params, lr=lr, weight_decay=self.weight_decay) + # Model was resumed, either due to a restart or a checkpoint + # being specified at the command line. + self.ema_params = [self._load_ema_parameters(rate) for rate in self.ema_rate] + else: + self.ema_params = [copy.deepcopy(self.master_params) for _ in range(len(self.ema_rate))] + + if th.cuda.is_available(): # DEBUG ** + self.use_ddp = True + print(dist_util.dev()) + # Distributed Data Parallel + self.ddp_model = DDP( + self.model, + device_ids=[dist_util.dev()], + output_device=dist_util.dev(), + broadcast_buffers=False, + bucket_cap_mb=128, + find_unused_parameters=False, + ) + else: + if dist.get_world_size() > 1: + logger.warn( + "Distributed training requires CUDA. " + "Gradients will not be synchronized properly!" + ) + self.use_ddp = False + self.ddp_model = self.model + + def _load_and_sync_parameters(self): + resume_checkpoint = find_resume_checkpoint() or self.resume_checkpoint + + if resume_checkpoint[-3:] == ".pt": + self.resume_step = parse_resume_step_from_filename(resume_checkpoint) + if dist.get_rank() == 0: + logger.log(f"loading model from checkpoint: {resume_checkpoint}...") + self.model.load_state_dict( + dist_util.load_state_dict( + actual_model_path(resume_checkpoint), map_location=dist_util.dev() + ) + ) + + dist_util.sync_params(self.model.parameters()) + + def _load_ema_parameters(self, rate): + ema_params = copy.deepcopy(self.master_params) + + main_checkpoint = find_resume_checkpoint() or self.resume_checkpoint + ema_checkpoint = find_ema_checkpoint(main_checkpoint, self.resume_step, rate) + if ema_checkpoint: + if dist.get_rank() == 0: + logger.log(f"loading EMA from checkpoint: {ema_checkpoint}...") + state_dict = dist_util.load_state_dict( + actual_model_path(ema_checkpoint), map_location=dist_util.dev() + ) + ema_params = self._state_dict_to_master_params(state_dict) + + dist_util.sync_params(ema_params) + return ema_params + + def _load_optimizer_state(self): + main_checkpoint = find_resume_checkpoint() or self.resume_checkpoint + if bf.exists(main_checkpoint): + logger.log(f"loading optimizer state from checkpoint: {main_checkpoint}") + state_dict = dist_util.load_state_dict( + actual_model_path(main_checkpoint), map_location=dist_util.dev() + ) + self.opt.load_state_dict(state_dict) + + def _setup_fp16(self): + self.master_params = make_master_params(self.model_params) + self.model.convert_to_fp16() + + def run_loop(self): + while not self.learning_steps or self.step + self.resume_step < self.learning_steps: + batch = next(self.data) + self.run_step(batch) + if self.step % self.log_interval == 0: + logger.dumpkvs() + if self.eval_data is not None and self.step % self.eval_interval == 0: + batch_eval = next(self.eval_data) + self.forward_only(batch_eval) + print("eval on validation set") + logger.dumpkvs() + if self.step > 0 and self.step % self.save_interval == 0: + self.save() + # Run for a finite amount of time in integration tests. + if os.environ.get("DIFFUSION_TRAINING_TEST", "") and self.step > 0: + return + self.step += 1 + # Save the last checkpoint if it wasn't already saved. + if (self.step - 1) % self.save_interval != 0: + self.save() + + def run_step( + self, + batch, + ): + self.forward_backward(batch) + if self.use_fp16: + self.optimize_fp16() + else: + self.optimize_normal() + self.log_step() + + def forward_only(self, batch): + with th.no_grad(): + zero_grad(self.model_params) + + for i in range(0, batch["sn_sp_repr"].shape[0], self.microbatch): + + mask = batch["mask"][i : i + self.microbatch].to(dist_util.dev()) + sn_sp_repr = batch["sn_sp_repr"][i : i + self.microbatch].to(dist_util.dev()) + sn_input_ids = batch["sn_input_ids"][i : i + self.microbatch].to(dist_util.dev()) + indices_pos_enc = batch["indices_pos_enc"][i : i + self.microbatch].to( + dist_util.dev() + ) + mask_sn_padding = batch["mask_sn_padding"][i : i + self.microbatch].to( + dist_util.dev() + ) + mask_transformer_att = batch["mask_transformer_att"][i : i + self.microbatch].to( + dist_util.dev() + ) + + last_batch = (i + self.microbatch) >= sn_sp_repr.shape[0] + t, weights = self.schedule_sampler.sample(sn_sp_repr.shape[0], dist_util.dev()) + compute_losses = functools.partial( + self.diffusion.training_losses, + self.ddp_model, + t, + sn_sp_repr, + mask, + sn_input_ids, + indices_pos_enc, + mask_sn_padding, + mask_transformer_att, + ) + + if last_batch or not self.use_ddp: + losses = compute_losses() + else: + with self.ddp_model.no_sync(): + losses = compute_losses() + + log_loss_dict( + self.diffusion, t, {f"eval_{k}": v * weights for k, v in losses.items()} + ) + + def forward_backward( + self, + batch, + ): + zero_grad(self.model_params) + + for i in range(0, batch["sn_sp_repr"].shape[0], self.microbatch): + + mask = batch["mask"][i : i + self.microbatch].to(dist_util.dev()) + sn_sp_repr = batch["sn_sp_repr"][i : i + self.microbatch].to(dist_util.dev()) + sn_input_ids = batch["sn_input_ids"][i : i + self.microbatch].to(dist_util.dev()) + indices_pos_enc = batch["indices_pos_enc"][i : i + self.microbatch].to(dist_util.dev()) + mask_sn_padding = batch["mask_sn_padding"][i : i + self.microbatch].to(dist_util.dev()) + mask_transformer_att = batch["mask_transformer_att"][i : i + self.microbatch].to( + dist_util.dev() + ) + + last_batch = (i + self.microbatch) >= sn_sp_repr.shape[0] + + # the indices in t are the number of noising steps; how many times is noise added to each instance + t, weights = self.schedule_sampler.sample(sn_sp_repr.shape[0], dist_util.dev()) + # print(micro_cond.keys()) + compute_losses = functools.partial( + self.diffusion.training_losses, + self.ddp_model, # the transformer model + t, # the number of times to add noise to each instance + sn_sp_repr, + mask, + sn_input_ids, + indices_pos_enc, + mask_sn_padding, + mask_transformer_att, + ) + + if last_batch or not self.use_ddp: + losses = compute_losses() + else: + with self.ddp_model.no_sync(): + losses = compute_losses() + + if isinstance(self.schedule_sampler, LossAwareSampler): + self.schedule_sampler.update_with_local_losses(t, losses["loss"].detach()) + + # weight the losses with what the schedule sampler returned + loss = (losses["loss"] * weights).mean() + log_loss_dict(self.diffusion, t, {k: v * weights for k, v in losses.items()}) + if self.use_fp16: + loss_scale = 2**self.lg_loss_scale + (loss * loss_scale).backward() + else: + loss.backward() + + def optimize_fp16(self): + if any(not th.isfinite(p.grad).all() for p in self.model_params): + self.lg_loss_scale -= 1 + logger.log(f"Found NaN, decreased lg_loss_scale to {self.lg_loss_scale}") + return + + model_grads_to_master_grads(self.model_params, self.master_params) + self.master_params[0].grad.mul_(1.0 / (2**self.lg_loss_scale)) + self._log_grad_norm() + self._anneal_lr() + self.opt.step() + for rate, params in zip(self.ema_rate, self.ema_params): + update_ema(params, self.master_params, rate=rate) + master_params_to_model_params(self.model_params, self.master_params) + self.lg_loss_scale += self.fp16_scale_growth + + def grad_clip(self): + # print('doing gradient clipping') + max_grad_norm = self.gradient_clipping # 3.0 + if hasattr(self.opt, "clip_grad_norm"): + # Some optimizers (like the sharded optimizer) have a specific way to do gradient clipping + self.opt.clip_grad_norm(max_grad_norm) + # else: + # assert False + # elif hasattr(self.model, "clip_grad_norm_"): + # # Some models (like FullyShardedDDP) have a specific way to do gradient clipping + # self.model.clip_grad_norm_(args.max_grad_norm) + else: + # Revert to normal clipping otherwise, handling Apex or full precision + th.nn.utils.clip_grad_norm_( + self.model.parameters(), # amp.master_params(self.opt) if self.use_apex else + max_grad_norm, + ) + + def optimize_normal(self): + if self.gradient_clipping > 0: + self.grad_clip() + + # log the gradient norm and the learning rate + self._log_grad_norm() + self._anneal_lr() + self.opt.step() + for rate, params in zip(self.ema_rate, self.ema_params): + update_ema(params, self.master_params, rate=rate) + + def _log_grad_norm(self): + sqsum = 0.0 + # cnt = 0 + for p in self.master_params: + # print(cnt, p) ## DEBUG + # print(cnt, p.grad) + # cnt += 1 + if p.grad is not None: + sqsum += (p.grad**2).sum().item() + logger.logkv_mean("grad_norm", np.sqrt(sqsum)) + + def _anneal_lr(self): + if not self.learning_steps: + return + frac_done = (self.step + self.resume_step) / self.learning_steps + lr = self.lr * (1 - frac_done) + for param_group in self.opt.param_groups: + param_group["lr"] = lr + + def log_step(self): + logger.logkv("step", self.step + self.resume_step) + logger.logkv("samples", (self.step + self.resume_step + 1) * self.global_batch) + if self.use_fp16: + logger.logkv("lg_loss_scale", self.lg_loss_scale) + + def save(self): + def save_checkpoint(rate, params): + state_dict = self._master_params_to_state_dict(params) + if dist.get_rank() == 0: + logger.log(f"saving model {rate}...") + if not rate: + filename = f"model{(self.step+self.resume_step):06d}.pt" + else: + filename = f"ema_{rate}_{(self.step+self.resume_step):06d}.pt" + print("writing to", bf.join(get_blob_logdir(), filename)) + print("writing to", bf.join(self.checkpoint_path, filename)) + # with bf.BlobFile(bf.join(get_blob_logdir(), filename), "wb") as f: + # th.save(state_dict, f) + with bf.BlobFile(bf.join(self.checkpoint_path, filename), "wb") as f: # DEBUG ** + th.save(state_dict, f) # save locally + # pass # save empty + + # save_checkpoint(0, self.master_params) + for rate, params in zip(self.ema_rate, self.ema_params): + save_checkpoint(rate, params) + + dist.barrier() + + def _master_params_to_state_dict(self, master_params): + if self.use_fp16: + master_params = unflatten_master_params( + list(self.model.parameters()), master_params # DEBUG ** + ) + state_dict = self.model.state_dict() + for i, (name, _value) in enumerate(self.model.named_parameters()): + assert name in state_dict + state_dict[name] = master_params[i] + return state_dict + + def _state_dict_to_master_params(self, state_dict): + params = [state_dict[name] for name, _ in self.model.named_parameters()] + if self.use_fp16: + return make_master_params(params) + else: + return params + + +def parse_resume_step_from_filename(filename): + """ + Parse filenames of the form path/to/modelNNNNNN.pt, where NNNNNN is the + checkpoint's number of steps. + """ + if filename[-3:] == ".pt": + return int(filename[-9:-3]) + else: + return 0 + + +def get_blob_logdir(): + return os.environ.get("DIFFUSION_BLOB_LOGDIR", logger.get_dir()) + + +def find_resume_checkpoint(): + # On your infrastructure, you may want to override this to automatically + # discover the latest checkpoint on your blob storage, etc. + return None + + +def find_ema_checkpoint(main_checkpoint, step, rate): + if main_checkpoint is None: + return None + filename = f"ema_{rate}_{(step):06d}.pt" + path = bf.join(bf.dirname(main_checkpoint), filename) + if bf.exists(path): + return path + return None + + +def log_loss_dict(diffusion, ts, losses): + for key, values in losses.items(): + logger.logkv_mean(key, values.mean().item()) + # Log the quantiles (four quartiles, in particular). + for sub_t, sub_loss in zip(ts.cpu().numpy(), values.detach().cpu().numpy()): + quartile = int(4 * sub_t / diffusion.num_timesteps) + logger.logkv_mean(f"{key}_q{quartile}", sub_loss) + + +def actual_model_path(model_path): + return model_path diff --git a/tests/test_imports.py b/tests/test_imports.py new file mode 100644 index 0000000000000000000000000000000000000000..a7a54f689ed7c6b53d7c1a6dff42ac85ed4f6bb0 --- /dev/null +++ b/tests/test_imports.py @@ -0,0 +1,112 @@ +import ast +import sys +import unittest +from pathlib import Path + +SCANDL2_ROOT = Path(__file__).resolve().parents[1] +PROJECT_ROOT = SCANDL2_ROOT.parent +EXCLUDED_DIRS = {"__pycache__", ".git", "tests"} +EXPECTED_IMPORT_FAILURES = { + ("ScanDL2/app.py", "import gradio as gr"), + ( + "ScanDL2/create_data.py", + "from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_emtec, process_emtec", + ), + ( + "ScanDL2/create_data.py", + "from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_bsc, process_bsc", + ), + ( + "ScanDL2/fix_dur_module/train_seq2seq.py", + "from ScanDL2.CONSTANTS import COMPLETE_FIXDUR_MODULE_TRAIN_PATH_BSC", + ), + ( + "ScanDL2/scandl_module/scripts/sp_run_train.py", + "from ScanDL2.CONSTANTS import (\n COMPLETE_SCANDL_MODULE_TRAIN_PATH_BSC", + ), + ("ScanDL2/scandl_module/original_scandl/utils/logger.py", "import tensorflow as tf"), + ( + "ScanDL2/scandl_module/original_scandl/utils/logger.py", + "from tensorflow.python import pywrap_tensorflow", + ), + ( + "ScanDL2/scandl_module/original_scandl/utils/logger.py", + "from tensorflow.core.util import event_pb2", + ), + ( + "ScanDL2/scandl_module/original_scandl/utils/logger.py", + "from tensorflow.python.util import compat", + ), +} + +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + + +class ScanDL2ImportTests(unittest.TestCase): + def test_import_lines(self): + failures = [] + + for file_path in _python_files(SCANDL2_ROOT): + source = file_path.read_text() + tree = ast.parse(source, filename=str(file_path)) + + for node in ast.walk(tree): + if not isinstance(node, (ast.Import, ast.ImportFrom)): + continue + if isinstance(node, ast.ImportFrom) and node.module == "__future__": + continue + + import_line = ast.get_source_segment(source, node) + try: + exec( + compile(import_line, str(file_path), "exec"), + _import_globals(file_path), + ) + except Exception as exc: + relative_path = str(file_path.relative_to(PROJECT_ROOT)) + if _is_expected_failure(relative_path, import_line): + continue + failures.append( + f"{relative_path}:{node.lineno}\n" + f"{import_line}\n" + f"{type(exc).__name__}: {exc}" + ) + + if failures: + self.fail("Failed import line(s):\n\n" + "\n\n".join(failures)) + + +def _python_files(root): + for file_path in root.rglob("*.py"): + if any(part in EXCLUDED_DIRS for part in file_path.parts): + continue + yield file_path + + +def _import_globals(file_path): + module_path = file_path.relative_to(PROJECT_ROOT).with_suffix("") + module_parts = module_path.parts + + if module_parts[-1] == "__init__": + module_name = ".".join(module_parts[:-1]) + package = module_name + else: + module_name = ".".join(module_parts) + package = ".".join(module_parts[:-1]) + + return { + "__name__": module_name, + "__package__": package, + } + + +def _is_expected_failure(relative_path, import_line): + return any( + relative_path == expected_path and import_line.startswith(expected_import) + for expected_path, expected_import in EXPECTED_IMPORT_FAILURES + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_model.py b/tests/test_model.py new file mode 100644 index 0000000000000000000000000000000000000000..f4397f55902010c8c76fcc6bb7632a7258b72ea4 --- /dev/null +++ b/tests/test_model.py @@ -0,0 +1,84 @@ +import sys +import unittest +from pprint import pprint +from pathlib import Path +from unittest.mock import patch + +SCANDL2_ROOT = Path(__file__).resolve().parents[1] +PROJECT_ROOT = SCANDL2_ROOT.parent +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + + +class ScanDL2SmokeTests(unittest.TestCase): + def test_sentence_model_runs_real_inference(self): + _skip_if_sentence_assets_are_missing() + + from ScanDL2.model import ScanDL2 + + try: + with patch.object(sys, "argv", [sys.argv[0]]): + model = ScanDL2(text_type="sentence", bsz=1, save=None, filename=None) + model.eval() + output = model(["The quick brown fox jumps."]) + except PermissionError as exc: + raise unittest.SkipTest(f"Real ScanDL2 smoke test needs socket access: {exc}") from exc + + print("\nScanDL2 output:") + pprint(output) + self.assertIsInstance(output, dict) + self.assertIn("predicted_sp_words", output) + self.assertIn("predicted_sp_ids", output) + self.assertIn("original_sn", output) + self.assertIn("predicted_fix_durs", output) + self.assertIn("unique_idx", output) + + def test_scandl_and_fixdur_modules_run_real_inference(self): + _skip_if_sentence_assets_are_missing() + + from ScanDL2.model import FixdurModule, ScanDLModule + + try: + with patch.object(sys, "argv", [sys.argv[0]]): + scandl_module = ScanDLModule(text_type="sentence", bsz=1) + scandl_output = scandl_module(texts=["The quick brown fox jumps."]) + + fixdur_module = FixdurModule(text_type="sentence", bsz=1) + fixdur_output = fixdur_module(scandl_module_output=scandl_output) + except PermissionError as exc: + raise unittest.SkipTest(f"Real ScanDL2 smoke test needs socket access: {exc}") from exc + + print("\nScanDLModule output:") + pprint(scandl_output) + print("\nFixdurModule output:") + pprint(fixdur_output) + + self.assertIsInstance(scandl_output, dict) + self.assertIn("predicted_sp_words", scandl_output) + self.assertIn("predicted_sp_ids", scandl_output) + self.assertIn("original_sn", scandl_output) + self.assertIn("unique_idx", scandl_output) + + self.assertIsInstance(fixdur_output, dict) + self.assertIn("predicted_sp_words", fixdur_output) + self.assertIn("predicted_sp_ids", fixdur_output) + self.assertIn("original_sn", fixdur_output) + self.assertIn("predicted_fix_durs", fixdur_output) + self.assertIn("unique_idx", fixdur_output) + + +def _skip_if_sentence_assets_are_missing(): + required_dirs = [ + SCANDL2_ROOT / "models" / "sentence" / "scandl-module", + SCANDL2_ROOT / "models" / "sentence" / "fixdur-module", + ] + missing = [path for path in required_dirs if not path.exists() or not any(path.iterdir())] + if missing: + raise unittest.SkipTest( + "Missing ScanDL2 sentence model assets: " + + ", ".join(str(path.relative_to(PROJECT_ROOT)) for path in missing) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/training_utils.py b/training_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..1576339adde14295d7c6280c617c121d48a75ecf --- /dev/null +++ b/training_utils.py @@ -0,0 +1,24 @@ +from ScanDL2.scandl_module.original_scandl.step_sample import ( + LossSecondMomentResampler, + UniformSampler, + create_named_schedule_sampler, +) +from ScanDL2.scandl_module.scripts.sp_train_util import TrainLoop +from ScanDL2.fix_dur_module.utils_data import ( + Seq2SeqDatasetHP, + prepare_seq2seq_data_hp, + split_train_val_data, +) +from ScanDL2.fix_dur_module.utils_train import EarlyStopping, train + +__all__ = [ + "train", + "TrainLoop", + "EarlyStopping", + "create_named_schedule_sampler", + "UniformSampler", + "LossSecondMomentResampler", + "split_train_val_data", + "Seq2SeqDatasetHP", + "prepare_seq2seq_data_hp", +] diff --git a/utils.py b/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..c9e9310f6c26368fea424eced090ba948604f285 --- /dev/null +++ b/utils.py @@ -0,0 +1,39 @@ +from ScanDL2.scandl_module.original_scandl.utils.dist_util import load_state_dict +from ScanDL2.fix_dur_module.utils_data import padding_and_mask_seq2seq, aggregate_input_embeddings +from ScanDL2.scandl2_utils import TextDataset +from ScanDL2.scandl2_utils import text_dataset_loader +from ScanDL2.scandl2_utils import FixdurDataset +from ScanDL2.scandl_module.scripts.sp_basic_utils import ( + create_model_and_diffusion, + load_defaults_config, + add_dict_to_argparser, + args_to_dict, +) + +from ScanDL2.fix_dur_module.utils_data import ( + get_embeddings_seq2seq, + prepare_seq2seq_data, + get_embeddings_seq2seq_hp, + prepare_seq2seq_data_hp, + Seq2SeqDataset, + Seq2SeqDatasetHP, +) + +from ScanDL2.fix_dur_module.scasim import scasim + +__all__ = [ + "padding_and_mask_seq2seq", + "aggregate_input_embeddings", + "TextDataset", + "text_dataset_loader", + "FixdurDataset", + "load_defaults_config", + "add_dict_to_argparser", + "args_to_dict", + "get_embeddings_seq2seq", + "prepare_seq2seq_data", + "Seq2SeqDataset", + "scasim", + "create_model_and_diffusion", + "load_state_dict", +]