Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
ast.hpp
Go to the documentation of this file.
1 #ifndef __STAN__GM__AST_HPP__
2 #define __STAN__GM__AST_HPP__
3 
4 #include <map>
5 #include <vector>
6 
7 #include <boost/variant/recursive_variant.hpp>
8 
9 namespace stan {
10 
11  namespace gm {
12 
15  struct nil { };
16 
17  // components of abstract syntax tree
18  struct array_literal;
19  struct assignment;
20  struct binary_op;
21  struct conditional_statement;
22  struct distribution;
23  struct double_var_decl;
24  struct double_literal;
25  struct expression;
26  struct for_statement;
27  struct fun;
28  struct identifier;
29  struct index_op;
30  struct int_literal;
31  struct inv_var_decl;
32  struct matrix_var_decl;
33  struct no_op_statement;
34  struct ordered_var_decl;
36  struct print_statement;
37  struct program;
38  struct range;
39  struct row_vector_var_decl;
40  struct sample;
41  struct simplex_var_decl;
42  struct unit_vector_var_decl;
43  struct statement;
44  struct statements;
45  struct unary_op;
46  struct variable;
47  struct variable_dims;
48  struct var_decl;
49  struct var_type;
50  struct vector_var_decl;
52 
53  // forward declarable enum hack (can't fwd-decl enum)
54  typedef int base_expr_type;
55  const int INT_T = 1;
56  const int DOUBLE_T = 2;
57  const int VECTOR_T = 3;
58  const int ROW_VECTOR_T = 4;
59  const int MATRIX_T = 5;
60  const int ILL_FORMED_T = 6;
61 
62  std::ostream& write_base_expr_type(std::ostream& o, base_expr_type type);
63 
64  struct expr_type {
66  size_t num_dims_;
67  expr_type();
68  expr_type(const base_expr_type base_type);
69  expr_type(const base_expr_type base_type,
70  size_t num_dims);
71  bool operator==(const expr_type& et) const;
72  bool operator!=(const expr_type& et) const;
73  bool is_primitive() const;
74  bool is_primitive_int() const;
75  bool is_primitive_double() const;
76  bool is_ill_formed() const;
77  base_expr_type type() const;
78  size_t num_dims() const;
79  };
80 
81  std::ostream& operator<<(std::ostream& o, const expr_type& et);
82 
84 
86  const expr_type& et2);
87 
88  typedef std::pair<expr_type, std::vector<expr_type> > function_signature_t;
89 
91  public:
92  static function_signatures& instance();
93  void add(const std::string& name,
94  const expr_type& result_type,
95  const std::vector<expr_type>& arg_types);
96  void add(const std::string& name,
97  const expr_type& result_type);
98  void add(const std::string& name,
99  const expr_type& result_type,
100  const expr_type& arg_type);
101  void add(const std::string& name,
102  const expr_type& result_type,
103  const expr_type& arg_type1,
104  const expr_type& arg_type2);
105  void add(const std::string& name,
106  const expr_type& result_type,
107  const expr_type& arg_type1,
108  const expr_type& arg_type2,
109  const expr_type& arg_type3);
110  void add(const std::string& name,
111  const expr_type& result_type,
112  const expr_type& arg_type1,
113  const expr_type& arg_type2,
114  const expr_type& arg_type3,
115  const expr_type& arg_type4);
116  void add(const std::string& name,
117  const expr_type& result_type,
118  const expr_type& arg_type1,
119  const expr_type& arg_type2,
120  const expr_type& arg_type3,
121  const expr_type& arg_type4,
122  const expr_type& arg_type5);
123  void add_nullary(const::std::string& name);
124  void add_unary(const::std::string& name);
125  void add_binary(const::std::string& name);
126  void add_ternary(const::std::string& name);
127  void add_quaternary(const::std::string& name);
128  int num_promotions(const std::vector<expr_type>& call_args,
129  const std::vector<expr_type>& sig_args);
130  expr_type get_result_type(const std::string& name,
131  const std::vector<expr_type>& args,
132  std::ostream& error_msgs);
133  private:
136  std::map<std::string, std::vector<function_signature_t> > sigs_map_;
137  static function_signatures* sigs_; // init below outside of class
138  };
139 
140  struct statements {
141  std::vector<var_decl> local_decl_;
142  std::vector<statement> statements_;
143  statements();
144  statements(const std::vector<var_decl>& local_decl,
145  const std::vector<statement>& stmts);
146  };
147 
148 
149  struct distribution {
150  std::string family_;
151  std::vector<expression> args_;
152  };
153 
154  struct expression_type_vis : public boost::static_visitor<expr_type> {
155  expr_type operator()(const nil& e) const;
156  expr_type operator()(const int_literal& e) const;
157  expr_type operator()(const double_literal& e) const;
158  expr_type operator()(const array_literal& e) const;
159  expr_type operator()(const variable& e) const;
160  expr_type operator()(const fun& e) const;
161  expr_type operator()(const index_op& e) const;
162  expr_type operator()(const binary_op& e) const;
163  expr_type operator()(const unary_op& e) const;
164  // template <typename T> expr_type operator()(const T& e) const;
165  };
166 
167 
168 
169 
170  struct expression;
171 
172  struct expression {
173  typedef boost::variant<boost::recursive_wrapper<nil>,
174  boost::recursive_wrapper<int_literal>,
175  boost::recursive_wrapper<double_literal>,
176  boost::recursive_wrapper<array_literal>,
177  boost::recursive_wrapper<variable>,
178  boost::recursive_wrapper<fun>,
179  boost::recursive_wrapper<index_op>,
180  boost::recursive_wrapper<binary_op>,
181  boost::recursive_wrapper<unary_op> >
183 
184  expression();
185  expression(const expression& e);
186 
187  // template <typename Expr> expression(const Expr& expr);
188  expression(const nil& expr);
189  expression(const int_literal& expr);
190  expression(const double_literal& expr);
191  expression(const array_literal& expr);
192  expression(const variable& expr);
193  expression(const fun& expr);
194  expression(const index_op& expr);
195  expression(const binary_op& expr);
196  expression(const unary_op& expr);
197  expression(const expression_t& expr_);
198 
199  expr_type expression_type() const;
200 
201  expression& operator+=(const expression& rhs);
202  expression& operator-=(const expression& rhs);
203  expression& operator*=(const expression& rhs);
204  expression& operator/=(const expression& rhs);
205 
207  };
208 
209  // struct contains_var : public boost::static_visitor<bool> {
210  // const variable_map& var_map_;
211  // contains_var(const variable_map& var_map);
212  // bool operator()(const nil& e) const;
213  // bool operator()(const int_literal& e) const;
214  // bool operator()(const double_literal& e) const;
215  // bool operator()(const array_literal& e) const;
216  // bool operator()(const variable& e) const;
217  // bool operator()(const fun& e) const;
218  // bool operator()(const index_op& e) const;
219  // bool operator()(const binary_op& e) const;
220  // bool operator()(const unary_op& e) const;
221  // };
222 
223  struct printable {
224  typedef boost::variant<boost::recursive_wrapper<std::string>,
225  boost::recursive_wrapper<expression> >
227 
228  printable();
230  printable(const std::string& msg);
232  printable(const printable& printable);
233 
235  };
236 
237  struct is_nil_op : public boost::static_visitor<bool> {
238  bool operator()(const nil& x) const;
239  bool operator()(const int_literal& x) const;
240  bool operator()(const double_literal& x) const;
241  bool operator()(const array_literal& x) const;
242  bool operator()(const variable& x) const;
243  bool operator()(const fun& x) const;
244  bool operator()(const index_op& x) const;
245  bool operator()(const binary_op& x) const;
246  bool operator()(const unary_op& x) const;
247 
248  // template <typename T>
249  // bool operator()(const T& x) const;
250  };
251 
252  bool is_nil(const expression& e);
253 
254  struct variable_dims {
255  std::string name_;
256  std::vector<expression> dims_;
257  variable_dims();
258  variable_dims(std::string const& name,
259  std::vector<expression> const& dims);
260  };
261 
262 
263  struct int_literal {
264  int val_;
266  int_literal();
267  int_literal(int val);
268  int_literal(const int_literal& il);
269  int_literal& operator=(const int_literal& il);
270  };
271 
272 
273  struct double_literal {
274  double val_;
276  double_literal();
277  double_literal(double val);
279  };
280 
281  struct array_literal {
282  std::vector<expression> args_;
284  array_literal();
285  array_literal(const std::vector<expression>& args);
287  };
288 
289  struct variable {
290  std::string name_;
292  variable();
293  variable(std::string name);
294  void set_type(const base_expr_type& base_type,
295  size_t num_dims);
296  };
297 
298  struct fun {
299  std::string name_;
300  std::vector<expression> args_;
302  fun();
303  fun(std::string const& name,
304  std::vector<expression> const& args);
305  void infer_type(); // FIXME: is this used anywhere?
306  };
307 
308  size_t total_dims(const std::vector<std::vector<expression> >& dimss);
309 
310  expr_type infer_type_indexing(const base_expr_type& expr_base_type,
311  size_t num_expr_dims,
312  size_t num_index_dims);
313 
315  size_t num_index_dims);
316 
317 
318  struct index_op {
320  std::vector<std::vector<expression> > dimss_;
322  index_op();
323  // vec of vec for e.g., e[1,2][3][4,5,6]
324  index_op(const expression& expr,
325  const std::vector<std::vector<expression> >& dimss);
326  void infer_type();
327  };
328 
329 
330  struct binary_op {
331  std::string op;
335  binary_op();
336  binary_op(const expression& left,
337  const std::string& op,
338  const expression& right);
339  };
340 
341  struct unary_op {
342  char op;
345  unary_op(char op,
346  expression const& subject);
347  };
348 
349  struct range {
352  range();
353  range(expression const& low,
354  expression const& high);
355  bool has_low() const;
356  bool has_high() const;
357  };
358 
359  typedef int var_origin;
360  const int data_origin = 1;
361  const int transformed_data_origin = 2;
362  const int parameter_origin = 3;
364  const int derived_origin = 5;
365  const int local_origin = 6;
366 
367 
368 
369  void print_var_origin(std::ostream& o, const var_origin& vo);
370 
371  struct base_var_decl {
372  std::string name_;
373  std::vector<expression> dims_;
375  base_var_decl();
376  base_var_decl(const base_expr_type& base_type);
377  base_var_decl(const std::string& name,
378  const std::vector<expression>& dims,
379  const base_expr_type& base_type);
380  };
381 
382  struct variable_map {
383  typedef std::pair<base_var_decl,var_origin> range_t;
384  std::map<std::string, range_t> map_;
385  bool exists(const std::string& name) const;
386  base_var_decl get(const std::string& name) const;
387  base_expr_type get_base_type(const std::string& name) const;
388  size_t get_num_dims(const std::string& name) const;
389  var_origin get_origin(const std::string& name) const;
390  void add(const std::string& name,
391  const base_var_decl& base_decl,
392  const var_origin& vo);
393  void remove(const std::string& name);
394  };
395 
396  struct int_var_decl : public base_var_decl {
398  int_var_decl();
399  int_var_decl(range const& range,
400  std::string const& name,
401  std::vector<expression> const& dims);
402  };
403 
404 
405  struct double_var_decl : public base_var_decl {
407  double_var_decl();
408  double_var_decl(range const& range,
409  std::string const& name,
410  std::vector<expression> const& dims);
411  };
412 
417  std::string const& name,
418  std::vector<expression> const& dims);
419  };
420 
421  struct simplex_var_decl : public base_var_decl {
424  simplex_var_decl(expression const& K,
425  std::string const& name,
426  std::vector<expression> const& dims);
427  };
428 
429  struct ordered_var_decl : public base_var_decl {
432  ordered_var_decl(expression const& K,
433  std::string const& name,
434  std::vector<expression> const& dims);
435  };
436 
441  std::string const& name,
442  std::vector<expression> const& dims);
443  };
444 
445  struct vector_var_decl : public base_var_decl {
448  vector_var_decl();
449  vector_var_decl(range const& range,
450  expression const& M,
451  std::string const& name,
452  std::vector<expression> const& dims);
453  };
454 
460  expression const& N,
461  std::string const& name,
462  std::vector<expression> const& dims);
463  };
464 
465  struct matrix_var_decl : public base_var_decl {
469  matrix_var_decl();
470  matrix_var_decl(range const& range,
471  expression const& M,
472  expression const& N,
473  std::string const& name,
474  std::vector<expression> const& dims);
475  };
476 
477 
478 
479 
480 
485  std::string const& name,
486  std::vector<expression> const& dims);
487  };
488 
489 
490 
495  std::string const& name,
496  std::vector<expression> const& dims);
497  };
498 
499 
500 
501  struct name_vis : public boost::static_visitor<std::string> {
502  name_vis();
503  std::string operator()(const nil& x) const;
504  std::string operator()(const int_var_decl& x) const;
505  std::string operator()(const double_var_decl& x) const;
506  std::string operator()(const vector_var_decl& x) const;
507  std::string operator()(const row_vector_var_decl& x) const;
508  std::string operator()(const matrix_var_decl& x) const;
509  std::string operator()(const simplex_var_decl& x) const;
510  std::string operator()(const unit_vector_var_decl& x) const;
511  std::string operator()(const ordered_var_decl& x) const;
512  std::string operator()(const positive_ordered_var_decl& x) const;
513  std::string operator()(const cov_matrix_var_decl& x) const;
514  std::string operator()(const corr_matrix_var_decl& x) const;
515  };
516 
517 
518 
519 
520  struct var_decl {
521  typedef boost::variant<boost::recursive_wrapper<nil>,
522  boost::recursive_wrapper<int_var_decl>,
523  boost::recursive_wrapper<double_var_decl>,
524  boost::recursive_wrapper<vector_var_decl>,
525  boost::recursive_wrapper<row_vector_var_decl>,
526  boost::recursive_wrapper<matrix_var_decl>,
527  boost::recursive_wrapper<simplex_var_decl>,
528  boost::recursive_wrapper<unit_vector_var_decl>,
529  boost::recursive_wrapper<ordered_var_decl>,
530  boost::recursive_wrapper<positive_ordered_var_decl>,
531  boost::recursive_wrapper<cov_matrix_var_decl>,
532  boost::recursive_wrapper<corr_matrix_var_decl> >
534 
536 
537  var_decl();
538 
539  // template <typename Decl>
540  // var_decl(Decl const& decl);
541  var_decl(const var_decl_t& decl);
542  var_decl(const nil& decl);
543  var_decl(const int_var_decl& decl);
544  var_decl(const double_var_decl& decl);
545  var_decl(const vector_var_decl& decl);
546  var_decl(const row_vector_var_decl& decl);
547  var_decl(const matrix_var_decl& decl);
548  var_decl(const simplex_var_decl& decl);
549  var_decl(const unit_vector_var_decl& decl);
550  var_decl(const ordered_var_decl& decl);
551  var_decl(const positive_ordered_var_decl& decl);
552  var_decl(const cov_matrix_var_decl& decl);
553  var_decl(const corr_matrix_var_decl& decl);
554 
555  std::string name() const;
556  };
557 
558  struct statement {
559  typedef boost::variant<boost::recursive_wrapper<nil>,
560  boost::recursive_wrapper<assignment>,
561  boost::recursive_wrapper<sample>,
562  boost::recursive_wrapper<statements>,
563  boost::recursive_wrapper<for_statement>,
564  boost::recursive_wrapper<conditional_statement>,
565  boost::recursive_wrapper<while_statement>,
566  boost::recursive_wrapper<print_statement>,
567  boost::recursive_wrapper<no_op_statement> >
569 
571 
572  statement();
573  statement(const statement_t& st);
574 
575  statement(const nil& st);
576  statement(const assignment& st);
577  statement(const sample& st);
578  statement(const statements& st);
579  statement(const for_statement& st);
580  statement(const conditional_statement& st);
581  statement(const while_statement& st);
582  statement(const print_statement& st);
583  statement(const no_op_statement& st);
584 
585  // template <typename Statement>
586  // statement(const Statement& statement);
587  };
588 
589  struct for_statement {
590  std::string variable_;
593  for_statement();
594  for_statement(std::string& variable,
595  range& range,
596  statement& stmt);
597  };
598 
599  // bodies may be 1 longer than conditions due to else
601  std::vector<expression> conditions_;
602  std::vector<statement> bodies_;
604  conditional_statement(const std::vector<expression>& conditions,
605  const std::vector<statement>& statements);
606  };
607 
611  while_statement();
612  while_statement(const expression& condition,
613  const statement& body);
614  };
615 
617  std::vector<printable> printables_;
618  print_statement();
619  print_statement(const std::vector<printable>& printables);
620  };
621 
622 
624  // no op, no data
625  };
626 
627 
628 
629  struct program {
630  std::vector<var_decl> data_decl_;
631  std::pair<std::vector<var_decl>,std::vector<statement> >
633  std::vector<var_decl> parameter_decl_;
634  std::pair<std::vector<var_decl>,std::vector<statement> >
637  std::pair<std::vector<var_decl>,std::vector<statement> > generated_decl_;
638  program();
639  program(const std::vector<var_decl>& data_decl,
640  const std::pair<std::vector<var_decl>,
641  std::vector<statement> >& derived_data_decl,
642  const std::vector<var_decl>& parameter_decl,
643  const std::pair<std::vector<var_decl>,
644  std::vector<statement> >& derived_decl,
645  const statement& st,
646  const std::pair<std::vector<var_decl>,
647  std::vector<statement> >& generated_decl);
648 
649 
650  };
651 
652  struct sample {
656  sample();
658  distribution& dist);
659  bool is_ill_formed() const;
660  };
661 
662  struct assignment {
663  variable_dims var_dims_; // lhs_var[dim0,...,dimN-1]
664  expression expr_; // = rhs
665  base_var_decl var_type_; // type of lhs_var
666  assignment();
667  assignment(variable_dims& var_dims,
668  expression& expr);
669  };
670 
671  // FIXME: is this next necessary dependency?
672  // from generator.hpp
673  void generate_expression(const expression& e, std::ostream& o);
674 
675  bool has_rng_suffix(const std::string& s);
676 
677 
678  }
679 }
680 
681 #endif

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