COURSE / SOURCE

submission.h

All lessons
Source filecode/day100-capstone-5/submission.h

This is the source used by the lesson and its recorded evidence. Compile commands and expected output live in the directory README.

// SPDX-License-Identifier: MIT
#pragma once

#include <cstddef>
#include <cstdint>

#include <cuda_runtime.h>

namespace day100 {

struct Dataset {
    const float* train_images;
    const std::int32_t* train_labels;
    const float* test_images;
    const std::int32_t* test_labels;
    int train_count;
    int test_count;
};

struct Model {
    float* w1;
    float* b1;
    float* w2;
    float* b2;
    float* velocity;
};

void solve_forward(const float* x, const float* w1, const float* b1,
                   const float* w2, const float* b2, float* hidden,
                   float* logits, int batch, cudaStream_t stream);

void solve_loss(const float* logits, const std::int32_t* labels,
                float* loss_out, float* dlogits, int batch,
                cudaStream_t stream);

void solve_backward(const float* x, const float* hidden,
                    const float* dlogits, const float* w1, const float* w2,
                    float* dw1, float* db1, float* dw2, float* db2, float* dx,
                    int batch, cudaStream_t stream);

void solve_step(float* weights, const float* gradients, float* velocity,
                float learning_rate, float momentum, std::size_t count,
                cudaStream_t stream);

void solve_train(const Dataset& data, Model& model, int epochs,
                 cudaStream_t stream);

}  // namespace day100