Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
statement_grammar_def.hpp
Go to the documentation of this file.
1 #ifndef __STAN__GM__PARSER__STATEMENT_GRAMMAR_DEF__HPP__
2 #define __STAN__GM__PARSER__STATEMENT_GRAMMAR_DEF__HPP__
3 
4 #include <cstddef>
5 #include <iomanip>
6 #include <iostream>
7 #include <istream>
8 #include <map>
9 #include <set>
10 #include <sstream>
11 #include <string>
12 #include <utility>
13 #include <vector>
14 #include <stdexcept>
15 
16 #include <boost/spirit/include/qi.hpp>
17 // FIXME: get rid of unused include
18 #include <boost/spirit/include/phoenix_core.hpp>
19 #include <boost/spirit/include/phoenix_function.hpp>
20 #include <boost/spirit/include/phoenix_fusion.hpp>
21 #include <boost/spirit/include/phoenix_object.hpp>
22 #include <boost/spirit/include/phoenix_operator.hpp>
23 #include <boost/spirit/include/phoenix_stl.hpp>
24 
25 #include <boost/lexical_cast.hpp>
26 #include <boost/fusion/include/adapt_struct.hpp>
27 #include <boost/fusion/include/std_pair.hpp>
28 #include <boost/config/warning_disable.hpp>
29 #include <boost/spirit/include/qi.hpp>
30 #include <boost/spirit/include/qi_numeric.hpp>
31 #include <boost/spirit/include/classic_position_iterator.hpp>
32 #include <boost/spirit/include/phoenix_core.hpp>
33 #include <boost/spirit/include/phoenix_function.hpp>
34 #include <boost/spirit/include/phoenix_fusion.hpp>
35 #include <boost/spirit/include/phoenix_object.hpp>
36 #include <boost/spirit/include/phoenix_operator.hpp>
37 #include <boost/spirit/include/phoenix_stl.hpp>
38 #include <boost/spirit/include/support_multi_pass.hpp>
39 #include <boost/tuple/tuple.hpp>
40 #include <boost/variant/apply_visitor.hpp>
41 #include <boost/variant/recursive_variant.hpp>
42 
43 #include <stan/gm/ast.hpp>
49 
51  (stan::gm::variable_dims, var_dims_)
52  (stan::gm::expression, expr_) )
53 
54 BOOST_FUSION_ADAPT_STRUCT(stan::gm::variable_dims,
55  (std::string, name_)
56  (std::vector<stan::gm::expression>, dims_) )
57 
58 BOOST_FUSION_ADAPT_STRUCT(stan::gm::distribution,
59  (std::string, family_)
60  (std::vector<stan::gm::expression>, args_) )
61 
62 BOOST_FUSION_ADAPT_STRUCT(stan::gm::for_statement,
63  (std::string, variable_)
64  (stan::gm::range, range_)
65  (stan::gm::statement, statement_) )
66 
67 BOOST_FUSION_ADAPT_STRUCT(stan::gm::print_statement,
68  (std::vector<stan::gm::printable>, printables_) )
69 
70 BOOST_FUSION_ADAPT_STRUCT(stan::gm::sample,
71  (stan::gm::expression, expr_)
72  (stan::gm::distribution, dist_)
73  (stan::gm::range, truncation_) )
74 
75 BOOST_FUSION_ADAPT_STRUCT(stan::gm::statements,
76  (std::vector<stan::gm::var_decl>, local_decl_)
77  (std::vector<stan::gm::statement>, statements_) )
78 
79 namespace stan {
80 
81  namespace gm {
82 
83  struct validate_assignment {
84  template <typename T1, typename T2, typename T3, typename T4>
85  struct result { typedef bool type; };
86 
87  bool operator()(assignment& a,
88  const var_origin& origin_allowed,
89  variable_map& vm,
90  std::stringstream& error_msgs) const {
91 
92  // validate existence
93  std::string name = a.var_dims_.name_;
94  if (!vm.exists(name)) {
95  error_msgs << "unknown variable in assignment"
96  << "; lhs variable=" << a.var_dims_.name_
97  << std::endl;
98  return false;
99  }
100 
101  // validate origin
102  var_origin lhs_origin = vm.get_origin(name);
103  if (lhs_origin != local_origin
104  && lhs_origin != origin_allowed) {
105  error_msgs << "attempt to assign variable in wrong block."
106  << " left-hand-side variable origin=";
107  print_var_origin(error_msgs,lhs_origin);
108  error_msgs << std::endl;
109  return false;
110  }
111 
112  // validate types
113  a.var_type_ = vm.get(name);
114  size_t lhs_var_num_dims = a.var_type_.dims_.size();
115  size_t num_index_dims = a.var_dims_.dims_.size();
116 
117  expr_type lhs_type = infer_type_indexing(a.var_type_.base_type_,
118  lhs_var_num_dims,
119  num_index_dims);
120 
121  if (lhs_type.is_ill_formed()) {
122  error_msgs << "too many indexes for variable "
123  << "; variable name = " << name
124  << "; num dimensions given = " << num_index_dims
125  << "; variable array dimensions = " << lhs_var_num_dims;
126  return false;
127  }
128  if (lhs_type.num_dims_ != a.expr_.expression_type().num_dims_) {
129  error_msgs << "mismatched dimensions on left- and right-hand side of assignment"
130  << "; left dims=" << lhs_type.num_dims_
131  << "; right dims=" << a.expr_.expression_type().num_dims_
132  << std::endl;
133  return false;
134  }
135 
136  base_expr_type lhs_base_type = lhs_type.base_type_;
137  base_expr_type rhs_base_type = a.expr_.expression_type().base_type_;
138  // int -> double promotion
139  bool types_compatible
140  = lhs_base_type == rhs_base_type
141  || ( lhs_base_type == DOUBLE_T && rhs_base_type == INT_T );
142  if (!types_compatible) {
143  error_msgs << "base type mismatch in assignment"
144  << "; left variable=" << a.var_dims_.name_
145  << "; left base type=";
146  write_base_expr_type(error_msgs,lhs_base_type);
147  error_msgs << "; right base type=";
148  write_base_expr_type(error_msgs,rhs_base_type);
149  error_msgs << std::endl;
150  return false;
151  }
152  return true;
153  }
154  };
155  boost::phoenix::function<validate_assignment> validate_assignment_f;
156 
157  struct validate_sample {
158  template <typename T1, typename T2>
159  struct result { typedef bool type; };
160 
161  bool is_double_return(const std::string& function_name,
162  const std::vector<expr_type>& arg_types,
163  std::ostream& error_msgs) const {
164  return function_signatures::instance()
165  .get_result_type(function_name,arg_types,error_msgs)
166  .is_primitive_double();
167  }
168  bool operator()(const sample& s,
169  std::ostream& error_msgs) const {
170  std::vector<expr_type> arg_types;
171  arg_types.push_back(s.expr_.expression_type());
172  for (size_t i = 0; i < s.dist_.args_.size(); ++i)
173  arg_types.push_back(s.dist_.args_[i].expression_type());
174  std::string function_name(s.dist_.family_);
175  function_name += "_log";
176  // expr_type result_type
177  // = function_signatures::instance()
178  // .get_result_type(function_name,arg_types,error_msgs);
179  // if (!result_type.is_primitive_double()) {
180  if (!is_double_return(function_name,arg_types,error_msgs)) {
181  error_msgs << "unknown distribution=" << s.dist_.family_ << std::endl;
182  return false;
183  }
184  if (s.truncation_.has_low()) {
185  std::vector<expr_type> arg_types_trunc(arg_types);
186  arg_types_trunc[0] = s.truncation_.low_.expression_type();
187  std::string function_name_cdf(s.dist_.family_);
188  function_name_cdf += "_cdf";
189  if (!is_double_return(function_name_cdf,arg_types_trunc,error_msgs)) {
190  error_msgs << "lower truncation not defined for specified arguments to "
191  << s.dist_.family_ << std::endl;
192  return false;
193  }
194  if (!is_double_return(function_name_cdf,arg_types,error_msgs)) {
195  error_msgs << "lower bound in truncation type does not match"
196  << " sampled variate in distribution's type"
197  << std::endl;
198  return false;
199  }
200  }
201  if (s.truncation_.has_high()) {
202  std::vector<expr_type> arg_types_trunc(arg_types);
203  arg_types_trunc[0] = s.truncation_.high_.expression_type();
204  std::string function_name_cdf(s.dist_.family_);
205  function_name_cdf += "_cdf";
206  if (!is_double_return(function_name_cdf,arg_types_trunc,error_msgs)) {
207  error_msgs << "upper truncation not defined for specified arguments to "
208  << s.dist_.family_ << std::endl;
209  return false;
210  }
211  if (!is_double_return(function_name_cdf,arg_types,error_msgs)) {
212  error_msgs << "upper bound in truncation type does not match"
213  << " sampled variate in distribution's type"
214  << std::endl;
215  return false;
216  }
217  }
218  return true;
219 
220  }
221  };
222  boost::phoenix::function<validate_sample> validate_sample_f;
223 
224  struct unscope_locals {
225  template <typename T1, typename T2>
226  struct result { typedef void type; };
227  void operator()(const std::vector<var_decl>& var_decls,
228  variable_map& vm) const {
229  for (size_t i = 0; i < var_decls.size(); ++i)
230  vm.remove(var_decls[i].name());
231  }
232  };
233  boost::phoenix::function<unscope_locals> unscope_locals_f;
234 
235  // struct add_conditional_condition {
236  // template <typename T1, typename T2, typename T3>
237  // struct result { typedef bool type; };
238  // bool operator()(conditional_statement& cs,
239  // const expression& e,
240  // std::stringstream& error_msgs) const {
241  // if (!e.expression_type().is_primitive()) {
242  // error_msgs << "conditions in if-else statement must be primitive int or real;"
243  // << " found type=" << e.expression_type() << std::endl;
244  // return false;
245  // }
246  // cs.conditions_.push_back(e);
247  // return true;
248  // }
249  // };
250  // boost::phoenix::function<add_conditional_condition> add_conditional_condition_f;
251 
252  // struct add_conditional_body {
253  // template <typename T1, typename T2>
254  // struct result { typedef void type; };
255  // void operator()(conditional_statement& cs,
256  // const statement& s) const {
257  // cs.bodies_.push_back(s);
258  // }
259  // };
260  // boost::phoenix::function<add_conditional_body> add_conditional_body_f;
261 
262  struct add_while_condition {
263  template <typename T1, typename T2, typename T3>
264  struct result { typedef bool type; };
265  bool operator()(while_statement& ws,
266  const expression& e,
267  std::stringstream& error_msgs) const {
268  if (!e.expression_type().is_primitive()) {
269  error_msgs << "conditions in while statement must be primitive int or real;"
270  << " found type=" << e.expression_type() << std::endl;
271  return false;
272  }
273  ws.condition_ = e;
274  return true;
275  }
276  };
277  boost::phoenix::function<add_while_condition> add_while_condition_f;
278 
279  struct add_while_body {
280  template <typename T1, typename T2>
281  struct result { typedef void type; };
282  void operator()(while_statement& ws,
283  const statement& s) const {
284  ws.body_ = s;
285  }
286  };
287  boost::phoenix::function<add_while_body> add_while_body_f;
288 
289  struct add_loop_identifier {
290  template <typename T1, typename T2, typename T3, typename T4>
291  struct result { typedef bool type; };
292  bool operator()(const std::string& name,
293  std::string& name_local,
294  variable_map& vm,
295  std::stringstream& error_msgs) const {
296  name_local = name;
297  if (vm.exists(name)) {
298  error_msgs << "ERROR: loop variable already declared."
299  << " variable name=\"" << name << "\"" << std::endl;
300  return false; // variable exists
301  }
302  vm.add(name,
303  base_var_decl(name,std::vector<expression>(),
304  INT_T),
305  local_origin); // loop var acts like local
306  return true;
307  }
308  };
309  boost::phoenix::function<add_loop_identifier> add_loop_identifier_f;
310 
311  struct remove_loop_identifier {
312  template <typename T1, typename T2>
313  struct result { typedef void type; };
314  void operator()(const std::string& name,
315  variable_map& vm) const {
316  vm.remove(name);
317  }
318  };
319  boost::phoenix::function<remove_loop_identifier> remove_loop_identifier_f;
320 
321  struct validate_int_expr2 {
322  template <typename T1, typename T2>
323  struct result { typedef bool type; };
324 
325  bool operator()(const expression& expr,
326  std::stringstream& error_msgs) const {
327  if (!expr.expression_type().is_primitive_int()) {
328  error_msgs << "expression denoting integer required; found type="
329  << expr.expression_type() << std::endl;
330  return false;
331  }
332  return true;
333  }
334  };
335  boost::phoenix::function<validate_int_expr2> validate_int_expr2_f;
336 
337  struct validate_allow_sample {
338  template <typename T1, typename T2>
339  struct result { typedef bool type; };
340 
341  bool operator()(const bool& allow_sample,
342  std::stringstream& error_msgs) const {
343  if (!allow_sample) {
344  error_msgs << "ERROR: sampling only allowed in model."
345  << std::endl;
346  return false;
347  }
348  return true;
349  }
350  };
351  boost::phoenix::function<validate_allow_sample> validate_allow_sample_f;
352 
353 
354  template <typename Iterator>
356  std::stringstream& error_msgs)
357  : statement_grammar::base_type(statement_r),
358  var_map_(var_map),
359  error_msgs_(error_msgs),
360  expression_g(var_map,error_msgs),
361  var_decls_g(var_map,error_msgs),
362  statement_2_g(var_map,error_msgs,*this)
363  {
364  using boost::spirit::qi::_1;
365  using boost::spirit::qi::char_;
366  using boost::spirit::qi::eps;
367  using boost::spirit::qi::lexeme;
368  using boost::spirit::qi::lit;
369  using boost::spirit::qi::_pass;
370  using boost::spirit::qi::_val;
371 
372  using boost::spirit::qi::labels::_a;
373  using boost::spirit::qi::labels::_r1;
374  using boost::spirit::qi::labels::_r2;
375 
376  // _r1 true if sample_r allowed (inherited)
377  // _r2 source of variables allowed for assignments
378  // set to true if sample_r are allowed
379  statement_r.name("statement");
380  statement_r
381  %= statement_seq_r(_r1,_r2)
382  | for_statement_r(_r1,_r2)
383  | while_statement_r(_r1,_r2)
384  | statement_2_g(_r1,_r2)
385  | print_statement_r(_r2)
386  | assignment_r(_r2)
387  [_pass
388  = validate_assignment_f(_1,_r2,boost::phoenix::ref(var_map_),
389  boost::phoenix::ref(error_msgs_))]
390  | sample_r(_r1,_r2) [_pass = validate_sample_f(_1,
391  boost::phoenix::ref(error_msgs_))]
392  | no_op_statement_r
393  ;
394 
395  // _r1, _r2 same as statement_r
396  statement_seq_r.name("sequence of statements");
397  statement_seq_r
398  %= lit('{')
399  > local_var_decls_r[_a = _1]
400  > *statement_r(_r1,_r2)
401  > lit('}')
402  > eps[unscope_locals_f(_a,boost::phoenix::ref(var_map_))]
403  ;
404 
405  local_var_decls_r
406  %= var_decls_g(false,local_origin); // - constants
407 
408  while_statement_r.name("while statement");
409  while_statement_r
410  = lit("while")
411  > lit('(')
412  > expression_g(_r2)
413  [_pass = add_while_condition_f(_val,_1,
414  boost::phoenix::ref(error_msgs_))]
415  > lit(')')
416  > statement_r(_r1,_r2)
417  [add_while_body_f(_val,_1)]
418  ;
419 
420 
421  // _r1, _r2 same as statement_r
422  for_statement_r.name("for statement");
423  for_statement_r
424  %= lit("for")
425  > lit('(')
426  > identifier_r [_pass
427  = add_loop_identifier_f(_1,_a,
428  boost::phoenix::ref(var_map_),
429  boost::phoenix::ref(error_msgs_))]
430  > lit("in")
431  > range_r(_r2)
432  > lit(')')
433  > statement_r(_r1,_r2)
434  > eps
435  [remove_loop_identifier_f(_a,boost::phoenix::ref(var_map_))];
436  ;
437 
438  print_statement_r.name("print statement");
439  print_statement_r
440  %= lit("print")
441  > lit('(')
442  > (printable_r(_r1) % ',')
443  > lit(')');
444 
445  printable_r.name("printable");
446  printable_r
447  %= printable_string_r
448  | expression_g(_r1);
449 
450  printable_string_r.name("printable quoted string");
451  printable_string_r
452  %= lit('"')
453  > lexeme[*char_("a-zA-Z0-9/~!@#$%^&*()`_+-={}|[]:;'<>?,./ ")]
454  > lit('"');
455 
456  identifier_r.name("identifier");
457  identifier_r
458  %= (lexeme[char_("a-zA-Z")
459  >> *char_("a-zA-Z0-9_.")]);
460 
461  range_r.name("range expression pair, colon");
462  range_r
463  %= expression_g(_r1)
464  [_pass = validate_int_expr2_f(_1,boost::phoenix::ref(error_msgs_))]
465  >> lit(':')
466  >> expression_g(_r1)
467  [_pass = validate_int_expr2_f(_1,boost::phoenix::ref(error_msgs_))];
468 
469  assignment_r.name("variable assignment by expression");
470  assignment_r
471  %= var_lhs_r(_r1)
472  >> lit("<-")
473  > expression_g(_r1)
474  > lit(';')
475  ;
476 
477  var_lhs_r.name("variable and array dimensions");
478  var_lhs_r
479  %= identifier_r
480  >> opt_dims_r(_r1);
481 
482  opt_dims_r.name("array dimensions (optional)");
483  opt_dims_r
484  %= - dims_r(_r1);
485 
486  dims_r.name("array dimensions");
487  dims_r
488  %= lit('[')
489  > (expression_g(_r1)
490  [_pass = validate_int_expr2_f(_1,boost::phoenix::ref(error_msgs_))]
491  % ',')
492  > lit(']')
493  ;
494 
495  // inherited _r1 = true if samples allowed as statements
496  sample_r.name("distribution of expression");
497  sample_r
498  %= expression_g(_r2)
499  >> lit('~')
500  > eps
501  [_pass
502  = validate_allow_sample_f(_r1,boost::phoenix::ref(error_msgs_))]
503  > distribution_r(_r2)
504  > -truncation_range_r(_r2)
505  > lit(';');
506 
507  distribution_r.name("distribution and parameters");
508  distribution_r
509  %= identifier_r
510  >> lit('(')
511  >> -(expression_g(_r1) % ',')
512  > lit(')');
513 
514  truncation_range_r.name("range pair");
515  truncation_range_r
516  %= lit('T')
517  > lit('[')
518  > -expression_g(_r1)
519  > lit(',')
520  > -expression_g(_r1)
521  > lit(']');
522 
523  no_op_statement_r.name("no op statement");
524  no_op_statement_r
525  %= lit(';') [_val = no_op_statement()]; // ok to re-use instance
526 
527  }
528 
529  }
530 }
531 #endif

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