library(data.table)
library(brms)



######
# script-wide parameters
######
dark <- "#23384F"
font <- "Inter"
fig_height <- 5
fig_width <- 7



######
# Data
######
d <- fread('https://soci620.netlify.app/data/delinquency.csv')


######
# Intercept-only norm vs categorical
######
cfreq <- mean(d$ever_cocaine, na.rm=TRUE)
cstd <- sqrt(var(d$ever_cocaine, na.rm=TRUE))

x <- seq(from=-0.2, to=1.2, length.out=600)
y <- dnorm(x,cfreq,cstd)
y <- y / max(y) / 2


svg(
    "img/09_norm_vs_categorical.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width, height=fig_height * 0.8
)
par(mar=c(2,1,0,1)+0.5, xpd = FALSE)
plot(
    NA, # no data for now
    main=NA, xlab=NA, ylab=NA,
    xaxt='n', yaxt='n', bty='n',yaxs='i',
    xlim=range(x),ylim=range(0,y*1.05,1),
    fg=dark,col=dark,
    col.lab=dark, cex.lab=1.5,
    font.lab=2, adj=1.0
)
    # horizontal axis
axis(
    1, at = c(0,1),
    lwd=2, col=dark,
    col.axis=dark, cex.axis=1.3,
    font.axis=2,lend='square'
)
# densities
lines(x, y, lwd=5, col=dark, lend='butt', lty='11')
# bars
lines(c(0,0),c(0,1-cfreq), lwd=18,lend='butt',col=dark)
lines(c(1,1),c(0,  cfreq), lwd=18,lend='butt',col=dark)
par(xpd=TRUE)
lines(range(x),c(0,0),lwd=2,col=dark,lend='butt')


dev.off()



######
# Generic Normal and Bernoulli
######

# Normal
x <- seq(from=-4, to=4, length.out=600)
y <- dnorm(x,0,1)
svg(
    "img/09_generic_normal_labeled.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width, height=fig_width
)
par(mar=c(2,0,0,0)+0.5, xpd = TRUE)
plot(
    NA, # no data for now
    main=NA, xlab=NA, ylab=NA,
    xaxt='n', yaxt='n', bty='n',
    xlim=range(x),ylim=range(0,y*1.05),
    fg=dark,col=dark,
    col.lab=dark, cex.lab=1.5,
    font.lab=2, adj=1.0
)
# densities
lines(x, y, lwd=8, col=dark, lend='butt')
# mu
lines(c(0,0),c(0,max(y)),lty='11',lend='butt',lwd=2.5,col=dark)
text(
    0,0,
    labels='μ',
    adj=c(0.5,1.5),col=dark,
    cex=2.0,font=4
)
# sigma
sx <- c(0,1.3)
sy <- rep(dnorm(sx[2],0,1),2)
lines(sx,sy,col=dark,lwd=2.5,lend='butt')
lines(rep(sx[1],2),sy + c(-.02,.02),col=dark,lwd=2.5,lend='butt')
lines(rep(sx[2],2),sy + c(-.02,.02),col=dark,lwd=2.5,lend='butt')
text(
    mean(sx),sy[1],
    labels='σ',
    adj=c(0.5,1.5),col=dark,
    cex=2.0,font=4
)

dev.off()

# Bernoulli
p <- .4
svg(
    "img/09_generic_bernoulli_labeled.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width, height=fig_width
)
par(mar=c(2,0,0,0)+0.5, xpd = TRUE)
plot(
    NA, # no data for now
    main=NA, xlab=NA, ylab=NA,
    xaxt='n', yaxt='n', bty='n',
    xlim=c(-.5,1.5),ylim=c(0,max(p,1-p)*1.05),
    fg=dark,col=dark,
    col.lab=dark, cex.lab=1.5,
    font.lab=2, adj=1.0
)
# pmf
lines(c(0,0), c(0,1-p), lwd=100, col=dark, lend='butt')
lines(c(1,1), c(0,p), lwd=100, col=dark, lend='butt')
# x labels
text(
    0,0,
    labels='x=0',
    adj=c(0.5,1.5),col=dark,
    cex=2.0,font=2
)
text(
    1,0,
    labels='x=1',
    adj=c(0.5,1.5),col=dark,
    cex=2.0,font=2
)
# in-bar labels
text(
    0,1-p,
    labels='1 – p',
    adj=c(0.5,1.5),col='white',
    cex=2.0,font=2
)
text(
    1,p,
    labels='p',
    adj=c(0.5,1.5),col='white',
    cex=2.0,font=2
)

