Di0nigi commited on
Commit
95456ed
·
verified ·
1 Parent(s): cc00ba3

First commit

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitignore +8 -0
  2. CITATION.cff +27 -0
  3. CONSTANTS.py +62 -0
  4. LICENSE +121 -0
  5. PATHS.py +4 -0
  6. README.md +212 -0
  7. __init__.py +27 -0
  8. app.py +185 -0
  9. config.json +53 -0
  10. config_bsc.json +54 -0
  11. config_emtec.json +54 -0
  12. create_data.py +167 -0
  13. fix_dur_module/__init__.py +0 -0
  14. fix_dur_module/__pycache__/__init__.cpython-313.pyc +0 -0
  15. fix_dur_module/__pycache__/model_seq2seq.cpython-313.pyc +0 -0
  16. fix_dur_module/__pycache__/scasim.cpython-313.pyc +0 -0
  17. fix_dur_module/__pycache__/utils_data.cpython-313.pyc +0 -0
  18. fix_dur_module/__pycache__/utils_train.cpython-313.pyc +0 -0
  19. fix_dur_module/model_seq2seq.py +89 -0
  20. fix_dur_module/scasim.py +185 -0
  21. fix_dur_module/train_seq2seq.py +283 -0
  22. fix_dur_module/utils_data.py +530 -0
  23. fix_dur_module/utils_train.py +195 -0
  24. handler.py +53 -0
  25. model.py +701 -0
  26. models/paragraph/fixdur-module/hyperparameters.json +1 -0
  27. models/paragraph/fixdur-module/min_max_scaler.pkl +3 -0
  28. models/paragraph/fixdur-module/seq2seq_fixdur.pt +3 -0
  29. models/paragraph/scandl-module/ema_0.9999_080000.pt +3 -0
  30. models/paragraph/scandl-module/training_args.json +52 -0
  31. models/sentence/fixdur-module/hyperparameters.json +1 -0
  32. models/sentence/fixdur-module/min_max_scaler.pkl +3 -0
  33. models/sentence/fixdur-module/seq2seq_fixdur.pt +3 -0
  34. models/sentence/scandl-module/ema_0.9999_080000.pt +3 -0
  35. models/sentence/scandl-module/training_args.json +52 -0
  36. requirements.txt +15 -0
  37. scandl2_utils.py +77 -0
  38. scandl_module/.DS_Store +0 -0
  39. scandl_module/__init__.py +4 -0
  40. scandl_module/__pycache__/__init__.cpython-313.pyc +0 -0
  41. scandl_module/original_scandl/__init__.py +0 -0
  42. scandl_module/original_scandl/__pycache__/__init__.cpython-313.pyc +0 -0
  43. scandl_module/original_scandl/__pycache__/sp_gaussian_diffusion.cpython-313.pyc +0 -0
  44. scandl_module/original_scandl/__pycache__/sp_rounding.cpython-313.pyc +0 -0
  45. scandl_module/original_scandl/__pycache__/sp_transformer_model.cpython-313.pyc +0 -0
  46. scandl_module/original_scandl/__pycache__/step_sample.cpython-313.pyc +0 -0
  47. scandl_module/original_scandl/config.json +52 -0
  48. scandl_module/original_scandl/sp_gaussian_diffusion.py +1183 -0
  49. scandl_module/original_scandl/sp_rounding.py +58 -0
  50. scandl_module/original_scandl/sp_transformer_model.py +167 -0
