Feature to optimize function expressions

This PR adds a new feature which enables additional validation/optimization of 
function expressions, gated by the `enableOptimizeFunctionExpressions` feature 
flag. When disabled, we actually revert the changes earlier in this stack, and 
do all our lowering of function expressions in AnalyzeFunctions. 

When the feature is enabled, we incrementally process function expressions in 
the various compilation stages, eg InferTypes infers into function expressions, 
ConstantPropagation propagates constants into function expressions, etc. Because 
this stage optimizes function expressions, in this mode codegen uses the HIR as 
the source rather than the original babel node. 

The feature is disabled by default so it has no impact on generated code. For 
now i've enabled the feature on just one test to demonstrate constant 
propagation into a function expression.
This commit is contained in:
Joe Savona
2023-06-08 18:05:41 -04:00
parent 4d7d2cf4dc
commit a8fff7cc5c
13 changed files with 194 additions and 53 deletions
@@ -7,6 +7,7 @@
import * as t from "@babel/types";
import invariant from "invariant";
import { ExternalFunction } from "../Entrypoint/Options";
import { log } from "../Utils/logger";
import {
DEFAULT_GLOBALS,
@@ -33,7 +34,6 @@ import {
ShapeRegistry,
addHook,
} from "./ObjectShape";
import { ExternalFunction } from "../Entrypoint/Options";
export type Hook = {
effectKind: Effect;
@@ -150,6 +150,15 @@ export type EnvironmentConfig = Partial<{
* }
*/
enableEmitFreeze: ExternalFunction | null;
/**
* When enabled, function expression codegen uses a subset of the compiler pipeline
* to transform and optimize their contents. When disabled, function expression
* codegen uses the original, un-transformed function body.
*
* Defaults to false (use the un-transformed function body).
*/
enableCodegenLoweredFunctionExpressions: boolean;
}>;
export class Environment {
@@ -165,6 +174,7 @@ export class Environment {
enableTreatHooksAsFunctions: boolean;
disableAllMemoization: boolean;
enableEmitFreeze: ExternalFunction | null;
enableCodegenLoweredFunctionExpressions: boolean;
#contextIdentifiers: Set<t.Identifier>;
@@ -208,6 +218,8 @@ export class Environment {
config?.enableTreatHooksAsFunctions ?? true;
this.disableAllMemoization = config?.disableAllMemoization ?? false;
this.enableEmitFreeze = config?.enableEmitFreeze ?? null;
this.enableCodegenLoweredFunctionExpressions =
config?.enableCodegenLoweredFunctionExpressions ?? false;
this.#contextIdentifiers = contextIdentifiers;
}
@@ -31,9 +31,11 @@ import { mapOptionalFallthroughs } from "./visitors";
export function mergeConsecutiveBlocks(fn: HIRFunction): void {
const merged = new MergedBlocks();
for (const [, block] of fn.body.blocks) {
for (const instr of block.instructions) {
if (instr.value.kind === "FunctionExpression") {
mergeConsecutiveBlocks(instr.value.loweredFunc);
if (fn.env.enableCodegenLoweredFunctionExpressions) {
for (const instr of block.instructions) {
if (instr.value.kind === "FunctionExpression") {
mergeConsecutiveBlocks(instr.value.loweredFunc);
}
}
}
@@ -13,9 +13,13 @@ import {
Identifier,
isRefValueType,
isUseRefType,
mergeConsecutiveBlocks,
Place,
ReactiveScopeDependency,
} from "../HIR";
import { constantPropagation } from "../Optimization";
import { eliminateRedundantPhi, enterSSA } from "../SSA";
import { inferTypes } from "../TypeInference";
import { logHIRFunction } from "../Utils/logger";
import { inferMutableRanges } from "./InferMutableRanges";
import inferReferenceEffects from "./InferReferenceEffects";
@@ -86,6 +90,14 @@ export default function analyseFunctions(func: HIRFunction): void {
}
function lower(func: HIRFunction): void {
if (!func.env.enableCodegenLoweredFunctionExpressions) {
mergeConsecutiveBlocks(func);
enterSSA(func);
eliminateRedundantPhi(func);
constantPropagation(func);
inferTypes(func);
}
analyseFunctions(func);
inferReferenceEffects(func, { isFunctionExpression: true });
inferMutableRanges(func);
@@ -7,6 +7,7 @@
import { isValidIdentifier } from "@babel/types";
import {
Environment,
GotoVariant,
HIRFunction,
IdentifierId,
@@ -47,7 +48,12 @@ import { eliminateRedundantPhi } from "../SSA";
* pass.
*/
export function constantPropagation(fn: HIRFunction): void {
const haveTerminalsChanged = applyConstantPropagation(fn);
const constants: Constants = new Map();
constantPropagationImpl(fn, constants);
}
function constantPropagationImpl(fn: HIRFunction, constants: Constants): void {
const haveTerminalsChanged = applyConstantPropagation(fn, constants);
if (haveTerminalsChanged) {
// If terminals have changed then blocks may have become newly unreachable.
// Re-run minification of the graph (incl reordering instruction ids)
@@ -80,7 +86,10 @@ export function constantPropagation(fn: HIRFunction): void {
}
}
function applyConstantPropagation(fn: HIRFunction): boolean {
function applyConstantPropagation(
fn: HIRFunction,
constants: Constants
): boolean {
// Track the set of identifiers which are used as dependencies for function expressions
// in order to avoid propagating these constants. This is necessary because the function
// itself will still reference the original value. If the dependency is propagated but the
@@ -99,8 +108,6 @@ function applyConstantPropagation(fn: HIRFunction): boolean {
}
let hasChanges = false;
const constants: Constants = new Map();
for (const [, block] of fn.body.blocks) {
// Initialize phi values if all operands have the same known constant value.
// Note that this analysis uses a single-pass only, so it will never fill in
@@ -137,11 +144,13 @@ function applyConstantPropagation(fn: HIRFunction): boolean {
continue;
}
const instr = block.instructions[i]!;
// Don't propagate constants used as function expression dependencies
if (functionDependencies.has(instr.lvalue.identifier.id)) {
continue;
if (!fn.env.enableCodegenLoweredFunctionExpressions) {
// Don't propagate constants used as function expression dependencies
if (functionDependencies.has(instr.lvalue.identifier.id)) {
continue;
}
}
const value = evaluateInstruction(constants, instr);
const value = evaluateInstruction(fn.env, constants, instr);
if (value !== null) {
constants.set(instr.lvalue.identifier.id, value);
}
@@ -180,6 +189,7 @@ function applyConstantPropagation(fn: HIRFunction): boolean {
}
function evaluateInstruction(
env: Environment,
constants: Constants,
instr: Instruction
): Constant | null {
@@ -350,8 +360,9 @@ function evaluateInstruction(
return placeValue;
}
case "FunctionExpression": {
// TODO: propagate constants in the outer scope into the function when traversing
constantPropagation(value.loweredFunc);
if (env.enableCodegenLoweredFunctionExpressions) {
constantPropagationImpl(value.loweredFunc, constants);
}
return null;
}
default: {
@@ -7,6 +7,7 @@
import * as t from "@babel/types";
import invariant from "invariant";
import { pruneUnusedLValues, pruneUnusedLabels, renameVariables } from ".";
import { CompilerError, ErrorSeverity } from "../CompilerError";
import { Environment } from "../HIR";
import {
@@ -30,8 +31,10 @@ import {
} from "../HIR/HIR";
import { printPlace } from "../HIR/PrintHIR";
import { eachPatternOperand } from "../HIR/visitors";
import { deadCodeElimination } from "../Optimization";
import { Err, Ok, Result } from "../Utils/Result";
import { assertExhaustive } from "../Utils/utils";
import { buildReactiveFunction } from "./BuildReactiveFunction";
export function codegenReactiveFunction(
fn: ReactiveFunction
@@ -957,7 +960,36 @@ function codegenInstructionValue(
break;
}
case "FunctionExpression": {
value = t.cloneNode(instrValue.expr, true, false);
if (cx.env.enableCodegenLoweredFunctionExpressions) {
const loweredFunc = instrValue.loweredFunc;
deadCodeElimination(loweredFunc);
const reactiveFunction = buildReactiveFunction(loweredFunc);
pruneUnusedLabels(reactiveFunction);
pruneUnusedLValues(reactiveFunction);
renameVariables(reactiveFunction);
const fn = codegenReactiveFunction(reactiveFunction).unwrap();
if (instrValue.expr.type === "ArrowFunctionExpression") {
let body: t.BlockStatement | t.Expression = fn.body;
if (body.body.length === 1) {
const stmt = body.body[0]!;
if (stmt.type === "ReturnStatement" && stmt.argument != null) {
body = stmt.argument;
}
}
value = t.arrowFunctionExpression(fn.params, body, fn.async);
} else {
value = t.functionExpression(
fn.id ??
(instrValue.name != null ? t.identifier(instrValue.name) : null),
fn.params,
fn.body,
fn.generator,
fn.async
);
}
} else {
value = t.cloneNode(instrValue.expr, true, false);
}
break;
}
case "TaggedTemplateExpression": {
@@ -99,7 +99,11 @@ export function eliminateRedundantPhi(fn: HIRFunction): void {
rewritePlace(instr.lvalue, rewrites);
// visit function expressions on first iteration of each block
if (!hasBackEdge && instr.value.kind === "FunctionExpression") {
if (
!hasBackEdge &&
instr.value.kind === "FunctionExpression" &&
fn.env.enableCodegenLoweredFunctionExpressions
) {
eliminateRedundantPhi(instr.value.loweredFunc);
}
}
@@ -248,11 +248,15 @@ function enterSSAImpl(
if (blockId === rootEntry) {
// NOTE: func.context should be empty for the root function
if (func.context.length !== 0) {
CompilerError.invariant(
`Expected function context to be empty for outer function declarations`,
func.loc
);
if (func.env.enableCodegenLoweredFunctionExpressions) {
if (func.context.length !== 0) {
CompilerError.invariant(
`Expected function context to be empty for outer function declarations`,
func.loc
);
}
} else {
func.context = func.context.map((p) => builder.defineContext(p));
}
func.params = func.params.map((p) => builder.definePlace(p));
}
@@ -261,7 +265,10 @@ function enterSSAImpl(
mapInstructionLValues(instr, (lvalue) => builder.definePlace(lvalue));
mapInstructionOperands(instr, (place) => builder.getPlace(place));
if (instr.value.kind === "FunctionExpression") {
if (
instr.value.kind === "FunctionExpression" &&
func.env.enableCodegenLoweredFunctionExpressions
) {
const loweredFunc = instr.value.loweredFunc;
const entry = loweredFunc.body.blocks.get(loweredFunc.body.entry)!;
invariant(
@@ -67,7 +67,11 @@ function apply(func: HIRFunction, unifier: Unifier): void {
}
const { lvalue, value } = instr;
lvalue.identifier.type = unifier.get(lvalue.identifier.type);
if (value.kind === "FunctionExpression") {
if (
value.kind === "FunctionExpression" &&
func.env.enableCodegenLoweredFunctionExpressions
) {
apply(value.loweredFunc, unifier);
}
}
@@ -247,7 +251,9 @@ function* generateInstructionTypes(
}
case "FunctionExpression": {
yield* generate(value.loweredFunc);
if (env.enableCodegenLoweredFunctionExpressions) {
yield* generate(value.loweredFunc);
}
break;
}
@@ -20,42 +20,41 @@ function Component(props) {
import { unstable_useMemoCache as useMemoCache } from "react";
function Component(props) {
const $ = useMemoCache(7);
const t0 = props.x;
const c_0 = $[0] !== t0;
let t1;
const c_0 = $[0] !== props.x;
let t0;
if (c_0) {
t1 = foo(t0);
$[0] = t0;
$[1] = t1;
t0 = foo(props.x);
$[0] = props.x;
$[1] = t0;
} else {
t1 = $[1];
t0 = $[1];
}
const x = t1;
const x = t0;
const c_2 = $[2] !== props;
const c_3 = $[3] !== x;
let t2;
let t1;
if (c_2 || c_3) {
t2 = function () {
t1 = function () {
const arr = [...bar(props)];
return arr.at(x);
};
$[2] = props;
$[3] = x;
$[4] = t2;
$[4] = t1;
} else {
t2 = $[4];
t1 = $[4];
}
const fn = t2;
const fn = t1;
const c_5 = $[5] !== fn;
let t3;
let t2;
if (c_5) {
t3 = fn();
t2 = fn();
$[5] = fn;
$[6] = t3;
$[6] = t2;
} else {
t3 = $[6];
t2 = $[6];
}
const fnResult = t3;
const fnResult = t2;
return fnResult;
}
@@ -0,0 +1,43 @@
## Input
```javascript
// @enableCodegenLoweredFunctionExpressions
function Component(props) {
const x = 42;
const onEvent = () => {
console.log(x);
};
return <Foo onEvent={onEvent} />;
}
```
## Code
```javascript
import { unstable_useMemoCache as useMemoCache } from "react"; // @enableCodegenLoweredFunctionExpressions
function Component(props) {
const $ = useMemoCache(2);
let t0;
if ($[0] === Symbol.for("react.memo_cache_sentinel")) {
t0 = () => {
console.log(42);
};
$[0] = t0;
} else {
t0 = $[0];
}
const onEvent = t0;
let t1;
if ($[1] === Symbol.for("react.memo_cache_sentinel")) {
t1 = <Foo onEvent={onEvent} />;
$[1] = t1;
} else {
t1 = $[1];
}
return t1;
}
```
@@ -0,0 +1,8 @@
// @enableCodegenLoweredFunctionExpressions
function Component(props) {
const x = 42;
const onEvent = () => {
console.log(x);
};
return <Foo onEvent={onEvent} />;
}
@@ -38,8 +38,8 @@ function Component(props) {
if (c_0) {
const allUrls = [];
const { media: t84, comments, urls } = post;
media = t84;
const { media: t85, comments, urls } = post;
media = t85;
const c_3 = $[3] !== comments.length;
let t0;
if (c_3) {
@@ -98,40 +98,44 @@ export async function compile(
let disableAllMemoization = false;
let validateRefAccessDuringRender = true;
let enableEmitFreeze = null;
let enableCodegenLoweredFunctionExpressions = false;
if (firstLine.indexOf("@forgetDirective") !== -1) {
enableOnlyOnUseForgetDirective = true;
}
if (firstLine.indexOf("@gating") !== -1) {
if (firstLine.includes("@gating")) {
gating = {
source: "ReactForgetFeatureFlag",
importSpecifierName: "isForgetEnabled_Fixtures",
};
}
if (firstLine.indexOf("@instrumentForget") !== -1) {
if (firstLine.includes("@instrumentForget")) {
instrumentForget = {
source: "react-forget-runtime",
importSpecifierName: "useRenderCounter",
};
}
if (firstLine.indexOf("@panicOnBailout false") !== -1) {
if (firstLine.includes("@panicOnBailout false")) {
panicOnBailout = false;
}
if (firstLine.indexOf("@memoizeJsxElements false") !== -1) {
if (firstLine.includes("@memoizeJsxElements false")) {
memoizeJsxElements = false;
}
if (firstLine.indexOf("@enableAssumeHooksFollowRulesOfReact true") !== -1) {
if (firstLine.includes("@enableAssumeHooksFollowRulesOfReact true")) {
enableAssumeHooksFollowRulesOfReact = true;
}
if (firstLine.indexOf("@enableTreatHooksAsFunctions false") !== -1) {
if (firstLine.includes("@enableTreatHooksAsFunctions false")) {
enableTreatHooksAsFunctions = false;
}
if (firstLine.indexOf("@disableAllMemoization true") !== -1) {
if (firstLine.includes("@disableAllMemoization true")) {
disableAllMemoization = true;
}
if (firstLine.indexOf("@validateRefAccessDuringRender false") !== -1) {
if (firstLine.includes("@validateRefAccessDuringRender false")) {
validateRefAccessDuringRender = false;
}
if (firstLine.indexOf("@enableEmitFreeze") !== -1) {
if (firstLine.includes("@enableCodegenLoweredFunctionExpressions")) {
enableCodegenLoweredFunctionExpressions = true;
}
if (firstLine.includes("@enableEmitFreeze")) {
enableEmitFreeze = {
source: "react-forget-runtime",
importSpecifierName: "makeReadOnly",
@@ -162,6 +166,7 @@ export async function compile(
validateRefAccessDuringRender,
validateFrozenLambdas: true,
enableEmitFreeze,
enableCodegenLoweredFunctionExpressions,
},
logger: null,
gating,