library(data.table)

test <- fread('../input/test_set_metadata.csv')
train <- fread('../input/training_set_metadata.csv')

with(train, table(is.na(distmod), target))

classes_id <- c(sort(unique(train$target)), 99)
classes <- sapply(classes_id, function(x) paste0('class_', x))
p <- length(classes)

weights <- c(1, 2, 1, 1, 1, 1, 1, 2, 1, 1, 1, 1, 1, 1, 2)

# init
submission <- data.table(object_id=test$object_id)
for (i in c(1:p)){
    submission[[classes[i]]] <- 0
}

a <- 0.04924697
# inner classes
inner_classes_id <- sort(union(unique(train[is.na(distmod),]$target), 99))
inner_classes <- sapply(inner_classes_id, function(x) paste0('class_', x))
inner_rows <- test[is.na(distmod), which=TRUE]
inner_cols <- which(classes_id %in% inner_classes_id)
print(inner_classes)

k <- sum(weights[inner_cols] == 1)
m <- sum(weights[inner_cols] == 2)

p1 <- 1/(2*a + k)
print(p1)
print(1 - k*p1 - 2*(m - 1)*p1 )
for (i in inner_cols){
    if (weights[i] == 1){
        submission[[classes[i]]][inner_rows] <- p1
    } else if (weights[i] == 2){
        submission[[classes[i]]][inner_rows] <- 1 - k*p1 - 2*(m - 1)*p1 
    }
}

# outer classes
outer_classes_id <- sort(union(unique(train[!is.na(distmod),]$target), 99))
outer_classes <- sapply(outer_classes_id, function(x) paste0('class_', x))
outer_rows <- test[!is.na(distmod), which=TRUE]
outer_cols <- which(classes_id %in% outer_classes_id)
print(outer_classes)

k <- sum(weights[outer_cols] == 1)
m <- sum(weights[outer_cols] == 2)

q1 <- 1/(k + 2*(m - 1) + 2*(1 - a))
print(q1)
print(1 - k*q1 - 2*(m - 1)*q1)
for (i in outer_cols){
    if (weights[i] == 1){
        submission[[classes[i]]][outer_rows] <- q1
    } else if ((weights[i] == 2) & (classes_id[i] != 99)){
        submission[[classes[i]]][outer_rows] <- 2*q1
    } else if ((weights[i] == 2) & (classes_id[i] == 99)){
        submission[[classes[i]]][outer_rows] <- 1 - k*q1 - 2*(m - 1)*q1    
    }
}

filename <- 'submission.csv'
fwrite(submission, filename, sep=',', quote=FALSE, row.names=FALSE)
zipfilename <- 'submission.zip'
zip(zipfilename, filename)
#zip(zipfilename, filename, flags='-m')