dev.off()



######
# Basic inverse logit
######
logit <- function(p){log(p/(1-p))}
inv_logit <- function(x){1/(1+exp(-x))}

x <- seq(from=-4.2, to=4.2, length.out=600)
y <- inv_logit(x)


svg(
    "img/09_basic_inv_logit.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width, height=fig_height 
)
par(mar=c(4,4,2,0)+.1, xpd = FALSE)
plot(
    NA, # no data for now
    main=NA, xlab=NA, ylab=NA,
    xaxt='n', yaxt='n', bty='n',yaxs='i',xaxs='i',
    xlim=range(x),ylim=c(0,1),
    fg=dark,col=dark,
    col.lab=dark, cex.lab=1.5,
    font.lab=2, adj=1.0
)
box(bty='O',lwd=3,col=dark,lend='square',ljoin='mitre')

# horizontal axis
axis(
    1,
    lwd=2, col=dark,
    col.axis=dark, cex.axis=1.3,
    font.axis=2,lend='square'
)
# vertical axis
axis(
    2,
    lwd=2, col=dark,
    col.axis=dark, cex.axis=1.3,
    font.axis=2,lend='square'
)
# labels
title(
    xlab="x",
    ylab="p",
    line=2.5,col=dark,
    col.lab=dark, cex.lab=1.5,
    font.lab=4, adj=1.0
)   


lines(x, y, lwd=5, col=dark, lend='butt', lty=1)

dev.off()




######
# Inverse logit priors
######

dlogitnorm <- function(x,mu,sigma){
    num <- exp((logit(x)-mu)^2 / (-2*sigma^2))
    den <- sigma * sqrt(2*pi) * x * (1-x)
    return(num/den)
}



svg(
    "img/09_normal_0_3.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width, height=fig_height
)

x <- seq(from=-10,to=10,length.out=700)
y <- dnorm(x,0,3)

par(mar=c(3,0,0,0)+0.5, xpd = TRUE)
plot(
    NA, # no data for now
    main=NA, xlab=NA, ylab=NA,
    xaxt='n', yaxt='n', bty='n',
    xlim=range(x),ylim=range(0,y*1.05),
    fg=dark,col=dark,
    col.lab=dark, cex.lab=1.5,
    font.lab=2, adj=1.0
)
# horizontal axis
axis(
    1, at=(-3:3)*3,
    lwd=2, col=dark,
    col.axis=dark, cex.axis=1.3,
    font.axis=2,lend='square'
)
# densities
lines(x, y, lwd=5, col=dark, lend='butt')
dev.off()



svg(
    "img/09_logitnormal_0_3.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width, height=fig_height
)

x <- seq(from=0,to=1,length.out=1800)
y <- dlogitnorm(x,0,3)
y[is.nan(y)] <- 0

par(mar=c(3,0,0,0)+0.5, xpd = TRUE)
plot(
    NA, # no data for now
    main=NA, xlab=NA, ylab=NA,
    xaxt='n', yaxt='n', bty='n',
    xlim=range(x),ylim=range(0,y*1.05),
    fg=dark,col=dark,
    col.lab=dark, cex.lab=1.5,
    font.lab=2, adj=1.0
)
# horizontal axis
axis(
    1,
    lwd=2, col=dark,
    col.axis=dark, cex.axis=1.3,
    font.axis=2,lend='square'
)
# densities
lines(x, y, lwd=5, col=dark, lend='butt')
dev.off()



######
# logit-transformation illustration
######

