Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions hkmc2/shared/src/main/scala/hkmc2/Config.scala
Original file line number Diff line number Diff line change
Expand Up @@ -163,6 +163,7 @@ object Config:
debugEta: Bool = false,
debugDpe: Bool = false,
debugDce: Bool = false,
logEffects: Bool = false,
):
def effectiveDebugEta: Bool = debug || debugEta
def effectiveDebugDpe: Bool = debug || debugDpe
Expand Down Expand Up @@ -541,6 +542,7 @@ object ConfigParser:
var debugEta = base.debugEta
var debugDpe = base.debugDpe
var debugDce = base.debugDce
var logEffects = base.logEffects
args.foreach:
case NamedArg("debug", value) =>
setFrom(value)(parseBool)(v => debug = v)
Expand All @@ -560,6 +562,8 @@ object ConfigParser:
setFrom(value)(parseBool)(v => debugDpe = v)
case NamedArg("debugDce", value) =>
setFrom(value)(parseBool)(v => debugDce = v)
case NamedArg("logEffects", value) =>
setFrom(value)(parseBool)(v => logEffects = v)
case other =>
unsupported(passName, other)
S(Config.FlowAnalysisConfig(
Expand All @@ -572,6 +576,7 @@ object ConfigParser:
debugEta,
debugDpe,
debugDce,
logEffects,
))
case _ =>
expect(s"${passName}(...)")(tree)
Expand Down
4 changes: 2 additions & 2 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/Block.scala
Original file line number Diff line number Diff line change
Expand Up @@ -1153,13 +1153,13 @@ sealed abstract class Result extends Located, HasErasedType:
ErrorReport(message -> loc :: Nil, source = Diagnostic.Source.Compilation)
this

