From c319448f73642c47f314a539a6de5d91519c8575 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=CE=BAeen?= <3han5chou7@gmail.com> Date: Sun, 28 Aug 2022 19:41:08 +0900 Subject: [PATCH 1/6] return unique polymorphic variables --- src/ast/typing.rs | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/src/ast/typing.rs b/src/ast/typing.rs index 987d289..03933bf 100644 --- a/src/ast/typing.rs +++ b/src/ast/typing.rs @@ -286,16 +286,23 @@ impl TyEnv { use Typing::*; match self.pool.pool.value_of(ty) { Variable(id) => vec![*id], - Char | Int | Real | OverloadedNum | OverloadedNumText => vec![], + Char | Int | Real | TyAbs(_, _) | OverloadedNum | OverloadedNumText => vec![], Fun(param, body) => { let mut ret = self.polymorphic_variables(param.clone()); ret.append(&mut self.polymorphic_variables(body.clone())); + ret.sort(); + ret.dedup(); + ret + } + Tuple(tys) => { + let mut ret = tys + .iter() + .flat_map(|ty| self.polymorphic_variables(ty.clone())) + .collect::>(); + ret.sort(); + ret.dedup(); ret } - Tuple(tys) => tys - .iter() - .flat_map(|ty| self.polymorphic_variables(ty.clone())) - .collect(), // currently Datatype will not be polymorphic Datatype(_) => vec![], } From cd155012ff4847a079476b4d6e80a9aba344eb7b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=CE=BAeen?= <3han5chou7@gmail.com> Date: Sun, 28 Aug 2022 21:09:06 +0900 Subject: [PATCH 2/6] polymorphic inference --- src/ast/case_simplify.rs | 4 +- src/ast/desugar.rs | 5 ++ src/ast/mod.rs | 10 +++ src/ast/pp.rs | 14 ++++ src/ast/rename.rs | 1 + src/ast/typing.rs | 145 ++++++++++++++++++++++++++++++--------- src/ast/util.rs | 8 +++ src/hir/ast2hir.rs | 5 +- 8 files changed, 156 insertions(+), 36 deletions(-) diff --git a/src/ast/case_simplify.rs b/src/ast/case_simplify.rs index c12db2d..b1a66ef 100644 --- a/src/ast/case_simplify.rs +++ b/src/ast/case_simplify.rs @@ -709,7 +709,9 @@ impl CaseSimplifyPass { ) -> bool { use Type::*; match ty { - Real | Variable(_) | Fun(_, _) => panic!("no way to pattern match against this type"), + Real | Variable(_) | Fun(_, _) | TyAbs(_, _) => { + panic!("no way to pattern match against this type") + } Char | Int => false, Tuple(_) => { // unlikely reachable, but writing incase it reaches. diff --git a/src/ast/desugar.rs b/src/ast/desugar.rs index 87268cc..4d055dc 100644 --- a/src/ast/desugar.rs +++ b/src/ast/desugar.rs @@ -187,6 +187,7 @@ impl Desugar { } => self.transform_externcall(span, module, fun, args, argty, retty), Fn { param, body } => self.transform_fn(span, param, body), App { fun, arg } => self.transform_app(span, fun, arg), + TyApp { fun, arg } => self.transform_tyapp(span, fun, arg), Case { cond, clauses } => self.transform_case(span, cond, clauses), Tuple { tuple } => self.transform_tuple(span, tuple), Constructor { arg, name } => self.transform_constructor(span, arg, name), @@ -284,6 +285,10 @@ impl Desugar { } } + fn transform_tyapp(&mut self, _: Span, fun: Symbol, arg: Vec) -> UntypedCoreExprKind { + ExprKind::TyApp { fun, arg } + } + fn transform_if( &mut self, span: Span, diff --git a/src/ast/mod.rs b/src/ast/mod.rs index a4d9bef..bf71e74 100644 --- a/src/ast/mod.rs +++ b/src/ast/mod.rs @@ -198,6 +198,10 @@ pub enum ExprKind< fun: Box>, arg: Box>, }, + TyApp { + fun: Symbol, + arg: Vec, + }, Case { cond: Box>, clauses: Vec<(Pattern, Expr)>, @@ -310,6 +314,7 @@ pub enum Type { Fun(Box, Box), Tuple(Vec), Datatype(Symbol), + TyAbs(Vec, Box), } #[derive(Debug, Clone, PartialEq)] @@ -476,6 +481,10 @@ impl CoreExpr { fun: fun.map_ty(f).boxed(), arg: arg.map_ty(f).boxed(), }, + TyApp { fun, arg } => TyApp { + fun, + arg: arg.into_iter().map(f).collect(), + }, Case { cond, clauses } => Case { cond: cond.map_ty(&mut *f).boxed(), clauses: clauses @@ -506,6 +515,7 @@ impl CoreExpr { ExternCall { .. } => false, Fn { .. } => true, App { .. } => false, + TyApp { .. } => true, // TODO: check compatibility with fn Case { .. } => false, Tuple { tuple } => tuple.iter().all(|e| e.is_value()), diff --git a/src/ast/pp.rs b/src/ast/pp.rs index fb6ff72..a6464a1 100644 --- a/src/ast/pp.rs +++ b/src/ast/pp.rs @@ -230,6 +230,16 @@ impl fmt next = next )?; } + TyApp { fun, arg } => { + write!(f, "({:indent$})@", fun, indent = indent,)?; + inter_iter! { + arg.iter(), + write!(f, ", ")?, + |ty| => { + write!(f,"{:indent$}", ty, indent = indent)?; + } + } + } Case { cond, clauses } => { let ind = nspaces(indent); write!(f, "case {:next$} of", cond, next = next)?; @@ -432,6 +442,10 @@ impl fmt::Display for Type { write!(f, ")")?; } Datatype(name) => write!(f, "{}", name)?, + TyAbs(vars, ty) => { + inter_iter!(vars.iter(), write!(f, ", ")?, |id| => { write!(f, "'{}", id)? }); + write!(f, " {}", ty)? + } } Ok(()) } diff --git a/src/ast/rename.rs b/src/ast/rename.rs index 496543e..ad7572a 100644 --- a/src/ast/rename.rs +++ b/src/ast/rename.rs @@ -161,6 +161,7 @@ impl<'a> Scope<'a> { } } } + TyAbs(_, ty) => self.rename_type(ty), } } } diff --git a/src/ast/typing.rs b/src/ast/typing.rs index 03933bf..7630f00 100644 --- a/src/ast/typing.rs +++ b/src/ast/typing.rs @@ -34,6 +34,7 @@ enum Typing { Fun(NodeId, NodeId), Tuple(Vec), Datatype(Symbol), + TyAbs(Vec, NodeId), OverloadedNum, OverloadedNumText, } @@ -60,6 +61,7 @@ fn conv_ty(pool: &UnificationPool, ty: Typing) -> Type { ), Tuple(tys) => Type::Tuple(tys.into_iter().map(|ty| resolve(pool, ty)).collect()), Datatype(type_id) => Type::Datatype(type_id), + TyAbs(vars, ty) => Type::TyAbs(vars, Box::new(resolve(pool, ty))), OverloadedNum => Type::Int, OverloadedNumText => Type::Int, } @@ -215,6 +217,35 @@ impl TypePool { ) -> std::result::Result { self.pool.try_unify_with(id1, id2, try_unify) } + + fn set(&mut self, old: NodeId, new: NodeId) { + self.pool + .try_unify_with::(old, new, |_, _, new| Ok(new)) + .unwrap(); + } + fn app_types(&mut self, ty: NodeId, params: &[TypeId], args: &[NodeId]) -> NodeId { + use Typing::*; + match self.pool.value_of(ty).clone() { + Variable(id) => match params.iter().position(|&p| p == id) { + Some(i) => args[i].clone(), + None => ty, + }, + Char | Int | Real | OverloadedNum | OverloadedNumText | Datatype(_) => ty, + Fun(f, a) => { + let f = self.app_types(f, params, args); + let a = self.app_types(a, params, args); + self.ty(Fun(f, a)) + } + Tuple(ts) => { + let ts = ts + .into_iter() + .map(|t| self.app_types(t, params, args)) + .collect(); + self.ty(Tuple(ts)) + } + TyAbs(_, _) => unreachable!(), + } + } } impl TypePool { @@ -329,19 +360,23 @@ impl TyEnv { .collect(), ), Type::Datatype(name) => Typing::Datatype(name), + Type::TyAbs(vars, ty) => { + let typing = self.convert(*ty); + Typing::TyAbs(vars, self.pool.ty(typing)) + } } } } impl TyEnv { - fn infer_ast(&mut self, ast: &Core) -> Result<()> { - for decl in ast.0.iter() { + fn infer_ast(&mut self, ast: &mut Core) -> Result<()> { + for decl in &mut ast.0 { self.infer_statement(decl)?; } Ok(()) } - fn infer_statement(&mut self, decl: &CoreDeclaration) -> Result<()> { + fn infer_statement(&mut self, decl: &mut CoreDeclaration) -> Result<()> { use Declaration::*; match decl { Datatype { .. } => Ok(()), @@ -355,22 +390,29 @@ impl TyEnv { self.infer_expr(expr)?; self.infer_pat(pattern.span(), pattern)?; self.unify(expr.span(), expr.ty(), pattern.ty())?; - if !rec { + if !*rec { for &(name, ty) in &names { self.insert(name.clone(), Clone::clone(ty)); } } let ty = expr.ty(); - match (&self.polymorphic_variables(ty)[..], expr.is_value()) { - ([], _) => Ok(()), + match (self.polymorphic_variables(ty), expr.is_value()) { + (v, _) if v.is_empty() => Ok(()), (_, false) => Err(TypeError::new( expr.span(), TypeErrorKind::PolymorphicExpression, )), - (_vars, true) => todo!(), + (vars, true) => { + println!("generalized {:?}", ty); + let cloned = self.pool.ty(self.pool.pool.value_of(ty).clone()); + let new_ty = self.pool.ty(Typing::TyAbs(vars, cloned)); + println!("new_ty {:?}", new_ty); + self.pool.set(ty, new_ty); + Ok(()) + } } } - LangItem { decl, .. } => self.infer_statement(decl.as_ref()), + LangItem { decl, .. } => self.infer_statement(decl.as_mut()), Local { binds, body } => { for b in binds { self.infer_statement(b)? @@ -384,7 +426,7 @@ impl TyEnv { } } - fn infer_expr(&mut self, expr: &CoreExpr) -> Result<()> { + fn infer_expr(&mut self, expr: &mut CoreExpr) -> Result<()> { use crate::ast::ExprKind::*; let int = self.pool.ty_int(); let real = self.pool.ty_real(); @@ -392,7 +434,8 @@ impl TyEnv { let overloaded_num = self.pool.ty_overloaded_num(); let overloaded_num_text = self.pool.ty_overloaded_num_text(); let ty = &expr.ty; - match &expr.inner { + let span = expr.span(); + match &mut expr.inner { Binds { binds, ret } => { for decl in binds { self.infer_statement(decl)?; @@ -406,48 +449,56 @@ impl TyEnv { match fun { Add | Sub | Mul => { assert!(args.len() == 2); - let l = &args[0]; - let r = &args[1]; + let (l, r) = match &mut args[..] { + [l, r] => (l, r), + _ => unreachable!(), + }; self.infer_expr(l)?; self.infer_expr(r)?; self.unify(r.span(), l.ty(), r.ty())?; self.unify(l.span(), l.ty(), overloaded_num)?; - self.unify(expr.span(), *ty, l.ty())?; + self.unify(span, *ty, l.ty())?; Ok(()) } Eq | Neq | Gt | Ge | Lt | Le => { assert!(args.len() == 2); - let l = &args[0]; - let r = &args[1]; + let (l, r) = match &mut args[..] { + [l, r] => (l, r), + _ => unreachable!(), + }; self.infer_expr(l)?; self.infer_expr(r)?; self.unify(r.span(), l.ty(), r.ty())?; self.unify(l.span(), l.ty(), overloaded_num_text)?; - self.unify(expr.span(), *ty, bool)?; + self.unify(span, *ty, bool)?; Ok(()) } Div | Mod => { assert!(args.len() == 2); - let l = &args[0]; - let r = &args[1]; + let (l, r) = match &mut args[..] { + [l, r] => (l, r), + _ => unreachable!(), + }; self.unify(l.span(), l.ty(), int)?; self.unify(r.span(), r.ty(), int)?; - self.unify(expr.span(), *ty, int)?; + self.unify(span, *ty, int)?; self.infer_expr(l)?; self.infer_expr(r)?; Ok(()) } Divf => { assert!(args.len() == 2); - let l = &args[0]; - let r = &args[1]; + let (l, r) = match &mut args[..] { + [l, r] => (l, r), + _ => unreachable!(), + }; self.unify(l.span(), l.ty(), real)?; self.unify(r.span(), r.ty(), real)?; - self.unify(expr.span(), *ty, real)?; + self.unify(span, *ty, real)?; self.infer_expr(l)?; self.infer_expr(r)?; Ok(()) @@ -463,7 +514,7 @@ impl TyEnv { ExternCall { args, argty, retty, .. } => { - for (arg, argty) in args.iter().zip(argty) { + for (arg, argty) in args.iter_mut().zip(argty) { self.infer_expr(arg)?; let argty = self.convert(argty.clone()); self.give(arg.span(), arg.ty(), argty)?; @@ -476,15 +527,20 @@ impl TyEnv { let param_ty = self.pool.tyvar(); self.insert(param.clone(), param_ty); self.infer_expr(body)?; - self.give(expr.span(), *ty, Typing::Fun(param_ty, body.ty()))?; + self.give(span, *ty, Typing::Fun(param_ty, body.ty()))?; Ok(()) } App { fun, arg } => { self.infer_expr(fun)?; self.infer_expr(arg)?; - self.give(expr.span(), fun.ty(), Typing::Fun(arg.ty(), *ty))?; + self.give(span, fun.ty(), Typing::Fun(arg.ty(), *ty))?; Ok(()) } + TyApp { .. } => { + // Do nothing + // This term is generated after typing + unreachable!() + } Case { cond, clauses } => { self.infer_expr(cond)?; for (pat, branch) in clauses { @@ -496,19 +552,40 @@ impl TyEnv { Ok(()) } Tuple { tuple } => { - self.infer_tuple(expr.span(), tuple, *ty)?; + self.infer_tuple(span, tuple, *ty)?; Ok(()) } Constructor { arg, name } => { - self.infer_constructor(expr.span(), name, arg, *ty)?; + self.infer_constructor(span, name, arg, *ty)?; Ok(()) } Symbol { name } => { - self.infer_symbol(expr.span(), name, *ty)?; + use std::iter; + let mut name = name.clone(); + self.infer_symbol(span, &mut name, *ty)?; + let (vars, bodyty) = match self.pool.pool.value_of(*ty) { + Typing::TyAbs(vars, bodyty) => (vars.clone(), bodyty.clone()), + // monomorphic path + _ => return Ok(()), + }; + // polymorphic path + let newtyvars = iter::repeat(()) + .take(vars.len()) + .map(|_| self.pool.tyvar()) + .collect::>(); + let newty = self.pool.app_types(bodyty, &vars, &newtyvars); + *expr = Expr { + ty: newty, + span: expr.span.clone(), + inner: TyApp { + fun: name, + arg: newtyvars, + }, + }; Ok(()) } Literal { value } => { - self.infer_literal(expr.span(), value, *ty)?; + self.infer_literal(span, value, *ty)?; Ok(()) } D(d) => match *d {}, @@ -519,15 +596,15 @@ impl TyEnv { &mut self, span: Span, sym: &Symbol, - arg: &Option>>, + arg: &mut Option>>, given: NodeId, ) -> Result<()> { match self.get(sym) { Some(ty) => { self.unify(span, ty, given)?; let arg_ty = self.symbol_table().get_argtype_of_constructor(sym); - if let (Some(arg), Some(arg_ty)) = (arg.clone(), arg_ty.cloned()) { - self.infer_expr(&arg)?; + if let (Some(mut arg), Some(arg_ty)) = (arg.clone(), arg_ty.cloned()) { + self.infer_expr(&mut arg)?; let arg_typing = self.convert(arg_ty); let arg_ty_id = self.pool.ty(arg_typing); self.unify(arg.span(), arg.ty(), arg_ty_id)?; @@ -623,7 +700,7 @@ impl TyEnv { fn infer_tuple( &mut self, span: Span, - tuple: &Vec>, + tuple: &mut Vec>, given: NodeId, ) -> Result<()> { use std::iter; @@ -631,7 +708,7 @@ impl TyEnv { .take(tuple.len()) .collect::>(); - for (e, t) in tuple.iter().zip(tys.iter()) { + for (e, t) in tuple.iter_mut().zip(tys.iter()) { self.infer_expr(e)?; self.unify(e.span(), e.ty(), *t)?; } diff --git a/src/ast/util.rs b/src/ast/util.rs index c2531dc..3e51c27 100644 --- a/src/ast/util.rs +++ b/src/ast/util.rs @@ -67,6 +67,7 @@ pub trait Traverse { } => self.traverse_externcall(span, module, fun, args, argty, retty), Fn { param, body } => self.traverse_fn(span, param, body), App { fun, arg } => self.traverse_app(span, fun, arg), + TyApp { fun, arg } => self.traverse_tyapp(span, fun, arg), Case { cond, clauses } => self.traverse_case(span, cond, clauses), Tuple { tuple } => self.traverse_tuple(span, tuple), Constructor { arg, name } => self.traverse_constructor(span, arg, name), @@ -116,6 +117,8 @@ pub trait Traverse { self.traverse_expr(arg); } + fn traverse_tyapp(&mut self, _: Span, _: &mut Symbol, _: &mut Vec) {} + fn traverse_case( &mut self, _: Span, @@ -260,6 +263,7 @@ pub trait Transform { } => self.transform_externcall(span, module, fun, args, argty, retty), Fn { param, body } => self.transform_fn(span, param, body), App { fun, arg } => self.transform_app(span, fun, arg), + TyApp { fun, arg } => self.transform_tyapp(span, fun, arg), Case { cond, clauses } => self.transform_case(span, cond, clauses), Tuple { tuple } => self.transform_tuple(span, tuple), Constructor { arg, name } => self.transform_constructor(span, arg, name), @@ -344,6 +348,10 @@ pub trait Transform { } } + fn transform_tyapp(&mut self, _: Span, fun: Symbol, arg: Vec) -> CoreExprKind { + ExprKind::TyApp { fun, arg } + } + fn transform_case( &mut self, _: Span, diff --git a/src/hir/ast2hir.rs b/src/hir/ast2hir.rs index 39a0496..1a20883 100644 --- a/src/hir/ast2hir.rs +++ b/src/hir/ast2hir.rs @@ -54,7 +54,7 @@ fn conv_ty(ty: ast::Type) -> HTy { Tuple(tys) => HTy::Tuple(tys.into_iter().map(conv_ty).collect()), Fun(arg, ret) => HTy::fun(conv_ty(*arg), conv_ty(*ret)), Datatype(name) => HTy::Datatype(name), - Variable(_) => panic!("polymorphism is not supported yet"), + TyAbs(_, _) | Variable(_) => panic!("polymorphism is not supported yet"), } } @@ -295,6 +295,9 @@ impl AST2HIRPass { } } E::App { fun, arg } => self.conv_expr(*fun).app1(conv_ty(ty), self.conv_expr(*arg)), + E::TyApp { fun, arg } => { + panic!("Function {} is not monomorphized with {:?}", fun, arg) + } E::Case { cond, clauses } => Expr::Case { ty: conv_ty(ty), expr: Box::new(self.conv_expr(*cond)), From 6f787346f84e9635495e9ed48b0b2b03010741f7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=CE=BAeen?= <3han5chou7@gmail.com> Date: Mon, 29 Aug 2022 00:17:48 +0900 Subject: [PATCH 3/6] implementing monomorphization --- Cargo.lock | 16 +++ Cargo.toml | 1 + src/ast/mod.rs | 4 +- src/ast/monomorphize.rs | 221 ++++++++++++++++++++++++++++++++++++++++ src/lib.rs | 1 + 5 files changed, 242 insertions(+), 1 deletion(-) create mode 100644 src/ast/monomorphize.rs diff --git a/Cargo.lock b/Cargo.lock index e579d7f..ce0b858 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -70,6 +70,12 @@ dependencies = [ "vec_map", ] +[[package]] +name = "either" +version = "1.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "90e5c1c8368803113bf0c9584fc495a58b86dc8a29edbf8fe877d21d9507e797" + [[package]] name = "env_logger" version = "0.7.1" @@ -107,6 +113,15 @@ dependencies = [ "quick-error", ] +[[package]] +name = "itertools" +version = "0.10.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a9a9d19fa1e79b6215ff29b9d6880b706147f16e9b1dbb1e4e5947b5b02bc5e3" +dependencies = [ + "either", +] + [[package]] name = "lazy_static" version = "1.4.0" @@ -364,6 +379,7 @@ version = "0.1.0" dependencies = [ "clap", "env_logger", + "itertools", "log", "nom", "nom_locate", diff --git a/Cargo.toml b/Cargo.toml index 33b96fc..8efa71b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -22,6 +22,7 @@ log = "0.4.8" env_logger = "0.7.1" regex = "1.3.7" wasm = { version = "0.1.2", package = "web-assembler" } +itertools = "0.10.3" [target.'cfg(target_arch = "wasm32")'.dependencies] diff --git a/src/ast/mod.rs b/src/ast/mod.rs index bf71e74..f1adecf 100644 --- a/src/ast/mod.rs +++ b/src/ast/mod.rs @@ -1,6 +1,7 @@ mod case_simplify; mod collect_langitems; mod desugar; +mod monomorphize; mod pp; mod rename; mod resolve_overload; @@ -11,6 +12,7 @@ mod var2constructor; pub use self::case_simplify::CaseSimplify; pub use self::collect_langitems::CollectLangItems; pub use self::desugar::Desugar; +pub use self::monomorphize::Monomorphize; pub use self::rename::Rename; pub use self::resolve_overload::ResolveOverload; pub use self::typing::Typer; @@ -305,7 +307,7 @@ pub struct SymbolTable { type TypeId = u64; -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Debug, Clone, PartialEq, Eq, Hash)] pub enum Type { Variable(TypeId), Char, diff --git a/src/ast/monomorphize.rs b/src/ast/monomorphize.rs new file mode 100644 index 0000000..de5ae67 --- /dev/null +++ b/src/ast/monomorphize.rs @@ -0,0 +1,221 @@ +use crate::ast::*; +use crate::id::Id; +use crate::prim::*; +use crate::Config; +use itertools::Itertools; +use std::collections::{HashMap, HashSet}; +use util::{Transform, Traverse}; + +#[derive(Debug)] +pub struct Monomorphize { + id: Id, +} + +impl Monomorphize { + pub fn new(id: Id) -> Self { + Self { id } + } +} + +#[derive(Debug, Default)] +struct InstanceCollector { + instance_table: HashMap>>, +} + +impl Traverse for InstanceCollector { + fn traverse_tyapp(&mut self, _: Span, fun: &mut Symbol, arg: &mut Vec) { + self.instance_table + .entry(fun.clone()) + .or_default() + .insert(arg.clone()); + } +} + +#[derive(Debug)] +struct Monomorphizer { + instance_table: HashMap>>, + name_table: HashMap<(Symbol, Vec), Symbol>, + id: Id, +} + +impl Monomorphizer { + fn new(instance_table: HashMap>>, id: Id) -> Self { + Self { + instance_table, + name_table: HashMap::new(), + id, + } + } + + fn instanciated_name(&mut self, name: Symbol, arg: Vec) -> Symbol { + // TODO: use Entry API + if let Some(n) = self.name_table.get(&(name.clone(), arg.clone())) { + return n.clone(); + } + let mut new_name = name.0.clone(); + new_name.push_str(&format!("{}", arg.iter().format("{}"))); + let id = self.id.next(); + let new_symbol = Symbol(new_name, id); + self.name_table.insert((name, arg), new_symbol.clone()); + new_symbol + } +} + +fn rewrite_type(ty: &mut Type, params: &[TypeId], args: &[Type]) { + use Type::*; + match ty { + Variable(id) => { + if let Some(index) = params.iter().position(|id2| id2 == id) { + *ty = args[index].clone(); + } + } + Char | Int | Real | Datatype(_) => (), + Fun(f, arg) => { + rewrite_type(&mut *f, params, args); + rewrite_type(&mut *arg, params, args); + } + Tuple(tys) => { + for ty in tys { + rewrite_type(&mut *ty, params, args); + } + } + TyAbs(_, _) => unreachable!(), + } +} + +#[derive(Debug)] +pub struct Instanciator<'a> { + params: &'a [TypeId], + args: &'a [Type], +} + +impl<'a> Traverse for Instanciator<'a> { + fn traverse_expr(&mut self, expr: &mut CoreExpr) { + rewrite_type(&mut expr.ty, self.params, self.args); + + // the same as default implementation + use crate::ast::ExprKind::*; + + let span = expr.span.clone(); + match &mut expr.inner { + Binds { binds, ret } => self.traverse_binds(span, binds, ret), + BuiltinCall { fun, args } => self.traverse_builtincall(span, fun, args), + ExternCall { + module, + fun, + args, + argty, + retty, + } => self.traverse_externcall(span, module, fun, args, argty, retty), + Fn { param, body } => self.traverse_fn(span, param, body), + App { fun, arg } => self.traverse_app(span, fun, arg), + TyApp { fun, arg } => self.traverse_tyapp(span, fun, arg), + Case { cond, clauses } => self.traverse_case(span, cond, clauses), + Tuple { tuple } => self.traverse_tuple(span, tuple), + Constructor { arg, name } => self.traverse_constructor(span, arg, name), + Symbol { name } => self.traverse_sym(span, name), + Literal { value } => self.traverse_lit(span, value), + D(_) => (), + } + } + + fn traverse_pattern(&mut self, pattern: &mut CorePattern) { + rewrite_type(&mut pattern.ty, self.params, self.args); + + // the same as default implementation + use PatternKind::*; + let span = pattern.span.clone(); + match &mut pattern.inner { + Constant { value } => self.traverse_pat_constant(span, value), + Char { value } => self.traverse_pat_char(span, value), + Constructor { name, arg } => self.traverse_pat_constructor(span, name, arg), + Tuple { tuple } => self.traverse_pat_tuple(span, tuple), + Variable { name } => self.traverse_pat_variable(span, name), + Wildcard {} => self.traverse_pat_wildcard(span), + D(d) => match *d {}, + } + } +} + +impl Transform for Monomorphizer { + fn transform_val( + &mut self, + rec: bool, + pattern: CorePattern, + expr: CoreExpr, + ) -> CoreDeclaration { + let pattern = self.transform_pattern(pattern); + let expr = self.transform_expr(expr); + + let binds = pattern + .binds() + .into_iter() + .map(|(n, _)| n) + .cloned() + .collect::>(); + if binds.iter().all(|b| !self.instance_table.contains_key(b)) { + return Declaration::Val { rec, pattern, expr }; + } + // currently only supports `val variable = expr` + assert_eq!(binds.len(), 1); + assert!(matches!(pattern.ty, Type::TyAbs(_, _))); + + for args in &self.instance_table[&binds[0]] { + let mut expr = expr.clone(); + let mut pattern = pattern.clone(); + match expr.ty { + Type::TyAbs(params, body) => { + expr.ty = *body; + Instanciator { + params: ¶ms, + args: &args, + } + .traverse_expr(&mut expr); + } + _ => (), + } + match pattern.ty { + Type::TyAbs(params, body) => { + pattern.ty = *body; + + Instanciator { + params: ¶ms, + args: &args, + } + .traverse_pattern(&mut pattern); + } + _ => (), + } + // TODO: insert into AST + Declaration::Val { rec, expr, pattern }; + } + unimplemented!() + } + + fn transform_tyapp(&mut self, _: Span, fun: Symbol, arg: Vec) -> CoreExprKind { + let name = self.instanciated_name(fun, arg); + ExprKind::Symbol { name } + } +} + +use crate::pass::Pass; +impl Pass for Monomorphize { + type Target = TypedCoreContext; + + fn trans(&mut self, context: TypedCoreContext, _: &Config) -> Result { + let mut ast = context.ast; + + let mut collector = InstanceCollector::default(); + collector.traverse_ast(&mut ast); + let mut monomorphizer = Monomorphizer::new(collector.instance_table, self.id.clone()); + let ast = monomorphizer.transform_ast(ast); + let symbol_table = context.symbol_table; + let lang_items = context.lang_items; + + Ok(Context { + symbol_table, + ast, + lang_items, + }) + } +} diff --git a/src/lib.rs b/src/lib.rs index 5fe93d4..8516233 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -38,6 +38,7 @@ pub fn compile_strings(inputs: Vec, config: &Config) -> Result, collect_langitems: ast::CollectLangItems::default(), typing: ast::Typer::default(), resolve_overload: ast::ResolveOverload::default(), + monomorphize: ast::Monomorphize::new(id.clone()), case_simplify: ast::CaseSimplify::new(id.clone()), ast_to_hir: hir::AST2HIR::new(id.clone()), constructor_to_enum: hir::ConstructorToEnum::default(), From 2a8042816729156ec8f0669dea4d5e0d983e7c9e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=CE=BAeen?= <3han5chou7@gmail.com> Date: Mon, 29 Aug 2022 01:15:40 +0900 Subject: [PATCH 4/6] implement monomorphization --- src/ast/monomorphize.rs | 44 +++++++++++++++++++++++++++-------------- src/ast/typing.rs | 2 -- src/ast/util.rs | 6 +++++- src/hir/ast2hir.rs | 2 +- 4 files changed, 35 insertions(+), 19 deletions(-) diff --git a/src/ast/monomorphize.rs b/src/ast/monomorphize.rs index de5ae67..9edd6a3 100644 --- a/src/ast/monomorphize.rs +++ b/src/ast/monomorphize.rs @@ -87,6 +87,7 @@ fn rewrite_type(ty: &mut Type, params: &[TypeId], args: &[Type]) { pub struct Instanciator<'a> { params: &'a [TypeId], args: &'a [Type], + m: &'a mut Monomorphizer, } impl<'a> Traverse for Instanciator<'a> { @@ -135,6 +136,11 @@ impl<'a> Traverse for Instanciator<'a> { D(d) => match *d {}, } } + fn traverse_pat_variable(&mut self, _: Span, name: &mut Symbol) { + if self.m.instance_table.contains_key(&name) { + *name = self.m.instanciated_name(name.clone(), self.args.to_vec()) + } + } } impl Transform for Monomorphizer { @@ -147,49 +153,57 @@ impl Transform for Monomorphizer { let pattern = self.transform_pattern(pattern); let expr = self.transform_expr(expr); + if !matches!(pattern.ty, Type::TyAbs(_, _)) { + return Declaration::Val { rec, pattern, expr }; + } + let binds = pattern .binds() .into_iter() .map(|(n, _)| n) .cloned() .collect::>(); - if binds.iter().all(|b| !self.instance_table.contains_key(b)) { - return Declaration::Val { rec, pattern, expr }; - } + // currently only supports `val variable = expr` assert_eq!(binds.len(), 1); - assert!(matches!(pattern.ty, Type::TyAbs(_, _))); - for args in &self.instance_table[&binds[0]] { + let mut ret = vec![]; + // TODO: no need to clone + for args in self.instance_table[&binds[0]].clone() { let mut expr = expr.clone(); let mut pattern = pattern.clone(); - match expr.ty { + match pattern.ty { Type::TyAbs(params, body) => { - expr.ty = *body; + pattern.ty = *body; + Instanciator { params: ¶ms, args: &args, + m: self, } - .traverse_expr(&mut expr); + .traverse_pattern(&mut pattern); } _ => (), } - match pattern.ty { - Type::TyAbs(params, body) => { - pattern.ty = *body; + match expr.ty { + Type::TyAbs(params, body) => { + expr.ty = *body; Instanciator { params: ¶ms, args: &args, + m: self, } - .traverse_pattern(&mut pattern); + .traverse_expr(&mut expr); } _ => (), } - // TODO: insert into AST - Declaration::Val { rec, expr, pattern }; + ret.push(Declaration::Val { rec, expr, pattern }); + } + Declaration::Local { + binds: vec![], + body: ret, } - unimplemented!() } fn transform_tyapp(&mut self, _: Span, fun: Symbol, arg: Vec) -> CoreExprKind { diff --git a/src/ast/typing.rs b/src/ast/typing.rs index 7630f00..2a5680c 100644 --- a/src/ast/typing.rs +++ b/src/ast/typing.rs @@ -403,10 +403,8 @@ impl TyEnv { TypeErrorKind::PolymorphicExpression, )), (vars, true) => { - println!("generalized {:?}", ty); let cloned = self.pool.ty(self.pool.pool.value_of(ty).clone()); let new_ty = self.pool.ty(Typing::TyAbs(vars, cloned)); - println!("new_ty {:?}", new_ty); self.pool.set(ty, new_ty); Ok(()) } diff --git a/src/ast/util.rs b/src/ast/util.rs index 3e51c27..c52faed 100644 --- a/src/ast/util.rs +++ b/src/ast/util.rs @@ -176,7 +176,11 @@ pub trait Traverse { _arg: &mut Option>>, ) { } - fn traverse_pat_tuple(&mut self, _: Span, _tuple: &mut Vec>) {} + fn traverse_pat_tuple(&mut self, _: Span, tuple: &mut Vec>) { + for t in tuple { + self.traverse_pattern(t) + } + } fn traverse_pat_variable(&mut self, _: Span, _value: &mut Symbol) {} fn traverse_pat_wildcard(&mut self, _: Span) {} } diff --git a/src/hir/ast2hir.rs b/src/hir/ast2hir.rs index 1a20883..c8bc2f0 100644 --- a/src/hir/ast2hir.rs +++ b/src/hir/ast2hir.rs @@ -54,7 +54,7 @@ fn conv_ty(ty: ast::Type) -> HTy { Tuple(tys) => HTy::Tuple(tys.into_iter().map(conv_ty).collect()), Fun(arg, ret) => HTy::fun(conv_ty(*arg), conv_ty(*ret)), Datatype(name) => HTy::Datatype(name), - TyAbs(_, _) | Variable(_) => panic!("polymorphism is not supported yet"), + TyAbs(_, _) | Variable(_) => panic!("polymorphism is not supported yet: {:?}", ty), } } From 588cb6439b0b341ae188dcf5264350f3d0485b9e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=CE=BAeen?= <3han5chou7@gmail.com> Date: Mon, 29 Aug 2022 01:16:04 +0900 Subject: [PATCH 5/6] add an example of polymorphic functions --- ml_example/polymorphic_function.sml | 3 +++ 1 file changed, 3 insertions(+) create mode 100644 ml_example/polymorphic_function.sml diff --git a/ml_example/polymorphic_function.sml b/ml_example/polymorphic_function.sml new file mode 100644 index 0000000..d3a8353 --- /dev/null +++ b/ml_example/polymorphic_function.sml @@ -0,0 +1,3 @@ +fun id x = x +val x = id 1 +val y = id false From d2be6b16dbd9bf18e06c0cc670f1419d8948ff7e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=CE=BAeen?= <3han5chou7@gmail.com> Date: Wed, 21 Sep 2022 16:21:31 +0900 Subject: [PATCH 6/6] correctly resolve polymorphisms --- ml_example/online.sml | 5 ++ ml_example/polymorphic_function.sml | 6 ++ ml_example/stringFromInt2digits.sml | 1 + src/ast/monomorphize.rs | 134 ++++++++++++++++++++++++++-- 4 files changed, 137 insertions(+), 9 deletions(-) create mode 100644 ml_example/online.sml create mode 100644 ml_example/stringFromInt2digits.sml diff --git a/ml_example/online.sml b/ml_example/online.sml new file mode 100644 index 0000000..c3fdfa4 --- /dev/null +++ b/ml_example/online.sml @@ -0,0 +1,5 @@ +fun fib 0 = 1 +| fib 1 = 1 +| fib n = fib (n - 1) + fib (n - 2) + +val () = printInt (fib 39) diff --git a/ml_example/polymorphic_function.sml b/ml_example/polymorphic_function.sml index d3a8353..97601cf 100644 --- a/ml_example/polymorphic_function.sml +++ b/ml_example/polymorphic_function.sml @@ -1,3 +1,9 @@ fun id x = x val x = id 1 val y = id false +val id2 = id +val x = id2 1 +val y = id2 false +val id3 = id2 +val x = id3 1 +val y = id3 false diff --git a/ml_example/stringFromInt2digits.sml b/ml_example/stringFromInt2digits.sml new file mode 100644 index 0000000..cfa5375 --- /dev/null +++ b/ml_example/stringFromInt2digits.sml @@ -0,0 +1 @@ +val s = stringFromInt 9 diff --git a/src/ast/monomorphize.rs b/src/ast/monomorphize.rs index 9edd6a3..c812e35 100644 --- a/src/ast/monomorphize.rs +++ b/src/ast/monomorphize.rs @@ -19,15 +19,120 @@ impl Monomorphize { #[derive(Debug, Default)] struct InstanceCollector { - instance_table: HashMap>>, + root_set: HashMap>>, + blocked: HashMap>>, + param_dependencies: HashMap>, + symbol_params: HashMap>, } impl Traverse for InstanceCollector { + fn traverse_val( + &mut self, + _: &mut bool, + pattern: &mut CorePattern, + expr: &mut CoreExpr, + ) { + let params = match &pattern.ty { + Type::TyAbs(params, _) => params.clone(), + _ => vec![], + }; + let mut binds = pattern + .binds() + .into_iter() + .map(|(n, _)| n) + .cloned() + .collect::>(); + // currently only supports `val variable = expr` + if binds.len() == 1 { + let name = binds.remove(0); + self.symbol_params.insert(name, params); + } + + // the original traverse_val + self.traverse_expr(expr); + self.traverse_pattern(pattern) + } + fn traverse_tyapp(&mut self, _: Span, fun: &mut Symbol, arg: &mut Vec) { - self.instance_table - .entry(fun.clone()) - .or_default() - .insert(arg.clone()); + if arg.iter().any(|t| matches! {t, Type::Variable(_)}) { + self.blocked + .entry(fun.clone()) + .or_default() + .insert(arg.clone()); + let params = arg.iter().filter_map(|t| match t { + Type::Variable(id) => Some(id), + _ => None, + }); + for param in params { + self.param_dependencies + .entry(param.clone()) + .or_default() + .insert(fun.clone()); + } + } else { + self.root_set + .entry(fun.clone()) + .or_default() + .insert(arg.clone()); + } + } +} + +impl InstanceCollector { + fn generate_instance_table(self) -> HashMap>> { + let mut ret: HashMap>> = HashMap::new(); + let InstanceCollector { + mut root_set, + mut blocked, + mut param_dependencies, + symbol_params, + } = self; + + // resolve all the parameter dependencies + while !root_set.is_empty() { + let mut next_root_set: HashMap>> = HashMap::new(); + for (name, args_set) in root_set { + let params = &symbol_params[&name]; + for args in args_set.clone() { + assert_eq!(params.len(), args.len()); + for (param, arg) in params.iter().zip(args) { + let dependencies = match param_dependencies.remove(param) { + Some(ds) => ds, + None => continue, + }; + for dep in dependencies { + // Because hash value may change, we cannot modify values in hashset. + // Thus, we remove onece and re-insert + let dep_args_set = blocked.remove(&dep).expect("must exist"); + let mut next_dep_args_set = HashSet::new(); + for mut dep_args in dep_args_set { + for dep_arg in &mut dep_args { + if dep_arg == &Type::Variable(*param) { + *dep_arg = arg.clone() + } + } + if dep_args.iter().all(|a| !matches! {a, Type::Variable(_)}) { + next_root_set + .entry(dep.clone()) + .or_default() + .insert(dep_args); + } else { + next_dep_args_set.insert(dep_args); + } + } + blocked.insert(dep, next_dep_args_set); + } + } + } + let set = ret.entry(name).or_default(); + for args in args_set { + set.insert(args); + } + } + root_set = next_root_set; + } + assert!(param_dependencies.is_empty()); + ret } } @@ -120,6 +225,12 @@ impl<'a> Traverse for Instanciator<'a> { } } + fn traverse_tyapp(&mut self, _: Span, _: &mut Symbol, args: &mut Vec) { + for arg in args { + rewrite_type(arg, self.params, self.args) + } + } + fn traverse_pattern(&mut self, pattern: &mut CorePattern) { rewrite_type(&mut pattern.ty, self.params, self.args); @@ -150,10 +261,9 @@ impl Transform for Monomorphizer { pattern: CorePattern, expr: CoreExpr, ) -> CoreDeclaration { - let pattern = self.transform_pattern(pattern); - let expr = self.transform_expr(expr); - if !matches!(pattern.ty, Type::TyAbs(_, _)) { + let pattern = self.transform_pattern(pattern); + let expr = self.transform_expr(expr); return Declaration::Val { rec, pattern, expr }; } @@ -198,8 +308,12 @@ impl Transform for Monomorphizer { } _ => (), } + let pattern = self.transform_pattern(pattern); + let expr = self.transform_expr(expr); + ret.push(Declaration::Val { rec, expr, pattern }); } + Declaration::Local { binds: vec![], body: ret, @@ -207,6 +321,7 @@ impl Transform for Monomorphizer { } fn transform_tyapp(&mut self, _: Span, fun: Symbol, arg: Vec) -> CoreExprKind { + // mark let name = self.instanciated_name(fun, arg); ExprKind::Symbol { name } } @@ -221,7 +336,8 @@ impl Pass for Monomorphize { let mut collector = InstanceCollector::default(); collector.traverse_ast(&mut ast); - let mut monomorphizer = Monomorphizer::new(collector.instance_table, self.id.clone()); + let mut monomorphizer = + Monomorphizer::new(collector.generate_instance_table(), self.id.clone()); let ast = monomorphizer.transform_ast(ast); let symbol_table = context.symbol_table; let lang_items = context.lang_items;