Skip to content

Commit fa90f47

Browse files
committed
feat(language): wasm cas simplification, symbolic differentiation, zero-gc reverse-mode ad tape, solvers bridge
1 parent 2f9946a commit fa90f47

6 files changed

Lines changed: 732 additions & 0 deletions

File tree

packages/language/src/bindings/javascript/bindings.ts

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2163,6 +2163,47 @@ export class LspFacade {
21632163
return this.exports.csg_op_difference ? this.exports.csg_op_difference(d1, d2) : Math.max(d1, -d2);
21642164
}
21652165

2166+
/** Simplifies an algebraic expression using CAS rewrite rules and constant folding in WASM. */
2167+
casSimplify(daePtr: number, exprId: number): number {
2168+
return this.exports.cas_export_simplify ? this.exports.cas_export_simplify(daePtr, exprId) : exprId;
2169+
}
2170+
2171+
/** Computes the exact symbolic derivative d(expr) / d(varId) in WASM. */
2172+
casDifferentiate(daePtr: number, exprId: number, targetVarId: number): number {
2173+
return this.exports.cas_export_differentiate
2174+
? this.exports.cas_export_differentiate(daePtr, exprId, targetVarId)
2175+
: 0;
2176+
}
2177+
2178+
/** Creates an Automatic Differentiation Tape instance in WASM. */
2179+
createAdTape(): number {
2180+
return this.exports.tape_create ? this.exports.tape_create() : 0;
2181+
}
2182+
2183+
/** Pushes an elementary operation node to the AD tape. */
2184+
tapePushOp(tapePtr: number, op: number, left: number, right: number, val: number): number {
2185+
return this.exports.tape_pushOp ? this.exports.tape_pushOp(tapePtr, op, left, right, val) : 0;
2186+
}
2187+
2188+
/** Runs the reverse-mode AD pass backwards from rootNode. */
2189+
tapeBackward(tapePtr: number, rootNode: number): void {
2190+
if (this.exports.tape_backward) {
2191+
this.exports.tape_backward(tapePtr, rootNode);
2192+
}
2193+
}
2194+
2195+
/** Retrieves the accumulated gradient for a node on the AD tape. */
2196+
tapeGetGrad(tapePtr: number, nodeIdx: number): number {
2197+
return this.exports.tape_getGrad ? this.exports.tape_getGrad(tapePtr, nodeIdx) : 0;
2198+
}
2199+
2200+
/** Resets the AD tape for the next evaluation pass. */
2201+
tapeReset(tapePtr: number): void {
2202+
if (this.exports.tape_reset) {
2203+
this.exports.tape_reset(tapePtr);
2204+
}
2205+
}
2206+
21662207
/** Formats/unparses the document AST using zero-GC AssemblyScript formatting rules. */
21672208
formatDocument(astRoot: number, preserveFormatting: boolean = false): string {
21682209
if (!this.exports.lsp_formatDocument || !this.exports.lsp_getBinaryBuffer) return "";

packages/language/src/codegen/parser.ts

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@ import {
88
arrayCode,
99
bltCode,
1010
builtins_mathCode,
11+
casCode,
1112
correspondenceCode,
1213
cursorCode,
1314
daeCode,
@@ -29,6 +30,7 @@ import {
2930
recoveryCode,
3031
recoveryConfigCode,
3132
stubCode,
33+
tapeCode,
3234
trigramCode,
3335
} from "../../build/src-gen/runtime-templates.js";
3436
import { generateAliasAnalysis } from "./alias.js";
@@ -1032,6 +1034,8 @@ export function generateParserTables(
10321034
{ filename: "ontology_projection.ts", content: ontology_projectionCode },
10331035
{ filename: "builtins_math.ts", content: builtins_mathCode },
10341036
{ filename: "flattener.ts", content: flattenerCode },
1037+
{ filename: "cas.ts", content: casCode },
1038+
{ filename: "tape.ts", content: tapeCode },
10351039
];
10361040

10371041
if (originalGrammar.typeSystem) {
@@ -1059,6 +1063,8 @@ export function generateParserTables(
10591063
code += extractExports(ontology_projectionCode, "./ontology_projection");
10601064
code += extractExports(builtins_mathCode, "./builtins_math");
10611065
code += extractExports(flattenerCode, "./flattener");
1066+
code += extractExports(casCode, "./cas");
1067+
code += extractExports(tapeCode, "./tape");
10621068

10631069
if (originalGrammar.cfgNodes || originalGrammar.analysis) {
10641070
let layoutContent = generateBlockLayoutConstants();
Lines changed: 204 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,204 @@
1+
/* eslint-disable */
2+
// @ts-nocheck
3+
import { DaeBuilder, ExprKind, BinOp, UnaryOp, EXPR_STRIDE, EXPR_KIND, EXPR_DATA1, EXPR_LEFT, EXPR_RIGHT } from "./dae";
4+
5+
/**
6+
* Computer Algebra System (CAS) & Symbolic Simplification Engine in WASM.
7+
* Implements algebraic rewrite rules, constant folding, and symbolic differentiation over DaeBuilder arena expressions.
8+
*/
9+
10+
export function cas_getRealValue(dae: DaeBuilder, exprId: u32): f64 {
11+
let offset = exprId * EXPR_STRIDE;
12+
let lo = dae.exprData.get(offset + EXPR_DATA1) as u32;
13+
let hi = dae.exprData.get(offset + EXPR_LEFT) as u32;
14+
let bits = ((hi as u64) << 32) | (lo as u64);
15+
return reinterpret<f64>(bits);
16+
}
17+
18+
export function cas_isZero(dae: DaeBuilder, exprId: u32): boolean {
19+
if (exprId >= dae.exprCount) return false;
20+
let offset = exprId * EXPR_STRIDE;
21+
let kind = dae.exprData.get(offset + EXPR_KIND);
22+
if (kind == ExprKind.IntLiteral) {
23+
return (dae.exprData.get(offset + EXPR_DATA1) as i32) == 0;
24+
}
25+
if (kind == ExprKind.RealLiteral) {
26+
return cas_getRealValue(dae, exprId) == 0.0;
27+
}
28+
return false;
29+
}
30+
31+
export function cas_isOne(dae: DaeBuilder, exprId: u32): boolean {
32+
if (exprId >= dae.exprCount) return false;
33+
let offset = exprId * EXPR_STRIDE;
34+
let kind = dae.exprData.get(offset + EXPR_KIND);
35+
if (kind == ExprKind.IntLiteral) {
36+
return (dae.exprData.get(offset + EXPR_DATA1) as i32) == 1;
37+
}
38+
if (kind == ExprKind.RealLiteral) {
39+
return cas_getRealValue(dae, exprId) == 1.0;
40+
}
41+
return false;
42+
}
43+
44+
export function cas_isConstant(dae: DaeBuilder, exprId: u32): boolean {
45+
if (exprId >= dae.exprCount) return false;
46+
let offset = exprId * EXPR_STRIDE;
47+
let kind = dae.exprData.get(offset + EXPR_KIND);
48+
return kind == ExprKind.IntLiteral || kind == ExprKind.RealLiteral;
49+
}
50+
51+
/**
52+
* Recursively simplifies an algebraic expression using rewrite rules.
53+
*/
54+
export function cas_simplify(dae: DaeBuilder, exprId: u32): u32 {
55+
if (exprId >= dae.exprCount) return exprId;
56+
let offset = exprId * EXPR_STRIDE;
57+
let kind = dae.exprData.get(offset + EXPR_KIND);
58+
59+
if (kind == ExprKind.Binary) {
60+
let op = dae.exprData.get(offset + EXPR_DATA1);
61+
let left = cas_simplify(dae, dae.exprData.get(offset + EXPR_LEFT));
62+
let right = cas_simplify(dae, dae.exprData.get(offset + EXPR_RIGHT));
63+
64+
// Constant folding if both operands are numeric constants
65+
if (cas_isConstant(dae, left) && cas_isConstant(dae, right)) {
66+
let vLeft = cas_getRealValue(dae, left);
67+
let vRight = cas_getRealValue(dae, right);
68+
if (op == BinOp.Add) return dae.addRealLiteral(vLeft + vRight);
69+
if (op == BinOp.Sub) return dae.addRealLiteral(vLeft - vRight);
70+
if (op == BinOp.Mul) return dae.addRealLiteral(vLeft * vRight);
71+
if (op == BinOp.Div && vRight != 0.0) return dae.addRealLiteral(vLeft / vRight);
72+
if (op == BinOp.Pow) return dae.addRealLiteral(Math.pow(vLeft, vRight));
73+
}
74+
75+
// Algebraic Rewrite Rules:
76+
if (op == BinOp.Add) {
77+
// x + 0 -> x
78+
if (cas_isZero(dae, right)) return left;
79+
// 0 + x -> x
80+
if (cas_isZero(dae, left)) return right;
81+
} else if (op == BinOp.Sub) {
82+
// x - 0 -> x
83+
if (cas_isZero(dae, right)) return left;
84+
// x - x -> 0
85+
if (left == right) return dae.addRealLiteral(0.0);
86+
} else if (op == BinOp.Mul) {
87+
// x * 0 -> 0 or 0 * x -> 0
88+
if (cas_isZero(dae, left) || cas_isZero(dae, right)) return dae.addRealLiteral(0.0);
89+
// x * 1 -> x
90+
if (cas_isOne(dae, right)) return left;
91+
// 1 * x -> x
92+
if (cas_isOne(dae, left)) return right;
93+
} else if (op == BinOp.Div) {
94+
// 0 / x -> 0
95+
if (cas_isZero(dae, left)) return dae.addRealLiteral(0.0);
96+
// x / 1 -> x
97+
if (cas_isOne(dae, right)) return left;
98+
// x / x -> 1 (when x != 0)
99+
if (left == right) return dae.addRealLiteral(1.0);
100+
} else if (op == BinOp.Pow) {
101+
// x ^ 0 -> 1
102+
if (cas_isZero(dae, right)) return dae.addRealLiteral(1.0);
103+
// x ^ 1 -> x
104+
if (cas_isOne(dae, right)) return left;
105+
}
106+
107+
return dae.addExpression(ExprKind.Binary, op, left, right);
108+
}
109+
110+
if (kind == ExprKind.Unary || kind == ExprKind.Negate) {
111+
let sub = cas_simplify(dae, dae.exprData.get(offset + EXPR_LEFT));
112+
if (cas_isConstant(dae, sub)) {
113+
let val = cas_getRealValue(dae, sub);
114+
return dae.addRealLiteral(-val);
115+
}
116+
// -(-x) -> x
117+
let subOffset = sub * EXPR_STRIDE;
118+
let subKind = dae.exprData.get(subOffset + EXPR_KIND);
119+
if (subKind == ExprKind.Negate || subKind == ExprKind.Unary) {
120+
return dae.exprData.get(subOffset + EXPR_LEFT);
121+
}
122+
return dae.addExpression(ExprKind.Negate, 0, sub, 0xffffffff);
123+
}
124+
125+
return exprId;
126+
}
127+
128+
/**
129+
* Computes exact symbolic derivative of an expression with respect to a variable ID: d(expr) / d(varId).
130+
*/
131+
export function cas_differentiate(dae: DaeBuilder, exprId: u32, targetVarId: u32): u32 {
132+
if (exprId >= dae.exprCount) return dae.addRealLiteral(0.0);
133+
let offset = exprId * EXPR_STRIDE;
134+
let kind = dae.exprData.get(offset + EXPR_KIND);
135+
136+
if (kind == ExprKind.Name) {
137+
let varId = dae.exprData.get(offset + EXPR_DATA1);
138+
// d(x) / dx = 1, d(y) / dx = 0
139+
return varId == targetVarId ? dae.addRealLiteral(1.0) : dae.addRealLiteral(0.0);
140+
}
141+
142+
if (kind == ExprKind.IntLiteral || kind == ExprKind.RealLiteral) {
143+
// d(const) / dx = 0
144+
return dae.addRealLiteral(0.0);
145+
}
146+
147+
if (kind == ExprKind.Binary) {
148+
let op = dae.exprData.get(offset + EXPR_DATA1);
149+
let u = dae.exprData.get(offset + EXPR_LEFT);
150+
let v = dae.exprData.get(offset + EXPR_RIGHT);
151+
let du = cas_differentiate(dae, u, targetVarId);
152+
let dv = cas_differentiate(dae, v, targetVarId);
153+
154+
if (op == BinOp.Add) {
155+
// d(u + v) = du + dv
156+
let add = dae.addExpression(ExprKind.Binary, BinOp.Add, du, dv);
157+
return cas_simplify(dae, add);
158+
}
159+
if (op == BinOp.Sub) {
160+
// d(u - v) = du - dv
161+
let sub = dae.addExpression(ExprKind.Binary, BinOp.Sub, du, dv);
162+
return cas_simplify(dae, sub);
163+
}
164+
if (op == BinOp.Mul) {
165+
// Product rule: d(u * v) = du * v + u * dv
166+
let t1 = dae.addExpression(ExprKind.Binary, BinOp.Mul, du, v);
167+
let t2 = dae.addExpression(ExprKind.Binary, BinOp.Mul, u, dv);
168+
let add = dae.addExpression(ExprKind.Binary, BinOp.Add, t1, t2);
169+
return cas_simplify(dae, add);
170+
}
171+
if (op == BinOp.Div) {
172+
// Quotient rule: d(u / v) = (du * v - u * dv) / (v ^ 2)
173+
let num1 = dae.addExpression(ExprKind.Binary, BinOp.Mul, du, v);
174+
let num2 = dae.addExpression(ExprKind.Binary, BinOp.Mul, u, dv);
175+
let num = dae.addExpression(ExprKind.Binary, BinOp.Sub, num1, num2);
176+
let den = dae.addExpression(ExprKind.Binary, BinOp.Mul, v, v);
177+
let div = dae.addExpression(ExprKind.Binary, BinOp.Div, num, den);
178+
return cas_simplify(dae, div);
179+
}
180+
}
181+
182+
if (kind == ExprKind.Negate || kind == ExprKind.Unary) {
183+
let u = dae.exprData.get(offset + EXPR_LEFT);
184+
let du = cas_differentiate(dae, u, targetVarId);
185+
let neg = dae.addExpression(ExprKind.Negate, 0, du, 0xffffffff);
186+
return cas_simplify(dae, neg);
187+
}
188+
189+
return dae.addRealLiteral(0.0);
190+
}
191+
192+
// ----------------------------------------------------------------------------
193+
// Standalone WASM Exports
194+
// ----------------------------------------------------------------------------
195+
196+
export function cas_export_simplify(daePtr: u32, exprId: u32): u32 {
197+
if (daePtr == 0) return exprId;
198+
return cas_simplify(changetype<DaeBuilder>(daePtr), exprId);
199+
}
200+
201+
export function cas_export_differentiate(daePtr: u32, exprId: u32, targetVarId: u32): u32 {
202+
if (daePtr == 0) return 0;
203+
return cas_differentiate(changetype<DaeBuilder>(daePtr), exprId, targetVarId);
204+
}

0 commit comments

Comments
 (0)