Skip to content

Commit 16a34ad

Browse files
committed
[query] no-sharing in matrix ir lowering
1 parent 1b19d49 commit 16a34ad

21 files changed

Lines changed: 944 additions & 1003 deletions

hail/hail/ir-gen/src/Main.scala

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -582,7 +582,8 @@ object Main {
582582
r += node("I64", in("x", att("Long"))).withTraits(Atom)
583583
r += node("F32", in("x", att("Float"))).withTraits(Atom)
584584
r += node("F64", in("x", att("Double"))).withTraits(Atom)
585-
r += node("Str", in("x", att("String"))).withTraits(Atom)
585+
// Making Str < Atom would lead to code bloat
586+
r += node("Str", in("x", att("String")))
586587
.withPreamble(
587588
"override def toString(): String = s\"\"\"Str(\"${StringEscapeUtils.escapeString(x)}\")\"\"\""
588589
): @nowarn("msg=possible missing interpolator")
@@ -671,7 +672,9 @@ object Main {
671672
)
672673

673674
r += node("MakeArray", in("args", child.*), _typ("TArray")).withCompanionExtension
674-
r += node("MakeStream", in("args", child.*), _typ("TStream"), mmPerElt).withCompanionExtension
675+
r += node("MakeStream", in("args", child.*), _typ("TStream"), mmPerElt)
676+
.typed("TStream")
677+
.withCompanionExtension
675678
r += node("ArrayRef", in("a", child), in("i", child), errorID)
676679
r += node(
677680
"ArraySlice",

hail/hail/src/is/hail/backend/driver/BatchQueryDriver.scala

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -113,7 +113,7 @@ object BatchQueryDriver extends HttpLikeRpc with Logging {
113113
tempFileManager = new OwningTempFileManager(env.fs),
114114
theHailClassLoader = env.hcl,
115115
flags = env.flags,
116-
irMetadata = new IrMetadata(),
116+
irMetadata = new IrMetadata,
117117
blockMatrixCache = ImmutableMap.empty,
118118
compileCache = ImmutableMap.empty,
119119
irCache = ImmutableMap.empty,

hail/hail/src/is/hail/backend/driver/Py4JQueryDriver.scala

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -346,7 +346,7 @@ final class Py4JQueryDriver(backend: Backend) extends Closeable with Logging {
346346
else new OwningTempFileManager(tmpFileManager.fs),
347347
theHailClassLoader = hcl,
348348
flags = flags,
349-
irMetadata = new IrMetadata(),
349+
irMetadata = new IrMetadata,
350350
blockMatrixCache = blockMatrixCache,
351351
compileCache = compiledCodeCache,
352352
irCache = irCache,

hail/hail/src/is/hail/expr/ir/Binds.scala

Lines changed: 26 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package is.hail.expr.ir
22

33
import is.hail.collection.FastSeq
44
import is.hail.collection.compat.immutable.ArraySeq
5+
import is.hail.expr.ir.Scope.AGG
56
import is.hail.expr.ir.defs._
67
import is.hail.types.tcoerce
78
import is.hail.types.virtual._
@@ -224,38 +225,40 @@ object Bindings {
224225
private def childEnvValue(ir: IR, i: Int): Bindings[Type] =
225226
ir match {
226227
case Block(bindings, _) =>
227-
val bindingsTypes = bindings.view.take(i).map(b => b.name -> b.value.typ).to(ArraySeq)
228+
val types = ArraySeq.newBuilder[(Name, Type)]
229+
types.sizeHint(i)
230+
228231
val eval = ArraySeq.newBuilder[Int]
232+
eval.sizeHint(i) // most likely binding in eval
233+
229234
val agg = ArraySeq.newBuilder[Int]
230235
val scan = ArraySeq.newBuilder[Int]
231-
for (k <- 0 until i) bindings(k) match {
232-
case Binding(_, _, Scope.EVAL) =>
233-
eval += k
234-
case Binding(_, _, Scope.AGG) =>
235-
agg += k
236-
case Binding(_, _, Scope.SCAN) =>
237-
scan += k
238-
}
239-
if (i < bindings.length) bindings(i).scope match {
240-
case Scope.EVAL =>
241-
Bindings(
242-
bindingsTypes,
243-
eval.result(),
244-
AggEnv.bindOrNoOp(agg.result()),
245-
AggEnv.bindOrNoOp(scan.result()),
246-
)
247-
case Scope.AGG =>
248-
Bindings(bindingsTypes, agg.result(), AggEnv.Promote, AggEnv.bindOrNoOp(scan.result()))
249-
case Scope.SCAN =>
250-
Bindings(bindingsTypes, scan.result(), AggEnv.bindOrNoOp(agg.result()), AggEnv.Promote)
236+
237+
for (k <- 0 until i) {
238+
val Binding(name, value, scope) = bindings(k)
239+
types += name -> value.typ
240+
scope match {
241+
case Scope.EVAL =>
242+
eval += k
243+
case Scope.AGG =>
244+
agg += k
245+
case Scope.SCAN =>
246+
scan += k
247+
}
251248
}
252-
else
249+
250+
if (i == bindings.length || bindings(i).scope == Scope.EVAL)
253251
Bindings(
254-
bindingsTypes,
252+
types.result(),
255253
eval.result(),
256254
AggEnv.bindOrNoOp(agg.result()),
257255
AggEnv.bindOrNoOp(scan.result()),
258256
)
257+
else if (bindings(i).scope == AGG)
258+
Bindings(types.result(), agg.result(), AggEnv.Promote, AggEnv.bindOrNoOp(scan.result()))
259+
else // SCAN
260+
Bindings(types.result(), scan.result(), AggEnv.bindOrNoOp(agg.result()), AggEnv.Promote)
261+
259262
case TailLoop(name, args, resultType, _) if i == args.length =>
260263
Bindings(
261264
args.map { case (name, ir) => name -> ir.typ } :+

0 commit comments

Comments
 (0)