Refactored AddTask definition for LSTM training

This commit is contained in:
Konstantin Sidorov
2017-07-31 12:36:22 +03:00
parent ff93ab7d7a
commit 9ca127a402
3 changed files with 47 additions and 30 deletions
@@ -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
+16 -18
View File
@@ -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;
}