Spaces:
Sleeping
Sleeping
| import sys | |
| from pathlib import Path | |
| sys.path.insert(0, str(Path(__file__).resolve().parent)) | |
| from jinja2 import Environment | |
| import gradio as gr | |
| from llm_query.llm_client import AzureClient | |
| from llm_query.span_marking.llm_client import AzureClient as SpanMarkingAzureClient | |
| from llm_query.span_marking.utils import extract_labels_from_marked | |
| from llm_query.utils import ( | |
| extract_edit_tooltips, | |
| extract_labels, | |
| join_natural, | |
| join_tokenized, | |
| tokenize_line, | |
| ) | |
| MODEL_NAME = "gpt-5.6-terra" | |
| AUX_SYNC_CONFIG_NAME = "config.v0.6.json" | |
| SPAN_MARKING_CONFIG_NAME = "config.span_marking.v0.6.json" | |
| DEFAULT_TEXT_A = ( | |
| "It is the fourth-largest city of the province, with a population of 118,450." | |
| ) | |
| DEFAULT_TEXT_B = ( | |
| "C'est la quatrième plus grande ville de Slovaquie avec une population de " | |
| "81 114 habitants en 2015." | |
| ) | |
| TEMPLATE_PATH = Path(__file__).parent / "result_template.html" | |
| TEMPLATE = Environment().from_string(TEMPLATE_PATH.read_text()) | |
| aux_sync_client: AzureClient | None = None | |
| span_marking_client: SpanMarkingAzureClient | None = None | |
| def get_aux_sync_client() -> AzureClient: | |
| global aux_sync_client | |
| if aux_sync_client is None: | |
| aux_sync_client = AzureClient( | |
| model_name=MODEL_NAME, | |
| config_name=AUX_SYNC_CONFIG_NAME, | |
| ) | |
| return aux_sync_client | |
| def get_span_marking_client() -> SpanMarkingAzureClient: | |
| global span_marking_client | |
| if span_marking_client is None: | |
| span_marking_client = SpanMarkingAzureClient( | |
| model_name=MODEL_NAME, | |
| config_name=SPAN_MARKING_CONFIG_NAME, | |
| ) | |
| return span_marking_client | |
| def label_to_highlight(label: int) -> int: | |
| return 10 if label == 1 else 0 | |
| def render_tokens( | |
| tokens: tuple[str, ...], | |
| labels: tuple[int, ...], | |
| tooltips: tuple[str, ...], | |
| ) -> str: | |
| token_labels = [] | |
| for index, token in enumerate(tokens): | |
| label = labels[index] if index < len(labels) else 0 | |
| tooltip = tooltips[index] if index < len(tooltips) else "" | |
| token_labels.append((token + " ", label_to_highlight(label), tooltip)) | |
| return TEMPLATE.render(token_labels=token_labels) | |
| def render_error(message: str) -> str: | |
| return f'<p style="color: #b45309;">{message}</p>' | |
| def empty_tooltips(tokens: tuple[str, ...]) -> tuple[str, ...]: | |
| return tuple("" for _ in tokens) | |
| def generate_diff(text_a: str, text_b: str): | |
| aux_client = get_aux_sync_client() | |
| marking_client = get_span_marking_client() | |
| text_a = join_tokenized(text_a) | |
| text_b = join_tokenized(text_b) | |
| tokens_a = tokenize_line(text_a) | |
| tokens_b = tokenize_line(text_b) | |
| aux_response_a = aux_client.query(text_a=text_a, text_b=text_b) | |
| aux_response_b = aux_client.query(text_a=text_b, text_b=text_a) | |
| if aux_response_a.edited_text_a is None: | |
| aux_html_a = render_error("Could not get an edited version for Text A.") | |
| edited_a = "" | |
| else: | |
| labels_a = extract_labels(text_a, aux_response_a.edited_text_a) | |
| tooltips_a = extract_edit_tooltips(text_a, aux_response_a.edited_text_a) | |
| aux_html_a = render_tokens(tokens_a, labels_a, tooltips_a) | |
| edited_a = join_natural(aux_response_a.edited_text_a) | |
| if aux_response_b.edited_text_a is None: | |
| aux_html_b = render_error("Could not get an edited version for Text B.") | |
| edited_b = "" | |
| else: | |
| labels_b = extract_labels(text_b, aux_response_b.edited_text_a) | |
| tooltips_b = extract_edit_tooltips(text_b, aux_response_b.edited_text_a) | |
| aux_html_b = render_tokens(tokens_b, labels_b, tooltips_b) | |
| edited_b = join_natural(aux_response_b.edited_text_a) | |
| span_response_a = marking_client.query(text_a=text_a, text_b=text_b) | |
| span_response_b = marking_client.query(text_a=text_b, text_b=text_a) | |
| if span_response_a.marked_text_a is None: | |
| span_html_a = render_error("Could not get a marked version for Text A.") | |
| marked_a = "" | |
| else: | |
| span_labels_a = extract_labels_from_marked( | |
| text_a, | |
| span_response_a.marked_text_a, | |
| tokenize=tokenize_line, | |
| ) | |
| span_html_a = render_tokens(tokens_a, span_labels_a, empty_tooltips(tokens_a)) | |
| marked_a = join_natural(span_response_a.marked_text_a) | |
| if span_response_b.marked_text_a is None: | |
| span_html_b = render_error("Could not get a marked version for Text B.") | |
| marked_b = "" | |
| else: | |
| span_labels_b = extract_labels_from_marked( | |
| text_b, | |
| span_response_b.marked_text_a, | |
| tokenize=tokenize_line, | |
| ) | |
| span_html_b = render_tokens(tokens_b, span_labels_b, empty_tooltips(tokens_b)) | |
| marked_b = join_natural(span_response_b.marked_text_a) | |
| return ( | |
| aux_html_a, | |
| aux_html_b, | |
| edited_a, | |
| edited_b, | |
| span_html_a, | |
| span_html_b, | |
| marked_a, | |
| marked_b, | |
| ) | |
| with gr.Blocks(title="Generative Semantic Diff") as demo: | |
| gr.Markdown("# Generative Semantic Diff") | |
| with gr.Row(): | |
| text_a = gr.Textbox( | |
| label="Text A", | |
| value=DEFAULT_TEXT_A, | |
| lines=2, | |
| ) | |
| text_b = gr.Textbox( | |
| label="Text B", | |
| value=DEFAULT_TEXT_B, | |
| lines=2, | |
| ) | |
| with gr.Row(): | |
| submit_btn = gr.Button(value="Generate Diff") | |
| gr.Markdown("## Auxiliary synchronization") | |
| with gr.Row(): | |
| with gr.Column(variant="panel"): | |
| aux_output_a = gr.HTML(label="Result for text A", show_label=True) | |
| with gr.Column(variant="panel"): | |
| aux_output_b = gr.HTML(label="Result for text B", show_label=True) | |
| with gr.Row(): | |
| edited_a = gr.Textbox(label="LLM-edited Text A", lines=2, interactive=False) | |
| edited_b = gr.Textbox(label="LLM-edited Text B", lines=2, interactive=False) | |
| gr.Markdown("## Span marking") | |
| with gr.Row(): | |
| with gr.Column(variant="panel"): | |
| span_output_a = gr.HTML(label="Result for text A", show_label=True) | |
| with gr.Column(variant="panel"): | |
| span_output_b = gr.HTML(label="Result for text B", show_label=True) | |
| with gr.Row(): | |
| marked_a = gr.Textbox(label="LLM-marked Text A", lines=2, interactive=False) | |
| marked_b = gr.Textbox(label="LLM-marked Text B", lines=2, interactive=False) | |
| submit_btn.click( | |
| fn=generate_diff, | |
| inputs=[text_a, text_b], | |
| outputs=[ | |
| aux_output_a, | |
| aux_output_b, | |
| edited_a, | |
| edited_b, | |
| span_output_a, | |
| span_output_b, | |
| marked_a, | |
| marked_b, | |
| ], | |
| ) | |
| demo.queue() | |
| demo.launch() | |