config = {"methods": ["fa", "vae", "lgminmcm", "ncfa"], "seeds": range(101)}
vae_config = {
    "num_epochs": 200,
    "lr": 1e-3,
    "latent_width": 1,
    "meas_width": 2,
    "num_hidden_layers": 1,
    "encoder_hidden_dim": 64,
}


def input_path(wc, outfile):
    return (
        f"data/{wc.dataset}/"
        + (f"seed={wc.seed}/" if wc.dataset == "simulated" else "")
        + outfile
    )


rule all:
    input:
        "results/table.tex",


rule simulate_data:
    output:
        dataset="data/simulated/seed={seed}/dataset.txt",
        latent_sample="data/simulated/seed={seed}/latent_sample.txt",
        graph="data/simulated/seed={seed}/graph.txt",
    params:
        num_latent=3,
        num_meas=10,
        edge_prob=0.3,
        samp_size=10000,
    script:
        "scripts/simulate_data.py"


rule causalchamber_data:
    output:
        dataset="data/causalchamber/dataset.txt",
        latent_sample="data/causalchamber/latent_sample.txt",
        graph="data/causalchamber/graph.txt",
    script:
        "scripts/causalchamber.py"


rule fit_and_eval_linear:
    input:
        dataset=lambda wildcards: input_path(wildcards, "dataset.txt"),
        latent_sample=lambda wildcards: input_path(wildcards, "latent_sample.txt"),
        graph=lambda wildcards: input_path(wildcards, "graph.txt"),
    output:
        result="results/{dataset}/seed={seed}/{method}_eval.txt",
        graph="results/{dataset}/seed={seed}/{method}_graph.txt",
    wildcard_constraints:
        dataset=r"(simulated|causalchamber)",
        method=r"(fa|lgminmcm)",
    params:
        method=lambda wildcards: wildcards.method,
    script:
        "scripts/linear.py"


rule fit_and_eval_nonlinear:
    input:
        dataset=lambda wildcards: input_path(wildcards, "dataset.txt"),
        latent_sample=lambda wildcards: input_path(wildcards, "latent_sample.txt"),
        graph=lambda wildcards: input_path(wildcards, "graph.txt"),
    output:
        result="results/{dataset}/seed={seed}/{method}_eval.txt",
        graph="results/{dataset}/seed={seed}/{method}_graph.txt",
    wildcard_constraints:
        dataset=r"(simulated|causalchamber)",
        method=r"(vae|ncfa)",
    params:
        method=lambda wildcards: wildcards.method,
        num_epochs=vae_config["num_epochs"],
        lr=vae_config["lr"],
        latent_width=vae_config["latent_width"],
        meas_width=vae_config["meas_width"],
        num_hidden_layers=vae_config["num_hidden_layers"],
        encoder_hidden_dim=vae_config["encoder_hidden_dim"],
    script:
        "scripts/nonlinear.py"


rule collect_results:
    input:
        expand(
            "results/{dataset}/seed={seed}/{method}_eval.txt",
            dataset=["simulated", "causalchamber"],
            seed=config["seeds"],
            method=config["methods"],
        ),
    output:
        "results/results.txt",
    script:
        "scripts/collect.py"


rule format_table:
    input:
        "results/results.txt",
    output:
        "results/table.tex",
    script:
        "scripts/format_table.py"
