import numpy as np
import pymc as pm


def model(data):
    time = data['Time']
    infected = data['Infected_Count']
    N = 50  # total population
    
    with pm.Model() as model:
        # Feature container for time
        time_data = pm.Data('Time', time, dims='obs_id')
        
        # Logistic growth parameters
        r = pm.HalfNormal('r', sigma=5)
        t0 = pm.Normal('t0', mu=1.0, sigma=0.5)
        
        # Deterministic logistic curve
        mu_infected = pm.Deterministic(
            'mu_infected',
            N / (1 + pm.math.exp(-r * (time_data - t0))),
            dims='obs_id'
        )
        
        # Overdispersion parameter
        sigma = pm.HalfNormal('sigma', sigma=5)
        
        # Negative Binomial likelihood
        y_obs = pm.NegativeBinomial(
            'y_obs',
            mu=mu_infected,
            sigma=sigma,
            observed=infected,
            dims='obs_id'
        )
    
    return model


def gen_model(observed_data):
    built_model = model({column: observed_data[column].to_numpy() for column in observed_data.columns})
    with built_model:
        trace = pm.sample(200, tune=200, target_accept=0.90, chains=1, cores=1, random_seed=42, idata_kwargs={"log_likelihood": True})
        posterior_predictive = pm.sample_posterior_predictive(trace, random_seed=314, return_inferencedata=False)
    return built_model, posterior_predictive, trace