Refactored AddTask definition for LSTM training
This commit is contained in:
@@ -36,13 +36,15 @@ class AddTask
|
||||
* @param labels The variable to store output sequences.
|
||||
* @param batchSize The dataset size.
|
||||
*/
|
||||
void Generate(arma::field<arma::colvec>& input,
|
||||
arma::field<arma::colvec>& labels,
|
||||
void Generate(arma::field<arma::mat>& input,
|
||||
arma::field<arma::mat>& labels,
|
||||
const size_t batchSize);
|
||||
|
||||
private:
|
||||
// Maximum binary length of numbers.
|
||||
size_t bitLen;
|
||||
|
||||
arma::field<arma::mat> Binarize(arma::field<arma::vec> data);
|
||||
};
|
||||
} // namespace tasks
|
||||
} // namespace augmented
|
||||
|
||||
@@ -37,12 +37,12 @@ AddTask::AddTask(const size_t bitLen) : bitLen(bitLen) {
|
||||
assert(bitLen > 0);
|
||||
}
|
||||
|
||||
void AddTask::Generate(arma::field<arma::colvec>& input,
|
||||
arma::field<arma::colvec>& labels,
|
||||
void AddTask::Generate(arma::field<arma::mat>& input,
|
||||
arma::field<arma::mat>& labels,
|
||||
const size_t batchSize)
|
||||
{
|
||||
input = arma::field<arma::colvec>(batchSize);
|
||||
labels = arma::field<arma::colvec>(batchSize);
|
||||
arma::field<arma::vec> vecInput = arma::field<arma::colvec>(batchSize);
|
||||
arma::field<arma::vec> vecLabels = arma::field<arma::colvec>(batchSize);
|
||||
for (size_t i = 0; i < batchSize; ++i) {
|
||||
// Random uniform length from [2..bitLen]
|
||||
size_t size_A = RandInt(2, bitLen + 1);
|
||||
@@ -50,18 +50,18 @@ void AddTask::Generate(arma::field<arma::colvec>& input,
|
||||
// Construct sequence of the form
|
||||
// (binary number with size_A bits) + '+'
|
||||
// + (binary number with size_B bits)
|
||||
input(i) = arma::randi<arma::colvec>(
|
||||
vecInput(i) = arma::randi<arma::colvec>(
|
||||
size_A + size_B + 1, arma::distr_param(0, 1));
|
||||
input(i).at(size_A) = 2; // special value for '+' delimiter
|
||||
vecInput(i).at(size_A) = 2; // special value for '+' delimiter
|
||||
int val_A = 0;
|
||||
for (size_t k = 0; k < size_A; ++k) {
|
||||
val_A <<= 1;
|
||||
val_A += input(i).at(k);
|
||||
val_A += vecInput(i).at(k);
|
||||
}
|
||||
int val_B = 0;
|
||||
for (size_t k = size_A+1; k < size_A+1+size_B; ++k) {
|
||||
val_B <<= 1;
|
||||
val_B += input(i).at(k);
|
||||
val_B += vecInput(i).at(k);
|
||||
}
|
||||
int tot = val_A + val_B;
|
||||
vector<int> binary_seq;
|
||||
@@ -70,13 +70,30 @@ void AddTask::Generate(arma::field<arma::colvec>& input,
|
||||
tot >>= 1;
|
||||
}
|
||||
size_t tot_len = binary_seq.size();
|
||||
labels(i) = arma::colvec(tot_len);
|
||||
vecLabels(i) = arma::colvec(tot_len);
|
||||
for (size_t j = 0; j < tot_len; ++j) {
|
||||
labels(i).at(j) = binary_seq[tot_len-j-1];
|
||||
vecLabels(i).at(j) = binary_seq[tot_len-j-1];
|
||||
}
|
||||
}
|
||||
input = Binarize(vecInput);
|
||||
labels = Binarize(vecLabels);
|
||||
}
|
||||
|
||||
arma::field<arma::mat> AddTask::Binarize(arma::field<arma::vec> data)
|
||||
{
|
||||
arma::field<arma::mat> procData(data.n_elem);
|
||||
for (size_t i = 0; i < data.n_elem; ++i) {
|
||||
arma::mat temp = arma::zeros(3, data.at(i).n_elem);
|
||||
for (size_t j = 0; j < data.at(i).n_elem; ++j) {
|
||||
int val = data.at(i).at(j);
|
||||
temp.at(val, j) = 1;
|
||||
}
|
||||
procData.at(i) = temp;
|
||||
}
|
||||
return procData;
|
||||
}
|
||||
|
||||
|
||||
} // namespace tasks
|
||||
} // namespace augmented
|
||||
} // namespace ann
|
||||
|
||||
@@ -60,7 +60,6 @@ class HardCodedCopyModel {
|
||||
}
|
||||
assert(oneCnt % zeroCnt == 0);
|
||||
nRepeats = oneCnt / zeroCnt;
|
||||
std::cerr << "Repeats: " << nRepeats << "\n";
|
||||
}
|
||||
void Predict(
|
||||
arma::mat& predictors,
|
||||
@@ -72,8 +71,6 @@ class HardCodedCopyModel {
|
||||
for (size_t i = 0; i < outputLen; ++i) {
|
||||
labels.at(seqLen+i) = predictors.at(2 * (i % seqLen));
|
||||
}
|
||||
std::cerr << "Predictors:\n" << predictors.t();
|
||||
std::cerr << "Labels:\n" << labels.t();
|
||||
}
|
||||
void Predict(
|
||||
arma::field<arma::mat>& predictors,
|
||||
@@ -134,19 +131,20 @@ class HardCodedSortModel {
|
||||
class HardCodedAddModel {
|
||||
public:
|
||||
HardCodedAddModel() {}
|
||||
void Train(arma::field<arma::colvec>& predictors,
|
||||
arma::field<arma::colvec>& labels)
|
||||
void Train(arma::field<arma::mat>& predictors,
|
||||
arma::field<arma::mat>& labels)
|
||||
{
|
||||
return;
|
||||
}
|
||||
void Predict(arma::colvec& predictors,
|
||||
arma::colvec& labels)
|
||||
void Predict(arma::mat& predictors,
|
||||
arma::mat& labels)
|
||||
{
|
||||
assert(predictors.n_rows == 3);
|
||||
int num_A = 0, num_B = 0;
|
||||
bool num = false; // true iff we have already seen the separating symbol
|
||||
auto len = predictors.n_elem;
|
||||
size_t len = predictors.n_cols;
|
||||
for (size_t i = 0; i < len; ++i) {
|
||||
auto digit = predictors.at(i);
|
||||
auto digit = arma::as_scalar(arma::find(1 == predictors.col(i), 1));
|
||||
if (digit != 0 && digit != 1)
|
||||
{
|
||||
// We should not see two separators
|
||||
@@ -175,16 +173,16 @@ class HardCodedAddModel {
|
||||
total >>= 1;
|
||||
}
|
||||
auto tot_len = binary_seq.size();
|
||||
labels = arma::colvec(tot_len);
|
||||
labels = arma::zeros(3, tot_len);
|
||||
for (size_t j = 0; j < tot_len; ++j) {
|
||||
labels.at(j) = binary_seq[tot_len-j-1];
|
||||
labels.at(binary_seq[tot_len-j-1], j) = 1;
|
||||
}
|
||||
}
|
||||
void Predict(
|
||||
arma::field<arma::colvec>& predictors,
|
||||
arma::field<arma::colvec>& labels) {
|
||||
arma::field<arma::mat>& predictors,
|
||||
arma::field<arma::mat>& labels) {
|
||||
auto sz = predictors.n_elem;
|
||||
labels = arma::field<arma::colvec>(sz);
|
||||
labels = arma::field<arma::mat>(sz);
|
||||
for (size_t i = 0; i < sz; ++i) {
|
||||
Predict(predictors.at(i), labels.at(i));
|
||||
}
|
||||
@@ -251,16 +249,16 @@ BOOST_AUTO_TEST_CASE(AddTaskTest) {
|
||||
bool ok = true;
|
||||
for (size_t bitLen = 2; bitLen <= 16; ++bitLen) {
|
||||
AddTask task(bitLen);
|
||||
arma::field<arma::colvec> trainPredictor, trainResponse;
|
||||
arma::field<arma::mat> trainPredictor, trainResponse;
|
||||
task.Generate(trainPredictor, trainResponse, 8);
|
||||
arma::field<arma::colvec> testPredictor, testResponse;
|
||||
arma::field<arma::mat> testPredictor, testResponse;
|
||||
task.Generate(testPredictor, testResponse, 8);
|
||||
HardCodedAddModel model;
|
||||
model.Train(trainPredictor, trainResponse);
|
||||
arma::field<arma::colvec> predResponse;
|
||||
arma::field<arma::mat> predResponse;
|
||||
model.Predict(testPredictor, predResponse);
|
||||
// A single failure is a failure.
|
||||
if (SequencePrecision<arma::colvec>(testResponse, predResponse) < 0.99) {
|
||||
if (SequencePrecision<arma::mat>(testResponse, predResponse) < 0.99) {
|
||||
ok = false;
|
||||
break;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user