Add routine to do code motion.

This commit is contained in:
Calvin Rose
2026-09-10 15:24:35 -05:00
parent 1f306c90e2
commit c821f6bb6f
6 changed files with 246 additions and 47 deletions
+1 -1
View File
@@ -58,7 +58,7 @@ LDFLAGS?=-rdynamic
LIBJANET_LDFLAGS?=$(LDFLAGS)
RUN:=$(RUN)
COMMON_CFLAGS:=-std=c99 -Wall -Wextra -Isrc/include -Isrc/conf -fvisibility=hidden -fPIC -fsanitize=address,undefined
COMMON_CFLAGS:=-std=c99 -Wall -Wextra -Isrc/include -Isrc/conf -fvisibility=hidden -fPIC
BOOT_CFLAGS:=-DJANET_BOOTSTRAP -DJANET_BUILD=$(JANET_BUILD) -O0 $(COMMON_CFLAGS) -g
BUILD_CFLAGS:=$(CFLAGS) $(COMMON_CFLAGS)
+162 -18
View File
@@ -109,12 +109,31 @@ const enum JanetInstructionType janet_instructions[JOP_INSTRUCTION_COUNT] = {
JINT_SSS /* JOP_CANCEL, */
};
static void rewrite_symbolmap(JanetFuncDef *def, uint32_t *pc_map) {
int32_t smout = 0;
for (int32_t i = 0; i < def->symbolmap_length; i++) {
JanetSymbolMap *sm = def->symbolmap + i;
int keep = 1;
/* Don't rewrite upvalue mappings */
if (sm->birth_pc < UINT32_MAX) {
sm->birth_pc = pc_map[sm->birth_pc];
sm->death_pc = pc_map[sm->death_pc];
/* entirely dead symbols can be removed from the symbol map. This can happen if a symbol is in dead code. */
if (sm->death_pc > (uint32_t) def->bytecode_length) sm->death_pc = (uint32_t) def->bytecode_length;
if (sm->birth_pc >= sm->death_pc) keep = 0;
}
/* Now shift if needed */
if (keep) def->symbolmap[smout++] = *sm;
}
def->symbolmap_length = smout;
}
/* Remove all noops while preserving jumps and debugging information.
* Useful as part of a filtering compiler pass. */
* Useful as part of a filtering compiler pass. Use this to remove bytecode. */
void janet_bytecode_remove_noops(JanetFuncDef *def) {
/* Get an instruction rewrite map so we can rewrite jumps */
uint32_t *pc_map = janet_smalloc(sizeof(uint32_t) * (1 + def->bytecode_length));
uint32_t *pc_map = janet_malloc(sizeof(uint32_t) * (1 + def->bytecode_length));
uint32_t new_bytecode_length = 0;
for (int32_t i = 0; i < def->bytecode_length; i++) {
uint32_t instr = def->bytecode[i];
@@ -166,25 +185,147 @@ void janet_bytecode_remove_noops(JanetFuncDef *def) {
}
/* Rewrite symbolmap */
int32_t smout = 0;
for (int32_t i = 0; i < def->symbolmap_length; i++) {
JanetSymbolMap *sm = def->symbolmap + i;
int keep = 1;
/* Don't rewrite upvalue mappings */
if (sm->birth_pc < UINT32_MAX) {
sm->birth_pc = pc_map[sm->birth_pc];
sm->death_pc = pc_map[sm->death_pc];
/* entirely dead symbols can be removed from the symbol map. This can happen if a symbol is in dead code. */
if (sm->birth_pc >= sm->death_pc) keep = 0;
}
/* Now shift if needed */
if (keep) def->symbolmap[smout++] = *sm;
}
def->symbolmap_length = smout;
rewrite_symbolmap(def, pc_map);
def->bytecode_length = new_bytecode_length;
def->bytecode = janet_realloc(def->bytecode, def->bytecode_length * sizeof(uint32_t));
janet_sfree(pc_map);
janet_free(pc_map);
}
/* Insert snippets of code into the bytecode while preserving symbol mapping and other
* internal structure. Sorts the chunks array increasing by the start field. Use
* this to add bytecode. */
void janet_bytecode_insert_chunks(JanetFuncDef *def, int32_t n_chunks, JanetBytecodeChunk *chunks) {
/* Sort chunks in increasing start order.
* Two chunks inserted at the same index should maintain order. */
for (int32_t i = 1; i < n_chunks; i++) {
JanetBytecodeChunk pivot = chunks[i];
int32_t j = i - 1;
while (j >= 0 && chunks[j].start > pivot.start) {
chunks[j + 1] = chunks[j];
j--;
}
chunks[j + 1] = pivot;
}
/* Calculate final length */
int32_t new_length = def->bytecode_length;
for (int32_t i = 0; i < n_chunks; i++) {
new_length += chunks[i].length;
}
uint32_t *new_bytecode = array_allocate(sizeof(uint32_t), new_length);
/* Update symbol map and rewrite jumps */
/* old pc -> new pc */
uint32_t *pc_map = janet_malloc(sizeof(uint32_t) * (1 + def->bytecode_length));
{
int32_t pc_cursor = 0;
int32_t j = 0;
for (int32_t i = 0; i < def->bytecode_length; i++) {
if (chunks[j].start == i) {
pc_cursor += chunks[j++].length;
}
pc_map[i] = pc_cursor++;
}
pc_map[def->bytecode_length] = pc_cursor;
}
/* Fix jumps */
for (int32_t i = 0; i < def->bytecode_length; i++) {
uint32_t instr = def->bytecode[i];
uint32_t opcode = instr & 0x7F;
int32_t old_jump_target = 0;
int32_t new_jump_target = 0;
int32_t new_location = pc_map[i];
switch (opcode) {
case JOP_NOOP:
continue;
case JOP_JUMP:
/* relative pc is in DS field of instruction */
old_jump_target = i + (((int32_t)instr) >> 8);
janet_assert(old_jump_target >= 0, "bounds");
janet_assert(old_jump_target < def->bytecode_length, "bounds");
new_jump_target = pc_map[old_jump_target];
def->bytecode[i] = (instr & 0xFF) | ((uint32_t)(new_jump_target - new_location) << 8);
break;
case JOP_JUMP_IF:
case JOP_JUMP_IF_NIL:
case JOP_JUMP_IF_NOT:
case JOP_JUMP_IF_NOT_NIL:
/* relative pc is in ES field of instruction */
old_jump_target = i + (((int32_t)instr) >> 16);
janet_assert(old_jump_target >= 0, "bounds");
janet_assert(old_jump_target < def->bytecode_length, "bounds");
new_jump_target = pc_map[old_jump_target];
def->bytecode[i] = (instr & 0xFFFF) | ((uint32_t)(new_jump_target - new_location) << 16);
break;
default:
break;
}
}
/* Copy bytecode */
uint32_t *read_cursor = def->bytecode;
uint32_t *write_cursor = new_bytecode;
for (int32_t i = 0; i <= n_chunks; i++) {
/* Rewrite original bytecode */
if (i && i < n_chunks) {
janet_assert(chunks[i].start >= chunks[i - 1].start, "chunk order");
}
int32_t last_start = i ? chunks[i - 1].start : 0;
int32_t this_start = (i == n_chunks) ? def->bytecode_length : chunks[i].start;
int32_t write_len = this_start - last_start;
janet_assert(write_len >= 0, "bad write_len");
if (write_len > 0) {
memcpy(write_cursor, read_cursor, sizeof(uint32_t) * (size_t) write_len);
read_cursor += write_len;
write_cursor += write_len;
}
if (i < n_chunks) {
int32_t chunk_len = chunks[i].length;
memcpy(write_cursor, chunks[i].bytecode, sizeof(uint32_t) * (size_t) chunk_len);
write_cursor += chunk_len;
}
}
/* Copy sourcemaps */
if (def->sourcemap) {
JanetSourceMapping *new_map = array_allocate(sizeof(JanetSourceMapping), new_length);
JanetSourceMapping *read_cursor = def->sourcemap;
JanetSourceMapping *write_cursor = new_map;
for (int32_t i = 0; i <= n_chunks; i++) {
int32_t last_start = i ? chunks[i - 1].start : 0;
int32_t this_start = (i == n_chunks) ? def->bytecode_length : chunks[i].start;
int32_t write_len = this_start - last_start;
janet_assert(write_len >= 0, "bad write_len");
if (write_len > 0) {
memcpy(write_cursor, read_cursor, sizeof(JanetSourceMapping) * (size_t) write_len);
read_cursor += write_len;
write_cursor += write_len;
}
if (i < n_chunks) {
int32_t chunk_len = chunks[i].length;
/* TODO - new source mapping */
for (int32_t j = 0; j < chunk_len; j++) {
write_cursor[j].line = -1;
write_cursor[j].column = -1;
}
write_cursor += chunk_len;
}
}
janet_free(def->sourcemap);
def->sourcemap = new_map;
}
/* Replace bytecode */
janet_free(def->bytecode);
def->bytecode = new_bytecode;
def->bytecode_length = new_length;
/* Rewrite symbolmap */
rewrite_symbolmap(def, pc_map);
janet_free(pc_map);
}
/* Remove redundant loads, moves and other instructions if possible and convert them to
@@ -364,6 +505,9 @@ void janet_bytecode_movopt(JanetFuncDef *def) {
case JOP_LOAD_TRUE:
case JOP_LOAD_FALSE:
case JOP_LOAD_SELF:
case JOP_MAKE_BUFFER:
case JOP_MAKE_STRING:
case JOP_MAKE_TABLE:
case JOP_MAKE_ARRAY:
case JOP_MAKE_TUPLE:
case JOP_MAKE_BRACKET_TUPLE: {
+1
View File
@@ -291,5 +291,6 @@ Shadowing janetc_shadowcheck(JanetCompiler *c, const uint8_t *sym);
void janet_bytecode_movopt(JanetFuncDef *def);
void janet_bytecode_remove_noops(JanetFuncDef *def);
void janet_bytecode_ovm_optimize(JanetFuncDef *def);
void janet_bytecode_insert_chunks(JanetFuncDef *def, int32_t n_chunks, JanetBytecodeChunk *chunks);
#endif
+3 -1
View File
@@ -1047,8 +1047,10 @@ static const uint8_t *unmarshal_one_def(
}
/* Validate */
if (janet_verify(def))
int status = 0;
if ((status = janet_verify(def))) {
janet_panic("funcdef has invalid bytecode");
}
/* Set def */
*out = def;
+72 -27
View File
@@ -325,7 +325,6 @@ static void janet_jump_threading(JanetFuncDef *def) {
recur = 1;
continue;
}
int32_t original_target = target;
while ((code[target] & 0x7F) == JOP_NOOP) { /* Skip noops */
target++;
}
@@ -340,9 +339,9 @@ static void janet_jump_threading(JanetFuncDef *def) {
}
uint32_t newcode;
if (is_branch) {
newcode = (code[i] & 0xFFFF) | (uint32_t)((target - i) << 16);
newcode = (code[i] & 0xFFFF) | ((uint32_t)(target - i) << 16);
} else {
newcode = (code[i] & 0xFF) | (uint32_t)((target - i) << 8);
newcode = (code[i] & 0xFF) | ((uint32_t)(target - i) << 8);
}
if (newcode != code[i]) recur = 1;
code[i] = newcode;
@@ -511,6 +510,9 @@ typedef struct {
* TODO - this is very slow and uses the naive "dense" analysis. However, Janet functions
* tend to be small and the overhead of creating and dealing with many extra nodes might
* not be all that helpful for our ISA. Sparse is probably better though.
*
* Also, iterate over basic blocks instead. There are implicit merges at these locations
* that we currently aren't handling.
*/
JanetTypeflowInstruction *janet_bytecode_lattice_types(JanetFuncDef *def, uint16_t *ret_types) {
@@ -542,6 +544,9 @@ JanetTypeflowInstruction *janet_bytecode_lattice_types(JanetFuncDef *def, uint16
size_t iterations = 0; /* debug counter */
/* While we have more states to visit, traverse them */
while (janet_v_count(state_stack)) {
if (iterations >= MAX_ITERATIONS) {
janet_eprintf("function %s has issue.\n", def->name);
}
janet_assert(iterations < MAX_ITERATIONS, "too many iterations. Check the code.");
iterations += 1;
int32_t pc = janet_v_last(state_stack);
@@ -748,10 +753,11 @@ JanetTypeflowInstruction *janet_bytecode_lattice_types(JanetFuncDef *def, uint16
uint16_t taken = invert ? false_t : true_t;
uint16_t not_taken = invert ? true_t : false_t;
uint16_t branches[2] = { taken, not_taken };
int32_t targets[2] = { pc + 1, pc + Ies };
int32_t targets[2] = { pc + Ies, pc + 1 };
for (int j = 0; j < 2; j++) {
target = targets[j];
if (!branches[j]) continue; /* impossible branch */
types[Ia] = branches[j];
target = targets[j];
/* Merge with existing state */
Janet oldwrapper = janet_table_get(states, janet_wrap_integer(target));
if (janet_checktype(oldwrapper, JANET_NIL)) {
@@ -1435,7 +1441,6 @@ static void bb_remove_redundant_writes(JanetFuncDef *def, JanetBB bb) {
* return to our function. */
switch (Iop) {
default:
fprintf(stderr, "opcode = %u\n", Iop);
janet_assert(0, "unhandled instruction");
continue;
case JOP_JUMP:
@@ -1542,7 +1547,7 @@ static void bb_remove_redundant_writes(JanetFuncDef *def, JanetBB bb) {
* Some of these instructions may have side effects if
* the inputs are abstracts. */
/* Loads that trite D */
/* Loads that write D */
case JOP_LOAD_NIL:
case JOP_LOAD_TRUE:
case JOP_LOAD_FALSE:
@@ -1768,30 +1773,70 @@ void janet_bytecode_ovm_optimize(JanetFuncDef *def) {
int32_t initial_length = def->bytecode_length;
total_before += initial_length;
/* Basic optimization */
janet_bytecode_movopt(def);
janet_jump_threading(def);
{
JanetBB *bbs = janet_basic_blocks(def);
for (int32_t i = 0; i < janet_v_count(bbs); i++) {
//janet_ovm_value_numbering(def, bbs[i]);
bb_remove_redundant_writes(def, bbs[i]); // slightly wrong
}
bb_dead_to_noop(def, bbs);
janet_v_free(bbs);
/*
JanetArray *before = janet_array(0);
for (int32_t i = 0; i < def->bytecode_length; i++) {
janet_array_push(before, janet_asm_decode_instruction(def->bytecode[i]));
}
janet_jump_threading(def);
*/
/* Basic optimization */
for (int i = 0; i < 2; i++) {
janet_bytecode_movopt(def);
janet_jump_threading(def);
{
JanetBB *bbs = janet_basic_blocks(def);
for (int32_t i = 0; i < janet_v_count(bbs); i++) {
//janet_ovm_value_numbering(def, bbs[i]);
bb_remove_redundant_writes(def, bbs[i]);
}
bb_dead_to_noop(def, bbs);
janet_v_free(bbs);
}
janet_bytecode_remove_noops(def);
}
/* Check chunk insertion just because */
uint32_t noops[3] = { JOP_NOOP, JOP_NOOP, JOP_NOOP };
JanetBytecodeChunk chunks[3];
chunks[0].start = 0;
chunks[1].start = (def->bytecode_length - 1) / 2;
chunks[2].start = def->bytecode_length - 1; /* no noops at end */
chunks[0].length = 3;
chunks[1].length = 3;
chunks[2].length = 3;
chunks[0].bytecode = noops;
chunks[1].bytecode = noops;
chunks[2].bytecode = noops;
janet_bytecode_insert_chunks(def, 3, chunks);
janet_bytecode_remove_noops(def);
/*
JanetArray *after = janet_array(0);
for (int32_t i = 0; i < def->bytecode_length; i++) {
janet_array_push(after, janet_asm_decode_instruction(def->bytecode[i]));
}
for (int32_t i = 0; i < after->count || i < before->count; i++) {
if (i < after->count) {
if (i < before->count) {
janet_eprintf("%4d: %Q -> %Q\n", i, before->data[i], after->data[i]);
} else {
janet_eprintf("%4d: -> %Q\n", i, after->data[i]);
}
} else {
janet_eprintf("%4d: %Q ->\n", i, before->data[i]);
}
}
*/
/* Lattice analysis */
//uint16_t rettypes = 0;
//JanetTypeflowInstruction *instrs = janet_bytecode_lattice_types(func->def, &rettypes);
//if (!instrs) janet_panic("function too complicated");
//JanetArray *ret = janet_array(func->def->bytecode_length + 1);
//for (int32_t i = 0; i < func->def->bytecode_length; i++) {
// Janet x = debug_lattice_types_instruction(instrs[i]);
// janet_array_push(ret, x);
//}
/*uint16_t rettypes = 0;*/
/*JanetTypeflowInstruction *instrs = janet_bytecode_lattice_types(def, &rettypes);*/
/*lattice_dead_to_noop(def, instrs);*/
/*janet_free(instrs);*/
/*janet_jump_threading(def);*/
/*janet_bytecode_remove_noops(def);*/
/* Info */
int32_t final_length = def->bytecode_length;
+7
View File
@@ -1121,6 +1121,13 @@ struct JanetSymbolMap {
const uint8_t *symbol;
};
/* Internal data structure for inserting chunks of bytecode */
typedef struct {
uint32_t *bytecode;
int32_t length;
int32_t start;
} JanetBytecodeChunk;
/* A function definition. Contains information needed to instantiate closures. */
struct JanetFuncDef {
JanetGCObject gc;