Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
stan_csv_reader.hpp
Go to the documentation of this file.
1 #ifndef __STAN__IO__STAN_CSV_READER_HPP__
2 #define __STAN__IO__STAN_CSV_READER_HPP__
3 
4 #include <istream>
5 #include <iostream>
6 #include <sstream>
7 #include <string>
8 #include <boost/algorithm/string.hpp>
9 #include <boost/lexical_cast.hpp>
10 #include <stan/math/matrix.hpp>
11 
12 namespace stan {
13  namespace io {
14 
15  // FIXME: should consolidate with the options from the command line in stan::gm
20 
21  std::string data;
22  std::string init;
25  size_t seed;
27  size_t chain_id;
28  size_t iter;
29  size_t warmup;
30  size_t thin;
35  double epsilon;
36  double epsilon_pm;
37  double delta;
38  double gamma;
39  };
40 
42  std::string sampler;
43  double step_size;
44  Eigen::MatrixXd step_size_multipliers;
45  };
46 
47  struct stan_csv {
49  Eigen::Matrix<std::string, Eigen::Dynamic, 1> header;
51  Eigen::MatrixXd samples;
52  };
53 
58  public:
64 
70 
71  static bool read_metadata(std::istream& in, stan_csv_metadata& metadata) {
72  std::stringstream ss;
73  std::string line;
74 
75  if (in.peek() != '#')
76  return false;
77  while (in.peek() == '#') {
78  std::getline(in, line);
79  ss << line << '\n';
80  }
81  ss.seekg(std::ios_base::beg);
82 
83  char comment;
84  std::string lhs;
85 
86  // skip first two lines
87  std::getline(ss, line);
88  std::getline(ss, line);
89 
90  while (ss.good()) {
91  ss >> comment;
92  std::getline(ss, lhs, '=');
93  boost::trim(lhs);
94  if (lhs.compare("") == 0) { // no-op
95  } else if (lhs.compare("stan_version_major") == 0) {
96  ss >> metadata.stan_version_major;
97  } else if (lhs.compare("stan_version_minor") == 0) {
98  ss >> metadata.stan_version_minor;
99  } else if (lhs.compare("stan_version_patch") == 0) {
100  ss >> metadata.stan_version_patch;
101  } else if (lhs.compare("data") == 0) {
102  ss >> metadata.data;
103  } else if (lhs.compare("init") == 0) {
104  std::getline(ss, metadata.init);
105  boost::trim(metadata.init);
106  ss.unget();
107  } else if (lhs.compare("append_samples") == 0) {
108  ss >> metadata.append_samples;
109  } else if (lhs.compare("save_warmup") == 0) {
110  ss >> metadata.save_warmup;
111  } else if (lhs.compare("seed") == 0) {
112  ss >> metadata.seed;
113  metadata.random_seed = false;
114  } else if (lhs.compare("chain_id") == 0) {
115  ss >> metadata.chain_id;
116  } else if (lhs.compare("iter") == 0) {
117  ss >> metadata.iter;
118  } else if (lhs.compare("warmup") == 0) {
119  ss >> metadata.warmup;
120  } else if (lhs.compare("thin") == 0) {
121  ss >> metadata.thin;
122  } else if (lhs.compare("equal_step_sizes") == 0) {
123  ss >> metadata.equal_step_sizes;
124  } else if (lhs.compare("nondiag_mass") == 0) {
125  ss >> metadata.nondiag_mass;
126  } else if (lhs.compare("leapfrog_steps") == 0) {
127  ss >> metadata.leapfrog_steps;
128  } else if (lhs.compare("max_treedepth") == 0) {
129  ss >> metadata.max_treedepth;
130  } else if (lhs.compare("epsilon") == 0) {
131  ss >> metadata.epsilon;
132  } else if (lhs.compare("epsilon_pm") == 0) {
133  ss >> metadata.epsilon_pm;
134  } else if (lhs.compare("delta") == 0) {
135  ss >> metadata.delta;
136  } else if (lhs.compare("gamma") == 0) {
137  ss >> metadata.gamma;
138  } else {
139  std::cout << "unused option: " << lhs << std::endl;
140  }
141  std::getline(ss, line);
142  }
143  if (ss.good() == true)
144  return false;
145  return true;
146  }
147 
148  static bool read_header(std::istream& in, Eigen::Matrix<std::string, Eigen::Dynamic, 1>& header) {
149  std::string line;
150 
151  if (in.peek() != 'l')
152  return false;
153  std::getline(in, line);
154  std::stringstream ss(line);
155 
156  header.resize(std::count(line.begin(), line.end(), ',') + 1);
157  int idx = 0;
158  while (ss.good()) {
159  std::string token;
160  std::getline(ss, token, ',');
161  boost::trim(token);
162 
163  int pos = token.find('.');
164  if (pos > 0) {
165  token.replace(pos, 1, "[");
166  std::replace(token.begin(), token.end(), '.', ',');
167  token += "]";
168  }
169  header(idx++) = token;
170  }
171  return true;
172  }
173 
174  static bool read_adaptation(std::istream& in, stan_csv_adaptation& adaptation) {
175  std::stringstream ss;
176  std::string line;
177  int lines = 0;
178 
179  if (in.peek() != '#' || in.good() == false)
180  return false;
181  while (in.peek() == '#') {
182  std::getline(in, line);
183  ss << line << '\n';
184  lines++;
185  }
186  ss.seekg(std::ios_base::beg);
187 
188  char comment;
189  // sampler
190  ss >> comment >> adaptation.sampler;
191  std::getline(ss, line);
192  // clean up sampler field
193  std::replace(adaptation.sampler.begin(),
194  adaptation.sampler.end(),
195  '(', ' ');
196  std::replace(adaptation.sampler.begin(),
197  adaptation.sampler.end(),
198  ')', ' ');
199  boost::trim(adaptation.sampler);
200 
201  // step size
202  ss >> comment;
203  std::getline(ss, line, '=');
204  boost::trim(line);
205  ss >> adaptation.step_size;
206  std::getline(ss, line);
207 
208  // parameter step size multipliers
209  std::getline(ss, line); // comment line
210  std::getline(ss, line); // step sizes
211  int rows = lines-3;
212  int cols = std::count(line.begin(), line.end(), ',') + 1;
213  adaptation.step_size_multipliers.resize(rows, cols);
214 
215  for (int row = 0; row < rows; row++) {
216  std::stringstream line_ss;
217  line_ss.str(line);
218  line_ss >> comment;
219  for (int col = 0; col < cols; col++) {
220  std::string token;
221  std::getline(line_ss, token, ',');
222  boost::trim(token);
223  adaptation.step_size_multipliers(row,col) = boost::lexical_cast<double>(token);
224  }
225  std::getline(ss, line); // step sizes
226  }
227  if (ss.good())
228  return false;
229  return true;
230  }
231 
232  static bool read_samples(std::istream& in, Eigen::MatrixXd& samples) {
233  std::stringstream ss;
234  std::string line;
235 
236  int rows = 0;
237  int cols = -1;
238 
239  if (in.peek() == '#' || in.good() == false)
240  return false;
241 
242  while (in.good()) {
243  bool comment_line = (in.peek() == '#');
244  std::getline(in, line);
245  if (!comment_line) {
246  ss << line << '\n';
247  int current_cols = std::count(line.begin(), line.end(), ',') + 1;
248  if (cols == -1) {
249  cols = current_cols;
250  } else if (cols != current_cols) {
251  std::cout << "Error: expected " << cols << " columns, but found "
252  << current_cols << " instead for row " << rows+1 << std::endl;
253  return false;
254  }
255  rows++;
256  }
257  in.peek();
258  }
259  ss.seekg(std::ios_base::beg);
260 
261  if (rows > 0) {
262  samples.resize(rows, cols);
263  char comma;
264  for (int row = 0; row < rows; row++) {
265  for (int col = 0; col < cols; col++) {
266  ss >> samples(row,col);
267  if (col != cols-1)
268  ss >> comma;
269  }
270  }
271  }
272  return true;
273  }
274 
279  static stan_csv parse(std::istream& in) {
280  stan_csv data;
281  if (!read_metadata(in, data.metadata)) {
282  std::cout << "Warning: non-fatal error reading metadata" << std::endl;
283  }
284  if (!read_header(in, data.header)) {
285  std::cout << "Error: error reading header" << std::endl;
286  throw std::invalid_argument("Error with header of input file in parse");
287  }
288  if (!read_adaptation(in, data.adaptation)) {
289  std::cout << "Warning: non-fatal error reading adapation data" << std::endl;
290  }
291  if (!read_samples(in, data.samples)) {
292  std::cout << "Warning: non-fatal error reading samples" << std::endl;
293  }
294  return data;
295  }
296 
297  };
298 
299  }
300 }
301 
302 #endif

     [ Stan Home Page ] © 2011–2013, Stan Development Team.