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'
{message}
' 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()