First commit
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitignore +8 -0
- CITATION.cff +27 -0
- CONSTANTS.py +62 -0
- LICENSE +121 -0
- PATHS.py +4 -0
- README.md +212 -0
- __init__.py +27 -0
- app.py +185 -0
- config.json +53 -0
- config_bsc.json +54 -0
- config_emtec.json +54 -0
- create_data.py +167 -0
- fix_dur_module/__init__.py +0 -0
- fix_dur_module/__pycache__/__init__.cpython-313.pyc +0 -0
- fix_dur_module/__pycache__/model_seq2seq.cpython-313.pyc +0 -0
- fix_dur_module/__pycache__/scasim.cpython-313.pyc +0 -0
- fix_dur_module/__pycache__/utils_data.cpython-313.pyc +0 -0
- fix_dur_module/__pycache__/utils_train.cpython-313.pyc +0 -0
- fix_dur_module/model_seq2seq.py +89 -0
- fix_dur_module/scasim.py +185 -0
- fix_dur_module/train_seq2seq.py +283 -0
- fix_dur_module/utils_data.py +530 -0
- fix_dur_module/utils_train.py +195 -0
- handler.py +53 -0
- model.py +701 -0
- models/paragraph/fixdur-module/hyperparameters.json +1 -0
- models/paragraph/fixdur-module/min_max_scaler.pkl +3 -0
- models/paragraph/fixdur-module/seq2seq_fixdur.pt +3 -0
- models/paragraph/scandl-module/ema_0.9999_080000.pt +3 -0
- models/paragraph/scandl-module/training_args.json +52 -0
- models/sentence/fixdur-module/hyperparameters.json +1 -0
- models/sentence/fixdur-module/min_max_scaler.pkl +3 -0
- models/sentence/fixdur-module/seq2seq_fixdur.pt +3 -0
- models/sentence/scandl-module/ema_0.9999_080000.pt +3 -0
- models/sentence/scandl-module/training_args.json +52 -0
- requirements.txt +15 -0
- scandl2_utils.py +77 -0
- scandl_module/.DS_Store +0 -0
- scandl_module/__init__.py +4 -0
- scandl_module/__pycache__/__init__.cpython-313.pyc +0 -0
- scandl_module/original_scandl/__init__.py +0 -0
- scandl_module/original_scandl/__pycache__/__init__.cpython-313.pyc +0 -0
- scandl_module/original_scandl/__pycache__/sp_gaussian_diffusion.cpython-313.pyc +0 -0
- scandl_module/original_scandl/__pycache__/sp_rounding.cpython-313.pyc +0 -0
- scandl_module/original_scandl/__pycache__/sp_transformer_model.cpython-313.pyc +0 -0
- scandl_module/original_scandl/__pycache__/step_sample.cpython-313.pyc +0 -0
- scandl_module/original_scandl/config.json +52 -0
- scandl_module/original_scandl/sp_gaussian_diffusion.py +1183 -0
- scandl_module/original_scandl/sp_rounding.py +58 -0
- 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
|