-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathMarkovTextPredictor.cpp
More file actions
135 lines (113 loc) · 3.95 KB
/
Copy pathMarkovTextPredictor.cpp
File metadata and controls
135 lines (113 loc) · 3.95 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
#include <fstream>
#include <future>
#include <iostream>
#include <optional>
#include <random>
#include <string>
#include <unordered_map>
#include <vector>
// Random generator variables
std::random_device rd; // Non-deterministic seed
std::mt19937 gen(rd()); // Mersenne Twister engine
// For atomic console output
static std::mutex coutMutex;
void printAtomicPredictorStart(int size) {
std::lock_guard<std::mutex> lock(coutMutex);
std::cout << "Predictor(size=" << size << "): Initializing..." << std::endl;
}
void printAtomicPredictorDone(int size) {
std::lock_guard<std::mutex> lock(coutMutex);
std::cout << "Predictor(size=" << size << "): Done." << std::endl;
}
class Predictor {
std::unordered_map<std::string, std::vector<char>> model_;
int size_;
int hits_ = 0;
std::future<void> fut;
public:
Predictor(const std::string& text, int size) : size_(size), hits_(0) {
auto buildModel = [this, &text, size]() {
printAtomicPredictorStart(size);
for (size_t i = 0; i + size < text.size(); ++i) {
std::string key = text.substr(i, size);
char next_char = text[i + size];
this->model_[key].push_back(next_char);
}
printAtomicPredictorDone(size);
};
fut = std::async(std::launch::async, buildModel);
}
std::optional<char> predictNextChar(const std::string& context) {
fut.wait(); // Ensure the model is built
// If the context is shorter than the model size, return nullopt
if (size_ > context.size()) return std::nullopt;
// Find the context in the model
std::string key = context.substr(context.size() - size_);
auto it = model_.find(key);
if (it == model_.end())
return std::nullopt;
// Get a random character from the vector
std::uniform_int_distribution<size_t> dist(0, it->second.size()-1);
size_t index = dist(gen);
hits_++;
return it->second[index];
}
int getHits() const { return hits_; }
std::future<void>& get_future() { return fut; }
};
class MarkovTextPredictor {
int max_length_;
std::vector<Predictor> predictors_; // Assuming Predictor is a defined class
public:
MarkovTextPredictor(const std::string& text, int max_length=15)
{
predictors_.reserve(max_length + 1);
for (int size = 0; size <= max_length; ++size) {
// Initialize each predictor with the text and size
predictors_.emplace_back(text, size);
}
// Wait for all predictors to finish building their models
for (auto& predictor : predictors_) {
predictor.get_future().wait();
}
}
void printStats() const
{
for (int size = 0; size < predictors_.size(); ++size) {
std::cout << "Predictor size " << size << " hits: " << predictors_[size].getHits() << std::endl;
}
}
char predictNextChar(const std::string& context)
{
// Iterate through predictors from largest to smallest
for (size_t size = predictors_.size()-1; ; --size) {
auto prediction = predictors_[size].predictNextChar(context);
if (prediction.has_value()) {
return prediction.value();
}
}
}
};
int main(int argc, char *argv[]) // Fix parameter names and types
{
if (argc < 2) {
std::cout << "Usage: MarkovTextPredictor <input_file> [<prompt>]\n";
return 1;
}
std::ifstream file(argv[1]);
if (!file) { // Check if file opened successfully
std::cout << "Error: Could not open file " << argv[1] << "\n";
return 1;
}
std::string text((std::istreambuf_iterator<char>(file)), std::istreambuf_iterator<char>());
MarkovTextPredictor predictor(text);
std::string prompt = argc > 2 ? argv[2] : "";
std::string output = prompt;
for (int i= 0; i < 1000; ++i) {
output += predictor.predictNextChar(output);
}
std::replace(output.begin(), output.end(), '\n', ' ');
std::cout << output << std::endl;
predictor.printStats();
return 0;
}