| #ifndef NEUROFLOW_SFT_HPP
|
| #define NEUROFLOW_SFT_HPP
|
|
|
| #include <cstddef>
|
| #include <random>
|
| #include <string>
|
| #include <vector>
|
|
|
| #include "adamw.hpp"
|
| #include "alignment_common.hpp"
|
| #include "causal_lm.hpp"
|
| #include "scheduler.hpp"
|
| #include "tokenizer.hpp"
|
|
|
| namespace neuroflow {
|
|
|
| struct SFTTrainConfig {
|
| std::string data_path;
|
| std::string ckpt_path;
|
| std::string tokenizer_path;
|
| std::string output_dir;
|
| float learning_rate = 2e-5f;
|
| int epochs = 3;
|
| size_t max_seq_len = 512;
|
| float warmup_ratio = 0.05f;
|
| float weight_decay = 0.01f;
|
| float grad_clip = 1.0f;
|
| float adam_beta1 = 0.9f;
|
| float adam_beta2 = 0.999f;
|
| float adam_eps = 1e-8f;
|
| size_t save_interval = 1000;
|
| size_t log_interval = 10;
|
| unsigned seed = 42;
|
| };
|
|
|
| class SFTDataLoader {
|
| public:
|
| SFTDataLoader(const std::string& jsonl_path, size_t max_samples = 0);
|
|
|
| bool has_next() const;
|
| SFTSample next();
|
| void reset();
|
| void shuffle(std::mt19937& rng);
|
| size_t total_samples() const { return samples_.size(); }
|
| size_t invalid_count() const { return invalid_count_; }
|
|
|
| private:
|
| std::vector<SFTSample> samples_;
|
| size_t cursor_ = 0;
|
| size_t invalid_count_ = 0;
|
| };
|
|
|
| struct MaskedCEOutput {
|
| float loss;
|
| Tensor logits_grad;
|
| size_t valid_token_count;
|
| };
|
|
|
| MaskedCEOutput masked_cross_entropy(const Tensor& logits, size_t vocab_size,
|
| const std::vector<size_t>& target_ids,
|
| const std::vector<float>& loss_mask);
|
|
|
| SFTTrainingTensors build_sft_training_tensors(const SFTSample& sample,
|
| BPETokenizer& tokenizer,
|
| size_t max_seq_len);
|
|
|
| class SFTTrainer {
|
| public:
|
| SFTTrainConfig config;
|
|
|
| SFTTrainer(const SFTTrainConfig& cfg);
|
|
|
| void train();
|
| float train_on_sample(const SFTSample& sample);
|
|
|
| private:
|
| std::unique_ptr<CausalLMHead> model_;
|
| std::unique_ptr<BPETokenizer> tokenizer_;
|
| std::unique_ptr<AdamW> optimizer_;
|
| std::unique_ptr<CosineScheduler> scheduler_;
|
| };
|
|
|
| }
|
|
|
| #endif |