import marimo

__generated_with = "0.25.1"
app = marimo.App(width="medium")


@app.cell
def _():
    import math
    import marimo as mo
    return math, mo


@app.cell
def _(mo):
    task_focus = mo.ui.slider(0, 1, step=0.05, value=0.5, label="Task focus", show_value=True)
    temperature = mo.ui.slider(0.05, 1, step=0.05, value=0.25, label="Temperature", show_value=True)
    mo.vstack([task_focus, temperature])
    return task_focus, temperature


@app.cell
def _(math, task_focus, temperature):
    acts = [
        ("Ask a question", 0.65, 0.65),
        ("Reflect", 0.35, 0.95),
        ("Offer advice", 0.95, 0.15),
    ]
    scores = [task_focus.value * task + (1 - task_focus.value) * social for _, task, social in acts]
    _weights = [math.exp((score - max(scores)) / temperature.value) for score in scores]
    probabilities = [weight / sum(_weights) for weight in _weights]
    return acts, probabilities, scores


@app.cell
def _(acts, math, mo, probabilities, scores):
    _rows = "\n".join(
        f"| {name} | {task:.2f} | {social:.2f} | {score:.2f} | {probability:.1%} |"
        for (name, task, social), score, probability in zip(acts, scores, probabilities)
    )
    _best = acts[max(range(len(acts)), key=lambda i: probabilities[i])][0]
    _entropy = -sum(p * math.log2(p) for p in probabilities)
    mo.md(
        "| Dialogue act | Task score | Social score | Combined score | Probability |\n"
        "| :--- | ---: | ---: | ---: | ---: |\n" + _rows
        + f"\n\n**Most likely act:** {_best}  \n**Policy entropy:** {_entropy:.2f} bits"
    )
    return


if __name__ == "__main__":
    app.run()
