import numpy as np
import pymc as pm

def model(data):
    """
    von Bertalanffy growth model for dugong length as a function of age.
    """
    age = data['age']
    length = data['length']
    
    with pm.Model() as dugong_model:
        # Data containers
        age_data = pm.Data("age", age, dims="obs_id")
        
        # Priors for von Bertalanffy parameters
        # L∞: asymptotic length, should be > max observed length
        L_inf = pm.HalfNormal("L_inf", sigma=3.0)
        
        # k: growth rate, positive
        k = pm.HalfNormal("k", sigma=2.0)
        
        # t0: theoretical age at zero length, can be negative
        t0 = pm.Normal("t0", mu=0.0, sigma=1.0)
        
        # Expected length
        mu = L_inf * (1 - pm.math.exp(-k * (age_data - t0)))
        
        # Observation model
        sigma = pm.HalfNormal("sigma", sigma=0.5)
        y_obs = pm.Normal("y_obs", mu=mu, sigma=sigma, observed=length, dims="obs_id")
    
    return dugong_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