logit_norm_plot <- function(sigma,mu=0,xlim=c(-20,20),quantiles=c(.25,.5,.75)){
    xgrid <- seq(from=min(xlim),to=max(xlim),length.out=1000)
    ygrid <- seq(from=0,to=1,length.out=2000)

    par(mar=c(3,3,3,3)+0.5, xpd = TRUE)
    layout(matrix(c(2,2,4,1,1,3,1,1,3),nrow=3,byrow=TRUE))

    # plot the inv logit
    plot(
        NA, # no data for now
        main=NA, xlab=NA, ylab=NA,
        xaxt='n', yaxt='n', bty='n',yaxs='i',xaxs='i',
        xlim=xlim,ylim=c(0,1),
        fg=dark,col=dark,
        col.lab=dark, cex.lab=1.5,
        font.lab=2, adj=1.0
    )
    box(bty='O',lwd=2,col=dark,lend='square',ljoin='mitre')

    # horizontal axis
    axis(
        1,
        lwd=2, col=dark,
        col.axis=dark, cex.axis=1.3,
        font.axis=2,lend='square'
    )
    # vertical axis
    axis(
        2,
        lwd=2, col=dark,
        col.axis=dark, cex.axis=1.3,
        font.axis=2,lend='square'
    )
    lines(xgrid,inv_logit(xgrid),lwd=5,col=dark,lend='butt')
    # quantiles
    for(x in qnorm(quantiles,mu,sigma)){
        lines(
            c(x,x),c(inv_logit(x),2),
            col=dark,lwd=.5,lend='butt')
        lines(
            c(x,max(xlim)*2),rep(inv_logit(x),2),
            col=dark,lwd=.5,lend='butt')

    }

    # plot the norm on top
    par(mar=c(0,3,3,3)+0.5, xpd = TRUE)
    plot(
        NA, # no data for now
        main=NA, xlab=NA, ylab=NA,
        xaxt='n', yaxt='n', bty='n',yaxs='i',xaxs='i',
        xlim=xlim,ylim=range(0,dnorm(xgrid,mu,sigma)),
        fg=dark,col=dark,
        col.lab=dark, cex.lab=1.5,
        font.lab=2, adj=1.0
    )
    lines(xgrid,dnorm(xgrid,mu,sigma),lwd=5,col=dark,lend='butt',lty='11')
    # quantiles
    for(x in qnorm(quantiles,mu,sigma)){
        lines(
            c(x,x),c(-2,dnorm(x,mu,sigma)),
            col=dark,lwd=.5,lend='butt')
    }

    # plot the logit norm on the right
    par(mar=c(3,0,3,3)+0.5, xpd = TRUE)
    lnvals <- dlogitnorm(ygrid,mu,sigma)
    lnvals[is.nan(lnvals)] <- 0
    plot(
        NA, # no data for now
        main=NA, xlab=NA, ylab=NA,
        xaxt='n', yaxt='n', bty='n',yaxs='i',xaxs='i',
        xlim=range(0,lnvals),ylim=c(0,1),
        fg=dark,col=dark,
        col.lab=dark, cex.lab=1.5,
        font.lab=2, adj=1.0
    )
    lines(lnvals,ygrid,lwd=5,col=dark,lend='butt',lty='11')
    # quantiles
    for(y in inv_logit(qnorm(quantiles,mu,sigma))){
        lines(
            c(-10,dlogitnorm(y,mu,sigma)),c(y,y),
            col=dark,lwd=.5,lend='butt')
    }


}

svg(
    "img/09_alphaprior_0_3.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width*1.8, height=fig_height*1.8
)
logit_norm_plot(3)
dev.off()


svg(
    "img/09_alphaprior_0_1.7.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width*1.8, height=fig_height*1.8
)
logit_norm_plot(1.7)
dev.off()


svg(
    "img/09_alphaprior_0_1.svg",
    bg="#FFFFFF00",
    family=font,
    width=fig_width*1.8, height=fig_height*1.8
)
logit_norm_plot(1)
dev.off()






######
# regression results
######

m_0 <- bf(ever_cocaine ~ 1, family=bernoulli)
p_0 <- prior_string("normal(0, 1.7)",class="Intercept")
f_0 <- brm(m_0, prior=p_0, data=d, chains=4, cores=4, iter=5000)