/* mayRaiseEffects indicates whether this call may raise effect (algebraic effect),
/* mayHaveEffects indicates whether this call may raise effect (algebraic effect),
* regardless of whether the check for effect is inserted or not.
* Note that the check for effect is inserted during HandlerLowering and setting this to true
* after handler is lowered does not have any effect on the code generation. */
case class CallMetadata(
isMlsFun: Bool,
mayRaiseEffects: Bool,
mayHaveEffects: Bool,
annotations: Ls[Annot],
):
lazy val explicitTailCall: Bool = annotations.exists(_.isInstanceOf[Annot.TailCall])
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1200,7 +1200,7 @@ class BlockSimplifier
val combined = Call(prefix.fun, (prefix.argss ::: argss).ne_!)(
CallMetadata(
prefix.metadata.isMlsFun,
prefix.metadata.mayRaiseEffects || c.metadata.mayRaiseEffects,
prefix.metadata.mayHaveEffects || c.metadata.mayHaveEffects,
prefix.metadata.annotations ++ c.metadata.annotations,
),
c.toLoc)
Expand Down
10 changes: 7 additions & 3 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/EtaExpansion.scala
Original file line number Diff line number Diff line change
Expand Up @@ -91,7 +91,7 @@ class EtaExpansionSolver(val constraintSolver: FlowConstraintSolver, tl: TraceLo
EtaTargets(a.paramCount, a.hasRestParam, a.prodFuns ++ b.prodFuns)
go(S(mergedRes))
else Nil
case UnknownProd => Nil
case UnknownProd | UnsafeEta => Nil
case _: Ctor => Nil
end go

Expand All @@ -101,7 +101,7 @@ class EtaExpansionSolver(val constraintSolver: FlowConstraintSolver, tl: TraceLo
case _ => false
then go(N)
else Nil
case UnknownProd => Nil
case UnknownProd | UnsafeEta => Nil
case _: Ctor => Nil
end funResShape

Expand All @@ -110,7 +110,11 @@ class EtaExpansionSolver(val constraintSolver: FlowConstraintSolver, tl: TraceLo
case N =>
val targets = EtaTargets(pf.params.size, pf.restParam.isDefined, Set.single(pf))
if !processing.contains(pf) then
val res = targets :: funResShape(pf.res)
// It'd be unsound if eta-expansion postpones the evaluation of a
// function if it may have or observe (side) effects.
// TODO: function divergence?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't consider divergence a side effect, but rather undefined behavior. A nonterminating program is incorrect, and rhe compiler may optimize its divergence away.

val safe = pf.effect.exists(EffectAnalysis.summarize(_) == EffectSummary.Pure)
val res = targets :: (if safe then funResShape(pf.res) else Nil)
cache(pf) = res
res
else
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -113,7 +113,11 @@ class FlowAnalysisBasedRewrite(
val body2 = withEliminatedParams(removed):
withEtaArgss(etaParams.map(_.args)):
applyFunBodyLikeBlock(body)
params2 -> body2
// The appended stage may be effectful even when the original producer was pure.
val annotations2 =
if etaParams.nonEmpty then fun.annotations.filterNot(_.isInstanceOf[Annot.Pure])
else fun.annotations
(params2, body2, annotations2)


// traversal
Expand All @@ -138,7 +142,9 @@ class FlowAnalysisBasedRewrite(
Return(etaCall(p).withLocOf(res2))
case c @ Call(fun, argss) =>
Return(
Call(fun, (argss ++ activeEtaArgss).ne_!)(c.metadata, c.toLoc))
Call(fun, (argss ++ activeEtaArgss).ne_!)(c.metadata.copy(
mayHaveEffects = true,
annotations = c.metadata.annotations.filterNot(_.isInstanceOf[Annot.Pure])), c.toLoc))
case _ =>
val tmp = TempSymbol(N, erasedType = N, "eta$res")
Scoped(
Expand Down Expand Up @@ -206,9 +212,9 @@ class FlowAnalysisBasedRewrite(
super.applyObjBody(defn)

override def applyFunDefn(fun: FunDefn): FunDefn =
val (params2, body2) = rewriteFunDefn(fun)
if (params2 is fun.params) && (body2 is fun.body) then fun
else FunDefn(fun.owner, fun.sym, fun.dSym, params2, body2)(fun.configOverride, fun.annotations)
val (params2, body2, annotations2) = rewriteFunDefn(fun)
if (params2 is fun.params) && (body2 is fun.body) && (annotations2 is fun.annotations) then fun
else FunDefn(fun.owner, fun.sym, fun.dSym, params2, body2)(fun.configOverride, annotations2)

override def applyLam(lam: Lambda): Lambda =
val lamId: ConcreteFunId = ConcreteId(lam.uid, instId)
Expand Down Expand Up @@ -236,9 +242,9 @@ class FlowAnalysisBasedRewrite(
new Rewriter(Nil):
override def applyFunDefn(fun: FunDefn): FunDefn =
rewrittenInPlace.get(fun.dSym) match
case S((params, body)) =>
if (params is fun.params) && (body is fun.body) then fun
else FunDefn(fun.owner, fun.sym, fun.dSym, params, body)(fun.configOverride, fun.annotations)
case S((params, body, annotations)) =>
if (params is fun.params) && (body is fun.body) && (annotations is fun.annotations) then fun
else FunDefn(fun.owner, fun.sym, fun.dSym, params, body)(fun.configOverride, annotations)
case N => super.applyFunDefn(fun)

def mkPolyFunCopy(
Expand All @@ -257,7 +263,7 @@ class FlowAnalysisBasedRewrite(
case _ => super.applyValue(v)(k)
end RefreshSymbol

val (rewrittenParams, rewrittenBody) = rewritten
val (rewrittenParams, rewrittenBody, rewrittenAnnotations) = rewritten
val refreshParamMap = MutMap.empty[Symbol, Symbol]
def refreshParam(p: Param): Param =
val newSym = new VarSymbol(Tree.Ident(p.sym.name), erasedType = p.sym.erasedType)
Expand All @@ -269,7 +275,7 @@ class FlowAnalysisBasedRewrite(
FunDefn(
N, bms, tSym, refreshedParams,
new RefreshSymbol(refreshParamMap.toMap).apply(rewrittenBody))(
original.configOverride, original.annotations)
original.configOverride, rewrittenAnnotations)
end mkPolyFunCopy

def otherNewFunDefns: Iterable[FunDefn] = Nil
Expand Down Expand Up @@ -314,7 +320,7 @@ object FlowAnalysisBasedRewrite:
)

val etaExpansionSolver =
if eta then new EtaExpansionSolver(
if eta && cfg.liftDefns.isDefined then new EtaExpansionSolver(
flowAnalysisRes, mkTl("eta-expansion > ", optCfg.effectiveDebugEta))
else NoEtaExpansion
val deadParamElimSolver =
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,7 @@ object HandlerLowering:

object EffectfulResult:
def unapply(r: Result)(using Config): Bool = r match
case c: Call if c.metadata.mayRaiseEffects => true
case c: Call if c.metadata.mayHaveEffects => true
case _: Instantiate if config.checkInstantiateEffect => true
case _ => false

Expand Down
8 changes: 6 additions & 2 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/Lifter.scala
Original file line number Diff line number Diff line change
Expand Up @@ -1088,14 +1088,18 @@ class Lifter(topLevelBlk: Block)(using State, Raise, Config):

val call = Call(fun.sym.asMemberRef(fun.dSym), args ne_:: Nil)(CallMetadata.mlsFunWithEffect, N)
val bod = Return(call)
// The new capture list shifts the stages described by affine annotations.
val annotations = fun.annotations.map:
case a @ Annot.Affine(n) => Annot.Affine(n + 1)(a.toLoc)
case a => a

FunDefn(
N,
auxSym,
auxDsym,
newPlists,
bod
)(N, if fun.noInline then fun.annotations else Annot.Inline()(N) :: fun.annotations)
)(N, if fun.noInline then annotations else Annot.Inline()(N) :: annotations)

private val aux = Lazy[Defn](mkAuxDefn)

Expand Down Expand Up @@ -1275,7 +1279,7 @@ class Lifter(topLevelBlk: Block)(using State, Raise, Config):
Call.raw(
flattenedSym.asMemberRef(flattenedDSym),
(formatArgs :: argss).ne_!
)(c.metadata.copy(isMlsFun = true, mayRaiseEffects = false), c.toLoc)
)(c.metadata.copy(isMlsFun = true, mayHaveEffects = false), c.toLoc)
if isTrivial then
if c.argss is argss then k(c)
else k(c.copy(argss = argss)(c.metadata, c.toLoc))
Expand Down
16 changes: 8 additions & 8 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/Lowering.scala
Original file line number Diff line number Diff line change
Expand Up @@ -458,27 +458,27 @@ class Lowering()(using Config, TL, Raise, State, Ctx, SymbolPrinter):
* trying to group as many as possible into a single one
* when they correspond to parameter lists of the same callee. */
def lowerMultiCall(fr: Path, isMlsFun: Bool, annotations: Ls[Annot], args: Ls[Term], loc: Opt[Loc])(k: Result => Block)(using LoweringCtx): Block =
def zipArgs(remainingParamss: Ls[ParamList], remainingArgss: Ls[Term], acc: Ls[Ls[Arg]], mayRaiseEffects: Bool): Block =
def zipArgs(remainingParamss: Ls[ParamList], remainingArgss: Ls[Term], acc: Ls[Ls[Arg]], mayHaveEffects: Bool): Block =
(remainingParamss, remainingArgss) match
case (ps :: remainingParams, args :: remainingArgs) =>
lowerArgs(args, expectedParamTypes(ps))(as => zipArgs(remainingParams, remainingArgs, as :: acc, mayRaiseEffects))
lowerArgs(args, expectedParamTypes(ps))(as => zipArgs(remainingParams, remainingArgs, as :: acc, mayHaveEffects))
case (Nil, Nil) =>
k(Call(fr, acc.reverse.ne_!)(CallMetadata(isMlsFun, mayRaiseEffects, annotations), loc))
k(Call(fr, acc.reverse.ne_!)(CallMetadata(isMlsFun, mayHaveEffects, annotations), loc))
case (Nil, args :: remainingArgss) =>
acc.reverse match
case Nil => lowerRemainingCalls(fr, args, remainingArgss, annotations, loc)(k)
case acc: NELs[Ls[Arg]] =>
val call = Call(fr, acc)(CallMetadata(isMlsFun, mayRaiseEffects, Nil), loc)
val call = Call(fr, acc)(CallMetadata(isMlsFun, mayHaveEffects, Nil), loc)
val tmp = loweringCtx.registerTempSymbol(N, erasedType = call.erasedValueType, "baseCall")
Assign(tmp, call, lowerRemainingCalls(tmp.asSimpleRef, args, remainingArgss, annotations, loc)(k))
case (_ :: _, Nil) =>
k(Call(fr, acc.reverse.ne_!)(CallMetadata(isMlsFun, mayRaiseEffects, annotations), loc))
k(Call(fr, acc.reverse.ne_!)(CallMetadata(isMlsFun, mayHaveEffects, annotations), loc))
fr.targetSymbol match
case S(fs: TermSymbol) =>
fs.defn match
case S(td: TermDefinition) =>
zipArgs(td.params, args, Nil, fs.mayRaiseEffects)
case _ => zipArgs(Nil, args, Nil, fs.mayRaiseEffects)
zipArgs(td.params, args, Nil, fs.mayHaveEffects)
case _ => zipArgs(Nil, args, Nil, fs.mayHaveEffects)
case _ => zipArgs(Nil, args, Nil, true)

def lowerRemainingCalls(base: Path, args: Term, remainingArgss: Ls[Term], annotations: Ls[Annot], loc: Opt[Loc])
Expand Down Expand Up @@ -679,7 +679,7 @@ class Lowering()(using Config, TL, Raise, State, Ctx, SymbolPrinter):
if isImplicitNullaryCall(td.tsym) then
return k(Call(
bs.asMemberRef(disamb.get).withLocOf(ref), Nil ne_:: Nil
)(CallMetadata(isMlsFun = true, mayRaiseEffects = true, annots), ref.toLoc))
)(CallMetadata(isMlsFun = true, mayHaveEffects = true, annots), ref.toLoc))
case S(td: TermDefinition) =>
td.tsym.owner match
case S(owner) =>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ abstract class FlowAnalysisSolverResult:
end FlowAnalysisSolverResult


type RewrittenFunDefn = (params: Ls[ParamList], body: Block)
type RewrittenFunDefn = (params: Ls[ParamList], body: Block, annotations: Ls[Annot])


abstract class PolyInstantiationRewrite(val constraintSolver: FlowConstraintSolver):
Expand Down
10 changes: 5 additions & 5 deletions hkmc2/shared/src/main/scala/hkmc2/codegen/deforest/Rewrite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -325,7 +325,7 @@ class DeforestRewriter(val solver: DeforestFusionSolver)(using Raise)
private class Rewriter(instId: InstantiationId) extends InstantiationRewriter(instId):

override def rewriteFunDefn(fun: FunDefn): RewrittenFunDefn =
fun.params -> applyBlock(fun.body)
(fun.params, applyBlock(fun.body), fun.annotations)

override def applyResult(r: Result)(k: Result => Block): Block =
r match
Expand Down Expand Up @@ -477,7 +477,7 @@ class DeforestRewriter(val solver: DeforestFusionSolver)(using Raise)
tSym: TermSymbol,
rewritten: RewrittenFunDefn,
): FunDefn =
val (rewrittenParams, rewrittenBody) = rewritten
val (rewrittenParams, rewrittenBody, rewrittenAnnotations) = rewritten
// refresh other local symbols: for funs, we can check existing scoped blocks and
// there is no need to add scoped blocks, because function bodies now already are scoped
val refreshParamMap = MutMap.empty[VarSymbol, VarSymbol]
Expand All @@ -492,7 +492,7 @@ class DeforestRewriter(val solver: DeforestFusionSolver)(using Raise)
val bodyWithCorrectSymbols = refreshExtractedBody(refreshParamMap.toMap, rewrittenBody)
FunDefn(
N, bms, tSym, refreshedParams,
bodyWithCorrectSymbols)(N, PrivateModifier :: original.annotations)
bodyWithCorrectSymbols)(N, PrivateModifier :: rewrittenAnnotations)

end mkPolyFunCopy

Expand Down Expand Up @@ -545,10 +545,10 @@ class DeforestRewriter(val solver: DeforestFusionSolver)(using Raise)
new Rewriter(Nil):
override def applyFunDefn(fun: FunDefn): FunDefn =
rewrittenInPlace.get(fun.dSym) match
case Some((params, rewrittenBody)) =>
case Some((params, rewrittenBody, annotations)) =>
// deforest never changes fun params in place
assert(params is fun.params)
FunDefn(fun.owner, fun.sym, fun.dSym, params, rewrittenBody)(fun.configOverride, fun.annotations)
FunDefn(fun.owner, fun.sym, fun.dSym, params, rewrittenBody)(fun.configOverride, annotations)
case None => super.applyFunDefn(fun)
// ====== end: implements the abstract members of `PolyInstantiationRewrite` ======
end DeforestRewriter
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
package hkmc2
package codegen
package flowAnalysis

import hkmc2.semantics.*
import hkmc2.utils.*, shorthands.*
import utils.*


enum EffectSummary:
case Pure, UnsafeEta


case class EffectAnalysisResult(
latentEffects: Map[ConcreteId[FunId], EffectSummary],
callEffects: Map[ConcreteId[ResultId], EffectSummary],
)


object EffectAnalysis:


def summarize(effectVar: StratVar): EffectSummary =
if effectVar.lowerBounds.forall(_.isInstanceOf[StratVar])
then EffectSummary.Pure
else EffectSummary.UnsafeEta

def apply(solver: FlowConstraintSolver)(using tl: TraceLogger): EffectAnalysisResult =
given fState: FlowAnalysis.State = solver.fState
given eState: Elaborator.State = solver.eState


val result = EffectAnalysisResult(
solver.functionEffectVars.iterator.map((id, effect) => id -> summarize(effect)).toMap,
solver.callEffectVars.iterator.map((id, effect) => id -> summarize(effect)).toMap,
)

if tl.doTrace then logResult(result)
result
end apply

private def logResult(result: EffectAnalysisResult)(using
tl: TraceLogger,
fState: FlowAnalysis.State,
eState: Elaborator.State,
): Unit =
def showRefSite(resultId: ResultId): Str =
resultId.getReferredFun match
case Some(fun) => s"${fun.nme}@$resultId"
case None => s"${resultId.getResult}@$resultId"

def showInstId(instId: InstantiationId): Str =
if instId.isEmpty then "<root>" else instId.map(showRefSite).mkString(".")

def showFunction(id: ConcreteId[FunId]): Str =
val name = id.exprId match
case (funSym: TermSymbol, whichParamList) => s"${funSym.nme}#$whichParamList"
case exprId: ResultId => s"lambda@$exprId"
s"function $name @ ${showInstId(id.instId)}"

def showPath(path: Path): Str = path match
case Value.SimpleRef(sym) => sym.nme
case Value.MemberRef(_, disamb) => disamb.nme
case Select(_, name) => name.name
case _ => "<dynamic>"

def showCall(id: ConcreteId[ResultId]): Str =
val call = id.exprId.getResult match
case Call(fun, _) => s"call ${showPath(fun)}@${id.exprId}"
case Instantiate(_, cls, _) => s"instantiate ${showPath(cls)}@${id.exprId}"
case other => s"call $other@${id.exprId}"
s"$call @ ${showInstId(id.instId)}"

tl.log(">>> effect-analysis results >>>")
result.latentEffects.iterator.map: (id, effect) =>
s"${showFunction(id)} -> $effect"
.toSeq.sorted.foreach(tl.log(_))
result.callEffects.iterator.map: (id, effect) =>
s"${showCall(id)} -> $effect"
.toSeq.sorted.foreach(tl.log(_))
tl.log("<<< effect-analysis results <<<")
end logResult
end EffectAnalysis
Loading
Loading