.gitignore ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ models
2
+ test_results
3
+
4
+ # Python bytecode/cache files
5
+ __pycache__/
6
+ **/__pycache__/
7
+ *.py[cod]
8
+ *$py.class
CITATION.cff ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ cff-version: 1.2.0
2
+ message: "If you use this work, please cite it as below."
3
+ title: "ScanDL 2.0: A Generative Model of Eye Movements in Reading Synthesizing Scanpaths and Fixation Durations"
4
+ authors:
5
+ - family-names: Bolliger
6
+ given-names: Lena S.
7
+ - family-names: Reich
8
+ given-names: David R.
9
+ - family-names: Jäger
10
+ given-names: Lena A.
11
+ date-released: 2025-05-01
12
+ journal: "Proceedings of the ACM on Human-Computer Interaction"
13
+ publisher: "Association for Computing Machinery"
14
+ location: "New York, NY, USA"
15
+ volume: "9"
16
+ issue: "ETRA5"
17
+ article-number: "5"
18
+ numpages: "30"
19
+ doi: "10.1145/3725830"
20
+ url: "https://doi.org/10.1145/3725830"
21
+ 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."
22
+ keywords:
23
+ - neural networks
24
+ - scanpath generation
25
+ - eye movements
26
+ - reading
27
+ - diffusion models
CONSTANTS.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ##### data paths #####
2
+
3
+ # TODO adapt your paths to folder that contains the celer and zuco folders
4
+ path_to_celer = "/data/lenbol/data/" # e.g., path_to_celer = 'data/' if 'data/celer/...'
5
+ path_to_zuco = "/data/lenbol/data/" # e.g., path_to_zuco = 'data/' if 'data/zuco/...'
6
+ path_to_emtec = "/data/lenbol/data/" # e.g., path_to_copco = 'data/' if 'data/copco/...'
7
+ path_to_bsc = "/data/lenbol/data/" # e.g., path_to_copco = 'data/' if 'data/BSC/...'
8
+
9
+ PATH_TO_FIX = f"{path_to_celer}CELER/data_v2.0/sent_fix.tsv"
10
+ PATH_TO_IA = f"{path_to_celer}CELER/data_v2.0/sent_ia.tsv"
11
+ SUB_METADATA_PATH = f"{path_to_celer}/CELER/participant_metadata/metadata.tsv"
12
+ PATH_TO_EMTEC_FIX = f"{path_to_emtec}/EMTeC/fixations_corrected.csv"
13
+ PATH_TO_EMTEC_STIM = f"{path_to_emtec}/EMTeC/stimuli.csv"
14
+ PATH_TO_BSC_WORD = f"{path_to_bsc}/BSC/BSC.Word.Info.v2.xlsx"
15
+ PATH_TO_BSC_FIX = f"{path_to_bsc}/BSC/BSC.EMD.txt"
16
+
17
+
18
+ ##### model paths #####
19
+
20
+ # TODO adapt your paths
21
+
22
+ # training of original ScanDL for modular use with seq2seq fixdur module
23
+ SCANDL_MODULE_TRAIN_PATH = ""
24
+ SCANDL_MODULE_INF_PATH = ""
25
+
26
+ # training of the fixation duration module
27
+ FIXDUR_MODULE_TRAIN_PATH = ""
28
+ FIXDUR_MODULE_INF_PATH = ""
29
+
30
+ # training and inference of the diffusion-only architecture
31
+ DIFFUSION_ONLY_TRAIN_PATH = ""
32
+ DIFFUSION_ONLY_INF_PATH = ""
33
+
34
+
35
+ # names for EMTeC
36
+ # training of original ScanDL for modular use with seq2seq fixdur module
37
+ SCANDL_MODULE_TRAIN_PATH_EMTEC = ""
38
+ SCANDL_MODULE_INF_PATH_EMTEC = ""
39
+
40
+ # training of the fixation duration module
41
+ FIXDUR_MODULE_TRAIN_PATH_EMTEC = ""
42
+ FIXDUR_MODULE_INF_PATH_EMTEC = ""
43
+
44
+
45
+ # names for BSC
46
+ # training of original ScanDL for modular use with seq2seq fixdur module
47
+ SCANDL_MODULE_TRAIN_PATH_BSC = ""
48
+ SCANDL_MODULE_INF_PATH_BSC = ""
49
+
50
+ # training of the fixation duration module
51
+ FIXDUR_MODULE_TRAIN_PATH_BSC = ""
52
+ FIXDUR_MODULE_INF_PATH_BSC = ""
53
+
54
+
55
+ # training of ScanDL 2.0 on all EMTeC data for paragraph-level ScanDL 2.0
56
+ COMPLETE_SCANDL_MODULE_TRAIN_PATH_EMTEC = ""
57
+ COMPLETE_FIXDUR_MODULE_TRAIN_PATH_EMTEC = ""
58
+
59
+
60
+ # training of ScanDL 2.0 on all CELER data for sentence-level ScanDL 2.0
61
+ COMPLETE_SCANDL_MODULE_TRAIN_PATH_CELER = ""
62
+ COMPLETE_FIXDUR_MODULE_TRAIN_PATH_CELER = ""
LICENSE ADDED
@@ -0,0 +1,121 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Creative Commons Legal Code
2
+
3
+ CC0 1.0 Universal
4
+
5
+ CREATIVE COMMONS CORPORATION IS NOT A LAW FIRM AND DOES NOT PROVIDE
6
+ LEGAL SERVICES. DISTRIBUTION OF THIS DOCUMENT DOES NOT CREATE AN
7
+ ATTORNEY-CLIENT RELATIONSHIP. CREATIVE COMMONS PROVIDES THIS
8
+ INFORMATION ON AN "AS-IS" BASIS. CREATIVE COMMONS MAKES NO WARRANTIES
9
+ REGARDING THE USE OF THIS DOCUMENT OR THE INFORMATION OR WORKS
10
+ PROVIDED HEREUNDER, AND DISCLAIMS LIABILITY FOR DAMAGES RESULTING FROM
11
+ THE USE OF THIS DOCUMENT OR THE INFORMATION OR WORKS PROVIDED
12
+ HEREUNDER.
13
+
14
+ Statement of Purpose
15
+
16
+ The laws of most jurisdictions throughout the world automatically confer
17
+ exclusive Copyright and Related Rights (defined below) upon the creator
18
+ and subsequent owner(s) (each and all, an "owner") of an original work of
19
+ authorship and/or a database (each, a "Work").
20
+
21
+ Certain owners wish to permanently relinquish those rights to a Work for
22
+ the purpose of contributing to a commons of creative, cultural and
23
+ scientific works ("Commons") that the public can reliably and without fear
24
+ of later claims of infringement build upon, modify, incorporate in other
25
+ works, reuse and redistribute as freely as possible in any form whatsoever
26
+ and for any purposes, including without limitation commercial purposes.
27
+ These owners may contribute to the Commons to promote the ideal of a free
28
+ culture and the further production of creative, cultural and scientific
29
+ works, or to gain reputation or greater distribution for their Work in
30
+ part through the use and efforts of others.
31
+
32
+ For these and/or other purposes and motivations, and without any
33
+ expectation of additional consideration or compensation, the person
34
+ associating CC0 with a Work (the "Affirmer"), to the extent that he or she
35
+ is an owner of Copyright and Related Rights in the Work, voluntarily
36
+ elects to apply CC0 to the Work and publicly distribute the Work under its
37
+ terms, with knowledge of his or her Copyright and Related Rights in the
38
+ Work and the meaning and intended legal effect of CC0 on those rights.
39
+
40
+ 1. Copyright and Related Rights. A Work made available under CC0 may be
41
+ protected by copyright and related or neighboring rights ("Copyright and
42
+ Related Rights"). Copyright and Related Rights include, but are not
43
+ limited to, the following:
44
+
45
+ i. the right to reproduce, adapt, distribute, perform, display,
46
+ communicate, and translate a Work;
47
+ ii. moral rights retained by the original author(s) and/or performer(s);
48
+ iii. publicity and privacy rights pertaining to a person's image or
49
+ likeness depicted in a Work;
50
+ iv. rights protecting against unfair competition in regards to a Work,
51
+ subject to the limitations in paragraph 4(a), below;
52
+ v. rights protecting the extraction, dissemination, use and reuse of data
53
+ in a Work;
54
+ vi. database rights (such as those arising under Directive 96/9/EC of the
55
+ European Parliament and of the Council of 11 March 1996 on the legal
56
+ protection of databases, and under any national implementation
57
+ thereof, including any amended or successor version of such
58
+ directive); and
59
+ vii. other similar, equivalent or corresponding rights throughout the
60
+ world based on applicable law or treaty, and any national
61
+ implementations thereof.
62
+
63
+ 2. Waiver. To the greatest extent permitted by, but not in contravention
64
+ of, applicable law, Affirmer hereby overtly, fully, permanently,
65
+ irrevocably and unconditionally waives, abandons, and surrenders all of
66
+ Affirmer's Copyright and Related Rights and associated claims and causes
67
+ of action, whether now known or unknown (including existing as well as
68
+ future claims and causes of action), in the Work (i) in all territories
69
+ worldwide, (ii) for the maximum duration provided by applicable law or
70
+ treaty (including future time extensions), (iii) in any current or future
71
+ medium and for any number of copies, and (iv) for any purpose whatsoever,
72
+ including without limitation commercial, advertising or promotional
73
+ purposes (the "Waiver"). Affirmer makes the Waiver for the benefit of each
74
+ member of the public at large and to the detriment of Affirmer's heirs and
75
+ successors, fully intending that such Waiver shall not be subject to
76
+ revocation, rescission, cancellation, termination, or any other legal or
77
+ equitable action to disrupt the quiet enjoyment of the Work by the public
78
+ as contemplated by Affirmer's express Statement of Purpose.
79
+
80
+ 3. Public License Fallback. Should any part of the Waiver for any reason
81
+ be judged legally invalid or ineffective under applicable law, then the
82
+ Waiver shall be preserved to the maximum extent permitted taking into
83
+ account Affirmer's express Statement of Purpose. In addition, to the
84
+ extent the Waiver is so judged Affirmer hereby grants to each affected
85
+ person a royalty-free, non transferable, non sublicensable, non exclusive,
86
+ irrevocable and unconditional license to exercise Affirmer's Copyright and
87
+ Related Rights in the Work (i) in all territories worldwide, (ii) for the
88
+ maximum duration provided by applicable law or treaty (including future
89
+ time extensions), (iii) in any current or future medium and for any number
90
+ of copies, and (iv) for any purpose whatsoever, including without
91
+ limitation commercial, advertising or promotional purposes (the
92
+ "License"). The License shall be deemed effective as of the date CC0 was
93
+ applied by Affirmer to the Work. Should any part of the License for any
94
+ reason be judged legally invalid or ineffective under applicable law, such
95
+ partial invalidity or ineffectiveness shall not invalidate the remainder
96
+ of the License, and in such case Affirmer hereby affirms that he or she
97
+ will not (i) exercise any of his or her remaining Copyright and Related
98
+ Rights in the Work or (ii) assert any associated claims and causes of
99
+ action with respect to the Work, in either case contrary to Affirmer's
100
+ express Statement of Purpose.
101
+
102
+ 4. Limitations and Disclaimers.
103
+
104
+ a. No trademark or patent rights held by Affirmer are waived, abandoned,
105
+ surrendered, licensed or otherwise affected by this document.
106
+ b. Affirmer offers the Work as-is and makes no representations or
107
+ warranties of any kind concerning the Work, express, implied,
108
+ statutory or otherwise, including without limitation warranties of
109
+ title, merchantability, fitness for a particular purpose, non
110
+ infringement, or the absence of latent or other defects, accuracy, or
111
+ the present or absence of errors, whether or not discoverable, all to
112
+ the greatest extent permissible under applicable law.
113
+ c. Affirmer disclaims responsibility for clearing rights of other persons
114
+ that may apply to the Work or any use thereof, including without
115
+ limitation any person's Copyright and Related Rights in the Work.
116
+ Further, Affirmer disclaims responsibility for obtaining any necessary
117
+ consents, permissions or other rights required for any use of the
118
+ Work.
119
+ d. Affirmer understands and acknowledges that Creative Commons is not a
120
+ party to this document and has no duty or obligation with respect to
121
+ this CC0 or use of the Work.
PATHS.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ SENT_SCANDL_MODULE = "ScanDL2/models/sentence/scandl-module/"
2
+ SENT_FIXDUR_MODULE = "ScanDL2/models/sentence/fixdur-module/"
3
+ PAR_SCANDL_MODULE = "ScanDL2/models/paragraph/scandl-module/"
4
+ PAR_FIXDUR_MODULE = "ScanDL2/models/paragraph/fixdur-module/"
README.md ADDED
@@ -0,0 +1,212 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ScanDL 2.0: A Generative Model of Eye Movements in Reading Synthesizing Scanpaths and Fixation Durations
2
+
3
+ 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.
4
+
5
+ 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.
6
+
7
+ ## Setup
8
+
9
+ 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.
10
+
11
+ ### Install requirements
12
+
13
+ The code uses PyTorch and Hugging Face libraries.
14
+
15
+ ```bash
16
+ python -m pip install -r ScanDL2/requirements.txt
17
+ ```
18
+
19
+ Install a PyTorch build appropriate for your platform, and install Gradio to use the web interface; neither is included in this requirements file.
20
+
21
+ ```bash
22
+ python -m pip install torch gradio
23
+ ```
24
+
25
+ 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.
26
+
27
+ ## Using pre-trained ScanDL 2.0
28
+
29
+ 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.
30
+
31
+ 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:
32
+
33
+ ```text
34
+ ScanDL2/models/
35
+ ├── sentence/
36
+ │ ├── scandl-module/
37
+ │ │ ├── ema_0.9999_080000.pt
38
+ │ │ └── training_args.json
39
+ │ └── fixdur-module/
40
+ │ ├── seq2seq_fixdur.pt
41
+ │ ├── hyperparameters.json
42
+ │ └── min_max_scaler.pkl
43
+ └── paragraph/
44
+ ├── scandl-module/ # same filenames as above
45
+ └── fixdur-module/ # same filenames as above
46
+ ```
47
+
48
+ ### Python example
49
+
50
+ ```python
51
+ import torch
52
+ from ScanDL2 import ScanDL2
53
+
54
+ model = ScanDL2(
55
+ text_type="sentence",
56
+ bsz=2,
57
+ save=None,
58
+ filename=None,
59
+ )
60
+ model.eval()
61
+ with torch.no_grad():
62
+ output = model(texts=["The quick brown fox jumps over the lazy dog."])
63
+
64
+ print(output)
65
+ ```
66
+
67
+ Set `text_type="paragraph"` to use the paragraph model.
68
+
69
+ ### Parameters
70
+
71
+ | Parameter | Default | Description |
72
+ | --- | --- | --- |
73
+ | `text_type` | `"sentence"` | Either `"sentence"` or `"paragraph"`; selects the corresponding pretrained modules |
74
+ | `bsz` | `2` | Inference batch size; adjust to available memory |
75
+ | `save` | `None` | Optional directory in which to save the output as JSON |
76
+ | `filename` | `None` | Optional output filename; defaults to `scandl2_outputs.json` when `save` is set |
77
+
78
+ 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.
79
+
80
+ ### Input and output
81
+
82
+ 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.
83
+
84
+ The returned dictionary contains:
85
+
86
+ | Key | Contents |
87
+ | --- | --- |
88
+ | `predicted_sp_words` | Predicted scanpaths as lists of words in fixation order |
89
+ | `predicted_sp_ids` | Corresponding word-position indices |
90
+ | `original_sn` | Each original input as a list of words |
91
+ | `predicted_fix_durs` | Predicted fixation durations in milliseconds |
92
+ | `unique_idx` | An identifier for each input within the current call |
93
+
94
+ 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.
95
+
96
+ ## Gradio interface
97
+
98
+ The added [app.py](app.py) provides a web interface to the existing model API.
99
+
100
+ ```bash
101
+ python -m ScanDL2.app
102
+ ```
103
+
104
+ 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.
105
+
106
+ The app uses port 7860 by default, configurable through `PORT`, and binds to `0.0.0.0`.
107
+
108
+ ## Endpoint handler
109
+
110
+ The added [handler.py](handler.py) defines an `EndpointHandler` adapter intended to accept requests with text in `inputs` and inference settings in `parameters`:
111
+
112
+ ```json
113
+ {
114
+ "inputs": ["The quick brown fox jumps over the lazy dog."],
115
+ "parameters": {
116
+ "text_type": "sentence",
117
+ "bsz": 2
118
+ }
119
+ }
120
+ ```
121
+
122
+ 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.
123
+
124
+ 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.
125
+
126
+ ## Training, inference, and evaluation
127
+
128
+ 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.
129
+
130
+ ### Download the data
131
+
132
+ - **CELER:** follow the instructions in the [dataset repository](https://github.com/berzak/celer).
133
+ - **ZuCo:** download from the [OSF repository](https://osf.io/q3zws/). The dataset requires substantial storage.
134
+ - **Beijing Sentence Corpus (BSC):** download from the [OSF repository](https://osf.io/vr3k8/).
135
+ - **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).
136
+
137
+ 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).
138
+
139
+ ### Preprocess the training and test data
140
+
141
+ 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.
142
+
143
+ 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.
144
+
145
+ ### ScanDL module
146
+
147
+ 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.
148
+
149
+ 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).
150
+
151
+ ### Fixation-duration module
152
+
153
+ 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`.
154
+
155
+ ### Sentence-level and paragraph-level training
156
+
157
+ 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.
158
+
159
+ ### Evaluation and the diffusion-only ablation
160
+
161
+ 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.
162
+
163
+ Local checks can be run from the project root:
164
+
165
+ ```bash
166
+ python -m unittest discover -s ScanDL2/tests
167
+ ```
168
+
169
+ 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.
170
+
171
+ ## Citation
172
+
173
+ ```bibtex
174
+ @article{bolliger2025scandl2,
175
+ author = {Bolliger, Lena S. and Reich, David R. and J\"{a}ger, Lena A.},
176
+ title = {ScanDL 2.0: A Generative Model of Eye Movements in Reading Synthesizing Scanpaths and Fixation Durations},
177
+ year = {2025},
178
+ issue_date = {May 2025},
179
+ publisher = {Association for Computing Machinery},
180
+ address = {New York, NY, USA},
181
+ volume = {9},
182
+ number = {ETRA5},
183
+ url = {https://doi.org/10.1145/3725830},
184
+ doi = {10.1145/3725830},
185
+ 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.},
186
+ journal = {Proceedings of the ACM on Human-Computer Interaction},
187
+ month = may,
188
+ articleno = {5},
189
+ numpages = {30},
190
+ keywords = {neural networks, scanpath generation, eye movements, reading, diffusion models}
191
+ }
192
+ ```
193
+
194
+ ## Related paper
195
+
196
+ 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/).
197
+
198
+ ```bibtex
199
+ @inproceedings{bolliger2023scandl,
200
+ title={ScanDL: A Diffusion Model for Generating Synthetic Scanpaths on Texts},
201
+ author={Bolliger, Lena S. and Reich, David R. and Haller, Patrick and Jakobi, Deborah N. and Prasse, Paul and Jäger, Lena A.},
202
+ booktitle={Proceedings of the 2023 Conference on Empirical Methods in Natural Language Processing},
203
+ year={2023},
204
+ pages={15513--15538},
205
+ doi={10.18653/v1/2023.emnlp-main.960},
206
+ url={https://aclanthology.org/2023.emnlp-main.960/}
207
+ }
208
+ ```
209
+
210
+ ## License
211
+
212
+ 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.
__init__.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from ScanDL2.model import ScanDL2
2
+ from ScanDL2.model import ScanDLModule
3
+ from ScanDL2.model import FixdurModule
4
+ from ScanDL2.scandl_module.original_scandl.sp_transformer_model import TransformerNetModel
5
+ from ScanDL2.fix_dur_module.model_seq2seq import Seq2SeqModel
6
+
7
+ from ScanDL2.scandl_module.original_scandl.sp_gaussian_diffusion import GaussianDiffusion
8
+ from ScanDL2.scandl_module.original_scandl.sp_gaussian_diffusion import SpacedDiffusion
9
+ from ScanDL2.fix_dur_module.model_seq2seq import Pooler
10
+ from ScanDL2.scandl_module.original_scandl.sp_rounding import denoised_fn_round
11
+
12
+ from ScanDL2 import utils
13
+ from ScanDL2 import training_utils
14
+
15
+ __all__ = [
16
+ "ScanDL2",
17
+ "ScanDL",
18
+ "FixdurModule",
19
+ "TransformerNetModel",
20
+ "ScanDLModule",
21
+ "GaussianDiffusion",
22
+ "SpacedDiffusion",
23
+ "Pooler",
24
+ "denoised_fn_round",
25
+ "utils",
26
+ "training_utils",
27
+ ]
app.py ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import json
4
+ import tempfile
5
+ from typing import Union, List, Dict, Any
6
+
7
+ import gradio as gr
8
+ import torch
9
+
10
+
11
+ from ScanDL2 import ScanDL2
12
+
13
+ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
14
+
15
+
16
+ _MODELS: Dict[str, ScanDL2] = {}
17
+
18
+
19
+ def get_model(text_type: str) -> ScanDL2:
20
+ """Load and cache a ScanDL2 model for the given text type."""
21
+ if text_type not in _MODELS:
22
+ # `save=None` -> we handle saving ourselves in the Gradio app
23
+ _MODELS[text_type] = ScanDL2(text_type=text_type, save=None)
24
+ return _MODELS[text_type]
25
+
26
+
27
+ def predict(
28
+ text: str,
29
+ text_type: str,
30
+ progress=gr.Progress(track_tqdm=True),
31
+ ) -> Dict[str, Any]:
32
+ """
33
+ Run ScanDL2 on the input text and return a structured result.
34
+ """
35
+ if text is None or text.strip() == "":
36
+ raise gr.Error("Please provide some input text.")
37
+
38
+ lines = [ln.strip() for ln in text.strip().split("\n") if ln.strip()]
39
+ if len(lines) == 0:
40
+ raise gr.Error("Input text is empty after cleaning.")
41
+
42
+ progress(0.05, desc="Loading ScanDL2 model...")
43
+ model = get_model(text_type)
44
+
45
+ progress(0.15, desc="Running ScanDL + FixDur modules...")
46
+ with torch.no_grad():
47
+ output = model(texts=lines)
48
+
49
+ return output
50
+
51
+
52
+ def format_output(output: Dict[str, Any]) -> str:
53
+ """Pretty-print the ScanDL2 output for display in the UI."""
54
+ if output is None:
55
+ return ""
56
+
57
+ n = len(output.get("original_sn", []))
58
+ lines: List[str] = []
59
+
60
+ for i in range(n):
61
+ original_sn = output["original_sn"][i]
62
+ sp_words = output["predicted_sp_words"][i]
63
+ sp_ids = output["predicted_sp_ids"][i]
64
+ fix_durs = output["predicted_fix_durs"][i]
65
+
66
+ lines.append(f"### Example {i + 1}")
67
+ lines.append("")
68
+ lines.append(f"**Original text:** {' '.join(original_sn)}")
69
+ lines.append("")
70
+
71
+ # Build a readable scanpath table
72
+ lines.append("| # | Word | Word index | Fixation duration (ms) |")
73
+ lines.append("|---|------|-----------|------------------------|")
74
+ for step, (w, wid, dur) in enumerate(zip(sp_words, sp_ids, fix_durs), start=1):
75
+ lines.append(f"| {step} | {w} | {wid} | {dur} |")
76
+ lines.append("")
77
+
78
+ return "\n".join(lines)
79
+
80
+
81
+ def format_json(output: Dict[str, Any]) -> str:
82
+ """Return the raw JSON string of the output."""
83
+ if output is None:
84
+ return "{}"
85
+ return json.dumps(output, indent=2, ensure_ascii=False)
86
+
87
+
88
+ DESCRIPTION = """
89
+ # ScanDL 2.0
90
+
91
+ **ScanDL 2.0** predicts human-like **eye-movement scanpaths** (which words are fixated, in what order)
92
+ and their **fixation durations** (in milliseconds) directly from text.
93
+
94
+ This Space wraps two jointly-trained modules:
95
+
96
+ 1. **ScanDL module** — a discrete diffusion model that generates fixation *locations* (a scanpath) over the input text.
97
+ 2. **FixDur module** — a sequence-to-sequence model that predicts the *duration* of each fixation.
98
+
99
+ ### How to use
100
+ 1. Paste your text in the box below.
101
+ - In **sentence** mode, each line is treated as a separate sentence.
102
+ - In **paragraph** mode, each line is treated as a separate paragraph.
103
+ 2. Choose the text type (must match the model checkpoint you want to use).
104
+ 3. Click **Run**.
105
+
106
+ ### Output
107
+ - A human-readable scanpath table with predicted fixation durations per word.
108
+ - The raw JSON output (fixated words, word indices, and durations).
109
+ """
110
+
111
+ EXAMPLES = [
112
+ [
113
+ "The quick brown fox jumps over the lazy dog.",
114
+ "sentence",
115
+ ],
116
+ [
117
+ "Researchers have long been interested in how humans process written language.\n"
118
+ "Eye-tracking studies reveal where and for how long readers fixate on words.",
119
+ "paragraph",
120
+ ],
121
+ ]
122
+
123
+
124
+ def build_demo() -> gr.Blocks:
125
+ with gr.Blocks(
126
+ title="ScanDL 2.0 — Eye-Movement Scanpath Prediction",
127
+ theme=gr.themes.Soft(),
128
+ ) as demo:
129
+ gr.Markdown(DESCRIPTION)
130
+
131
+ with gr.Row():
132
+ with gr.Column(scale=3):
133
+ text_in = gr.Textbox(
134
+ label="Input text",
135
+ placeholder="Paste a sentence or paragraph here...",
136
+ lines=8,
137
+ )
138
+ text_type_in = gr.Radio(
139
+ choices=["sentence", "paragraph"],
140
+ value="sentence",
141
+ label="Text type",
142
+ info="Must match an available ScanDL2 checkpoint.",
143
+ )
144
+ with gr.Row():
145
+ run_btn = gr.Button("Run", variant="primary")
146
+ clear_btn = gr.Button("Clear")
147
+
148
+ with gr.Column(scale=4):
149
+ table_out = gr.Markdown(
150
+ label="Predicted scanpath",
151
+ value="_Results will appear here._",
152
+ )
153
+ json_out = gr.Code(
154
+ label="Raw JSON output",
155
+ language="json",
156
+ value="{}",
157
+ )
158
+
159
+ gr.Examples(examples=EXAMPLES, inputs=[text_in, text_type_in])
160
+
161
+ def _run(text, text_type):
162
+ output = predict(text, text_type)
163
+ return format_output(output), format_json(output)
164
+
165
+ run_btn.click(
166
+ fn=_run,
167
+ inputs=[text_in, text_type_in],
168
+ outputs=[table_out, json_out],
169
+ )
170
+ clear_btn.click(
171
+ fn=lambda: ("", "sentence", "_Results will appear here._", "{}"),
172
+ inputs=None,
173
+ outputs=[text_in, text_type_in, table_out, json_out],
174
+ )
175
+
176
+ return demo
177
+
178
+
179
+ if __name__ == "__main__":
180
+ demo = build_demo()
181
+ demo.queue(max_size=16).launch(
182
+ server_name="0.0.0.0",
183
+ server_port=int(os.environ.get("PORT", 7860)),
184
+ show_error=True,
185
+ )
config.json ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "lr": 0.0001,
3
+ "batch_size": 128,
4
+ "microbatch": 64,
5
+ "learning_steps": 80000,
6
+ "log_interval": 50,
7
+ "save_interval": 5000,
8
+ "eval_interval": 500,
9
+ "ema_rate": "0.9999",
10
+ "resume_checkpoint": "none",
11
+ "schedule_sampler": "lossaware",
12
+ "diffusion_steps": 2000,
13
+ "noise_schedule": "sqrt",
14
+ "timestep_respacing": "",
15
+ "vocab": "bert",
16
+ "use_plm_init": "no",
17
+ "vocab_size": 0,
18
+ "config_name": "bert-base-cased",
19
+ "gpt_config_name": "gpt2",
20
+ "notes": "folder-notes",
21
+ "data_dir": "processed_data",
22
+ "dataset": "dataset-name",
23
+ "checkpoint_path": "checkpoint-path/test-run",
24
+ "seq_len": 128,
25
+ "hidden_t_dim": 128,
26
+ "hidden_dim": 256,
27
+ "dropout": 0.1,
28
+ "use_fp16": false,
29
+ "fp16_scale_growth": 0.001,
30
+ "seed": 102,
31
+ "gradient_clipping": -1.0,
32
+ "weight_decay": 0.0,
33
+ "learn_sigma": false,
34
+ "use_kl": false,
35
+ "predict_xstart": true,
36
+ "rescale_timesteps": true,
37
+ "rescale_learned_sigmas": false,
38
+ "sigma_small": false,
39
+ "emb_scale_factor": 1.0,
40
+ "num_transformer_layers": 12,
41
+ "num_transformer_heads": 8,
42
+ "one_noise_step": true,
43
+ "mask_padding": false,
44
+ "celer_only_L1": true,
45
+ "data_split_criterion": "scanpath",
46
+ "corpus": "celer",
47
+ "inference": "none",
48
+ "n_folds": 5,
49
+ "ablation_type": "none",
50
+ "nll_in_loss": false,
51
+ "load_from_checkpoint": false,
52
+ "load_train_data": "-"
53
+ }
config_bsc.json ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "lr": 0.0001,
3
+ "batch_size": 128,
4
+ "microbatch": 64,
5
+ "learning_steps": 80000,
6
+ "log_interval": 50,
7
+ "save_interval": 5000,
8
+ "eval_interval": 500,
9
+ "ema_rate": "0.9999",
10
+ "resume_checkpoint": "none",
11
+ "schedule_sampler": "lossaware",
12
+ "diffusion_steps": 2000,
13
+ "noise_schedule": "sqrt",
14
+ "timestep_respacing": "",
15
+ "vocab": "bert",
16
+ "use_plm_init": "no",
17
+ "vocab_size": 0,
18
+ "config_name": "bert-base-chinese",
19
+ "gpt_config_name": "benjamin/gpt2-wechsel-chinese",
20
+ "notes": "folder-notes",
21
+ "data_dir": "processed_data_bsc",
22
+ "dataset": "dataset-name",
23
+ "checkpoint_path": "checkpoint-path/test-run",
24
+ "seq_len": 68,
25
+ "hidden_t_dim": 68,
26
+ "hidden_dim": 256,
27
+ "dropout": 0.1,
28
+ "use_fp16": false,
29
+ "fp16_scale_growth": 0.001,
30
+ "seed": 102,
31
+ "gradient_clipping": -1.0,
32
+ "weight_decay": 0.0,
33
+ "learn_sigma": false,
34
+ "use_kl": false,
35
+ "predict_xstart": true,
36
+ "rescale_timesteps": true,
37
+ "rescale_learned_sigmas": false,
38
+ "sigma_small": false,
39
+ "emb_scale_factor": 1.0,
40
+ "num_transformer_layers": 12,
41
+ "num_transformer_heads": 8,
42
+ "one_noise_step": true,
43
+ "mask_padding": false,
44
+ "celer_only_L1": true,
45
+ "data_split_criterion": "scanpath",
46
+ "corpus": "bsc",
47
+ "inference": "none",
48
+ "n_folds": 5,
49
+ "ablation_type": "none",
50
+ "nll_in_loss": false,
51
+ "load_from_checkpoint": false,
52
+ "load_train_data": "-"
53
+ }
54
+
config_emtec.json ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "lr": 0.0001,
3
+ "batch_size": 128,
4
+ "microbatch": 64,
5
+ "learning_steps": 80000,
6
+ "log_interval": 50,
7
+ "save_interval": 5000,
8
+ "eval_interval": 500,
9
+ "ema_rate": "0.9999",
10
+ "resume_checkpoint": "none",
11
+ "schedule_sampler": "lossaware",
12
+ "diffusion_steps": 2000,
13
+ "noise_schedule": "sqrt",
14
+ "timestep_respacing": "",
15
+ "vocab": "bert",
16
+ "use_plm_init": "no",
17
+ "vocab_size": 0,
18
+ "config_name": "bert-base-cased",
19
+ "gpt_config_name": "gpt2",
20
+ "notes": "folder-notes",
21
+ "data_dir": "processed_data_emtec",
22
+ "dataset": "dataset-name",
23
+ "checkpoint_path": "checkpoint-path/test-run",
24
+ "seq_len": 352,
25
+ "hidden_t_dim": 352,
26
+ "hidden_dim": 256,
27
+ "dropout": 0.1,
28
+ "use_fp16": false,
29
+ "fp16_scale_growth": 0.001,
30
+ "seed": 102,
31
+ "gradient_clipping": -1.0,
32
+ "weight_decay": 0.0,
33
+ "learn_sigma": false,
34
+ "use_kl": false,
35
+ "predict_xstart": true,
36
+ "rescale_timesteps": true,
37
+ "rescale_learned_sigmas": false,
38
+ "sigma_small": false,
39
+ "emb_scale_factor": 1.0,
40
+ "num_transformer_layers": 12,
41
+ "num_transformer_heads": 8,
42
+ "one_noise_step": true,
43
+ "mask_padding": false,
44
+ "celer_only_L1": true,
45
+ "data_split_criterion": "scanpath",
46
+ "corpus": "emtec",
47
+ "inference": "none",
48
+ "n_folds": 5,
49
+ "ablation_type": "none",
50
+ "nll_in_loss": false,
51
+ "load_from_checkpoint": false,
52
+ "load_train_data": "-"
53
+ }
54
+
create_data.py ADDED
@@ -0,0 +1,167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Create the data for training ScanDL on all data.
3
+ """
4
+
5
+ import argparse
6
+ import os
7
+ import json
8
+ import numpy as np
9
+ import pandas as pd
10
+ import sys
11
+
12
+
13
+ from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import (
14
+ load_celer,
15
+ load_celer_speakers,
16
+ process_celer,
17
+ )
18
+ from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import (
19
+ load_zuco,
20
+ process_zuco,
21
+ get_kfold,
22
+ get_kfold_indices_combined,
23
+ )
24
+ from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_emtec, process_emtec
25
+ from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import load_bsc, process_bsc
26
+ from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import flatten_data, unflatten_data
27
+ from transformers import set_seed, BertTokenizerFast
28
+
29
+ sys.path.append("./")
30
+ sys.path.append("../")
31
+
32
+
33
+ def create_argparser() -> argparse.ArgumentParser:
34
+ parser = argparse.ArgumentParser()
35
+ parser.add_argument(
36
+ "--folder-name",
37
+ type=str,
38
+ default="processed_data_all",
39
+ help="Name of the folder to save the processed data in.",
40
+ )
41
+ parser.add_argument(
42
+ "--max-fix-dur",
43
+ type=int,
44
+ help="max fixatino duration value. greater fixation durations are replaced with this value.",
45
+ default=999,
46
+ )
47
+ parser.add_argument(
48
+ "--data",
49
+ type=str,
50
+ choices=["celer", "emtec", "bsc"],
51
+ required=True,
52
+ )
53
+ defaults = dict()
54
+ defaults.update(load_defaults_config(parser.parse_args()))
55
+
56
+ add_dict_to_argparser(parser, defaults)
57
+ return parser
58
+
59
+
60
+ def load_defaults_config(args):
61
+ """
62
+ Load defaults for training args.
63
+ """
64
+ if args.data == "emtec":
65
+ config_name = "config_emtec.json"
66
+ elif args.data == "bsc":
67
+ config_name = "config_bsc.json"
68
+ else:
69
+ config_name = "config.json"
70
+ with open(f"diffusion_only/scandl_diff_dur/{config_name}", "r") as f:
71
+ return json.load(f)
72
+
73
+
74
+ def add_dict_to_argparser(parser, default_dict):
75
+ for k, v in default_dict.items():
76
+ v_type = type(v)
77
+ if v is None:
78
+ v_type = str
79
+ elif isinstance(v, bool):
80
+ v_type = str2bool
81
+ parser.add_argument(f"--{k}", default=v, type=v_type)
82
+
83
+
84
+ def str2bool(v):
85
+ """
86
+ https://stackoverflow.com/questions/15008758/parsing-boolean-values-with-argparse
87
+ """
88
+ if isinstance(v, bool):
89
+ return v
90
+ if v.lower() in ("yes", "true", "t", "y", "1"):
91
+ return True
92
+ elif v.lower() in ("no", "false", "f", "n", "0"):
93
+ return False
94
+ else:
95
+ raise argparse.ArgumentTypeError("boolean value expected")
96
+
97
+
98
+ def main():
99
+
100
+ base_folder_name = "scandl2_pkg"
101
+
102
+ print("Loading argument parser...")
103
+ args = create_argparser().parse_args()
104
+ set_seed(args.seed)
105
+
106
+ if args.data == "celer":
107
+
108
+ tokenizer = BertTokenizerFast.from_pretrained(args.config_name)
109
+ data_path = args.folder_name + "_celer"
110
+ if not os.path.exists(os.path.join(base_folder_name, data_path)):
111
+ os.makedirs(os.path.join(base_folder_name, data_path))
112
+
113
+ # load Celer data
114
+ word_info_df, eyemovement_df = load_celer()
115
+ reader_list = load_celer_speakers(only_native_speakers=args.celer_only_L1)
116
+ sn_list = np.unique(
117
+ word_info_df[word_info_df["list"].isin(reader_list)].sentenceid.values
118
+ ).tolist()
119
+
120
+ data, splitting_IDs_dict = process_celer(
121
+ sn_list=sn_list,
122
+ reader_list=reader_list,
123
+ word_info_df=word_info_df,
124
+ eyemovement_df=eyemovement_df,
125
+ tokenizer=tokenizer,
126
+ args=args,
127
+ inference="cv",
128
+ max_fix_dur=args.max_fix_dur,
129
+ )
130
+ flattened_data = flatten_data(data)
131
+ flattened_data = np.array(flattened_data, dtype=object).tolist()
132
+ train_data = unflatten_data(flattened_data=flattened_data, split="train")
133
+ train_data.save_to_disk(os.path.join(base_folder_name, data_path))
134
+
135
+ elif args.data == "bsc":
136
+
137
+ raise NotImplementedError("BSC data not implemented yet.")
138
+
139
+ elif args.data == "emtec":
140
+
141
+ tokenizer = BertTokenizerFast.from_pretrained(args.config_name)
142
+ data_path = args.folder_name + "_emtec"
143
+ if not os.path.exists(os.path.join(base_folder_name, data_path)):
144
+ os.makedirs(os.path.join(base_folder_name, data_path))
145
+
146
+ # load EMTeC data
147
+ print("Loading EMTeC data...")
148
+ fixations_df, stimuli_df = load_emtec()
149
+ data, splitting_IDs_dict = process_emtec(
150
+ fixations_df=fixations_df,
151
+ stimuli_df=stimuli_df,
152
+ tokenizer=tokenizer,
153
+ args=args,
154
+ inference="cv",
155
+ max_fix_dur=args.max_fix_dur,
156
+ )
157
+ flattened_data = flatten_data(data)
158
+ flattened_data = np.array(flattened_data, dtype=object).tolist()
159
+ train_data = unflatten_data(flattened_data=flattened_data, split="train")
160
+ train_data.save_to_disk(os.path.join(base_folder_name, data_path))
161
+
162
+ else:
163
+ raise NotImplementedError("Data not implemented yet.")
164
+
165
+
166
+ if __name__ == "__main__":
167
+ raise SystemExit(main())
fix_dur_module/__init__.py ADDED
File without changes
fix_dur_module/__pycache__/__init__.cpython-313.pyc ADDED
Binary file (181 Bytes). View file
 
fix_dur_module/__pycache__/model_seq2seq.cpython-313.pyc ADDED
Binary file (4.16 kB). View file
 
fix_dur_module/__pycache__/scasim.cpython-313.pyc ADDED
Binary file (7.43 kB). View file
 
fix_dur_module/__pycache__/utils_data.cpython-313.pyc ADDED
Binary file (17.8 kB). View file
 
fix_dur_module/__pycache__/utils_train.cpython-313.pyc ADDED
Binary file (6.16 kB). View file
 
fix_dur_module/model_seq2seq.py ADDED
@@ -0,0 +1,89 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ from transformers.models.bert.modeling_bert import BertEncoder
4
+ from typing import Optional
5
+
6
+
7
+ class Seq2SeqModel(nn.Module):
8
+
9
+ def __init__(
10
+ self,
11
+ config,
12
+ output_dim,
13
+ num_linear,
14
+ dropout,
15
+ ):
16
+ super().__init__()
17
+ self.config = config
18
+ self.output_dim = output_dim
19
+ self.encoder = BertEncoder(config)
20
+ self.pooler = Pooler(config)
21
+
22
+ layers_list = list()
23
+ for i in range(num_linear):
24
+ layers_list.append(nn.Linear(config.hidden_size, config.hidden_size))
25
+ layers_list.append(nn.ReLU())
26
+ layers_list.append(nn.Dropout(dropout))
27
+ self.ff = nn.Sequential(*layers_list)
28
+ # self.ff = nn.Linear(config.hidden_size, config.hidden_size)
29
+ self.ff_out = nn.Linear(config.hidden_size, output_dim)
30
+
31
+ def _invert_attention_mask(self, attention_mask):
32
+ if attention_mask.dim() == 3:
33
+ extended_attention_mask = attention_mask[:, None, :, :]
34
+ elif attention_mask.dim() == 2:
35
+ extended_attention_mask = attention_mask[:, None, None, :]
36
+ extended_attention_mask = (1.0 - extended_attention_mask) * torch.finfo(torch.float32).min
37
+ return extended_attention_mask
38
+
39
+ def forward(
40
+ self,
41
+ sp_embeddings,
42
+ attention_mask: Optional[torch.Tensor] = None,
43
+ output_attentions: Optional[bool] = None,
44
+ ):
45
+ # get the extended attention mask
46
+ # zeros and ones are inverted such that what is not maked is 0 and what is masked is -inf
47
+ if attention_mask is not None:
48
+ attention_mask = self._invert_attention_mask(attention_mask)
49
+ encoder_outputs = self.encoder(
50
+ sp_embeddings,
51
+ attention_mask=attention_mask,
52
+ output_attentions=output_attentions,
53
+ )
54
+ else:
55
+ encoder_outputs = self.encoder(
56
+ sp_embeddings,
57
+ output_attentions=output_attentions,
58
+ )
59
+
60
+ last_hidden_state = encoder_outputs.last_hidden_state
61
+
62
+ # pool the encoder output: the hidden state of the CLS token is passed through another linear layer
63
+ pooled_output = self.pooler(last_hidden_state)
64
+
65
+ # map to the output dimension
66
+ out = self.ff(pooled_output)
67
+ out = self.ff_out(out)
68
+
69
+ if output_attentions:
70
+ attentions = encoder_outputs.attentions
71
+ return out, attentions
72
+
73
+ else:
74
+ return out
75
+
76
+
77
+ class Pooler(nn.Module):
78
+ def __init__(self, config):
79
+ super().__init__()
80
+ self.dense = nn.Linear(config.hidden_size, config.hidden_size)
81
+ self.activation = nn.Tanh()
82
+
83
+ def forward(self, hidden_states):
84
+ # pool the output by taking the hidden state of the first token (the CLS token)
85
+ # and pass it through another linear layer wtih tanh activation
86
+ cls_out = hidden_states[:, 0]
87
+ pooled_output = self.dense(cls_out)
88
+ pooled_output = self.activation(pooled_output)
89
+ return pooled_output
fix_dur_module/scasim.py ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Script that implements the scanpath similarity metric ScaSim by
3
+ Von der Malsburg, Titus, and Shravan Vasishth.
4
+ "What is the scanpath signature of syntactic reanalysis?."
5
+ Journal of Memory and Language 65.2 (2011): 109-127.
6
+ """
7
+
8
+ from __future__ import annotations
9
+ from math import pi, sin, cos, acos
10
+ import numpy as np
11
+ from typing import List, Tuple, Optional, Any, Union
12
+
13
+
14
+ # only need 0, 2 of s/t due to word index instead of x/y location
15
+ def scasim(
16
+ s: List[Tuple[int, int, Union[int, float]]],
17
+ t: List[Tuple[int, int, Union[int, float]]],
18
+ modulator: Optional[float] = 0.83,
19
+ normalize: Optional[str] = None, # fixations, durations, None
20
+ ) -> float:
21
+ """
22
+ Calculate the similarity between two scanpaths s and t.
23
+ :param s: scanpath s, consisting of fixation locations (word indices) and fixation durations
24
+ :param t: scanpath t, consisting of fixation locations (word indices) and fixation durations
25
+ :param modulator: modulator specifies how spatial distances between fixations are assessed. When set to 0, any spatial divergence of two
26
+ compared scanpaths is penalized independently of its degree. When set to 1, the scanpaths are compared only with respect to their
27
+ temporal patterns. The default value approximates the sensitivity to spatial distance found in the human visual system.
28
+ :param normalize: if 'fixations', the similarity score is normalized by the number of fixations in the two scanpaths. If 'durations',
29
+ the similarity score is normalized by the sum of fixation durations in the two scanpaths. If None, no normalization is applied.
30
+
31
+ :return: similarity between scanpaths s and t
32
+ """
33
+ m, n = len(s), len(t)
34
+ d = [list(map(lambda i: 0, range(n + 1))) for _ in range(m + 1)]
35
+
36
+ # sum of fixation durations of the two scanpaths
37
+ s_fixdur_sum = sum([fix[2] for fix in s])
38
+ t_fixdur_sum = sum([fix[2] for fix in t])
39
+ # number of fixations in the two scanpaths
40
+ s_nfix = len(s)
41
+ t_nfix = len(t)
42
+
43
+ acc = 0
44
+ # sequence alignment
45
+ # loop over fixations in scanpath s:
46
+ for fix_i in range(1, m + 1):
47
+ acc += s[fix_i - 1][2]
48
+ d[fix_i][0] = acc
49
+
50
+ # loop over fixations in scanpath t:
51
+ acc = 0
52
+ for fix_j in range(1, n + 1):
53
+ acc += t[fix_j - 1][2]
54
+ d[0][fix_j] = acc
55
+
56
+ # Compute similarity:
57
+ for fix_i in range(n):
58
+ for fix_j in range(m):
59
+ # calculating angle between fixation targets:
60
+ slon = s[fix_j][0] / (180 / pi) # longitude (x-axis)
61
+ tlon = t[fix_i][0] / (180 / pi)
62
+ slat = s[fix_j][1] / (180 / pi) # latitude (y-axis)
63
+ tlat = t[fix_i][1] / (180 / pi)
64
+
65
+ angle = acos(sin(slat) * sin(tlat) + cos(slat) * cos(tlat) * cos(slon - tlon)) * (
66
+ 180 / pi
67
+ )
68
+
69
+ # approximation of cortical magnification:
70
+ mixer = modulator**angle
71
+
72
+ # cost for substitution:
73
+ cost = abs(t[fix_i][2] - s[fix_j][2]) * mixer + (t[fix_i][2] + s[fix_j][2]) * (
74
+ 1.0 - mixer
75
+ )
76
+
77
+ # select optimal edit operation
78
+ ops = (
79
+ d[fix_j][fix_i + 1] + s[fix_j][2],
80
+ d[fix_j + 1][fix_i] + t[fix_i][2],
81
+ d[fix_j][fix_i] + cost,
82
+ )
83
+
84
+ # mi = which_min(*ops)
85
+ mi = np.argmin(ops)
86
+
87
+ d[fix_j + 1][fix_i + 1] = ops[mi]
88
+
89
+ result = d[-1][-1]
90
+ if normalize == "fixations":
91
+ result /= s_nfix + t_nfix
92
+ elif normalize == "durations":
93
+ result /= s_fixdur_sum + t_fixdur_sum
94
+
95
+ return result
96
+
97
+
98
+ def main():
99
+
100
+ predicted_sp_ids = [
101
+ [0, 1, 2, 3, 4, 5, 6, 8, 10, 10, 10, 11],
102
+ [0, 1, 1, 2, 4, 5, 7, 4, 3, 7, 1, 8],
103
+ [0, 1, 2, 4, 4, 5, 7, 7, 8, 9],
104
+ ]
105
+ original_sp_ids = [
106
+ [0, 1, 2, 4, 2, 3, 5, 6, 8, 9, 10, 4, 11],
107
+ [0, 1, 2, 4, 3, 8],
108
+ [0, 1, 6, 8, 9],
109
+ ]
110
+ predicted_fix_durs = [
111
+ [69, 52, 374, 374, 374, 374, 256, 423, 423, 423, 423, 188],
112
+ [69, 52, 374, 384, 374, 374, 423, 423, 423, 423, 52, 423],
113
+ [69, 52, 374, 502, 502, 374, 423, 423, 374, 374],
114
+ ]
115
+ original_fix_durs = [
116
+ [0, 208, 232, 197, 314, 151, 219, 308, 195, 280, 260, 102],
117
+ [0, 192, 182, 297, 134],
118
+ [0, 195, 130, 101],
119
+ ]
120
+
121
+ # remove last element in each sublist of list for predicted_sp_ids, original_sp_ids, and predicted_fix_durs
122
+ # these are the pad tokens and they are not contained in original_fix_durs
123
+ predicted_sp_ids = [sublist[:-1] for sublist in predicted_sp_ids]
124
+ original_sp_ids = [sublist[:-1] for sublist in original_sp_ids]
125
+ predicted_fix_durs = [sublist[:-1] for sublist in predicted_fix_durs]
126
+
127
+ # create dummy y values for original_sp_ids and predicted_sp_ids
128
+ dummy_y_original_sp_ids = [[1] * len(sublist) for sublist in original_sp_ids]
129
+ dummy_y_predicted_sp_ids = [[1] * len(sublist) for sublist in predicted_sp_ids]
130
+
131
+ # zip together the predicted_sp_ids and predicted_fix_durs lists as list of list of tuples
132
+ predicted_sp = list(
133
+ map(
134
+ lambda x, y, z: list(zip(x, y, z)),
135
+ predicted_sp_ids,
136
+ dummy_y_predicted_sp_ids,
137
+ predicted_fix_durs,
138
+ )
139
+ )
140
+ # zip together the original_sp_ids and original_fix_durs lists as list of list of tuples
141
+ original_sp = list(
142
+ map(
143
+ lambda x, y, z: list(zip(x, y, z)),
144
+ original_sp_ids,
145
+ dummy_y_original_sp_ids,
146
+ original_fix_durs,
147
+ )
148
+ )
149
+
150
+ s1 = predicted_sp[0]
151
+ t1 = original_sp[0]
152
+ sim1 = scasim(s=s1, t=t1)
153
+
154
+ s2 = predicted_sp[1]
155
+ t2 = original_sp[1]
156
+ sim2 = scasim(s=s2, t=t2)
157
+
158
+ s3 = predicted_sp[2]
159
+ t3 = original_sp[2]
160
+ sim3 = scasim(s=s3, t=t3)
161
+
162
+ # normalize by fixations
163
+ sim10 = scasim(s=s1, t=t1, normalize="fixations")
164
+ sim11 = scasim(s=s2, t=t2, normalize="fixations")
165
+ sim12 = scasim(s=s3, t=t3, normalize="fixations")
166
+
167
+ # normalize by durations
168
+ sim13 = scasim(s=s1, t=t1, normalize="durations")
169
+ sim14 = scasim(s=s2, t=t2, normalize="durations")
170
+ sim15 = scasim(s=s3, t=t3, normalize="durations")
171
+
172
+ print("normalize by fixations")
173
+ print("sim10:", sim10)
174
+ print("sim11:", sim11)
175
+ print("sim12:", sim12)
176
+ print("normalize by durations")
177
+ print("sim13:", sim13)
178
+ print("sim14:", sim14)
179
+ print("sim15:", sim15)
180
+
181
+ breakpoint()
182
+
183
+
184
+ if __name__ == "__main__":
185
+ raise SystemExit(main())
fix_dur_module/train_seq2seq.py ADDED
@@ -0,0 +1,283 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ The training script for training the fixation duration module.
3
+ """
4
+
5
+ import joblib
6
+ import sys
7
+ import json
8
+ import os
9
+ from typing import Dict
10
+ from argparse import ArgumentParser
11
+
12
+ from transformers import GPT2TokenizerFast, GPT2LMHeadModel, GPT2Model, AutoConfig, BertModel
13
+ from transformers.models.bert.modeling_bert import BertEncoder, BertPooler
14
+ from transformers import AdamW, get_linear_schedule_with_warmup
15
+
16
+ import numpy as np
17
+ import torch
18
+ import torch.nn as nn
19
+ from torch.utils.data import Dataset, DataLoader
20
+ from datasets import load_from_disk, DatasetDict
21
+ from sklearn.preprocessing import MinMaxScaler
22
+
23
+ from ScanDL2.fix_dur_module.utils_data import (
24
+ prepare_seq2seq_data,
25
+ get_embeddings_seq2seq,
26
+ Seq2SeqDataset,
27
+ split_train_val_data,
28
+ )
29
+ from ScanDL2.fix_dur_module.model_seq2seq import Seq2SeqModel
30
+ from ScanDL2.fix_dur_module.utils_train import EarlyStopping, train
31
+
32
+ sys.path.append("./")
33
+ sys.path.append("../")
34
+ sys.path.append("../../")
35
+
36
+ from ScanDL2.CONSTANTS import (
37
+ COMPLETE_FIXDUR_MODULE_TRAIN_PATH_BSC,
38
+ COMPLETE_FIXDUR_MODULE_TRAIN_PATH_CELER,
39
+ COMPLETE_FIXDUR_MODULE_TRAIN_PATH_EMTEC,
40
+ )
41
+
42
+
43
+ def get_parser() -> ArgumentParser:
44
+ parser = ArgumentParser()
45
+ parser.add_argument(
46
+ "--max-length",
47
+ type=int,
48
+ default=128,
49
+ help="The maximum sequence length.",
50
+ )
51
+ parser.add_argument(
52
+ "--num-heads",
53
+ type=int,
54
+ default=12,
55
+ help="The number of attention heads in the Transformer encoder.",
56
+ )
57
+ parser.add_argument(
58
+ "--num-layers",
59
+ type=int,
60
+ default=12,
61
+ help="The number of layers in the Transformer encoder.",
62
+ )
63
+ parser.add_argument(
64
+ "--num-linear",
65
+ type=int,
66
+ default=8,
67
+ help="The number of linear layers.",
68
+ )
69
+ parser.add_argument(
70
+ "--bsz",
71
+ type=int,
72
+ default=128,
73
+ help="The batch size.",
74
+ )
75
+ parser.add_argument(
76
+ "--dropout",
77
+ type=float,
78
+ default=0.5,
79
+ help="The dropout rate.",
80
+ )
81
+ parser.add_argument(
82
+ "--num-epochs",
83
+ type=int,
84
+ default=400,
85
+ )
86
+ parser.add_argument(
87
+ "--sp-pad-token",
88
+ type=int,
89
+ default=127,
90
+ help="the padding token appended to the sp, usually seq_len-1",
91
+ )
92
+ parser.add_argument(
93
+ "--use-attention-mask",
94
+ action="store_true",
95
+ help="Whether to use the attention mask in the Transformer encoder.",
96
+ )
97
+ parser.add_argument(
98
+ "--data",
99
+ type=str,
100
+ required=True,
101
+ choices=["emtec", "bsc", "celer"],
102
+ help="The dataset to train on.",
103
+ )
104
+ return parser
105
+
106
+
107
+ def main():
108
+
109
+ args = get_parser().parse_args()
110
+
111
+ max_length = args.max_length
112
+ output_attentions = False
113
+ learning_rate = 1e-4
114
+ num_epochs = args.num_epochs
115
+ patience = 25
116
+ normalize = True
117
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
118
+
119
+ if args.data == "emtec":
120
+ path_save_model = COMPLETE_FIXDUR_MODULE_TRAIN_PATH_EMTEC
121
+ path_to_data = "processed_data_all_emtec"
122
+ elif args.data == "bsc":
123
+ raise NotImplementedError("Training on BSC data is not yet implemented.")
124
+ path_save_model = COMPLETE_FIXDUR_MODULE_TRAIN_PATH_BSC
125
+ path_to_data = "processed_data_all_bsc"
126
+ elif args.data == "celer":
127
+ path_save_model = COMPLETE_FIXDUR_MODULE_TRAIN_PATH_CELER
128
+ path_to_data = "processed_data_all_celer"
129
+ else:
130
+ raise ValueError("Unknown dataset.")
131
+
132
+ if not os.path.exists(path_save_model):
133
+ os.makedirs(path_save_model)
134
+ model_name = "seq2seq_fixdur.pt"
135
+
136
+ hypeparameters = {
137
+ "num_heads": args.num_heads,
138
+ "num_layers": args.num_layers,
139
+ "num_linear": args.num_linear,
140
+ "bsz": args.bsz,
141
+ "dropout": args.dropout,
142
+ "use_attention_mask": args.use_attention_mask,
143
+ }
144
+ with open(os.path.join(path_save_model, "hyperparameters.json"), "w") as f:
145
+ json.dump(hypeparameters, f)
146
+
147
+ # load GPT-2 and GPT-2 tokenizer to get the contextualized embeddings
148
+ if args.data == "bsc":
149
+ raise NotImplementedError("Training on BSC data is not yet implemented.")
150
+ gpt_config_name = "benjamin/gpt2-wechsel-chinese"
151
+ else:
152
+ gpt_config_name = "gpt2"
153
+
154
+ tokenizer = GPT2TokenizerFast.from_pretrained(gpt_config_name, add_prefix_space=True)
155
+ gpt2_model = GPT2Model.from_pretrained(gpt_config_name)
156
+ tokenizer.pad_token = tokenizer.eos_token
157
+ # freeze parameters
158
+ for param in gpt2_model.parameters():
159
+ param.requires_grad = False
160
+
161
+ # load BERT config (for model architecture) and BERT model (for embeddings of CLS and PAD tokens)
162
+
163
+ if args.data == "bsc":
164
+ raise NotImplementedError("Training on BSC data is not yet implemented.")
165
+ bert_config_name = "bert-base-chinese"
166
+ else:
167
+ bert_config_name = "bert-base-cased"
168
+ config = AutoConfig.from_pretrained(bert_config_name)
169
+ bert_embeddings = BertModel.from_pretrained(bert_config_name).embeddings.word_embeddings
170
+ # freeze parameters
171
+ for param in bert_embeddings.parameters():
172
+ param.requires_grad = False
173
+
174
+ # change the parameters in the config
175
+ config.num_attention_heads = args.num_heads
176
+ config.num_hidden_layers = args.num_layers
177
+
178
+ # training
179
+ print("--- load and prepare data ...")
180
+ train_data = load_from_disk(os.path.join("scandl2_pkg", path_to_data, "train"))
181
+ new_data = DatasetDict()
182
+ new_data["train"] = train_data
183
+
184
+ # prepare the data for training
185
+ data = prepare_seq2seq_data(
186
+ data=new_data,
187
+ tokenizer=tokenizer,
188
+ gpt2_model=gpt2_model,
189
+ bert_embeddings=bert_embeddings,
190
+ aggregate="mean",
191
+ max_length=max_length,
192
+ sp_pad_token=args.sp_pad_token,
193
+ )
194
+
195
+ fix_dur_colname = "fix_durs"
196
+
197
+ if normalize:
198
+ min_max_scaler = MinMaxScaler()
199
+ fix_durs = [t.cpu().detach().numpy() for t in data["fix_durs"]]
200
+ flattened = np.concatenate(fix_durs).reshape(-1, 1)
201
+ # fit the scaler on the training data
202
+ min_max_scaler.fit(flattened)
203
+ # normalize the fixation durations
204
+ flattened_normalized = min_max_scaler.transform(flattened)
205
+ # reshape
206
+ split_indices = [len(t) for t in fix_durs]
207
+ normalized_data = np.split(flattened_normalized.flatten(), np.cumsum(split_indices)[:-1])
208
+ # convert back to tensors
209
+ normalized_tensors = [torch.tensor(t) for t in normalized_data]
210
+ data["fix_durs_normalized"] = normalized_tensors
211
+ # save the scaler (needed for inference)
212
+ joblib.dump(min_max_scaler, os.path.join(path_save_model, "min_max_scaler.pkl"))
213
+ fix_dur_colname = "fix_durs_normalized"
214
+
215
+ # split data into train and val data (val data for early stopping)
216
+ train_data, val_data = split_train_val_data(
217
+ data=data,
218
+ val_size=0.1,
219
+ )
220
+
221
+ # create dataset and dataloader
222
+ train_dataset = Seq2SeqDataset(
223
+ data=train_data,
224
+ normalize=normalize,
225
+ )
226
+ val_dataset = Seq2SeqDataset(
227
+ data=val_data,
228
+ normalize=normalize,
229
+ )
230
+ train_loader = DataLoader(
231
+ train_dataset,
232
+ batch_size=args.bsz,
233
+ shuffle=True,
234
+ )
235
+ val_loader = DataLoader(
236
+ val_dataset,
237
+ batch_size=args.bsz,
238
+ shuffle=False,
239
+ )
240
+
241
+ # model, loss, optimizer, scheduler, early stopping
242
+
243
+ model = Seq2SeqModel(
244
+ config=config,
245
+ output_dim=max_length,
246
+ num_linear=args.num_linear,
247
+ dropout=args.dropout,
248
+ )
249
+ model.to(device)
250
+ criterion = nn.MSELoss(reduction="mean")
251
+ optimizer = AdamW(model.parameters(), lr=learning_rate)
252
+ early_stopping = EarlyStopping(
253
+ patience=patience,
254
+ path=os.path.join(path_save_model, model_name),
255
+ )
256
+
257
+ num_training_steps = len(train_loader) * num_epochs
258
+ num_warmup_steps = int(0.05 * num_training_steps)
259
+ scheduler = get_linear_schedule_with_warmup(
260
+ optimizer,
261
+ num_warmup_steps=num_warmup_steps,
262
+ num_training_steps=num_training_steps,
263
+ )
264
+
265
+ # training
266
+ train(
267
+ model=model,
268
+ num_epochs=num_epochs,
269
+ train_loader=train_loader,
270
+ val_loader=val_loader,
271
+ criterion=criterion,
272
+ optimizer=optimizer,
273
+ early_stopping=early_stopping,
274
+ scheduler=scheduler,
275
+ device=device,
276
+ fix_dur_colname=fix_dur_colname,
277
+ output_attentions=output_attentions,
278
+ use_attention_mask=args.use_attention_mask,
279
+ )
280
+
281
+
282
+ if __name__ == "__main__":
283
+ raise SystemExit(main())
fix_dur_module/utils_data.py ADDED
@@ -0,0 +1,530 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Utils for the fixation duration module.
3
+ """
4
+
5
+ import torch
6
+ from datasets import load_from_disk
7
+ from typing import Dict, Any, List, Optional, Union
8
+ import transformers
9
+ from torch.utils.data import Dataset
10
+ import datasets
11
+ from tqdm import tqdm
12
+ import random
13
+
14
+ # def get_input_embeddings(
15
+ # data_instance: Dict[str, Any],
16
+ # tokenizer: transformers.GPT2TokenizerFast,
17
+ # gpt2_model: transformers.GPT2Model,
18
+ # aggregate: str = 'mean', # 'mean', 'sum'
19
+ # #max_length: int = 128,
20
+ # ):
21
+ # # dummy code for now
22
+ # sn_repr_len = data_instance['sn_repr_len']
23
+
24
+ # # the sentence
25
+ # sn_words = data_instance['words_for_mapping'].split()
26
+ # while sn_words[-1] == '[PAD]':
27
+ # sn_words.pop()
28
+
29
+ # # the scanpath
30
+ # # remove the CLS token (already have a SEP token at the end of the sentence and the two will beconcatenated)
31
+ # # the SEP token can stay
32
+ # sp_ids = data_instance['sn_sp_repr'][sn_repr_len:][1:]
33
+ # # cut off the trialing pad tokens
34
+ # while sp_ids[-1] == 127:
35
+ # sp_ids.pop()
36
+ # # get the scanpath as fixated words
37
+ # sp_words = list()
38
+ # for sp_id in sp_ids:
39
+ # sp_words.append(sn_words[sp_id])
40
+
41
+ # # TODO doesn't make sense to have CLS sn SEP sp SEP for auto-regressive model. BOS token
42
+ # # join sentence and scanpath into strings and concatenate them
43
+ # sn = ' '.join(sn_words)
44
+ # sp = ' '.join(sp_words)
45
+ # sn_sp = sn + ' ' + sp
46
+
47
+ # encoded = tokenizer.encode_plus(
48
+ # sn_sp,
49
+ # add_special_tokens=False,
50
+ # return_tensors='pt',
51
+ # return_attention_mask=True,
52
+ # )
53
+ # word_ids = torch.Tensor(encoded.word_ids())
54
+
55
+ # last_hidden = gpt2_model(encoded.input_ids).last_hidden_state
56
+
57
+ # # aggregate the embeddings to word-level
58
+ # embeddings = aggregate_input_embeddings(
59
+ # embeddings=last_hidden,
60
+ # word_ids=word_ids,
61
+ # aggregate=aggregate,
62
+ # )
63
+
64
+ # return embeddings, word_ids, sn_sp
65
+
66
+
67
+ def get_embeddings_seq2seq(
68
+ data_instance: Dict[str, Any],
69
+ tokenizer: transformers.GPT2TokenizerFast,
70
+ gpt2_model: transformers.GPT2Model,
71
+ bert_embeddings: torch.nn.Embedding,
72
+ instance_idx: int,
73
+ aggregate: str = "mean", # 'mean', 'sum'
74
+ max_length: int = 128,
75
+ sp_pad_token: int = 127,
76
+ ):
77
+ """
78
+ Get the embeddings of the scanpath (fixated words) from the encoder.
79
+ :param data_instance: the data instance from the dataset.
80
+ :param tokenizer: the tokenizer.
81
+ :param gpt2_model: the GPT2 model.
82
+ :param bert_embeddings: the BERT embeddings.
83
+ :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings.
84
+ :return: the embeddings of the scanpath, the padded fixation durations, and the attention mask.
85
+ """
86
+
87
+ sn_repr_len = data_instance["sn_repr_len"]
88
+
89
+ # the sentence
90
+ sn_words = data_instance["words_for_mapping"].split()
91
+ while sn_words[-1] == "[PAD]":
92
+ sn_words.pop()
93
+ # remove the CLS and SEP tokens
94
+ sn_words = sn_words[1:-1]
95
+ # chinese characters are one string in a list
96
+ if sp_pad_token == 67: # chinese pad token
97
+ sn_words = list(sn_words[0])
98
+
99
+ # the scanpath
100
+ # remove the CLS token
101
+ sp_ids = data_instance["sn_sp_repr"][sn_repr_len:][1:]
102
+ # cut off the trailing pad tokens
103
+ while sp_ids[-1] == sp_pad_token:
104
+ sp_ids.pop()
105
+ # remove the SEP token
106
+ sp_ids = sp_ids[:-1]
107
+
108
+ # make the scanpath ids start from 0 for re-ordering of the embeddings
109
+ sp_ids = [sp_id - 1 for sp_id in sp_ids]
110
+
111
+ # get the scanpath as fixated words
112
+ sp_words = list()
113
+ try:
114
+ for sp_id in sp_ids:
115
+ sp_words.append(sn_words[sp_id])
116
+ except:
117
+ print(f"Error at index {instance_idx}")
118
+ # breakpoint()
119
+ return None, None, None
120
+
121
+ # get the fixation durations
122
+ fix_durs = data_instance["sn_sp_fix_dur"][sn_repr_len + 1 :]
123
+ while fix_durs[-1] == 0:
124
+ fix_durs.pop()
125
+ # convert to tensor
126
+ fix_durs = torch.Tensor(fix_durs)
127
+
128
+ # get the sentence encoding
129
+ sn_enc = tokenizer.encode_plus(
130
+ sn_words,
131
+ add_special_tokens=False,
132
+ return_tensors="pt",
133
+ is_split_into_words=True,
134
+ )
135
+ sn_word_ids = torch.Tensor(sn_enc.word_ids())
136
+
137
+ # get the embeddings
138
+ with torch.no_grad():
139
+ last_hidden = gpt2_model(sn_enc.input_ids).last_hidden_state
140
+
141
+ # aggregate the embeddings to word-level
142
+ sn_embeddings = aggregate_input_embeddings(
143
+ embeddings=last_hidden,
144
+ word_ids=sn_word_ids,
145
+ aggregate=aggregate,
146
+ )
147
+
148
+ # convert sp_ids to tensor
149
+ sp_ids = torch.Tensor(sp_ids).long()
150
+
151
+ # re-order the embeddings as scanpath
152
+ sp_embeddings = sn_embeddings[:, sp_ids, :]
153
+
154
+ # pad the embeddings and fixation durations to max input length
155
+ # and get the attention mask
156
+ sp_embeddings_padded, fix_durs_padded, attention_mask = padding_and_mask_seq2seq(
157
+ sp_embeddings=sp_embeddings,
158
+ fix_durs=fix_durs,
159
+ bert_embeddings=bert_embeddings,
160
+ max_length=max_length,
161
+ )
162
+
163
+ return sp_embeddings_padded.squeeze(0), fix_durs_padded, attention_mask.squeeze(0)
164
+
165
+
166
+ def padding_and_mask_seq2seq(
167
+ sp_embeddings: torch.Tensor,
168
+ bert_embeddings: torch.nn.Embedding,
169
+ max_length: int,
170
+ fix_durs: Optional[torch.Tensor] = None,
171
+ inference: Optional[bool] = None,
172
+ ):
173
+ """
174
+ Add the BERT CLS token to the beginning of the scanpath embedding (needed for pooler output).
175
+ Pad the scanpath embeddings and fixation durations to max input lenght.
176
+ Use the PAD token embedding for padding.
177
+ """
178
+ # get the embedding for the pad token
179
+ pad_emb = bert_embeddings(torch.Tensor([0]).long())
180
+ cls_emb = bert_embeddings(torch.Tensor([101]).long())
181
+
182
+ # prepend the cls emb to the sp_embeddings
183
+ sp_embeddings = torch.cat((cls_emb.unsqueeze(0), sp_embeddings), dim=1)
184
+
185
+ # pad the embeddings
186
+ current_length = sp_embeddings.size(1)
187
+ padding_needed = max_length - current_length
188
+ pad_tensor = pad_emb.unsqueeze(0).expand(1, padding_needed, -1)
189
+ sp_embeddings_padded = torch.cat((sp_embeddings, pad_tensor), dim=1)
190
+
191
+ # create attention mask
192
+ sp_mask = torch.ones((1, current_length), dtype=torch.long)
193
+ pad_mask = torch.zeros((1, padding_needed), dtype=torch.long)
194
+ attention_mask = torch.cat((sp_mask, pad_mask), dim=1)
195
+
196
+ if inference:
197
+ return sp_embeddings_padded, attention_mask
198
+
199
+ # prepend 0 to the fixation durations because the first word is the CLS token
200
+ fix_durs = torch.cat((torch.Tensor([0]), fix_durs), dim=0)
201
+
202
+ # pad the fixation durations
203
+ fix_dur_pad = torch.zeros(padding_needed)
204
+ fix_durs_padded = torch.cat((fix_durs, fix_dur_pad), dim=0)
205
+
206
+ return sp_embeddings_padded, fix_durs_padded, attention_mask
207
+
208
+
209
+ def aggregate_input_embeddings(
210
+ embeddings: torch.Tensor,
211
+ word_ids: torch.Tensor,
212
+ aggregate: str = "mean", # 'mean', 'sum'
213
+ ):
214
+ """
215
+ Aggregate the embeddings that are input to the fixation module to word-level.
216
+ :param embeddings: the last hidden state (contextualised embeddings) of the sentence-scanpath concatenation
217
+ when passed through the GPT2 model.
218
+ :param word_ids: the word ids of the sentence-scanpath concatenation.
219
+ :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings.
220
+ :return: the aggregated word embeddings.
221
+ """
222
+ # get the unique indices and inverse
223
+ unique_indices, inverse_indices = torch.unique(word_ids, return_inverse=True)
224
+
225
+ # sum the tensor along the dimension 1 (sequence length) for the same word ids
226
+ summed_tensor = torch.zeros((1, unique_indices.size(0), embeddings.size(2)))
227
+ summed_tensor = summed_tensor.scatter_add(
228
+ 1, inverse_indices.unsqueeze(0).unsqueeze(-1).expand_as(embeddings), embeddings
229
+ )
230
+
231
+ if aggregate == "sum":
232
+ return summed_tensor
233
+
234
+ elif aggregate == "mean":
235
+
236
+ # count the occurrences of each word id (how many sub-words per word)
237
+ counts = torch.zeros(unique_indices.size(0)).scatter_add(
238
+ 0, inverse_indices, torch.ones_like(inverse_indices, dtype=torch.float)
239
+ )
240
+
241
+ # average the summed tensor
242
+ averaged_tensor = summed_tensor / counts.view(1, -1, 1)
243
+ return averaged_tensor
244
+
245
+
246
+ class Seq2SeqDataset(Dataset):
247
+ def __init__(
248
+ self,
249
+ data: Dict[str, torch.Tensor],
250
+ normalize: Optional[bool] = None,
251
+ inference: Optional[bool] = None,
252
+ ):
253
+ super().__init__()
254
+ self.data = data
255
+ self.normalize = normalize
256
+ self.inference = inference
257
+
258
+ def __len__(self):
259
+ return len(self.data["sp_embeddings"])
260
+
261
+ def __getitem__(self, idx):
262
+ if self.inference:
263
+ sample = {
264
+ "sp_embeddings": self.data["sp_embeddings"][idx],
265
+ "attention_masks": self.data["attention_masks"][idx],
266
+ }
267
+ return sample
268
+ else:
269
+ sample = {
270
+ "sp_embeddings": self.data["sp_embeddings"][idx],
271
+ "attention_masks": self.data["attention_masks"][idx],
272
+ "fix_durs": self.data["fix_durs"][idx],
273
+ }
274
+ if self.normalize:
275
+ sample["fix_durs_normalized"] = self.data["fix_durs_normalized"][idx]
276
+ return sample
277
+
278
+
279
+ def prepare_seq2seq_data(
280
+ data: datasets.DatasetDict,
281
+ tokenizer: transformers.GPT2TokenizerFast,
282
+ gpt2_model: transformers.GPT2Model,
283
+ bert_embeddings: torch.nn.Embedding,
284
+ aggregate: str = "mean",
285
+ max_length: int = 128,
286
+ sp_pad_token: int = 127,
287
+ ):
288
+ """
289
+ Prepare the data for training the fixation duration module.
290
+ :param data: the dataset.
291
+ :param tokenizer: the tokenizer.
292
+ :param gpt2_model: the GPT2 model.
293
+ :param bert_embeddings: the BERT embeddings.
294
+ :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings.
295
+ :param max_length: the maximum input length.
296
+ :return: the data for training the fixation duration module.
297
+ """
298
+ data_dict = {
299
+ "sp_embeddings": [],
300
+ "attention_masks": [],
301
+ "fix_durs": [],
302
+ }
303
+
304
+ for idx, instance in tqdm(enumerate(data["train"])):
305
+
306
+ sp_embeddings, fix_durs, attention_mask = get_embeddings_seq2seq(
307
+ data_instance=instance,
308
+ tokenizer=tokenizer,
309
+ gpt2_model=gpt2_model,
310
+ bert_embeddings=bert_embeddings,
311
+ instance_idx=idx,
312
+ aggregate=aggregate,
313
+ max_length=max_length,
314
+ sp_pad_token=sp_pad_token,
315
+ )
316
+ if sp_embeddings is None:
317
+ continue
318
+
319
+ data_dict["sp_embeddings"].append(sp_embeddings)
320
+ data_dict["attention_masks"].append(attention_mask)
321
+ data_dict["fix_durs"].append(fix_durs)
322
+
323
+ return data_dict
324
+
325
+
326
+ def split_train_val_data(
327
+ data: Dict[str, List[torch.Tensor]],
328
+ val_size: float = 0.1,
329
+ ):
330
+ """
331
+ Split the train data into train and validation data.
332
+ :param data: the data.
333
+ :param val_size: the size of the validation data.
334
+ :return: the train and validation data.
335
+ """
336
+ num_samples = len(next(iter(data.values())))
337
+ # shuffle the indices
338
+ indices = list(range(num_samples))
339
+ random.shuffle(indices)
340
+
341
+ # compute the split point
342
+ split_point = int(num_samples * val_size)
343
+ train_indices = indices[split_point:]
344
+ val_indices = indices[:split_point]
345
+
346
+ train_data = {key: [value[i] for i in train_indices] for key, value in data.items()}
347
+ val_data = {key: [value[i] for i in val_indices] for key, value in data.items()}
348
+
349
+ return train_data, val_data
350
+
351
+
352
+ def get_embeddings_seq2seq_hp(
353
+ sn_repr_len: int,
354
+ sn_words: List[str],
355
+ sp_ids: List[int],
356
+ tokenizer: transformers.GPT2TokenizerFast,
357
+ gpt2_model: transformers.GPT2Model,
358
+ bert_embeddings: torch.nn.Embedding,
359
+ aggregate: str = "mean",
360
+ max_length: int = 128,
361
+ sp_pad_token: int = 127,
362
+ ):
363
+ """
364
+ Get the embeddings of the scanpath (fixated words) from the encoder.
365
+ :param sn_repr_len: the length of the sentence representation.
366
+ :param sn_words: the words of the sentence.
367
+ :param sp_ids: the scanpath ids.
368
+ :param tokenizer: the tokenizer.
369
+ :param gpt2_model: the GPT2 model.
370
+ :param bert_embeddings: the BERT embeddings.
371
+ :param aggregate: the aggregation method, either summing or averaging the sub-word embeddings.
372
+ :return: the embeddings of the scanpath, the padded fixation durations, and the attention mask.
373
+ """
374
+
375
+ pad_idx = [i for i, word in enumerate(sn_words) if word == "[PAD]"]
376
+ sep_idx = [sn_words.index("[SEP]")]
377
+ all_remove_idx = [0] # for CLS
378
+ all_remove_idx += sep_idx
379
+ all_remove_idx += pad_idx
380
+
381
+ # get rid of trailing pad tokens in sentence
382
+ while sn_words[-1] == "[PAD]":
383
+ sn_words.pop()
384
+ # get rid of the CLS and SEP tokens
385
+ sn_words = sn_words[1:-1]
386
+
387
+ # the scanpath
388
+ # get rid of predicted CLS, SEP and wrongly predicted PAD tokens (will throw error)
389
+ sp_ids = [sp_id for sp_id in sp_ids if sp_id not in all_remove_idx]
390
+
391
+ # make the scanpath ids start from 0 for re-ordering of the embeddings
392
+ sp_ids = [sp_id - 1 for sp_id in sp_ids]
393
+
394
+ # get the scanpath as fixated words
395
+ sp_words = list()
396
+ for sp_id in sp_ids:
397
+ sp_words.append(sn_words[sp_id])
398
+
399
+ # get the sentence encoding
400
+ sn_enc = tokenizer.encode_plus(
401
+ sn_words,
402
+ add_special_tokens=False,
403
+ return_tensors="pt",
404
+ is_split_into_words=True,
405
+ )
406
+ sn_word_ids = torch.Tensor(sn_enc.word_ids())
407
+
408
+ # get the embeddings
409
+ with torch.no_grad():
410
+ last_hidden = gpt2_model(sn_enc.input_ids).last_hidden_state
411
+
412
+ # aggregate the embeddings to word-level
413
+ sn_embeddings = aggregate_input_embeddings(
414
+ embeddings=last_hidden,
415
+ word_ids=sn_word_ids,
416
+ aggregate=aggregate,
417
+ )
418
+
419
+ # convert sp_ids to tensor
420
+ sp_ids = torch.Tensor(sp_ids).long()
421
+
422
+ # re-order the embeddings as scanpath
423
+ sp_embeddings = sn_embeddings[:, sp_ids, :]
424
+
425
+ # pad the embeddings to max input length and get the attention mask
426
+ sp_embeddings_padded, attention_mask = padding_and_mask_seq2seq(
427
+ sp_embeddings=sp_embeddings,
428
+ bert_embeddings=bert_embeddings,
429
+ max_length=max_length,
430
+ inference=True,
431
+ )
432
+
433
+ return sp_embeddings_padded.squeeze(0), attention_mask.squeeze(0)
434
+
435
+
436
+ def prepare_seq2seq_data_hp(
437
+ scandl_output: Dict[str, Any],
438
+ tokenizer: transformers.GPT2TokenizerFast,
439
+ gpt2_model: transformers.GPT2Model,
440
+ bert_embeddings: torch.nn.Embedding,
441
+ aggregate: str = "mean",
442
+ max_length: int = 128,
443
+ sp_pad_token: int = 127,
444
+ ):
445
+ """
446
+ Prepare the scandl output for inference of the hyper-parameter search of the Seq2Seq fixation duration model.
447
+ :param scandl_output: the ScanDL output.
448
+ :return: the data for inference.
449
+ """
450
+ data_dict = {
451
+ "sp_embeddings": [],
452
+ "attention_masks": [],
453
+ "original_fix_durs": [],
454
+ "predicted_sp_ids": [],
455
+ "reader_ids": [],
456
+ "sn_ids": [],
457
+ }
458
+
459
+ for idx in tqdm(range(len(scandl_output["predicted_sp_ids"]))):
460
+
461
+ sn_repr_len = scandl_output["sn_repr_len"][idx]
462
+ if sp_pad_token == 67:
463
+ # for Chinese: make sure the words are split correctly (chinese characters have no whitespace)
464
+ sn_words = scandl_output["words_for_mapping"][idx].split()
465
+ sn_words = [sn_words[0]] + list(sn_words[1]) + sn_words[2:]
466
+ else:
467
+ sn_words = scandl_output["words_for_mapping"][idx].split()
468
+ sp_ids = scandl_output["predicted_sp_ids"][idx]
469
+
470
+ try:
471
+ sp_embeddings, attention_mask = get_embeddings_seq2seq_hp(
472
+ sn_repr_len=sn_repr_len,
473
+ sn_words=sn_words,
474
+ sp_ids=sp_ids,
475
+ tokenizer=tokenizer,
476
+ gpt2_model=gpt2_model,
477
+ bert_embeddings=bert_embeddings,
478
+ aggregate=aggregate,
479
+ max_length=max_length,
480
+ sp_pad_token=sp_pad_token,
481
+ )
482
+
483
+ # get the original fixation durations
484
+ fix_durs = scandl_output["sn_sp_fix_dur"][idx][sn_repr_len:]
485
+ while fix_durs[-1] == 0:
486
+ fix_durs.pop()
487
+ fix_durs.append(0)
488
+
489
+ data_dict["sp_embeddings"].append(sp_embeddings)
490
+ data_dict["attention_masks"].append(attention_mask)
491
+ data_dict["original_fix_durs"].append(str(fix_durs))
492
+ data_dict["predicted_sp_ids"].append(str(sp_ids))
493
+ data_dict["reader_ids"].append(scandl_output["reader_ids"][idx])
494
+ data_dict["sn_ids"].append(scandl_output["sn_ids"][idx])
495
+ except:
496
+ print(f"Error at index {idx}")
497
+ continue
498
+
499
+ return data_dict
500
+
501
+
502
+ class Seq2SeqDatasetHP(Dataset):
503
+ def __init__(
504
+ self,
505
+ data: Dict[str, Union[torch.Tensor, Any]],
506
+ ):
507
+ super().__init__()
508
+ self.data = data
509
+
510
+ def __len__(self):
511
+ return len(self.data["sp_embeddings"])
512
+
513
+ def __getitem__(self, idx):
514
+ sample = {
515
+ "sp_embeddings": self.data["sp_embeddings"][idx],
516
+ "attention_masks": self.data["attention_masks"][idx],
517
+ #'predicted_sp_words': self.data['predicted_sp_words'][idx],
518
+ #'original_sp_words': self.data['original_sp_words'][idx],
519
+ "predicted_sp_ids": self.data["predicted_sp_ids"][idx],
520
+ # 'original_sp_ids': self.data['original_sp_ids'][idx],
521
+ # 'original_sn': self.data['original_sn'][idx],
522
+ "sn_ids": self.data["sn_ids"][idx],
523
+ "reader_ids": self.data["reader_ids"][idx],
524
+ # 'sn_repr_len': self.data['sn_repr_len'][idx],
525
+ # 'words_for_mapping': self.data['words_for_mapping'][idx],
526
+ # 'sn_sp_repr': self.data['sn_sp_repr'][idx],
527
+ # 'sn_sp_fix_dur': self.data['sn_sp_fix_dur'][idx],
528
+ "original_fix_durs": self.data["original_fix_durs"][idx],
529
+ }
530
+ return sample
fix_dur_module/utils_train.py ADDED
@@ -0,0 +1,195 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import numpy as np
4
+ import transformers
5
+
6
+ from typing import Optional
7
+
8
+
9
+ class EarlyStopping:
10
+ def __init__(
11
+ self,
12
+ patience: int,
13
+ path: str,
14
+ delta: Optional[int] = 0,
15
+ ):
16
+ self.patience = patience
17
+ self.delta = delta
18
+ self.best_score = None
19
+ self.early_stop = False
20
+ self.counter = 0
21
+ self.best_loss = np.inf
22
+ self.path = path
23
+
24
+ def __call__(
25
+ self,
26
+ val_loss,
27
+ model,
28
+ ):
29
+ score = -val_loss
30
+
31
+ if self.best_score is None:
32
+ self.best_score = score
33
+ self.save_checkpoint(val_loss, model)
34
+ elif score < self.best_score + self.delta:
35
+ self.counter += 1
36
+ print(f"EarlyStopping counter: {self.counter} out of {self.patience}")
37
+ if self.counter >= self.patience:
38
+ self.early_stop = True
39
+ else:
40
+ self.best_score = score
41
+ self.save_checkpoint(val_loss, model)
42
+ self.counter = 0
43
+
44
+ def save_checkpoint(
45
+ self,
46
+ val_loss,
47
+ model,
48
+ ):
49
+ """Saves model when validation loss decreases."""
50
+ print(
51
+ f"Validation loss decreased ({self.best_loss:.6f} --> {val_loss:.6f}). Saving model..."
52
+ )
53
+ torch.save(model.state_dict(), self.path)
54
+ self.best_loss = val_loss
55
+
56
+
57
+ def train(
58
+ model,
59
+ num_epochs: int,
60
+ train_loader: torch.utils.data.DataLoader,
61
+ val_loader: torch.utils.data.DataLoader,
62
+ criterion: nn.MSELoss,
63
+ optimizer: transformers.AdamW,
64
+ early_stopping: EarlyStopping,
65
+ scheduler: transformers.get_linear_schedule_with_warmup,
66
+ device: torch.device,
67
+ fix_dur_colname: str,
68
+ output_attentions: Optional[bool] = None,
69
+ use_attention_mask: Optional[bool] = None,
70
+ ):
71
+ """
72
+ Train loop to train the Seq2Seq model.
73
+ :param model: the model to train
74
+ :param num_epochs: number of epochs to train
75
+ :param train_loader: the training data loader
76
+ :param val_loader: the validation data loader
77
+ :param criterion: the loss function (MSE Loss)
78
+ :param optimizer: the optimizer (AdamW)
79
+ :param early_stopping: the early stopping object
80
+ :param scheduler: the learning rate scheduler
81
+ :param device: the device to train on
82
+ :param fix_dur_colname: the name of the column containing the fixations durations
83
+ """
84
+ for epoch in range(num_epochs):
85
+
86
+ model.train()
87
+
88
+ for batch_idx, train_batch in enumerate(train_loader):
89
+
90
+ optimizer.zero_grad()
91
+
92
+ sp_embeddings = train_batch["sp_embeddings"].to(device)
93
+ attention_mask = train_batch["attention_masks"].to(device)
94
+ fix_durs = train_batch[fix_dur_colname].to(device)
95
+
96
+ # forward pass
97
+ if use_attention_mask:
98
+
99
+ if output_attentions:
100
+
101
+ out, _ = model(
102
+ sp_embeddings=sp_embeddings,
103
+ attention_mask=attention_mask,
104
+ output_attentions=output_attentions,
105
+ )
106
+ else:
107
+
108
+ out = model(
109
+ sp_embeddings=sp_embeddings,
110
+ attention_mask=attention_mask,
111
+ output_attentions=output_attentions,
112
+ )
113
+ else:
114
+
115
+ if output_attentions:
116
+
117
+ out, _ = model(
118
+ sp_embeddings=sp_embeddings,
119
+ output_attentions=output_attentions,
120
+ )
121
+ else:
122
+ out = model(
123
+ sp_embeddings=sp_embeddings,
124
+ output_attentions=output_attentions,
125
+ )
126
+
127
+ # train_loss = criterion(out, fix_durs)
128
+ # mask the padding in the loss computation
129
+ loss_mask = (fix_durs != 0).float()
130
+ # train_loss = criterion(out * loss_mask, fix_durs * loss_mask)
131
+ train_loss = criterion(out, fix_durs)
132
+
133
+ train_loss.backward()
134
+ optimizer.step()
135
+ scheduler.step()
136
+
137
+ print(f"\t epoch {epoch+1}, batch {batch_idx+1}, loss: {train_loss.item():.4f}")
138
+
139
+ # validation
140
+
141
+ model.eval()
142
+ val_loss = 0.0
143
+
144
+ with torch.no_grad():
145
+
146
+ for val_batch in val_loader:
147
+
148
+ sp_embeddings = val_batch["sp_embeddings"].to(device)
149
+ attention_mask = val_batch["attention_masks"].to(device)
150
+ fix_durs = val_batch["fix_durs"].to(device)
151
+
152
+ if use_attention_mask:
153
+
154
+ if output_attentions:
155
+ out, attentions = model(
156
+ sp_embeddings=sp_embeddings,
157
+ attention_mask=attention_mask,
158
+ output_attentions=output_attentions,
159
+ )
160
+ else:
161
+ out = model(
162
+ sp_embeddings=sp_embeddings,
163
+ attention_mask=attention_mask,
164
+ output_attentions=output_attentions,
165
+ )
166
+ else:
167
+
168
+ # forward pass
169
+ if output_attentions:
170
+ out, attentions = model(
171
+ sp_embeddings=sp_embeddings,
172
+ output_attentions=output_attentions,
173
+ )
174
+
175
+ else:
176
+ out = model(
177
+ sp_embeddings=sp_embeddings,
178
+ output_attentions=output_attentions,
179
+ )
180
+
181
+ val_loss_mask = (fix_durs != 0).float()
182
+ # val_loss += criterion(out * val_loss_mask, fix_durs * val_loss_mask).item()
183
+ val_loss += criterion(out, fix_durs).item()
184
+
185
+ # average the losses
186
+ val_loss /= len(val_loader)
187
+ train_loss /= len(train_loader)
188
+
189
+ print(f"Epoch {epoch+1}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}")
190
+
191
+ # check for early stopping
192
+ early_stopping(val_loss, model)
193
+ if early_stopping.early_stop:
194
+ print("Early stopping")
195
+ break
handler.py ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+ from ScanDL2 import ScanDL2
4
+
5
+
6
+ class EndpointHandler:
7
+ def __init__(self, path: str = ""):
8
+
9
+ self.models = {
10
+ "sentence": ScanDL2(
11
+ text_type="sentence",
12
+ bsz=2,
13
+ save=None,
14
+ filename=None,
15
+ ),
16
+ "paragraph": ScanDL2(
17
+ text_type="paragraph",
18
+ bsz=2,
19
+ save=None,
20
+ filename=None,
21
+ ),
22
+ }
23
+
24
+ for m in self.models.values():
25
+ # m.to(self.device)
26
+ m.eval()
27
+
28
+ def __call__(self, data):
29
+
30
+ inputs = data.get("inputs", data)
31
+
32
+ parameters = data.get("parameters", {})
33
+
34
+ text_type = parameters.get("text_type", "sentence")
35
+ model = self.models[text_type]
36
+ bsz = parameters.get("bsz", 2)
37
+
38
+ if model.scandl_module.args.batch_size != bsz:
39
+ model.scandl_module.args.batch_size = bsz
40
+ model.fixdur_module.bsz = bsz
41
+ model.fixdur_module.args["bsz"] = bsz
42
+
43
+ if isinstance(inputs, str):
44
+ texts = [inputs]
45
+ elif isinstance(inputs, list):
46
+ texts = inputs
47
+ else:
48
+ raise ValueError("'inputs' must be a string or list of strings.")
49
+
50
+ with torch.no_grad():
51
+ output = model(texts=texts)
52
+
53
+ return output
model.py ADDED
@@ -0,0 +1,701 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import time
4
+ import json
5
+ import joblib
6
+ import argparse
7
+
8
+ from functools import partial
9
+
10
+ from typing import Union, List, Dict, Optional, Any
11
+
12
+ import torch
13
+ import torch.nn as nn
14
+ import torch.distributed as dist
15
+ from torch.utils.data import DataLoader
16
+
17
+ import numpy as np
18
+ import pandas as pd
19
+
20
+ from tqdm import tqdm
21
+
22
+ from transformers import (
23
+ set_seed,
24
+ BertTokenizerFast,
25
+ GPT2TokenizerFast,
26
+ GPT2LMHeadModel,
27
+ GPT2Model,
28
+ AutoConfig,
29
+ BertModel,
30
+ )
31
+ from transformers.models.bert.modeling_bert import BertEncoder
32
+ from datasets import DatasetDict
33
+ from datasets import Dataset as Dataset2
34
+
35
+ from ScanDL2.scandl_module.original_scandl.sp_rounding import denoised_fn_round
36
+ from ScanDL2.scandl_module.original_scandl.utils import dist_util, logger
37
+ from ScanDL2.scandl_module.original_scandl.utils.nn import *
38
+
39
+ from ScanDL2.scandl2_utils import text_dataset_loader, FixdurDataset
40
+
41
+ from ScanDL2.scandl_module.scripts.sp_load_celer_zuco import _collate_batch_helper
42
+ from ScanDL2.scandl_module.scripts.sp_basic_utils import (
43
+ load_defaults_config,
44
+ create_model_and_diffusion,
45
+ add_dict_to_argparser,
46
+ args_to_dict,
47
+ )
48
+
49
+ from ScanDL2.fix_dur_module.model_seq2seq import Seq2SeqModel
50
+ from ScanDL2.fix_dur_module.utils_data import aggregate_input_embeddings, padding_and_mask_seq2seq
51
+
52
+ from ScanDL2.PATHS import (
53
+ SENT_SCANDL_MODULE,
54
+ SENT_FIXDUR_MODULE,
55
+ PAR_SCANDL_MODULE,
56
+ PAR_FIXDUR_MODULE,
57
+ )
58
+
59
+
60
+ class ScanDL2(nn.Module):
61
+ def __init__(
62
+ self,
63
+ text_type: str = "sentence", # sentence, paragraph
64
+ bsz: Optional[int] = 2,
65
+ save: Optional[str] = None,
66
+ filename: Optional[str] = None,
67
+ ):
68
+ super(ScanDL2, self).__init__()
69
+
70
+ self.save = save
71
+ self.filename = filename
72
+
73
+ # initialize the ScanDL module and the Fixdur Module
74
+ self.scandl_module = ScanDLModule(
75
+ text_type=text_type,
76
+ bsz=bsz,
77
+ )
78
+ self.fixdur_module = FixdurModule(
79
+ text_type=text_type,
80
+ bsz=bsz,
81
+ )
82
+
83
+ def forward(
84
+ self,
85
+ texts: Union[str, List[str]],
86
+ ):
87
+ # check if the input is in the correct format
88
+ self._validate_inputs(texts=texts)
89
+
90
+ # get the fixation location predictions from the ScanDL module
91
+ scandl_module_output = self.scandl_module(texts=texts)
92
+
93
+ # get the fixation duration predictions from the Fixdur module
94
+ fixdur_module_output = self.fixdur_module(scandl_module_output=scandl_module_output)
95
+
96
+ if self.save is not None:
97
+ filename = self.filename if self.filename is not None else f"scandl2_outputs.json"
98
+ filename = f"{filename}.json" if not filename.endswith(".json") else filename
99
+ self._save_results(results=fixdur_module_output, filename=filename)
100
+
101
+ return fixdur_module_output
102
+
103
+ def _save_results(self, results, filename):
104
+ if not os.path.exists(self.save):
105
+ os.makedirs(self.save)
106
+ with open(os.path.join(self.save, filename), "w") as f:
107
+ json.dump(results, f)
108
+ print(f"--- ScanDL 2.0 outputs saved to {os.path.join(self.save, filename)}.")
109
+
110
+ def _validate_inputs(self, texts: Union[str, List[str]]) -> None:
111
+ if not isinstance(texts, (str, list)) or (
112
+ isinstance(texts, list) and not all(isinstance(t, str) for t in texts)
113
+ ):
114
+ raise TypeError("Invalid input: 'texts' must be of type 'str' or 'List[str]'.")
115
+
116
+
117
+ class ScanDLModule(nn.Module):
118
+
119
+ def __init__(
120
+ self,
121
+ text_type: str, # sentence, paragraph
122
+ bsz: int,
123
+ ):
124
+ super(ScanDLModule, self).__init__()
125
+
126
+ base_path = os.path.dirname(__file__)
127
+ if text_type == "paragraph":
128
+ self.path_to_config = os.path.join(base_path, "config_emtec.json")
129
+ self.path_to_scandl_module = PAR_SCANDL_MODULE
130
+ elif text_type == "sentence":
131
+ self.path_to_config = os.path.join(base_path, "config.json")
132
+ self.path_to_scandl_module = SENT_SCANDL_MODULE
133
+ else:
134
+ raise NotImplementedError(f"Text type {text_type} not implemented.")
135
+
136
+ # get the args
137
+ self.args = self._get_args()
138
+ self.args.batch_size = bsz
139
+
140
+ # seting up the environment
141
+ dist_util.setup_dist()
142
+ logger.configure()
143
+ self.world_size = dist.get_world_size() or 1
144
+ self.rank = dist.get_rank() or 0
145
+ # set_seed(self.args.seed2)
146
+
147
+ # load the tokenizer
148
+ self.tokenizer = self._load_tokenizer()
149
+
150
+ # load the ScanDL module and the Diffusion
151
+ self.scandl_module, self.diffusion = self._load_scandl_module(
152
+ path_to_scandl_module=self.path_to_scandl_module
153
+ )
154
+ self.sn_sp_repr_embedding = self._get_sn_sp_repr_emb()
155
+
156
+ def forward(
157
+ self,
158
+ texts: Union[str, List[str]],
159
+ ) -> Dict[str, Union[List[List[str]], List[List[int]], List[str]]]:
160
+
161
+ data_loader = self._preprocess_text(texts=texts)
162
+
163
+ predicted_sp_words, predicted_sp_ids = [], []
164
+ original_sn = []
165
+
166
+ print("\t\t### ScanDL Module generates fixation locations ...")
167
+
168
+ unique_idx = list()
169
+ idx_ctr = 0
170
+
171
+ for batch_idx, batch in tqdm(enumerate(data_loader)):
172
+
173
+ mask = batch["mask"].to(dist_util.dev())
174
+ sn_sp_repr = batch["sn_sp_repr"].to(dist_util.dev())
175
+ sn_input_ids = batch["sn_input_ids"].to(dist_util.dev())
176
+ indices_pos_enc = batch["indices_pos_enc"].to(dist_util.dev())
177
+ sn_repr_len = batch["sn_repr_len"].to(dist_util.dev())
178
+ words_for_mapping = batch["words_for_mapping"]
179
+
180
+ sn_sp_emb, pos_enc, sn_input_ids_emb = self.scandl_module.get_embeds(
181
+ sn_sp_repr=sn_sp_repr,
182
+ sn_input_ids=sn_input_ids,
183
+ indices_pos_enc=indices_pos_enc,
184
+ )
185
+
186
+ x_start = sn_sp_emb
187
+ noise = torch.randn_like(x_start)
188
+ mask = torch.broadcast_to(mask.unsqueeze(dim=-1), x_start.shape).to(dist_util.dev())
189
+ x_noised = torch.where(mask == 0, x_start, noise)
190
+
191
+ self.args.use_ddim = False
192
+ step_gap = 1
193
+
194
+ sample_fn = (
195
+ self.diffusion.p_sample_loop
196
+ if not self.args.use_ddim
197
+ else self.diffusion.ddim_sample_loop
198
+ )
199
+
200
+ sample_shape = (x_start.shape[0], self.args.seq_len, self.args.hidden_dim)
201
+ subwords = [self.tokenizer.convert_ids_to_tokens(i) for i in sn_input_ids]
202
+
203
+ samples = sample_fn(
204
+ model=self.scandl_module,
205
+ shape=sample_shape,
206
+ noise=x_noised,
207
+ sn_input_ids_emb=sn_input_ids_emb,
208
+ pos_enc=pos_enc,
209
+ mask_sn_padding=None,
210
+ mask_transformer_att=None,
211
+ clip_denoised=self.args.clip_denoised,
212
+ denoised_fn=partial(denoised_fn_round, self.args, self.sn_sp_repr_embedding),
213
+ model_kwargs=None,
214
+ top_p=self.args.top_p,
215
+ clamp_step=self.args.clamp_step,
216
+ clamp_first=self.args.clamp_first_bool,
217
+ mask=mask,
218
+ x_start=x_start,
219
+ gap=step_gap,
220
+ )
221
+ sample = samples[-1]
222
+
223
+ logits = self.scandl_module.get_logits(sample)
224
+ cands = torch.topk(logits, k=1, dim=-1)
225
+
226
+ for instance_idx, (pred_seq, orig_words, sn_len) in enumerate(
227
+ zip(cands.indices, words_for_mapping, sn_repr_len)
228
+ ):
229
+ pred_seq_sp = pred_seq[sn_len:]
230
+ words_split = orig_words.split()
231
+ predicted_sp = [words_split[i] for i in pred_seq_sp]
232
+ pred_sp_ids = [e.item() for e in pred_seq_sp]
233
+
234
+ # cut off trailing pad tokens
235
+ while len(predicted_sp) > 1 and predicted_sp[-1] == "[PAD]":
236
+ predicted_sp.pop()
237
+ while len(pred_sp_ids) > 1 and pred_sp_ids[-1] == self.args.seq_len - 1:
238
+ pred_sp_ids.pop()
239
+ while len(words_split) > 1 and words_split[-1] == "[PAD]":
240
+ words_split.pop()
241
+
242
+ # remove CLS and SEP tokens from predictions
243
+ if predicted_sp[0] == "[CLS]":
244
+ predicted_sp = predicted_sp[1:]
245
+ pred_sp_ids = pred_sp_ids[1:]
246
+ if predicted_sp[-1] == "[SEP]":
247
+ predicted_sp = predicted_sp[:-1]
248
+ pred_sp_ids = pred_sp_ids[:-1]
249
+ words_split = words_split[1:-1]
250
+
251
+ # filter out erroneously predicted PAD tokens (they will raise an error in the fixdur module)
252
+ pred_sp_ids, predicted_sp = self._remove_special_tokens(
253
+ predicted_sp_ids=pred_sp_ids,
254
+ predicted_sp_words=predicted_sp,
255
+ token="[PAD]",
256
+ )
257
+ pred_sp_ids, predicted_sp = self._remove_special_tokens(
258
+ predicted_sp_ids=pred_sp_ids,
259
+ predicted_sp_words=predicted_sp,
260
+ token="[CLS]",
261
+ )
262
+ pred_sp_ids, predicted_sp = self._remove_special_tokens(
263
+ predicted_sp_ids=pred_sp_ids,
264
+ predicted_sp_words=predicted_sp,
265
+ token="[SEP]",
266
+ )
267
+
268
+ predicted_sp_words.append(predicted_sp)
269
+ predicted_sp_ids.append(pred_sp_ids)
270
+ original_sn.append(words_split)
271
+
272
+ idx_ctr += 1
273
+ unique_idx.append(idx_ctr)
274
+
275
+ predictions = {
276
+ "predicted_sp_words": predicted_sp_words,
277
+ "predicted_sp_ids": predicted_sp_ids,
278
+ "original_sn": original_sn,
279
+ "unique_idx": unique_idx,
280
+ }
281
+ return predictions
282
+
283
+ def _remove_special_tokens(
284
+ self,
285
+ predicted_sp_ids: List[int],
286
+ predicted_sp_words: List[str],
287
+ token: str, # '[CLS]' or '[SEP]' or '[PAD]'
288
+ ):
289
+ filtered_sp_ids, filtered_sp_words = [], []
290
+ for sp_word, sp_id in zip(predicted_sp_words, predicted_sp_ids):
291
+ if sp_word != token:
292
+ filtered_sp_ids.append(sp_id)
293
+ filtered_sp_words.append(sp_word)
294
+ return filtered_sp_ids, filtered_sp_words
295
+
296
+ def _preprocess_text(
297
+ self,
298
+ texts: Union[str, List[str]],
299
+ ):
300
+ data = {
301
+ "mask": [],
302
+ "sn_sp_repr": [],
303
+ "sn_input_ids": [],
304
+ "indices_pos_enc": [],
305
+ "words_for_mapping": [],
306
+ "sn_repr_len": [],
307
+ }
308
+
309
+ if isinstance(texts, str):
310
+ texts = [texts]
311
+
312
+ for sn_idx, sn in enumerate(texts):
313
+
314
+ if sn.startswith("[CLS]") and sn.endswith("[SEP]"):
315
+ sn = sn
316
+ elif sn.startswith("[CLS]"):
317
+ sn = sn + " [SEP]"
318
+ elif sn.endswith("[SEP]"):
319
+ sn = "[CLS] " + sn
320
+ else:
321
+ sn = "[CLS] " + sn + " [SEP]"
322
+
323
+ encoded_sn = self.tokenizer.encode_plus(
324
+ sn.split(),
325
+ add_special_tokens=False,
326
+ padding=False,
327
+ return_attention_mask=False,
328
+ is_split_into_words=True,
329
+ truncation=False,
330
+ )
331
+
332
+ if len(encoded_sn) > self.args.seq_len / 2:
333
+ print(f"Sentence {sn} is too long. Continue.")
334
+
335
+ sn_word_ids = encoded_sn.word_ids()
336
+ sn_input_ids = encoded_sn["input_ids"]
337
+
338
+ sn_sp_repr = sn_word_ids
339
+
340
+ mask = [0] * len(sn_word_ids)
341
+ indices_pos_enc = list(range(0, len(sn_word_ids))) + list(
342
+ range(0, self.args.seq_len - len(sn_word_ids))
343
+ )
344
+ words_for_mapping = sn.split() + (self.args.seq_len - len(sn.split())) * ["[PAD]"]
345
+
346
+ data["mask"].append(mask)
347
+ data["sn_sp_repr"].append(sn_sp_repr)
348
+ data["sn_input_ids"].append(sn_input_ids)
349
+ data["indices_pos_enc"].append(indices_pos_enc)
350
+ data["words_for_mapping"].append(" ".join(words_for_mapping))
351
+ data["sn_repr_len"].append(len(sn_word_ids))
352
+
353
+ # padding
354
+ data["mask"] = _collate_batch_helper(
355
+ examples=data["mask"],
356
+ pad_token_id=1,
357
+ max_length=self.args.seq_len,
358
+ )
359
+ data["sn_sp_repr"] = _collate_batch_helper(
360
+ examples=data["sn_sp_repr"],
361
+ pad_token_id=self.args.seq_len - 1,
362
+ max_length=self.args.seq_len,
363
+ )
364
+ data["sn_input_ids"] = _collate_batch_helper(
365
+ examples=data["sn_input_ids"],
366
+ pad_token_id=self.tokenizer.pad_token_id,
367
+ max_length=self.args.seq_len,
368
+ )
369
+
370
+ split = "inference"
371
+ dataset = Dataset2.from_dict(data)
372
+ dataset_dict = DatasetDict()
373
+ dataset_dict[split] = dataset
374
+ data_loader = text_dataset_loader(
375
+ data=dataset_dict,
376
+ data_args=self.args,
377
+ split=split,
378
+ deterministic=True,
379
+ )
380
+ return data_loader
381
+
382
+ def _load_scandl_module(
383
+ self,
384
+ path_to_scandl_module: str,
385
+ ):
386
+ logger.log("### Loading ScanDL Diffusion Module ...")
387
+ scandl_module, diffusion = create_model_and_diffusion(
388
+ **args_to_dict(self.args, load_defaults_config(config_path=self.path_to_config).keys())
389
+ )
390
+ # TODO Name scandl module, not model
391
+ scandl_module.load_state_dict(
392
+ dist_util.load_state_dict(
393
+ os.path.join(self.path_to_scandl_module, "ema_0.9999_080000.pt"), map_location="cpu"
394
+ )
395
+ )
396
+ pytorch_total_params = sum(p.numel() for p in scandl_module.parameters())
397
+ logger.log(f"### Total number of parameters: {pytorch_total_params}")
398
+ scandl_module.eval().requires_grad_(False).to(dist_util.dev())
399
+ return scandl_module, diffusion
400
+
401
+ def _get_sn_sp_repr_emb(self):
402
+ sn_sp_repr_embedding = nn.Embedding(
403
+ num_embeddings=self.args.hidden_t_dim,
404
+ embedding_dim=self.args.hidden_dim,
405
+ _weight=self.scandl_module.sn_sp_repr_embedding.weight.clone().cpu(),
406
+ )
407
+ return sn_sp_repr_embedding
408
+
409
+ def _get_args(self):
410
+ args = self._get_parser().parse_args()
411
+ # load the training arguments
412
+ with open(os.path.join(self.path_to_scandl_module, "training_args.json")) as f:
413
+ training_args = json.load(f)
414
+ training_args["batch_size"] = args.batch_size
415
+ args.__dict__.update(training_args)
416
+ if args.clamp_first == "yes":
417
+ args.clamp_first_bool = True
418
+ else:
419
+ args.clamp_first_bool = False
420
+ # TODO self.args.clamp_first_bool as argument
421
+ # set mask_padding to False
422
+ args.mask_padding = False
423
+ return args
424
+
425
+ def _load_tokenizer(self):
426
+ tokenizer = BertTokenizerFast.from_pretrained(self.args.config_name)
427
+ self.args.vocab_size = tokenizer.vocab_size
428
+ return tokenizer
429
+
430
+ def _get_parser(self) -> argparse.ArgumentParser:
431
+ defaults = dict(
432
+ model_path="",
433
+ step=0,
434
+ out_dir="",
435
+ top_p=0,
436
+ clamp_first="yes",
437
+ test_set_sns="mixed",
438
+ atten_vis=False,
439
+ notes="-",
440
+ tsne_vis=False,
441
+ sp_vis=False,
442
+ no_inst=0,
443
+ atten_vis_sp=False,
444
+ load_ids="-",
445
+ load_test_data="-",
446
+ setting="-",
447
+ fold=0,
448
+ )
449
+ decode_defaults = dict(
450
+ split="valid",
451
+ clamp_step=0,
452
+ seed2=105,
453
+ clip_denoised=False,
454
+ )
455
+
456
+ defaults.update(load_defaults_config(config_path=self.path_to_config))
457
+ defaults.update(decode_defaults)
458
+ parser = argparse.ArgumentParser()
459
+ add_dict_to_argparser(parser, defaults)
460
+
461
+ return parser
462
+
463
+
464
+ class FixdurModule(nn.Module):
465
+
466
+ def __init__(
467
+ self,
468
+ text_type: str, # sentence, paragraph
469
+ bsz: Optional[int] = 2,
470
+ ):
471
+ super(FixdurModule, self).__init__()
472
+
473
+ base_path = os.path.dirname(__file__)
474
+ if text_type == "paragraph":
475
+ self.path_to_config = os.path.join(base_path, "config_emtec.json")
476
+ self.path_to_fixdur_module = PAR_FIXDUR_MODULE
477
+ elif text_type == "sentence":
478
+ self.path_to_config = os.path.join(base_path, "config.json")
479
+ self.path_to_fixdur_module = SENT_FIXDUR_MODULE
480
+ else:
481
+ raise NotImplementedError(f"Text type {text_type} not implemented.")
482
+
483
+ self.bsz = bsz
484
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
485
+
486
+ self.config = load_defaults_config(config_path=self.path_to_config)
487
+ self.args = self._get_args(config=self.config)
488
+ self.hyperparameters = self._get_hyperparams(
489
+ path_to_fixdur_module=self.path_to_fixdur_module
490
+ )
491
+
492
+ # load GPT-2 model and tokenizer, and BERT embeddings
493
+ self.gpt2_model, self.tokenizer, self.bert_embeddings = self._load_gpt_and_bert(
494
+ config=self.config
495
+ )
496
+
497
+ # load the Fixdur module and the MinMax Scaler
498
+ self.fixdur_module = self._load_fixdur_module()
499
+ self.scaler = self._load_scaler()
500
+
501
+ def _get_args(self, config: Dict[str, Any]) -> Dict[str, Any]:
502
+ args = {
503
+ "max_length": config["seq_len"],
504
+ "normalize": True,
505
+ "output_attentions": False,
506
+ "bsz": self.bsz,
507
+ "corpus": config["corpus"],
508
+ "sp_pad_token": config["seq_len"] - 1,
509
+ }
510
+ return args
511
+
512
+ def _get_hyperparams(self, path_to_fixdur_module: str) -> Dict[str, Any]:
513
+ with open(os.path.join(path_to_fixdur_module, "hyperparameters.json")) as f:
514
+ return json.load(f)
515
+
516
+ def _load_gpt_and_bert(self, config: Dict[str, Any]):
517
+ """
518
+ Load GPT-2 and GPT-2 tokenizer to get the contextualized embeddings.
519
+ Load BERT model (for embeddings of CLS and PAD tokens)
520
+ """
521
+ # GPT-2
522
+ gpt_config_name = config["gpt_config_name"]
523
+ tokenizer = GPT2TokenizerFast.from_pretrained(gpt_config_name, add_prefix_space=True)
524
+ gpt2_model = GPT2Model.from_pretrained(gpt_config_name)
525
+ tokenizer.pad_token = tokenizer.eos_token
526
+ # freeze parameters
527
+ for param in gpt2_model.parameters():
528
+ param.requires_grad = False
529
+
530
+ # BERT
531
+ bert_config_name = config["config_name"]
532
+ bert_embeddings = BertModel.from_pretrained(bert_config_name).embeddings.word_embeddings
533
+ # freeze parameters
534
+ for param in bert_embeddings.parameters():
535
+ param.requires_grad = False
536
+
537
+ return gpt2_model, tokenizer, bert_embeddings
538
+
539
+ def _load_fixdur_module(self):
540
+ fixdur_module_config = AutoConfig.from_pretrained("bert-base-cased")
541
+ fixdur_module_config.num_attention_heads = self.hyperparameters["num_heads"]
542
+ fixdur_module_config.num_hidden_layers = self.hyperparameters["num_layers"]
543
+ fixdur_module = Seq2SeqModel(
544
+ config=fixdur_module_config,
545
+ output_dim=self.args["max_length"],
546
+ num_linear=self.hyperparameters["num_linear"],
547
+ dropout=self.hyperparameters["dropout"],
548
+ )
549
+ fixdur_module.load_state_dict(
550
+ torch.load(
551
+ os.path.join(self.path_to_fixdur_module, "seq2seq_fixdur.pt"),
552
+ map_location=self.device,
553
+ )
554
+ )
555
+ fixdur_module.eval()
556
+ fixdur_module.to(self.device)
557
+ return fixdur_module
558
+
559
+ def _load_scaler(self):
560
+ scaler = joblib.load(os.path.join(self.path_to_fixdur_module, "min_max_scaler.pkl"))
561
+ return scaler
562
+
563
+ def _prepare_data(
564
+ self,
565
+ scandl_module_output: Dict[str, Union[List[List[str]], List[List[int]], List[str]]],
566
+ ):
567
+ data_dict = {
568
+ "sp_embeddings": [],
569
+ "attention_mask": [],
570
+ "unique_idx": [],
571
+ }
572
+ for idx in range(len(scandl_module_output["predicted_sp_words"])):
573
+
574
+ sn_words = scandl_module_output["original_sn"][idx]
575
+ sp_ids = scandl_module_output["predicted_sp_ids"][idx]
576
+ unique_id = scandl_module_output["unique_idx"][idx]
577
+
578
+ # make the scanpath ids start at 0 for re-ordering of the embeddings
579
+ sp_ids = [i - 1 for i in sp_ids]
580
+
581
+ sp_words = scandl_module_output["predicted_sp_words"][idx]
582
+
583
+ # get the sentence encoding
584
+ sn_enc = self.tokenizer(
585
+ sn_words,
586
+ add_special_tokens=False,
587
+ return_tensors="pt",
588
+ is_split_into_words=True,
589
+ )
590
+ sn_word_ids = torch.Tensor(sn_enc.word_ids())
591
+
592
+ # get the embeddings
593
+ with torch.no_grad():
594
+ last_hidden = self.gpt2_model(sn_enc.input_ids).last_hidden_state
595
+
596
+ # aggregate the embeddings to word level
597
+ sn_embeddings = aggregate_input_embeddings(
598
+ embeddings=last_hidden,
599
+ word_ids=sn_word_ids,
600
+ aggregate="mean",
601
+ )
602
+
603
+ # convert sp_ids to tensor
604
+ sp_ids = torch.Tensor(sp_ids).long()
605
+
606
+ # re-order the embeddings as scanpath
607
+ try:
608
+ sp_embeddings = sn_embeddings[:, sp_ids, :]
609
+ except:
610
+ breakpoint()
611
+
612
+ # pad the embeddings to max input length and get the attentino mask
613
+ sp_embeddings_padded, attention_mask = padding_and_mask_seq2seq(
614
+ sp_embeddings=sp_embeddings,
615
+ bert_embeddings=self.bert_embeddings,
616
+ max_length=self.args["max_length"],
617
+ inference=True,
618
+ )
619
+
620
+ data_dict["sp_embeddings"].append(sp_embeddings_padded)
621
+ data_dict["attention_mask"].append(attention_mask)
622
+ data_dict["unique_idx"].append(unique_id)
623
+
624
+ return data_dict
625
+
626
+ def forward(
627
+ self,
628
+ scandl_module_output: Dict[str, Union[List[List[str]], List[List[int]], List[str]]],
629
+ ) -> Dict[str, Union[List[List[str]], List[List[int]], List[str], List[List[float]]]]:
630
+
631
+ output_dict = {
632
+ "predicted_sp_words": [],
633
+ "predicted_sp_ids": [],
634
+ "original_sn": [],
635
+ "predicted_fix_durs": [],
636
+ "unique_idx": [],
637
+ }
638
+
639
+ data_df = pd.DataFrame(scandl_module_output)
640
+
641
+ data_dict = self._prepare_data(
642
+ scandl_module_output=scandl_module_output,
643
+ )
644
+ dataset = FixdurDataset(data=data_dict)
645
+ data_loader = DataLoader(
646
+ dataset,
647
+ batch_size=self.bsz,
648
+ shuffle=False,
649
+ )
650
+
651
+ print("\t\t### FixDur Module generates fixation durations ...")
652
+ for batch_idx, batch in tqdm(enumerate(data_loader)):
653
+
654
+ sp_embeddings = batch["sp_embeddings"].squeeze(1).to(self.device)
655
+ attention_mask = batch["attention_mask"].squeeze(1).to(self.device)
656
+ unique_indices = batch["unique_idx"]
657
+
658
+ out = self.fixdur_module(
659
+ sp_embeddings=sp_embeddings,
660
+ attention_mask=attention_mask,
661
+ output_attentions=self.args["output_attentions"],
662
+ )
663
+
664
+ # scale the output back to the original range
665
+ out_transformed = self.scaler.inverse_transform(out.detach().cpu().numpy())
666
+ out_transformed_rounded = np.round(out_transformed, 2)
667
+
668
+ # iterate over the individual predictions
669
+ for out_idx, out_instance in enumerate(out_transformed_rounded):
670
+
671
+ predicted_fix_durs = out_instance
672
+ unique_idx = unique_indices[out_idx].item()
673
+
674
+ # find predicted_sp_words, predicted_sp_ids, and original_sn in data_df conditioned on unique_idx
675
+ predicted_sp_words = data_df.loc[
676
+ data_df["unique_idx"] == unique_idx, "predicted_sp_words"
677
+ ].values[0]
678
+ predicted_sp_ids = data_df.loc[
679
+ data_df["unique_idx"] == unique_idx, "predicted_sp_ids"
680
+ ].values[0]
681
+ original_sn = data_df.loc[
682
+ data_df["unique_idx"] == unique_idx, "original_sn"
683
+ ].values[0]
684
+
685
+ sp_len = len(predicted_sp_ids)
686
+
687
+ # cut off the predicted_fix_durs to the length of the scanpath
688
+ # the predicted fixation durations still contain predictions for the CLS and SEP token as well
689
+ pred_fix_durs = predicted_fix_durs[: sp_len + 2].tolist()[1:-1]
690
+ pred_fix_durs = [round(d, 2) for d in pred_fix_durs]
691
+
692
+ # add to output_dict
693
+ output_dict["predicted_sp_words"].append(predicted_sp_words)
694
+ output_dict["predicted_sp_ids"].append(predicted_sp_ids)
695
+ output_dict["original_sn"].append(original_sn)
696
+ output_dict["predicted_fix_durs"].append(pred_fix_durs)
697
+ output_dict["unique_idx"].append(unique_idx)
698
+
699
+ print(f"fixdur original sn: {original_sn}")
700
+
701
+ return output_dict
models/paragraph/fixdur-module/hyperparameters.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {"num_heads": 12, "num_layers": 12, "num_linear": 8, "bsz": 48, "dropout": 0.5, "use_attention_mask": true}
models/paragraph/fixdur-module/min_max_scaler.pkl ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0b234e10b46971a6d87694bdaa150483394cb952ebf984ab1c93801caaec927c
3
+ size 667
models/paragraph/fixdur-module/seq2seq_fixdur.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f4cef30a1665cb4d09c5622b2bd91622f8a12ab9d5cdc1c1f6da347ffd4d8885
3
+ size 362646514
models/paragraph/scandl-module/ema_0.9999_080000.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:eea605cedc021bc437f965febf00ab01fba073ca9b0d955ba282af99876ac236
3
+ size 616139530
models/paragraph/scandl-module/training_args.json ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "checkpoint_path": "/data/lenbol/projects/ScanDL-fix-dur/complete/EMTeC/scandl-module/checkpoint-path",
3
+ "vocab": "bert",
4
+ "use_plm_init": "no",
5
+ "lr": 0.0001,
6
+ "batch_size": 64,
7
+ "microbatch": 64,
8
+ "diffusion_steps": 2000,
9
+ "noise_schedule": "sqrt",
10
+ "schedule_sampler": "lossaware",
11
+ "seq_len": 352,
12
+ "resume_checkpoint": "none",
13
+ "hidden_t_dim": 352,
14
+ "seed": 101,
15
+ "hidden_dim": 256,
16
+ "learning_steps": 80000,
17
+ "save_interval": 5000,
18
+ "notes": "-",
19
+ "data_split_criterion": "reader",
20
+ "num_transformer_layers": 12,
21
+ "num_transformer_heads": 8,
22
+ "corpus": "emtec",
23
+ "inference": "cv",
24
+ "load_train_data": "processed_data_all_emtec",
25
+ "log_interval": 50,
26
+ "eval_interval": 500,
27
+ "ema_rate": "0.9999",
28
+ "timestep_respacing": "",
29
+ "vocab_size": 28996,
30
+ "config_name": "bert-base-cased",
31
+ "data_dir": "processed_data",
32
+ "dataset": "dataset-name",
33
+ "dropout": 0.1,
34
+ "use_fp16": false,
35
+ "fp16_scale_growth": 0.001,
36
+ "gradient_clipping": -1.0,
37
+ "weight_decay": 0.0,
38
+ "learn_sigma": false,
39
+ "use_kl": false,
40
+ "predict_xstart": true,
41
+ "rescale_timesteps": true,
42
+ "rescale_learned_sigmas": false,
43
+ "sigma_small": false,
44
+ "emb_scale_factor": 1.0,
45
+ "one_noise_step": true,
46
+ "mask_padding": false,
47
+ "celer_only_L1": true,
48
+ "n_folds": 5,
49
+ "ablation_type": "none",
50
+ "nll_in_loss": false,
51
+ "load_from_checkpoint": false
52
+ }
models/sentence/fixdur-module/hyperparameters.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {"num_heads": 12, "num_layers": 12, "num_linear": 8, "bsz": 128, "dropout": 0.5, "use_attention_mask": true}
models/sentence/fixdur-module/min_max_scaler.pkl ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9cea38a7b6b840f0ea990707db171ae219ecf0796d1824e57be935cd66b16dd8
3
+ size 667
models/sentence/fixdur-module/seq2seq_fixdur.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:75e490252c13d83cc099e1a2bc1c63bd75373e17fd29229c58fa10a9e9cc00e5
3
+ size 361957490
models/sentence/scandl-module/ema_0.9999_080000.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1109935957d7a54efcdd0979dcdeb2a4cde0ae1e50d6274a5ed4ce16e301b8a9
3
+ size 612809098
models/sentence/scandl-module/training_args.json ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "checkpoint_path": "/data/lenbol/projects/ScanDL-fix-dur/complete/CELER/scandl-module/checkpoint-path",
3
+ "vocab": "bert",
4
+ "use_plm_init": "no",
5
+ "lr": 0.0001,
6
+ "batch_size": 64,
7
+ "microbatch": 64,
8
+ "diffusion_steps": 2000,
9
+ "noise_schedule": "sqrt",
10
+ "schedule_sampler": "lossaware",
11
+ "seq_len": 128,
12
+ "resume_checkpoint": "none",
13
+ "hidden_t_dim": 128,
14
+ "seed": 101,
15
+ "hidden_dim": 256,
16
+ "learning_steps": 80000,
17
+ "save_interval": 5000,
18
+ "notes": "-",
19
+ "data_split_criterion": "reader",
20
+ "num_transformer_layers": 12,
21
+ "num_transformer_heads": 8,
22
+ "corpus": "celer",
23
+ "inference": "cv",
24
+ "load_train_data": "processed_data_all_celer",
25
+ "log_interval": 50,
26
+ "eval_interval": 500,
27
+ "ema_rate": "0.9999",
28
+ "timestep_respacing": "",
29
+ "vocab_size": 28996,
30
+ "config_name": "bert-base-cased",
31
+ "data_dir": "processed_data",
32
+ "dataset": "dataset-name",
33
+ "dropout": 0.1,
34
+ "use_fp16": false,
35
+ "fp16_scale_growth": 0.001,
36
+ "gradient_clipping": -1.0,
37
+ "weight_decay": 0.0,
38
+ "learn_sigma": false,
39
+ "use_kl": false,
40
+ "predict_xstart": true,
41
+ "rescale_timesteps": true,
42
+ "rescale_learned_sigmas": false,
43
+ "sigma_small": false,
44
+ "emb_scale_factor": 1.0,
45
+ "one_noise_step": true,
46
+ "mask_padding": false,
47
+ "celer_only_L1": true,
48
+ "n_folds": 5,
49
+ "ablation_type": "none",
50
+ "nll_in_loss": false,
51
+ "load_from_checkpoint": false
52
+ }
requirements.txt ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ blobfile==2.0.1
2
+ datasets==2.14.7
3
+ huggingface-hub==0.17.3
4
+ joblib
5
+ matplotlib>=3.7,<3.9
6
+ numpy==1.23.5
7
+ openpyxl==3.0.10
8
+ pandas==1.5.3
9
+ scikit-learn==1.6.1
10
+ seaborn==0.12.2
11
+ textdistance
12
+ tqdm==4.66.4
13
+ transformers==4.34.1
14
+ wandb==0.14.0
15
+ setuptools==68.2.2 #for compatibility with wandb 0.14.0
scandl2_utils.py ADDED
@@ -0,0 +1,77 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import pandas as pd
2
+ import numpy as np
3
+ from tqdm import tqdm
4
+ import os
5
+ import random
6
+ import torch
7
+ from torch.utils.data import Dataset, DataLoader
8
+ from typing import Dict, Union, Any, Optional, List
9
+
10
+
11
+ class TextDataset(Dataset):
12
+
13
+ def __init__(
14
+ self,
15
+ dataset,
16
+ data_args,
17
+ split, # 'train', 'test', 'val'
18
+ ):
19
+ super().__init__()
20
+ self.dataset = dataset
21
+ self.length = len(self.dataset[split])
22
+ self.data_args = data_args
23
+ self.split = split
24
+
25
+ def __len__(self):
26
+ return self.length
27
+
28
+ def __getitem__(self, idx):
29
+ sample = {
30
+ "mask": np.array(self.dataset[self.split][idx]["mask"]),
31
+ "sn_sp_repr": np.array(self.dataset[self.split][idx]["sn_sp_repr"]),
32
+ "sn_input_ids": np.array(self.dataset[self.split][idx]["sn_input_ids"]),
33
+ "indices_pos_enc": np.array(self.dataset[self.split][idx]["indices_pos_enc"]),
34
+ "sn_repr_len": np.array(self.dataset[self.split][idx]["sn_repr_len"]),
35
+ "words_for_mapping": self.dataset[self.split][idx]["words_for_mapping"],
36
+ }
37
+ return sample
38
+
39
+
40
+ def text_dataset_loader(
41
+ data,
42
+ data_args,
43
+ split: str,
44
+ deterministic: bool = False,
45
+ ):
46
+ dataset = TextDataset(
47
+ dataset=data,
48
+ data_args=data_args,
49
+ split=split,
50
+ )
51
+ data_loader = DataLoader(
52
+ dataset,
53
+ batch_size=data_args.batch_size,
54
+ shuffle=not deterministic,
55
+ num_workers=0,
56
+ )
57
+ return iter(data_loader)
58
+
59
+
60
+ class FixdurDataset(Dataset):
61
+ def __init__(
62
+ self,
63
+ data: Dict[str, Union[torch.Tensor, Any]],
64
+ ):
65
+ super().__init__()
66
+ self.data = data
67
+
68
+ def __len__(self):
69
+ return len(self.data["sp_embeddings"])
70
+
71
+ def __getitem__(self, idx):
72
+ sample = {
73
+ "sp_embeddings": self.data["sp_embeddings"][idx],
74
+ "attention_mask": self.data["attention_mask"][idx],
75
+ "unique_idx": self.data["unique_idx"][idx],
76
+ }
77
+ return sample
scandl_module/.DS_Store ADDED
Binary file (6.15 kB). View file
 
scandl_module/__init__.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ from .original_scandl import utils as utils
2
+ from . import original_scandl as original_scandl
3
+
4
+ __all__ = ["utils", "original_scandl"]
scandl_module/__pycache__/__init__.cpython-313.pyc ADDED
Binary file (294 Bytes). View file
 
scandl_module/original_scandl/__init__.py ADDED
File without changes
scandl_module/original_scandl/__pycache__/__init__.cpython-313.pyc ADDED
Binary file (196 Bytes). View file
 
scandl_module/original_scandl/__pycache__/sp_gaussian_diffusion.cpython-313.pyc ADDED
Binary file (43.3 kB). View file
 
scandl_module/original_scandl/__pycache__/sp_rounding.cpython-313.pyc ADDED
Binary file (3.51 kB). View file
 
scandl_module/original_scandl/__pycache__/sp_transformer_model.cpython-313.pyc ADDED
Binary file (6.81 kB). View file
 
scandl_module/original_scandl/__pycache__/step_sample.cpython-313.pyc ADDED
Binary file (9.98 kB). View file
 
scandl_module/original_scandl/config.json ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "lr": 0.0001,
3
+ "batch_size": 128,
4
+ "microbatch": 64,
5
+ "learning_steps": 80000,
6
+ "log_interval": 50,
7
+ "save_interval": 5000,
8
+ "eval_interval": 500,
9
+ "ema_rate": "0.9999",
10
+ "resume_checkpoint": "none",
11
+ "schedule_sampler": "lossaware",
12
+ "diffusion_steps": 2000,
13
+ "noise_schedule": "sqrt",
14
+ "timestep_respacing": "",
15
+ "vocab": "bert",
16
+ "use_plm_init": "no",
17
+ "vocab_size": 0,
18
+ "config_name": "bert-base-cased",
19
+ "notes": "folder-notes",
20
+ "data_dir": "processed_data",
21
+ "dataset": "dataset-name",
22
+ "checkpoint_path": "checkpoint-path/test-run",
23
+ "seq_len": 128,
24
+ "hidden_t_dim": 128,
25
+ "hidden_dim": 256,
26
+ "dropout": 0.1,
27
+ "use_fp16": false,
28
+ "fp16_scale_growth": 0.001,
29
+ "seed": 102,
30
+ "gradient_clipping": -1.0,
31
+ "weight_decay": 0.0,
32
+ "learn_sigma": false,
33
+ "use_kl": false,
34
+ "predict_xstart": true,
35
+ "rescale_timesteps": true,
36
+ "rescale_learned_sigmas": false,
37
+ "sigma_small": false,
38
+ "emb_scale_factor": 1.0,
39
+ "num_transformer_layers": 12,
40
+ "num_transformer_heads": 8,
41
+ "one_noise_step": true,
42
+ "mask_padding": false,
43
+ "celer_only_L1": true,
44
+ "data_split_criterion": "scanpath",
45
+ "corpus": "celer",
46
+ "inference": "none",
47
+ "n_folds": 5,
48
+ "ablation_type": "none",
49
+ "nll_in_loss": false,
50
+ "load_from_checkpoint": false,
51
+ "load_train_data": "-"
52
+ }
scandl_module/original_scandl/sp_gaussian_diffusion.py ADDED
@@ -0,0 +1,1183 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ This code is adapted from Gong et al.'s 2023 DiffuSeq Model: https://github.com/Shark-NLP/DiffuSeq
3
+ """
4
+
5
+ import math
6
+ import numpy as np
7
+ import torch as th
8
+ import sys
9
+ import os
10
+ import torch.nn
11
+
12
+ from .utils.nn import mean_flat
13
+
14
+ sys.path.append(".")
15
+
16
+
17
+ def get_named_beta_schedule(schedule_name, num_diffusion_timesteps):
18
+ """
19
+ Get a pre-defined beta schedule for the given name.
20
+
21
+ The beta schedule library consists of beta schedules which remain similar
22
+ in the limit of num_diffusion_timesteps.
23
+ Beta schedules may be added, but should not be removed or changed once
24
+ they are committed to maintain backwards compatibility.
25
+ """
26
+ if schedule_name == "linear":
27
+ # Linear schedule from Ho et al, extended to work for any number of
28
+ # diffusion steps.
29
+ scale = 1000 / num_diffusion_timesteps
30
+ beta_start = scale * 0.0001
31
+ beta_end = scale * 0.02
32
+ return np.linspace(beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64)
33
+ elif schedule_name == "cosine":
34
+ return betas_for_alpha_bar(
35
+ num_diffusion_timesteps,
36
+ lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2,
37
+ )
38
+ elif schedule_name == "sqrt":
39
+ return betas_for_alpha_bar(
40
+ num_diffusion_timesteps,
41
+ lambda t: 1 - np.sqrt(t + 0.0001),
42
+ )
43
+ elif schedule_name == "trunc_cos":
44
+ return betas_for_alpha_bar_left(
45
+ num_diffusion_timesteps,
46
+ lambda t: np.cos((t + 0.1) / 1.1 * np.pi / 2) ** 2,
47
+ )
48
+ elif schedule_name == "trunc_lin":
49
+ scale = 1000 / num_diffusion_timesteps
50
+ beta_start = scale * 0.0001 + 0.01
51
+ beta_end = scale * 0.02 + 0.01
52
+ return np.linspace(beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64)
53
+ elif schedule_name == "pw_lin":
54
+ scale = 1000 / num_diffusion_timesteps
55
+ beta_start = scale * 0.0001 + 0.01
56
+ beta_mid = scale * 0.0001 # scale * 0.02
57
+ beta_end = scale * 0.02
58
+ first_part = np.linspace(beta_start, beta_mid, 10, dtype=np.float64)
59
+ second_part = np.linspace(
60
+ beta_mid, beta_end, num_diffusion_timesteps - 10, dtype=np.float64
61
+ )
62
+ return np.concatenate([first_part, second_part])
63
+ else:
64
+ raise NotImplementedError(f"unknown beta schedule: {schedule_name}")
65
+
66
+
67
+ def betas_for_alpha_bar_left(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
68
+ """
69
+ Create a beta schedule that discretizes the given alpha_t_bar function, but shifts towards left interval starting from 0
70
+ which defines the cumulative product of (1-beta) over time from t = [0,1].
71
+
72
+ :param num_diffusion_timesteps: the number of betas to produce.
73
+ :param alpha_bar: a lambda that takes an argument t from 0 to 1 and
74
+ produces the cumulative product of (1-beta) up to that
75
+ part of the diffusion process.
76
+ :param max_beta: the maximum beta to use; use values lower than 1 to
77
+ prevent singularities.
78
+ """
79
+ betas = []
80
+ betas.append(min(1 - alpha_bar(0), max_beta))
81
+ for i in range(num_diffusion_timesteps - 1):
82
+ t1 = i / num_diffusion_timesteps
83
+ t2 = (i + 1) / num_diffusion_timesteps
84
+ betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
85
+ return np.array(betas)
86
+
87
+
88
+ def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
89
+ """
90
+ Create a beta schedule that discretizes the given alpha_t_bar function,
91
+ which defines the cumulative product of (1-beta) over time from t = [0,1].
92
+
93
+ :param num_diffusion_timesteps: the number of betas to produce.
94
+ :param alpha_bar: a lambda that takes an argument t from 0 to 1 and
95
+ produces the cumulative product of (1-beta) up to that
96
+ part of the diffusion process.
97
+ :param max_beta: the maximum beta to use; use values lower than 1 to
98
+ prevent singularities.
99
+ """
100
+ betas = []
101
+ for i in range(num_diffusion_timesteps):
102
+ t1 = i / num_diffusion_timesteps
103
+ t2 = (i + 1) / num_diffusion_timesteps
104
+ betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
105
+ return np.array(betas)
106
+
107
+
108
+ class GaussianDiffusion:
109
+ """
110
+ Utilities for training and sampling diffusion models.
111
+ """
112
+
113
+ def __init__(
114
+ self,
115
+ *,
116
+ betas,
117
+ predict_xstart,
118
+ rescale_learned_sigmas,
119
+ learn_sigmas,
120
+ sigma_small,
121
+ use_kl,
122
+ one_noise_step,
123
+ nll_in_loss,
124
+ mask_padding,
125
+ rescale_timesteps=False,
126
+ ):
127
+ self.rescale_timesteps = rescale_timesteps
128
+ self.predict_xstart = predict_xstart
129
+ self.rescale_learned_sigmas = rescale_learned_sigmas
130
+ self.learn_sigmas = learn_sigmas
131
+ self.sigma_small = sigma_small
132
+ self.use_kl = use_kl
133
+ self.one_noise_step = one_noise_step
134
+ self.nll_in_loss = nll_in_loss
135
+ self.mask_padding = mask_padding
136
+
137
+ # Use float64 for accuracy.
138
+ betas = np.array(betas, dtype=np.float64) # shape [diffusion_steps]
139
+ self.betas = betas
140
+ assert len(betas.shape) == 1, "betas must be 1-D"
141
+ assert (betas > 0).all() and (betas <= 1).all()
142
+
143
+ self.num_timesteps = int(betas.shape[0])
144
+
145
+ alphas = 1.0 - betas
146
+
147
+ self.alphas_cumprod = np.cumprod(alphas, axis=0) # will approximate 0
148
+ self.alphas_cumprod_prev = np.append(
149
+ 1.0, self.alphas_cumprod[:-1]
150
+ ) # shifted one to the right
151
+ self.alphas_cumprod_next = np.append(
152
+ self.alphas_cumprod[1:], 0.0
153
+ ) # shifted one to the left
154
+ assert self.alphas_cumprod_prev.shape == (self.num_timesteps,)
155
+
156
+ # calculations for diffusion q(x_t | x_{t-1}) and others
157
+ self.sqrt_alphas_cumprod = np.sqrt(self.alphas_cumprod)
158
+ self.sqrt_one_minus_alphas_cumprod = np.sqrt(1.0 - self.alphas_cumprod)
159
+ self.log_one_minus_alphas_cumprod = np.log(1.0 - self.alphas_cumprod)
160
+ self.sqrt_recip_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod)
161
+ self.sqrt_recipm1_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod - 1)
162
+
163
+ # calculations for posterior q(x_{t-1} | x_t, x_0)
164
+ self.posterior_variance = (
165
+ betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
166
+ )
167
+ # log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain.
168
+ self.posterior_log_variance_clipped = np.log(
169
+ np.append(self.posterior_variance[1], self.posterior_variance[1:])
170
+ )
171
+ self.posterior_mean_coef1 = (
172
+ betas * np.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
173
+ )
174
+ self.posterior_mean_coef2 = (
175
+ (1.0 - self.alphas_cumprod_prev) * np.sqrt(alphas) / (1.0 - self.alphas_cumprod)
176
+ )
177
+
178
+ self.mapping_func = None # implement in train main()
179
+ self.add_mask_noise = False # TODO
180
+
181
+ def training_losses(self, model, *args, **kwargs):
182
+ self.model = model
183
+ return self.training_losses_seq2seq(model, *args, **kwargs)
184
+
185
+ def _predict_xstart_from_eps(self, x_t, t, eps):
186
+ assert x_t.shape == eps.shape
187
+ return (
188
+ _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t
189
+ - _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps
190
+ )
191
+
192
+ def _predict_eps_from_xstart(self, x_t, t, pred_xstart):
193
+ return (
194
+ _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - pred_xstart
195
+ ) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
196
+
197
+ def _scale_timesteps(self, t):
198
+ if self.rescale_timesteps:
199
+ return t.float() * (1000.0 / self.num_timesteps)
200
+ return t
201
+
202
+ def q_mean_variance(self, x_start, t):
203
+ """
204
+ Get the distribution q(x_t | x_0).
205
+
206
+ :param x_start: the [N x C x ...] tensor of noiseless inputs.
207
+ :param t: the number of diffusion steps (minus 1). Here, 0 means one step.
208
+ :return: A tuple (mean, variance, log_variance), all of x_start's shape.
209
+ """
210
+ mean = _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
211
+ variance = _extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape)
212
+ log_variance = _extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape)
213
+ return mean, variance, log_variance
214
+
215
+ def q_sample(self, x_start, t, noise=None, mask=None):
216
+ """
217
+ Diffuse the data for a given number of diffusion steps.
218
+
219
+ In other words, sample from q(x_t | x_0).
220
+
221
+ :param x_start: the initial data batch.
222
+ :param t: the number of diffusion steps (minus 1). Here, 0 means one step.
223
+ :param noise: if specified, the split-out normal noise.
224
+ :param mask: anchoring masked position
225
+ :return: A noisy version of x_start.
226
+ """
227
+ if noise is None:
228
+ noise = th.randn_like(x_start)
229
+
230
+ assert noise.shape == x_start.shape
231
+ x_t = (
232
+ _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start # mu * x_0
233
+ + _extract_into_tensor(
234
+ self.sqrt_one_minus_alphas_cumprod, t, x_start.shape
235
+ ) # sd * noise
236
+ * noise
237
+ )
238
+
239
+ if mask == None:
240
+ return x_t
241
+ else:
242
+ mask = th.broadcast_to(mask.unsqueeze(dim=-1), x_start.shape)
243
+ return th.where(mask == 0, x_start, x_t)
244
+
245
+ def q_posterior_mean_variance(self, x_start, x_t, t):
246
+ """
247
+ Compute the mean and variance of the diffusion posterior:
248
+ q(x_{t-1} | x_t, x_0)
249
+
250
+ """
251
+ assert x_start.shape == x_t.shape
252
+ posterior_mean = (
253
+ _extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start
254
+ + _extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t
255
+ )
256
+ posterior_variance = _extract_into_tensor(self.posterior_variance, t, x_t.shape)
257
+ posterior_log_variance_clipped = _extract_into_tensor(
258
+ self.posterior_log_variance_clipped, t, x_t.shape
259
+ )
260
+
261
+ assert (
262
+ posterior_mean.shape[0]
263
+ == posterior_variance.shape[0]
264
+ == posterior_log_variance_clipped.shape[0]
265
+ == x_start.shape[0]
266
+ )
267
+ return posterior_mean, posterior_variance, posterior_log_variance_clipped
268
+
269
+ def p_mean_variance(
270
+ self,
271
+ model,
272
+ x,
273
+ sn_input_ids_emb,
274
+ pos_enc,
275
+ mask_sn_padding,
276
+ mask_transformer_att,
277
+ t,
278
+ clip_denoised=True,
279
+ denoised_fn=None,
280
+ model_kwargs=None,
281
+ subwords_list=None,
282
+ atten_vis=None,
283
+ atten_vis_fn=None,
284
+ atten_vis_path=None,
285
+ batch_idx=None,
286
+ rank=None,
287
+ atten_vis_sp=None,
288
+ ):
289
+ """
290
+ Apply the model to get p(x_{t-1} | x_t), as well as a prediction of
291
+ the initial x, x_0.
292
+
293
+ :param model: the model, which takes a signal and a batch of timesteps
294
+ as input.
295
+ :param x: the [N x C x ...] tensor at time t; the noised input where the embedded condition/sn was not noised
296
+ and the target/sp was replaced with Gaussian noise; at time step t
297
+ :param sn_input_ids_emb: the BERT embeddings
298
+ :param pos_enc: the positional embeddings
299
+ :param mask_sn_padding:
300
+ :param mask_transformer_att: the attention mask for the transformer
301
+ :param t: a 1-D Tensor of timesteps.
302
+ :param clip_denoised: if True, clip the denoised signal into [-1, 1].
303
+ :param denoised_fn: if not None, a function which applies to the x_start prediction before it is used to sample.
304
+ Applies before clip_denoised.
305
+ :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning.
306
+ :param subwords_list: list of list containing the subwortokens of each instance (for attention visualization)
307
+ :param atten_vis: bool: if True, attention is visualized (heatmaps)
308
+ :param atten_vis_fn: the attention visualization function
309
+ :param atten_vis_path: the path where to save the heatmaps to
310
+ :param batch_idx: index of the current batch
311
+ :paran rank: if parallel processing, the GPU index
312
+ :param atten_vis_sp: bool: save attention scores of the last timestep
313
+ :return: a dict with the following keys:
314
+ - 'mean': the model mean output.
315
+ - 'variance': the model variance output.
316
+ - 'log_variance': the log of 'variance'.
317
+ - 'pred_xstart': the prediction for x_0.
318
+ """
319
+ if model_kwargs is None:
320
+ model_kwargs = {}
321
+
322
+ B, C = x.size(0), x.size(-1)
323
+ assert t.shape == (B,)
324
+
325
+ if not atten_vis and not atten_vis_sp:
326
+ model_output = model(
327
+ x=x,
328
+ ts=self._scale_timesteps(t),
329
+ sn_input_ids_emb=sn_input_ids_emb,
330
+ pos_enc=pos_enc,
331
+ attention_mask=mask_transformer_att,
332
+ **model_kwargs,
333
+ )
334
+ else:
335
+ # visualise the attention (heatmaps)
336
+ if atten_vis:
337
+ model_output, attention_scores = model(
338
+ x=x,
339
+ ts=self._scale_timesteps(t),
340
+ sn_input_ids_emb=sn_input_ids_emb,
341
+ pos_enc=pos_enc,
342
+ attention_mask=mask_transformer_att,
343
+ atten_vis=atten_vis,
344
+ **model_kwargs,
345
+ )
346
+ # visualise attention for last denoising step
347
+ if t[0] % 200 == 0:
348
+
349
+ atten_vis_fn(
350
+ attention_scores=attention_scores,
351
+ subwords_list=subwords_list,
352
+ batch_idx=batch_idx,
353
+ path_to_dir=atten_vis_path,
354
+ denoising_step=t[0].item(),
355
+ aggregate=True,
356
+ rank=rank,
357
+ )
358
+ if atten_vis_sp:
359
+
360
+ model_output, attention_scores = model(
361
+ x=x,
362
+ ts=self._scale_timesteps(t),
363
+ sn_input_ids_emb=sn_input_ids_emb,
364
+ pos_enc=pos_enc,
365
+ attention_mask=mask_transformer_att,
366
+ atten_vis=True,
367
+ **model_kwargs,
368
+ )
369
+ if t[0] == 0:
370
+
371
+ out_path_heatmaps_sp = os.path.join(atten_vis_path, "heatmaps_sps")
372
+ if not os.path.exists(out_path_heatmaps_sp):
373
+ os.makedirs(out_path_heatmaps_sp)
374
+ filename = f"att_scores_rank{rank}_batch{batch_idx}.pt"
375
+ path_to_file = os.path.join(out_path_heatmaps_sp, filename)
376
+ torch.save(attention_scores, path_to_file)
377
+
378
+ model_variance = np.append(self.posterior_variance[1], self.betas[1:])
379
+ model_log_variance = np.log(np.append(self.posterior_variance[1], self.betas[1:]))
380
+
381
+ model_variance = _extract_into_tensor(model_variance, t, x.shape)
382
+ model_log_variance = _extract_into_tensor(model_log_variance, t, x.shape)
383
+
384
+ # The denoised_fn is applied to x_start (the model output) before it is used for sampling
385
+ def process_xstart(x):
386
+ """here x is the model output"""
387
+ if denoised_fn is not None:
388
+ # print(denoised_fn)
389
+ x = denoised_fn(x, t)
390
+ if clip_denoised:
391
+ return x.clamp(-1, 1)
392
+ return x
393
+
394
+ if self.predict_xstart:
395
+ # the denoised fn is applied to the model output
396
+ pred_xstart = process_xstart(model_output)
397
+ else:
398
+ ### model is used to predict eps
399
+ pred_xstart = process_xstart(
400
+ self._predict_xstart_from_eps(x_t=x, t=t, eps=model_output)
401
+ )
402
+
403
+ # this is the mean of the posterior distribution q(x_{t-1} | x_t, x_0), estimated from x_t, which is the noised
404
+ # input, and pred_xstart, which is what the model predicted to be x_0 from the noised input x_noised/x_t
405
+ model_mean, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t)
406
+
407
+ assert model_mean.shape == model_log_variance.shape == pred_xstart.shape == x.shape
408
+ return {
409
+ "mean": model_mean,
410
+ "variance": model_variance,
411
+ "log_variance": model_log_variance,
412
+ "pred_xstart": pred_xstart,
413
+ }
414
+
415
+ def p_sample(
416
+ self,
417
+ model,
418
+ x,
419
+ sn_input_ids_emb,
420
+ pos_enc,
421
+ mask_sn_padding,
422
+ mask_transformer_att,
423
+ t,
424
+ clip_denoised=True,
425
+ denoised_fn=None,
426
+ model_kwargs=None,
427
+ top_p=None,
428
+ mask=None,
429
+ x_start=None,
430
+ subwords_list=None,
431
+ atten_vis=None,
432
+ atten_vis_fn=None,
433
+ atten_vis_path=None,
434
+ batch_idx=None,
435
+ rank=None,
436
+ atten_vis_sp=None,
437
+ ):
438
+ """
439
+ Sample x_{t-1} from the model at the given timestep.
440
+
441
+ :param model: the model to sample from; the transformer model that learned the denoising
442
+ :param x: the current tensor at x_{t-1}.
443
+ :param sn_input_ids_emb: the BERT embeddings
444
+ :param pos_enc: the positional embeddings
445
+ :param mask_sn_padding:
446
+ :param mask_transformer_att: the attention mask for the transformer
447
+ :param t: the value of t, starting at 0 for the first diffusion step.
448
+ :param clip_denoised: if True, clip the x_start prediction to [-1, 1].
449
+ :param denoised_fn: if not None, a function which applies to the x_start prediction before it is used to sample.
450
+ :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning.
451
+ :param top_p:
452
+ :param mask: anchoring masked position to x_start
453
+ :param x_start:
454
+ :param subwords_list: list of list containing the subwortokens of each instance (for attention visualization)
455
+ :param atten_vis: bool: if True, attention is visualized (heatmaps)
456
+ :param atten_vis_fn: the attention visualization function
457
+ :param atten_vis_path: the path where to save the heatmaps to
458
+ :param batch_idx: index of the current batch
459
+ :paran rank: if parallel processing, the GPU index
460
+ :param atten_vis_sp: bool: save attention scores of the last timestep
461
+ :return: a dict containing the following keys:
462
+ - 'sample': a random sample from the model.
463
+ - 'pred_xstart': a prediction of x_0.
464
+ """
465
+ out = self.p_mean_variance(
466
+ model=model,
467
+ x=x,
468
+ sn_input_ids_emb=sn_input_ids_emb,
469
+ pos_enc=pos_enc,
470
+ mask_sn_padding=mask_sn_padding,
471
+ mask_transformer_att=mask_transformer_att,
472
+ t=t,
473
+ clip_denoised=clip_denoised,
474
+ denoised_fn=denoised_fn,
475
+ model_kwargs=model_kwargs,
476
+ subwords_list=subwords_list,
477
+ atten_vis=atten_vis,
478
+ atten_vis_fn=atten_vis_fn,
479
+ atten_vis_path=atten_vis_path,
480
+ batch_idx=batch_idx,
481
+ rank=rank,
482
+ atten_vis_sp=atten_vis_sp,
483
+ )
484
+
485
+ if top_p is not None and top_p > 0:
486
+ # print('top_p sampling')
487
+ noise = th.randn_like(x)
488
+ replace_mask = th.abs(noise) > top_p
489
+ while replace_mask.any():
490
+ noise[replace_mask] = th.randn_like(noise[replace_mask])
491
+ replace_mask = th.abs(noise) > top_p
492
+ assert (th.abs(noise) <= top_p).all()
493
+
494
+ else:
495
+ noise = th.randn_like(x)
496
+
497
+ nonzero_mask = (
498
+ (t != 0).float().view(-1, *([1] * (len(x.shape) - 1)))
499
+ ) # no noise when t == 0
500
+
501
+ sample = out["mean"] + nonzero_mask * th.exp(0.5 * out["log_variance"]) * noise
502
+
503
+ if mask == None:
504
+ pass
505
+ else:
506
+ # the original embedding for the sn, and the predicted sample for the sp
507
+ sample = th.where(mask == 0, x_start, sample)
508
+
509
+ return {
510
+ "sample": sample,
511
+ "pred_xstart": out["pred_xstart"],
512
+ "greedy_mean": out["mean"],
513
+ "out": out,
514
+ }
515
+
516
+ def p_sample_loop(
517
+ self,
518
+ model,
519
+ shape,
520
+ noise=None,
521
+ sn_input_ids_emb=None,
522
+ pos_enc=None,
523
+ mask_sn_padding=None,
524
+ mask_transformer_att=None,
525
+ clip_denoised=True,
526
+ denoised_fn=None,
527
+ model_kwargs=None,
528
+ device=None,
529
+ progress=False,
530
+ top_p=None,
531
+ clamp_step=None,
532
+ clamp_first=None,
533
+ mask=None,
534
+ x_start=None,
535
+ subwords_list=None,
536
+ atten_vis=None,
537
+ atten_vis_fn=None,
538
+ atten_vis_path=None,
539
+ batch_idx=None,
540
+ gap=1,
541
+ rank=None,
542
+ atten_vis_sp=None,
543
+ ):
544
+ """
545
+ Generate samples from the model.
546
+
547
+ :param model: the transformer model that was trained to learn the denoising
548
+ :param shape: the shape of the samples, (N, C, H, W).
549
+ :param noise: the Gaussian noise that should be denoised at inference (the replaced word ID emb)
550
+ :param sn_input_ids_emb: the BERT embeddings
551
+ :param pos_enc: the positional embeddings
552
+ :param mask_sn_padding:
553
+ :param mask_transformer_att: the attention mask for the transformer
554
+ :param clip_denoised: if True, clip x_start predictions to [-1, 1].
555
+ :param denoised_fn: if not None, a function which applies to the x_start prediction before it is used to sample.
556
+ :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning.
557
+ :param device: if specified, the device to create the samples on. If not specified, use a model parameter's device.
558
+ :param progress: if True, show a tqdm progress bar.
559
+ :param top_p:
560
+ :param clamp_step: in clamp_first mode, choose end clamp step, otherwise starting clamp step
561
+ :param clamp_first: bool, clamp_first mode
562
+ :param mask: anchoring masked position to x_start
563
+ :param x_start: the word ID embedding before replaced by noise
564
+ :param subwords_list: list of list containing the subwortokens of each instance (for attention visualization)
565
+ :param atten_vis: bool: if True, attention is visualized (heatmaps)
566
+ :param atten_vis_fn: the attention visualization function
567
+ :param atten_vis_path: the path where to save the heatmaps to
568
+ :param batch_idx: index of the current batch
569
+ :param gap:
570
+ :paran rank: if parallel processing, the GPU index
571
+ :param atten_vis_sp: bool: save attention scores of the last timestep
572
+ :return: a non-differentiable batch of samples.
573
+ """
574
+ final = []
575
+ for sample in self.p_sample_loop_progressive(
576
+ model,
577
+ shape,
578
+ noise=noise,
579
+ sn_input_ids_emb=sn_input_ids_emb,
580
+ pos_enc=pos_enc,
581
+ mask_sn_padding=mask_sn_padding,
582
+ mask_transformer_att=mask_transformer_att,
583
+ clip_denoised=clip_denoised,
584
+ denoised_fn=denoised_fn,
585
+ model_kwargs=model_kwargs,
586
+ device=device,
587
+ progress=progress,
588
+ top_p=top_p,
589
+ clamp_step=clamp_step,
590
+ clamp_first=clamp_first,
591
+ mask=mask,
592
+ x_start=x_start,
593
+ subwords_list=subwords_list,
594
+ atten_vis=atten_vis,
595
+ atten_vis_fn=atten_vis_fn,
596
+ atten_vis_path=atten_vis_path,
597
+ batch_idx=batch_idx,
598
+ rank=rank,
599
+ atten_vis_sp=atten_vis_sp,
600
+ ):
601
+ final.append(sample["sample"])
602
+ return final
603
+
604
+ def p_sample_loop_progressive(
605
+ self,
606
+ model,
607
+ shape,
608
+ noise=None,
609
+ sn_input_ids_emb=None,
610
+ pos_enc=None,
611
+ mask_sn_padding=None,
612
+ mask_transformer_att=None,
613
+ clip_denoised=True,
614
+ denoised_fn=None,
615
+ model_kwargs=None,
616
+ device=None,
617
+ progress=False,
618
+ top_p=None,
619
+ clamp_step=None,
620
+ clamp_first=None,
621
+ mask=None,
622
+ x_start=None,
623
+ subwords_list=None,
624
+ atten_vis=None,
625
+ atten_vis_fn=None,
626
+ atten_vis_path=None,
627
+ batch_idx=None,
628
+ rank=None,
629
+ atten_vis_sp=None,
630
+ ):
631
+ """
632
+ Generate samples from the model and yield intermediate samples from
633
+ each timestep of diffusion.
634
+
635
+ Arguments are the same as p_sample_loop().
636
+ Returns a generator over dicts, where each dict is the return value of
637
+ p_sample().
638
+ """
639
+ if device is None:
640
+ device = next(model.parameters()).device
641
+ assert isinstance(shape, (tuple, list))
642
+
643
+ # noise/sample_x is the input that was noised: the concatenated sn-sp embedding where the sp was completely
644
+ # replaced with Gaussian noise from the standard normal distribution
645
+ if noise is not None:
646
+ sample_x = noise
647
+ else:
648
+ sample_x = th.randn(*shape, device=device)
649
+
650
+ # the number of diffusion steps in reverse order
651
+ indices = list(range(self.num_timesteps))[::-1]
652
+
653
+ if progress:
654
+ # Lazy import so that we don't depend on tqdm.
655
+ from tqdm.auto import tqdm
656
+
657
+ indices = tqdm(indices)
658
+
659
+ # denoising from the number of diffusion steps T to t=0
660
+ for i in indices: # from T to 0
661
+
662
+ t = th.tensor([i] * shape[0], device=device)
663
+ if not clamp_first:
664
+ if i > clamp_step:
665
+ denoised_fn_cur = None
666
+ else:
667
+ denoised_fn_cur = denoised_fn
668
+ else:
669
+ if i >= clamp_step:
670
+ denoised_fn_cur = denoised_fn
671
+ else:
672
+ denoised_fn_cur = None
673
+
674
+ with th.no_grad():
675
+ out = self.p_sample(
676
+ model=model,
677
+ x=sample_x,
678
+ sn_input_ids_emb=sn_input_ids_emb,
679
+ pos_enc=pos_enc,
680
+ mask_sn_padding=mask_sn_padding,
681
+ mask_transformer_att=mask_transformer_att,
682
+ t=t,
683
+ clip_denoised=clip_denoised,
684
+ denoised_fn=denoised_fn_cur,
685
+ model_kwargs=model_kwargs,
686
+ top_p=top_p,
687
+ mask=mask,
688
+ subwords_list=subwords_list,
689
+ x_start=x_start,
690
+ atten_vis=atten_vis,
691
+ atten_vis_fn=atten_vis_fn,
692
+ atten_vis_path=atten_vis_path,
693
+ batch_idx=batch_idx,
694
+ rank=rank,
695
+ atten_vis_sp=atten_vis_sp,
696
+ )
697
+ yield out
698
+ sample_x = out["sample"]
699
+
700
+ def _get_x_start(self, x_start_mean, std):
701
+ """
702
+ Word embedding projection from {Emb(w)} to {x_0}
703
+ :param x_start_mean: word embedding
704
+ :return: x_0
705
+ """
706
+ noise = th.randn_like(x_start_mean)
707
+ assert noise.shape == x_start_mean.shape
708
+ # print(x_start_mean.device, noise.device)
709
+ return x_start_mean + std * noise
710
+
711
+ def _token_discrete_loss(self, x_t, get_logits, input_ids, mask=None, truncate=False, t=None):
712
+ """
713
+ the loss of -log p(w|z_0)
714
+ :param x_start_mean: word embedding
715
+ :return: x_0
716
+ """
717
+ reshaped_x_t = x_t
718
+ logits = get_logits(reshaped_x_t) # shape [microbatch size, seq_len, vocabulary]
719
+ # print(logits.shape)
720
+ loss_fct = th.nn.CrossEntropyLoss(reduction="none")
721
+ decoder_nll = loss_fct(logits.view(-1, logits.size(-1)), input_ids.view(-1)).view(
722
+ input_ids.shape
723
+ )
724
+ if mask != None:
725
+ decoder_nll *= mask
726
+ # print(decoder_nll.shape)
727
+ if mask != None:
728
+ decoder_nll = decoder_nll.sum(dim=-1) / mask.sum(dim=-1)
729
+ else:
730
+ decoder_nll = decoder_nll.mean(dim=-1)
731
+
732
+ return decoder_nll
733
+
734
+ def _x0_helper(self, model_output, x, t):
735
+
736
+ if self.predict_xstart:
737
+ pred_xstart = model_output
738
+ pred_prev, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t)
739
+
740
+ else: # predict eps
741
+ pred_xstart = self._predict_xstart_from_eps(x_t=x, t=t, eps=model_output)
742
+
743
+ pred_prev, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t)
744
+
745
+ return {"pred_xprev": pred_prev, "pred_xstart": pred_xstart}
746
+
747
+ def training_losses_seq2seq(
748
+ self,
749
+ model, # the transformer model
750
+ t, # the number of noise adding steps for each instance in the microbatch
751
+ sn_sp_repr,
752
+ mask,
753
+ sn_input_ids,
754
+ indices_pos_enc,
755
+ mask_sn_padding,
756
+ mask_transformer_att,
757
+ noise=None,
758
+ ):
759
+ """
760
+ Compute training losses for a single timestep.
761
+
762
+ :param model: the transformer model
763
+ :param t: a batch of timestep indices.
764
+ :param sn_sp_repr: the word IDs of sn and sp
765
+ :param mask: masking the sn
766
+ :param sn_input_ids: the tokenizer input IDs
767
+ :param indices_pos_enc: the indices for pos enc
768
+ :param_mask_sn_padding:
769
+ :param mask_transformer_att: the transformer att
770
+ :param model_kwargs: if not None, a dict of extra keyword arguments to pass to the model. This can be used for conditioning.
771
+ :param noise: if specified, the specific Gaussian noise to try to remove.
772
+ :return: a dict with the key "loss" containing a tensor of shape [N].
773
+ Some mean or variance settings may also have other keys.
774
+ """
775
+
776
+ microbatch_size, seq_len = sn_sp_repr.shape
777
+
778
+ # get the word ID embedding, BERT embedding, positional embedding
779
+ sn_sp_emb, pos_enc, sn_input_ids_emb = model.model.module.get_embeds(
780
+ sn_sp_repr=sn_sp_repr,
781
+ sn_input_ids=sn_input_ids,
782
+ indices_pos_enc=indices_pos_enc,
783
+ )
784
+
785
+ # get the standard deviation, shape [microbatch, args.seq_len, hidden_size=768]
786
+ std = _extract_into_tensor(
787
+ self.sqrt_one_minus_alphas_cumprod, th.tensor([0]).to(sn_sp_emb.device), sn_sp_emb.shape
788
+ )
789
+
790
+ # map sn_sp_emb to x_start, which is a one-step noised sn_sp_emb (in paper it's z_0)
791
+ if (
792
+ self.one_noise_step
793
+ ): # this should always be true actually (without it performance is bad)
794
+ x_start = self._get_x_start(sn_sp_emb, std)
795
+ else:
796
+ x_start = sn_sp_emb
797
+
798
+ # sample noise in the same shape as our input
799
+ if noise is None:
800
+ noise = th.randn_like(x_start)
801
+
802
+ # get the noised sample x_t, which is still of shape [microbatch, args.seq_len, hidden_size=768]
803
+ # the condition/sn is not noised (hence the input mask)
804
+ # each instance in the microbatch receives a different amount of noise (t noising steps, as given in vector t)
805
+ x_t = self.q_sample(
806
+ x_start=x_start,
807
+ t=t,
808
+ noise=noise,
809
+ mask=mask,
810
+ )
811
+
812
+ terms = {}
813
+
814
+ target = x_start
815
+
816
+ # model_output is of shape [microbatch, args.seq_len, emb_dim=768]
817
+ model_output = model(
818
+ x=x_t,
819
+ ts=self._scale_timesteps(t),
820
+ sn_input_ids_emb=sn_input_ids_emb,
821
+ pos_enc=pos_enc,
822
+ attention_mask=mask_transformer_att,
823
+ )
824
+ assert model_output.shape == target.shape == x_start.shape
825
+
826
+ # Loss 1: Mean Squared Error (MSE) (L_{VLB})
827
+ terms["mse"] = mean_flat((target - model_output) ** 2)
828
+ model_out_x_start = self._x0_helper(model_output, x_t, t)["pred_xstart"]
829
+ t0_mask = (
830
+ t == 0
831
+ ) # mask that says true for every instance where no noise was received, i.e. t=0
832
+ # MSE between the model output and the embedded input before the one noise step
833
+ t0_loss = mean_flat((sn_sp_emb - model_out_x_start) ** 2)
834
+ # update the MSE between the model output and the one-step noised input embeddings with the MSE between the
835
+ # model output and the embeddings before the one noise step wherever there was no noise received in the noising
836
+ # process (i.e., wherever t was 0)
837
+ terms["mse"] = th.where(t0_mask, t0_loss, terms["mse"])
838
+
839
+ # Loss 2: L_{round}
840
+ out_mean, _, _ = self.q_mean_variance(
841
+ x_start, th.LongTensor([self.num_timesteps - 1]).to(x_start.device)
842
+ )
843
+ tT_loss = mean_flat(out_mean**2)
844
+
845
+ # for the NLL losses, we need to convert the model output into logits
846
+ get_logits = model.model.module.get_logits
847
+
848
+ # Loss 3: L_{EMB}
849
+ # compute the NLL between the one-noised embeddings and the initial representation (word IDs)
850
+ # embedding regularisation
851
+ decoder_nll = self._token_discrete_loss(x_start, get_logits, sn_sp_repr)
852
+
853
+ # unused Loss
854
+ terms["nll"] = self._token_discrete_loss(model_output, get_logits, sn_sp_repr, mask=mask)
855
+
856
+ # combined loss
857
+ if self.nll_in_loss: # should be False; model performance drops if nll included
858
+ terms["loss"] = terms["mse"] + tT_loss + decoder_nll + terms["nll"]
859
+ else:
860
+ terms["loss"] = terms["mse"] + tT_loss + decoder_nll
861
+
862
+ return terms
863
+
864
+ def ddim_sample(
865
+ self,
866
+ model,
867
+ x,
868
+ t,
869
+ clip_denoised=True,
870
+ denoised_fn=None,
871
+ model_kwargs=None,
872
+ eta=0.0,
873
+ langevin_fn=None,
874
+ mask=None,
875
+ x_start=None,
876
+ ):
877
+ """
878
+ Sample x_{t-1} from the model using DDIM.
879
+
880
+ Same usage as p_sample().
881
+ """
882
+ out = self.p_mean_variance(
883
+ model,
884
+ x,
885
+ t,
886
+ clip_denoised=clip_denoised,
887
+ denoised_fn=denoised_fn,
888
+ model_kwargs=model_kwargs,
889
+ )
890
+ # Usually our model outputs epsilon, but we re-derive it
891
+ # in case we used x_start or x_prev prediction.
892
+ eps = self._predict_eps_from_xstart(x, t, out["pred_xstart"])
893
+ alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape)
894
+ alpha_bar_prev = _extract_into_tensor(self.alphas_cumprod_prev, t, x.shape)
895
+ sigma = (
896
+ eta
897
+ * th.sqrt((1 - alpha_bar_prev) / (1 - alpha_bar))
898
+ * th.sqrt(1 - alpha_bar / alpha_bar_prev)
899
+ )
900
+ # Equation 12.
901
+ noise = th.randn_like(x)
902
+ mean_pred = (
903
+ out["pred_xstart"] * th.sqrt(alpha_bar_prev)
904
+ + th.sqrt(1 - alpha_bar_prev - sigma**2) * eps
905
+ )
906
+ nonzero_mask = (
907
+ (t != 0).float().view(-1, *([1] * (len(x.shape) - 1)))
908
+ ) # no noise when t == 0
909
+ # print(sigma.mean())
910
+ sample = mean_pred + nonzero_mask * sigma * noise
911
+ if langevin_fn:
912
+ print(t.shape)
913
+ sample = langevin_fn(sample, mean_pred, sigma, self.alphas_cumprod_prev[t[0]], t, x)
914
+
915
+ if mask == None:
916
+ pass
917
+ else:
918
+ sample = th.where(mask == 0, x_start, sample)
919
+
920
+ return {"sample": sample, "pred_xstart": out["pred_xstart"]}
921
+
922
+ def ddim_reverse_sample(
923
+ self,
924
+ model,
925
+ x,
926
+ t,
927
+ clip_denoised=True,
928
+ denoised_fn=None,
929
+ model_kwargs=None,
930
+ eta=0.0,
931
+ ):
932
+ """
933
+ Sample x_{t+1} from the model using DDIM reverse ODE.
934
+ """
935
+ assert eta == 0.0, "Reverse ODE only for deterministic path"
936
+ out = self.p_mean_variance(
937
+ model,
938
+ x,
939
+ t,
940
+ clip_denoised=clip_denoised,
941
+ denoised_fn=denoised_fn,
942
+ model_kwargs=model_kwargs,
943
+ )
944
+ # Usually our model outputs epsilon, but we re-derive it
945
+ # in case we used x_start or x_prev prediction.
946
+ eps = (
947
+ _extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x.shape) * x
948
+ - out["pred_xstart"]
949
+ ) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x.shape)
950
+ alpha_bar_next = _extract_into_tensor(self.alphas_cumprod_next, t, x.shape)
951
+
952
+ # Equation 12. reversed
953
+ mean_pred = out["pred_xstart"] * th.sqrt(alpha_bar_next) + th.sqrt(1 - alpha_bar_next) * eps
954
+
955
+ return {"sample": mean_pred, "pred_xstart": out["pred_xstart"]}
956
+
957
+ def ddim_sample_loop(
958
+ self,
959
+ model,
960
+ shape,
961
+ noise=None,
962
+ clip_denoised=True,
963
+ denoised_fn=None,
964
+ model_kwargs=None,
965
+ device=None,
966
+ progress=False,
967
+ top_p=None,
968
+ clamp_step=None,
969
+ clamp_first=None,
970
+ mask=None,
971
+ x_start=None,
972
+ gap=1,
973
+ ):
974
+ """
975
+ Generate samples from the model using DDIM.
976
+ :param gap: compute ddim sampling for each {gap} step
977
+
978
+ Same usage as p_sample_loop().
979
+ """
980
+ final = []
981
+ for sample in self.ddim_sample_loop_progressive(
982
+ model,
983
+ shape,
984
+ noise=noise,
985
+ clip_denoised=clip_denoised,
986
+ denoised_fn=denoised_fn,
987
+ model_kwargs=model_kwargs,
988
+ device=device,
989
+ progress=progress,
990
+ mask=mask,
991
+ x_start=x_start,
992
+ gap=gap,
993
+ ):
994
+ final.append(sample["sample"])
995
+ return final
996
+
997
+ def ddim_sample_loop_progressive(
998
+ self,
999
+ model,
1000
+ shape,
1001
+ noise=None,
1002
+ clip_denoised=True,
1003
+ denoised_fn=None,
1004
+ model_kwargs=None,
1005
+ device=None,
1006
+ progress=False,
1007
+ eta=0.0,
1008
+ langevin_fn=None,
1009
+ mask=None,
1010
+ x_start=None,
1011
+ gap=1,
1012
+ ):
1013
+ """
1014
+ Use DDIM to sample from the model and yield intermediate samples from
1015
+ each timestep of DDIM.
1016
+
1017
+ Same usage as p_sample_loop_progressive().
1018
+ """
1019
+ if device is None:
1020
+ device = next(model.parameters()).device
1021
+ assert isinstance(shape, (tuple, list))
1022
+ if noise is not None:
1023
+ sample_x = noise
1024
+ else:
1025
+ sample_x = th.randn(*shape, device=device)
1026
+ indices = list(range(self.num_timesteps))[::-1][::gap]
1027
+
1028
+ if progress:
1029
+ # Lazy import so that we don't depend on tqdm.
1030
+ from tqdm.auto import tqdm
1031
+
1032
+ indices = tqdm(indices)
1033
+
1034
+ for i in indices:
1035
+ t = th.tensor([i] * shape[0], device=device)
1036
+ with th.no_grad():
1037
+ out = self.ddim_sample(
1038
+ model,
1039
+ sample_x,
1040
+ t,
1041
+ clip_denoised=clip_denoised,
1042
+ denoised_fn=denoised_fn,
1043
+ model_kwargs=model_kwargs,
1044
+ mask=mask,
1045
+ x_start=x_start,
1046
+ )
1047
+ yield out
1048
+ sample_x = out["sample"]
1049
+
1050
+
1051
+ def _extract_into_tensor(arr, timesteps, broadcast_shape):
1052
+ """
1053
+ Extract values from a 1-D numpy array for a batch of indices.
1054
+
1055
+ :param arr: the 1-D numpy array.
1056
+ :param timesteps: a tensor of indices into the array to extract.
1057
+ :param broadcast_shape: a larger shape of K dimensions with the batch
1058
+ dimension equal to the length of timesteps.
1059
+ :return: a tensor of shape [batch_size, 1, ...] where the shape has K dims.
1060
+ """
1061
+ res = th.from_numpy(arr).to(device=timesteps.device)[timesteps].float()
1062
+ while len(res.shape) < len(broadcast_shape):
1063
+ res = res[..., None]
1064
+ return res.expand(broadcast_shape)
1065
+
1066
+
1067
+ def space_timesteps(num_timesteps, section_counts):
1068
+ """
1069
+ Create a list of timesteps to use from an original diffusion process,
1070
+ given the number of timesteps we want to take from equally-sized portions
1071
+ of the original process.
1072
+
1073
+ For example, if there's 300 timesteps and the section counts are [10,15,20]
1074
+ then the first 100 timesteps are strided to be 10 timesteps, the second 100
1075
+ are strided to be 15 timesteps, and the final 100 are strided to be 20.
1076
+
1077
+ If the stride is a string starting with "ddim", then the fixed striding
1078
+ from the DDIM paper is used, and only one section is allowed.
1079
+
1080
+ :param num_timesteps: the number of diffusion steps in the original
1081
+ process to divide up.
1082
+ :param section_counts: either a list of numbers, or a string containing
1083
+ comma-separated numbers, indicating the step count
1084
+ per section. As a special case, use "ddimN" where N
1085
+ is a number of steps to use the striding from the
1086
+ DDIM paper.
1087
+ :return: a set of diffusion steps from the original process to use.
1088
+ """
1089
+ if isinstance(section_counts, str):
1090
+ if section_counts.startswith("ddim"):
1091
+ desired_count = int(section_counts[len("ddim") :])
1092
+ for i in range(1, num_timesteps):
1093
+ if len(range(0, num_timesteps, i)) == desired_count:
1094
+ return set(range(0, num_timesteps, i))
1095
+ raise ValueError(f"cannot create exactly {num_timesteps} steps with an integer stride")
1096
+ section_counts = [int(x) for x in section_counts.split(",")]
1097
+ size_per = num_timesteps // len(section_counts)
1098
+ extra = num_timesteps % len(section_counts)
1099
+ start_idx = 0
1100
+ all_steps = []
1101
+ for i, section_count in enumerate(section_counts):
1102
+ size = size_per + (1 if i < extra else 0)
1103
+ if size < section_count:
1104
+ raise ValueError(f"cannot divide section of {size} steps into {section_count}")
1105
+ if section_count <= 1:
1106
+ frac_stride = 1
1107
+ else:
1108
+ frac_stride = (size - 1) / (section_count - 1)
1109
+ cur_idx = 0.0
1110
+ taken_steps = []
1111
+ for _ in range(section_count):
1112
+ taken_steps.append(start_idx + round(cur_idx))
1113
+ cur_idx += frac_stride
1114
+ all_steps += taken_steps
1115
+ start_idx += size
1116
+ return set(all_steps)
1117
+
1118
+
1119
+ class SpacedDiffusion(GaussianDiffusion):
1120
+ """
1121
+ A diffusion process which can skip steps in a base diffusion process.
1122
+
1123
+ :param use_timesteps: a collection (sequence or set) of timesteps from the
1124
+ original diffusion process to retain.
1125
+ :param kwargs: the kwargs to create the base diffusion process.
1126
+ """
1127
+
1128
+ def __init__(self, use_timesteps, **kwargs):
1129
+ self.use_timesteps = set(use_timesteps)
1130
+ self.timestep_map = []
1131
+ self.original_num_steps = len(kwargs["betas"])
1132
+
1133
+ # print(kwargs.keys())
1134
+ base_diffusion = GaussianDiffusion(**kwargs) # pylint: disable=missing-kwoa
1135
+ last_alpha_cumprod = 1.0
1136
+ new_betas = []
1137
+ for i, alpha_cumprod in enumerate(base_diffusion.alphas_cumprod):
1138
+ if i in self.use_timesteps:
1139
+ new_betas.append(1 - alpha_cumprod / last_alpha_cumprod)
1140
+ last_alpha_cumprod = alpha_cumprod
1141
+ self.timestep_map.append(i)
1142
+ kwargs["betas"] = np.array(new_betas)
1143
+ super().__init__(**kwargs)
1144
+
1145
+ def p_mean_variance(self, model, *args, **kwargs): # pylint: disable=signature-differs
1146
+ # print('called p_mean_var')
1147
+ return super().p_mean_variance(self._wrap_model(model), *args, **kwargs)
1148
+
1149
+ def training_losses(self, model, *args, **kwargs): # pylint: disable=signature-differs
1150
+ # print('called training_losses')
1151
+ return super().training_losses(self._wrap_model(model), *args, **kwargs)
1152
+
1153
+ def _wrap_model(self, model):
1154
+ if isinstance(model, _WrappedModel):
1155
+ return model
1156
+ return _WrappedModel(
1157
+ model, self.timestep_map, self.rescale_timesteps, self.original_num_steps
1158
+ )
1159
+
1160
+ def _scale_timesteps(self, t):
1161
+ # Scaling is done by the wrapped model.
1162
+ return t
1163
+
1164
+
1165
+ class _WrappedModel:
1166
+ def __init__(self, model, timestep_map, rescale_timesteps, original_num_steps):
1167
+ self.model = model
1168
+ self.timestep_map = timestep_map
1169
+ self.rescale_timesteps = rescale_timesteps
1170
+ self.original_num_steps = original_num_steps
1171
+
1172
+ def __call__(self, x, ts, **kwargs):
1173
+ # print(ts)
1174
+ map_tensor = th.tensor(self.timestep_map, device=ts.device, dtype=ts.dtype)
1175
+ new_ts = map_tensor[ts]
1176
+ # print(new_ts)
1177
+ if self.rescale_timesteps:
1178
+ new_ts = new_ts.float() * (1000.0 / self.original_num_steps)
1179
+ # temp = self.model(x, new_ts, **kwargs)
1180
+ # print(temp.shape)
1181
+ # return temp
1182
+ # print(new_ts)
1183
+ return self.model(x, new_ts, **kwargs)
scandl_module/original_scandl/sp_rounding.py ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+
4
+
5
+ def get_knn(model_emb, text_emb, dist="cos"):
6
+ if dist == "cos":
7
+ adjacency = model_emb @ text_emb.transpose(1, 0).to(model_emb.device)
8
+ elif dist == "l2":
9
+ adjacency = model_emb.unsqueeze(1).expand(-1, text_emb.size(0), -1) - text_emb.unsqueeze(
10
+ 0
11
+ ).expand(model_emb.size(0), -1, -1)
12
+ adjacency = -torch.norm(adjacency, dim=-1)
13
+ topk_out = torch.topk(adjacency, k=6, dim=0)
14
+ return topk_out.values, topk_out.indices
15
+
16
+
17
+ def get_efficient_knn(sn_sp_repr_embedding_weight, text_emb):
18
+ """
19
+ :param sn_sp_repr_embedding_weight:
20
+ :param text_emb:
21
+ """
22
+ emb_norm = (sn_sp_repr_embedding_weight**2).sum(-1).view(-1, 1)
23
+ text_emb_t = torch.transpose(text_emb.view(-1, text_emb.size(-1)), 0, 1)
24
+ arr_norm = (text_emb**2).sum(-1).view(-1, 1)
25
+ dist = (
26
+ emb_norm
27
+ + arr_norm.cpu().transpose(0, 1)
28
+ - 2.0 * torch.mm(sn_sp_repr_embedding_weight, text_emb_t.cpu())
29
+ ) # (vocab, d) x (d, bsz*seqlen)
30
+ dist = torch.clamp(dist, 0.0, np.inf)
31
+ topk_out = torch.topk(-dist, k=1, dim=0)
32
+ return topk_out.values, topk_out.indices
33
+
34
+
35
+ def denoised_fn_round(args, sn_sp_repr_embedding, text_emb, t):
36
+ """
37
+ :param sn_sp_repr_embedding: the weights/parameter of the embedding layer that embeds the concatenated word IDs
38
+ :param text_emb: the model output at denoising step t; the transformer received the noise as input; this is the pred.
39
+ shape [batch size, args.seq_len, hidden_dim=768]
40
+ :param t: the current time step, shape [batch size] (same t for each instance in the batch)
41
+ """
42
+ sn_sp_repr_embedding_weight = sn_sp_repr_embedding.weight
43
+ old_shape = text_emb.shape
44
+ old_device = text_emb.device
45
+
46
+ if len(text_emb.shape) > 2:
47
+ text_emb = text_emb.reshape(-1, text_emb.size(-1))
48
+ else:
49
+ text_emb = text_emb
50
+
51
+ text_emb.to(sn_sp_repr_embedding_weight.device)
52
+
53
+ val, indices = get_efficient_knn(
54
+ sn_sp_repr_embedding_weight=sn_sp_repr_embedding_weight, text_emb=text_emb
55
+ )
56
+ rounded_tokens = indices[0]
57
+ new_embeds = sn_sp_repr_embedding(rounded_tokens).view(old_shape).to(old_device)
58
+ return new_embeds
scandl_module/original_scandl/sp_transformer_model.py ADDED
@@ -0,0 +1,167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import AutoConfig
2
+ from transformers.models.bert.modeling_bert import BertEncoder, BertModel
3
+ import torch
4
+ import torch as th
5
+ import torch.nn as nn
6
+ from typing import Optional
7
+
8
+ from .utils.nn import (
9
+ SiLU,
10
+ linear,
11
+ timestep_embedding,
12
+ )
13
+
14
+
15
+ class TransformerNetModel(nn.Module):
16
+ """
17
+ The ScanDL transformer.
18
+ """
19
+
20
+ def __init__(
21
+ self,
22
+ input_dims,
23
+ output_dims,
24
+ hidden_t_dim,
25
+ num_transformer_layers,
26
+ num_transformer_heads,
27
+ one_noise_step,
28
+ mask_padding,
29
+ dropout=0,
30
+ config=None,
31
+ config_name="bert-base-uncased",
32
+ vocab_size=None,
33
+ init_pretrained="no",
34
+ logits_mode=1,
35
+ ):
36
+ super().__init__()
37
+
38
+ if config is None:
39
+ config = AutoConfig.from_pretrained(config_name)
40
+ config.hidden_dropout_prob = dropout
41
+ config.num_hidden_layers = num_transformer_layers
42
+ config.num_attention_heads = num_transformer_heads
43
+ config.hidden_size = input_dims
44
+
45
+ self.input_dims = input_dims
46
+ self.hidden_t_dim = hidden_t_dim
47
+ self.output_dims = output_dims
48
+ self.dropout = dropout
49
+ self.logits_mode = logits_mode
50
+
51
+ self.mask_padding = mask_padding
52
+ self.one_noise_step = one_noise_step
53
+
54
+ self.bert_for_embedding = BertModel.from_pretrained(config_name)
55
+ # freeze BERT parameters (so that embeddings are freezed)
56
+ for param in self.bert_for_embedding.parameters():
57
+ param.requires_grad = False
58
+
59
+ self.sn_sp_repr_embedding = nn.Embedding(self.hidden_t_dim, self.input_dims)
60
+
61
+ self.positional_encoding = nn.Embedding(self.hidden_t_dim, self.input_dims)
62
+ self.sn_input_ids_embedding = self.bert_for_embedding.embeddings.word_embeddings
63
+ # additional linear layer needed if hidden is not 768 to map from pretrained BERT embeddings to other dim
64
+ if self.input_dims != 768:
65
+ self.proj_bert_emb = nn.Linear(768, self.input_dims)
66
+
67
+ self.lm_head = nn.Linear(self.input_dims, self.hidden_t_dim)
68
+ with torch.no_grad():
69
+ self.lm_head.weight = self.sn_sp_repr_embedding.weight
70
+
71
+ time_embed_dim = hidden_t_dim * 4
72
+ self.time_embed = nn.Sequential(
73
+ linear(hidden_t_dim, time_embed_dim),
74
+ SiLU(),
75
+ linear(time_embed_dim, config.hidden_size),
76
+ )
77
+
78
+ self.input_transformers = BertEncoder(config)
79
+ self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
80
+ self.dropout = nn.Dropout(config.hidden_dropout_prob)
81
+
82
+ def get_embeds(
83
+ self,
84
+ sn_sp_repr,
85
+ sn_input_ids,
86
+ indices_pos_enc,
87
+ ):
88
+
89
+ sn_sp_emb = self.sn_sp_repr_embedding(sn_sp_repr)
90
+ pos_enc = self.positional_encoding(indices_pos_enc)
91
+ if self.input_dims == 768:
92
+ sn_input_ids_emb = self.sn_input_ids_embedding(sn_input_ids)
93
+ else:
94
+ sn_input_ids_emb_bert_embs = self.sn_input_ids_embedding(sn_input_ids)
95
+ sn_input_ids_emb = self.proj_bert_emb(sn_input_ids_emb_bert_embs)
96
+ return sn_sp_emb, pos_enc, sn_input_ids_emb
97
+
98
+ def get_logits(self, model_output):
99
+ if self.logits_mode == 1:
100
+ return self.lm_head(model_output)
101
+ elif self.logits_mode == 2: # standard cosine similarity
102
+ raise NotImplementedError(
103
+ "standard cosine similarity not yet implemented for sp model output."
104
+ )
105
+ else:
106
+ raise NotImplementedError
107
+
108
+ def forward(
109
+ self,
110
+ x, # x_t
111
+ ts,
112
+ sn_input_ids_emb,
113
+ pos_enc,
114
+ attention_mask: Optional[torch.tensor] = None,
115
+ atten_vis: Optional[bool] = False,
116
+ ):
117
+ """
118
+ Apply the model to an input batch.
119
+
120
+ :param x: the noised input ID embeddings
121
+ :param ts: a 1-D batch of timesteps.
122
+ :param sn_input_ids_emb: the BERT embeddings
123
+ :param pos_enc: the positional embeddings
124
+ :param attention_mask: the attention mask (only given during training, not during inference)
125
+ :atten_vis: visualise attention
126
+ """
127
+ # timestep embedding
128
+ emb_t = self.time_embed(timestep_embedding(ts, self.hidden_t_dim))
129
+
130
+ # add the input x_t, the positional encoding pos_enc, the word ID embedding word_id_emb, and the timestep emb
131
+ emb_inputs = x + pos_enc + sn_input_ids_emb + emb_t.unsqueeze(1).expand(-1, x.size(1), -1)
132
+
133
+ # pipe through dropout and layer normalisation
134
+ emb_inputs = self.dropout(self.LayerNorm(emb_inputs))
135
+
136
+ if self.mask_padding:
137
+ if attention_mask == None:
138
+ raise ValueError("padding should be masked, but no attention mask given.")
139
+
140
+ extended_attention_mask = attention_mask[:, None, None, :]
141
+
142
+ if atten_vis:
143
+ model_out = self.input_transformers(
144
+ emb_inputs, attention_mask=extended_attention_mask, output_attentions=True
145
+ )
146
+ input_trans_hidden_states = model_out.last_hidden_state
147
+ attention_scores = model_out.attentions
148
+
149
+ else:
150
+ input_trans_hidden_states = self.input_transformers(
151
+ emb_inputs, attention_mask=extended_attention_mask
152
+ ).last_hidden_state
153
+
154
+ else:
155
+ if atten_vis:
156
+ model_out = self.input_transformers(emb_inputs, output_attentions=True)
157
+ input_trans_hidden_states = model_out.last_hidden_state
158
+ attention_scores = model_out.attentions
159
+ else:
160
+ input_trans_hidden_states = self.input_transformers(emb_inputs).last_hidden_state
161
+
162
+ h = input_trans_hidden_states
163
+ h = h.type(x.dtype)
164
+ if atten_vis:
165
+ return h, attention_scores
166
+ else:
167
+ return h