File size: 12,345 Bytes
67a2b13
 
 
 
 
 
 
 
 
 
 
 
 
 
ddb3f9b
 
 
67a2b13
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22e4917
 
 
 
67a2b13
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22e4917
 
 
67a2b13
 
 
 
 
 
 
22e4917
 
 
67a2b13
22e4917
67a2b13
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22e4917
 
 
 
 
67a2b13
 
 
 
 
 
 
22e4917
 
67a2b13
 
 
 
22e4917
 
 
67a2b13
 
 
22e4917
 
 
 
 
67a2b13
 
 
 
 
ddb3f9b
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
import argparse
import html
import time

from extend import spacy_component  # this is needed to register the spacy component

import spacy
import streamlit as st
from annotated_text import annotation
from classy.scripts.model.demo import tabbed_navigation
from classy.utils.streamlit import get_md_200_random_color_generator


def main(
    model_checkpoint_path: str,
    default_inventory_path: str,
    cuda_device: int,
):
    # setup examples
    examples = [
        "Italy beat England and won Euro 2021.",
        "Japan began the defence of their Asian Cup title with a lucky 2-1 win against Syria in a Group C championship match on Friday.",
        "The project was coded in Java.",
    ]

    # css rules
    st.write(
        """
            <style type="text/css">
                a {
                    text-decoration: none !important;
                }
            </style>
        """,
        unsafe_allow_html=True,
    )

    # setup header
    st.markdown(
        "<h1 style='text-align: center;'>ExtEnD: Extractive Entity Disambiguation</h1>",
        unsafe_allow_html=True,
    )
    st.write(
        """
            <div align="center">
                <a href="https://sunglasses-ai.github.io/classy/">
                    <img alt="Python" style="height: 3em; margin: 0 1em" src="">
                </a>
                <a href="https://spacy.io/" tyle="text-decoration: none">
                    <img alt="spaCy" style="height: 3em; margin: 0 1em;" src="">
                </a>
            </div> 
        """,
        unsafe_allow_html=True,
    )

    def model_demo():
        @st.cache(allow_output_mutation=True)
        def load_resources(inventory_path):

            # load nlp
            nlp = spacy.load("en_core_web_sm")
            extend_config = dict(
                checkpoint_path=model_checkpoint_path,
                mentions_inventory_path=inventory_path,
                device=cuda_device,
                tokens_per_batch=10_000,
            )
            nlp.add_pipe("extend", after="ner", config=extend_config)

            # mock call to load resources
            nlp(examples[0])

            # return
            return nlp

        # read input
        placeholder = st.selectbox(
            "Examples",
            options=examples,
            index=0,
        )
        input_text = st.text_area("Input text to entity-disambiguate", placeholder)

        # custom inventory
        uploaded_inventory_path = st.file_uploader(
            "[Optional] Upload custom inventory (tsv file, mention \\t desc1 \\t desc2 \\t)",
            accept_multiple_files=False,
            type=["tsv"],
        )
        if uploaded_inventory_path is not None:
            inventory_path = f"data/inventories/{uploaded_inventory_path.name}"
            with open(inventory_path, "wb") as f:
                f.write(uploaded_inventory_path.getbuffer())
        else:
            inventory_path = default_inventory_path

        # load model and color generator
        nlp = load_resources(inventory_path)
        color_generator = get_md_200_random_color_generator()

        if st.button("Disambiguate", key="classify"):

            # tag sentence
            time_start = time.perf_counter()
            doc = nlp(input_text)
            time_end = time.perf_counter()

            # extract entities
            entities = {}
            for ent in doc.ents:
                if ent._.disambiguated_entity is not None:
                    entities[ent.start_char] = (
                        ent.start_char,
                        ent.end_char,
                        ent.text,
                        ent._.disambiguated_entity,
                    )

            # create annotated html components

            annotated_html_components = []

            assert all(any(t.idx == _s for t in doc) for _s in entities)
            it = iter(list(doc))
            while True:
                try:
                    t = next(it)
                except StopIteration:
                    break
                if t.idx in entities:
                    _start, _end, _text, _entity = entities[t.idx]
                    while t.idx + len(t) != _end:
                        t = next(it)
                    annotated_html_components.append(
                        str(annotation(*(_text, _entity, color_generator())))
                    )
                else:
                    annotated_html_components.append(str(html.escape(t.text)))

            st.markdown(
                "\n".join(
                    [
                        "<div>",
                        *annotated_html_components,
                        "<p></p>"
                        f'<div style="text-align: right"><p style="color: gray">Time: {(time_end - time_start):.2f}s</p></div>'
                        "</div>",
                    ]
                ),
                unsafe_allow_html=True,
            )

    def hiw():
        st.markdown("ExtEnD frames Entity Disambiguation as a text extraction problem:")
        st.image(
            "data/repo-assets/extend_formulation.png", caption="ExtEnD Formulation"
        )
        st.markdown(
            """            
            Given the sentence *After a long fight Superman saved Metropolis*, where *Superman* is the mention
            to disambiguate, ExtEnD first concatenates the descriptions of all the possible candidates of *Superman* in the
            inventory and then selects the span whose description best suits the mention in its context.

            To convert this task to end2end entity linking, as we do in *Model demo*, we leverage spaCy 
            (more specifically, its NER) and run ExtEnD on each named entity spaCy identifies 
            (if the corresponding mention is contained in the inventory).
        """
        )

    def abstract():
        st.write(
            """
            Local models for Entity Disambiguation (ED) have today become extremely powerful, in most part thanks to the advent of large pre-trained language models. However, despite their significant performance achievements, most of these approaches frame ED through classification formulations that have intrinsic limitations, both computationally and from a modeling perspective. In contrast with this trend, here we propose EXTEND, a novel local formulation for ED where we frame this task as a text extraction problem, and present two Transformer-based architectures that implement it. Based on experiments in and out of domain, and training over two different data regimes, we find our approach surpasses all its competitors in terms of both data efficiency and raw performance. EXTEND outperforms its alternatives by as few as 6 F 1 points on the more constrained of the two data regimes and, when moving to the other higher-resourced regime, sets a new state of the art on 4 out of 6 benchmarks under consideration, with average improvements of 0.7 F 1 points overall and 1.1 F 1 points out of domain. In addition, to gain better insights from our results, we also perform a fine-grained evaluation of our performances on different classes of label frequency, along with an ablation study of our architectural choices and an error analysis. We release our code and models for research purposes at https:// github.com/SapienzaNLP/extend.

            Link to full paper: https://www.researchgate.net/publication/359392427_ExtEnD_Extractive_Entity_Disambiguation 
        """
        )

    tabs = dict(
        model=("Model demo", model_demo),
        hiw=("How it works", hiw),
        abstract=("Abstract", abstract),
    )

    tabbed_navigation(tabs, "model")


if __name__ == "__main__":
    main(
        "experiments/extend-longformer-large/2021-10-22/09-11-39/checkpoints/best.ckpt",
        "data/inventories/aida.tsv",
        cuda_device=-1,
    )