# How to use this file

# You are expected to run the whole file at one time, follow the instructions below

# Step 0: Uncomment and run the first time you run this code to install dependencies
# using Pkg
# Pkg.add("MLDatasets")

# Step 1: Include this file in your code by writing
# include("loocv.jl")

# Step 2: Train a predictor theta_hat by calling
# theta_hat = train_theta_hat()

# Step 3: Test your predictor on the test set by calling
# error = test_theta_hat(theta_hat)
# The error returned is the mean-squared error.
#
#IF THIS CODE DOES NOT WORK FOR YOU, THERE IS AN ALTNERATE METHOD AT THE BOTTOM    

using MLDatasets # imports MNIST
using LinearAlgebra

function train_theta_hat()
    trainset = MNIST()
    # trainset.features is a 28 x 28 x 60000 matrix of data
    # trainset.targets is a 60000 vector of labels

    # reshape the data into a matrix where each ROW is a sample and each COLUMN is a feature
    w,d,N = size(trainset.features)
    A = reshape(trainset.features, (w*d, N))'
    println("A matrix has size ", size(A))
    y = trainset.targets

    theta_hat = pinv(A)*y; # this is the linear regression step
    return theta_hat
end

function test_theta_hat(theta_hat)
    testset = MNIST()

    w,d,N = size(testset.features)    
    A = reshape(testset.features, (w*d, N))'
    
    # y_hat is the predicted values
    y_hat = A*theta_hat
    y_true = testset.targets

    # Since y_hat is supposed to predict the digits 0-9, we'll round y_hat to the nearest integer
    y_hat = round.(y_hat)

    # the mean-squared error
    return sum((y_hat .- y_true).^2)./N
end


# IF FOR SOME REASON YOU CANNOT DOWNLOAD MNIST() THROUGH MLDatasets:
# Use the alternate steps below

# To set up the alternate code:
# Download and unzip mnist.zip from the julia_files directory.

# The first time you run this code, install dependencies with
# using Pkg; 
#Pkg.add("CSV"); 
#Pkg.add("DataFrames")

# You do not need this alternate code if the code above works on your machine.

#= # delete this line if you need to uncomment the alternate method

using CSV
using DataFrames

function train_theta_hat_alternate()
    A = Matrix(DataFrame(CSV.File("mnist_data.csv")))
    y = Matrix(DataFrame(CSV.File("mnist_labels.csv")))
    # A is a 28*28 x 60000 matrix of data
    # y is a 60000 vector of labels

    # reshape the data into a matrix where each ROW is a sample and each COLUMN is a feature
    wd,N = size(A)
    println("A matrix has size ", size(A))

    theta_hat = pinv(A)*y; # this is the linear regression step
    return theta_hat
end

function test_theta_hat_alternate(theta_hat)
    A = Matrix(DataFrame(CSV.File("mnist_data.csv")))
    y_true = Matrix(DataFrame(CSV.File("mnist_labels.csv")))
    # A is a 28*28 x 60000 matrix of data
    # y_true is a 60000 vector of labels
    N = length(y)

    # y_hat is the predicted values
    y_hat = A*theta_hat

    # Since y_hat is supposed to predict the digits 0-9, we'll round y_hat to the nearest integer
    y_hat = round.(y_hat)

    # the mean-squared error
    return sum((y_hat .- y_true).^2)./N
end

=# # delete this line if you need to uncomment the alternate method
