Harbor

branch main
showing the latest snapshot on main
emit_wgsl.odin 16.4 KB · Plain text
shader/emit_wgsl.odin 0644 Raw
package shader

import "core:fmt"
import "core:strings"

// WGSL backend — emits from IR_Module

Wgsl_Emitter :: struct {
	w:           Writer,
	module:      ^IR_Module,
	current_fn:  ^IR_Function,
	diagnostics: [dynamic]Diagnostic,
}

emit_wgsl :: proc(module: ^IR_Module, allocator := context.allocator) -> (string, []Diagnostic) {
	e := Wgsl_Emitter{
		w           = writer_init(allocator),
		module      = module,
		diagnostics = make([dynamic]Diagnostic, allocator),
	}

	// Specialization constants (WGSL override declarations)
	for &sc in module.spec_constants {
		write_line(&e.w, "@id(", sc.spec_id, ") override ", sc.name, ": ", resolved_type_to_wgsl(sc.type), " = ", ir_const_value_to_string(sc.default_value), ";")
	}
	if len(module.spec_constants) > 0 {
		write_line(&e.w, "")
	}

	// Structs (for I/O and uniforms — skip compute)
	for &fn in module.functions {
		if !fn.is_entry do continue
		if fn.stage == .Compute do continue
		emit_wgsl_io_structs(&e, &fn)
	}

	// Shared variables (workgroup memory)
	for sv in module.shared_vars {
		write_line(&e.w, "var<workgroup> ", sv.name, ": ", resolved_type_to_wgsl(sv.type), ";")
	}
	if len(module.shared_vars) > 0 {
		write_line(&e.w, "")
	}

	// Bindings
	for &b in module.bindings {
		emit_wgsl_binding(&e, &b)
	}
	if len(module.bindings) > 0 {
		write_line(&e.w, "")
	}

	// Functions
	for &fn in module.functions {
		emit_wgsl_function(&e, &fn)
		write_line(&e.w, "")
	}

	return writer_to_string(e.w), e.diagnostics[:]
}

// -- I/O Structs --

@(private = "file")
emit_wgsl_io_structs :: proc(e: ^Wgsl_Emitter, fn: ^IR_Function) {
	// Input struct
	if len(fn.inputs) > 0 {
		// Find the param type name
		input_name := len(fn.params) > 0 ? wgsl_struct_name(fn.params[0].type) : "Input"
		write_line(&e.w, "struct ", input_name, " {")
		indent(&e.w)
		// Use write/write_line carefully - WGSL uses { } so avoid write_fmt
		for io in fn.inputs {
			for _ in 0 ..< e.w.indent do write(&e.w, "\t")
			if io.builtin != "" {
				write(&e.w, "@builtin(", wgsl_builtin_name(io.builtin), ") ", io.name, ": ", resolved_type_to_wgsl(io.type))
			} else {
				write(&e.w, "@location(", io.location, ") ", io.name, ": ", resolved_type_to_wgsl(io.type))
			}
			write(&e.w, ",\n")
		}
		dedent(&e.w)
		write_line(&e.w, "}")
		write_line(&e.w, "")
	}

	// Output struct
	if len(fn.outputs) > 0 {
		output_name := fn.return_type != nil ? wgsl_struct_name(fn.return_type) : "Output"
		write_line(&e.w, "struct ", output_name, " {")
		indent(&e.w)
		for io in fn.outputs {
			for _ in 0 ..< e.w.indent do write(&e.w, "\t")
			if io.builtin != "" {
				write(&e.w, "@builtin(", wgsl_builtin_name(io.builtin), ") ", io.name, ": ", resolved_type_to_wgsl(io.type))
			} else {
				write(&e.w, "@location(", io.location, ") ", io.name, ": ", resolved_type_to_wgsl(io.type))
			}
			write(&e.w, ",\n")
		}
		dedent(&e.w)
		write_line(&e.w, "}")
		write_line(&e.w, "")
	}
}

// -- Bindings --

@(private = "file")
emit_wgsl_binding :: proc(e: ^Wgsl_Emitter, b: ^IR_Binding) {
	switch b.kind {
	case .Texture:
		write_line(&e.w, "@group(", b.group, ") @binding(", b.binding_num, ") var ", b.name, ": texture_2d<f32>;")
	case .Sampler:
		write_line(&e.w, "@group(", b.group, ") @binding(", b.binding_num, ") var ", b.name, ": sampler;")
	case .Uniform, .Buffer:
		if b.struct_ref != nil {
			write_line(&e.w, "struct ", b.struct_ref.name, " {")
			indent(&e.w)
			for f in b.struct_ref.fields {
				write_line(&e.w, f.name, ": ", resolved_type_to_wgsl(f.type), ",")
			}
			dedent(&e.w)
			write_line(&e.w, "}")
			write_line(&e.w, "")
		}

		addr_space := b.kind == .Uniform ? "uniform" : "storage, read_write"
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		write(&e.w, "@group(", b.group, ") @binding(", b.binding_num, ") var<", addr_space, "> ", b.name, ": ")
		if b.struct_ref != nil {
			write(&e.w, b.struct_ref.name)
		} else {
			write(&e.w, resolved_type_to_wgsl(b.type))
		}
		write(&e.w, ";\n")
	case .Push_Constant:
		// WGSL has no push constants — emit as uniform fallback
		if b.struct_ref != nil {
			write_line(&e.w, "struct ", b.struct_ref.name, " {")
			indent(&e.w)
			for f in b.struct_ref.fields {
				write_line(&e.w, f.name, ": ", resolved_type_to_wgsl(f.type), ",")
			}
			dedent(&e.w)
			write_line(&e.w, "}")
			write_line(&e.w, "")
		}
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		write(&e.w, "var<push_constant> ", b.name, ": ")
		if b.struct_ref != nil {
			write(&e.w, b.struct_ref.name)
		} else {
			write(&e.w, resolved_type_to_wgsl(b.type))
		}
		write(&e.w, ";\n")
	}
}

// -- Functions --

@(private = "file")
emit_wgsl_function :: proc(e: ^Wgsl_Emitter, fn: ^IR_Function) {
	e.current_fn = fn

	if fn.is_entry {
		emit_wgsl_entry_point(e, fn)
	} else {
		emit_wgsl_helper_function(e, fn)
	}

	e.current_fn = nil
}

@(private = "file")
emit_wgsl_entry_point :: proc(e: ^Wgsl_Emitter, fn: ^IR_Function) {
	// Entry point attribute
	for _ in 0 ..< e.w.indent do write(&e.w, "\t")
	#partial switch fn.stage {
	case .Vertex:   write(&e.w, "@vertex\n")
	case .Fragment: write(&e.w, "@fragment\n")
	case .Compute:
		write(&e.w, "@compute @workgroup_size(", fn.workgroup_size[0])
		if fn.workgroup_size[1] > 0 do write(&e.w, ", ", fn.workgroup_size[1])
		if fn.workgroup_size[2] > 0 do write(&e.w, ", ", fn.workgroup_size[2])
		write(&e.w, ")\n")
	}

	// Compute: builtin params, void return
	if fn.stage == .Compute {
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		write(&e.w, "fn ", fn.name, "(")
		first := true
		for io in fn.inputs {
			if io.builtin == "" do continue
			if !first do write(&e.w, ", ")
			first = false
			write(&e.w, "@builtin(", builtin_to_wgsl(io.builtin), ") ", io.name, ": ", resolved_type_to_wgsl(io.type))
		}
		write(&e.w, ") {\n")
		indent(&e.w)
		emit_wgsl_stmts(e, fn.body[:])
		dedent(&e.w)
		write_line(&e.w, "}")
		return
	}

	// Function signature
	input_name := len(fn.params) > 0 ? wgsl_struct_name(fn.params[0].type) : "Input"
	output_name := fn.return_type != nil ? wgsl_struct_name(fn.return_type) : "Output"

	for _ in 0 ..< e.w.indent do write(&e.w, "\t")
	write(&e.w, "fn ", fn.name, "(")
	if len(fn.params) > 0 {
		write(&e.w, fn.params[0].name, ": ", input_name)
	}
	write(&e.w, ") -> ", output_name, " {\n")
	indent(&e.w)

	if len(fn.outputs) > 0 {
		write_line(&e.w, "var __luma_output: ", output_name, ";")
	}
	emit_wgsl_stmts(e, fn.body[:])
	if len(fn.outputs) > 0 {
		write_line(&e.w, "return __luma_output;")
	}

	dedent(&e.w)
	write_line(&e.w, "}")
}

@(private = "file")
emit_wgsl_helper_function :: proc(e: ^Wgsl_Emitter, fn: ^IR_Function) {
	for _ in 0 ..< e.w.indent do write(&e.w, "\t")
	write(&e.w, "fn ", fn.name, "(")
	for p, i in fn.params {
		if i > 0 do write(&e.w, ", ")
		write(&e.w, p.name, ": ", resolved_type_to_wgsl(p.type))
	}
	write(&e.w, ") -> ", resolved_type_to_wgsl(fn.return_type), " {\n")
	indent(&e.w)

	emit_wgsl_stmts(e, fn.body[:])

	dedent(&e.w)
	write_line(&e.w, "}")
}

// -- Statements --

@(private = "file")
emit_wgsl_stmts :: proc(e: ^Wgsl_Emitter, stmts: []IR_Stmt) {
	for stmt in stmts {
		emit_wgsl_stmt(e, stmt)
	}
}

@(private = "file")
emit_wgsl_stmt :: proc(e: ^Wgsl_Emitter, stmt: IR_Stmt) {
	switch s in stmt {
	case ^IR_Let:
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		write(&e.w, "let ", s.name, " = ")
		emit_wgsl_expr(e, s.value)
		write(&e.w, ";\n")

	case ^IR_Assign:
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		emit_wgsl_expr(e, s.target)
		write(&e.w, " = ")
		emit_wgsl_expr(e, s.value)
		write(&e.w, ";\n")

	case ^IR_Return:
		if s.value != nil {
			for _ in 0 ..< e.w.indent do write(&e.w, "\t")
			write(&e.w, "return ")
			emit_wgsl_expr(e, s.value)
			write(&e.w, ";\n")
		} else {
			write_line(&e.w, "return;")
		}

	case ^IR_Store_Output:
		// Handled by entry point emitter; if we get here in a non-entry context, emit as-is
		fn := e.current_fn
		if fn != nil && s.io_index >= 0 && s.io_index < len(fn.outputs) {
			io := fn.outputs[s.io_index]
			for _ in 0 ..< e.w.indent do write(&e.w, "\t")
			write(&e.w, "__luma_output.", io.name, " = ")
			emit_wgsl_expr(e, s.value)
			write(&e.w, ";\n")
		}

	case ^IR_If:
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		write(&e.w, "if (")
		emit_wgsl_expr(e, s.condition)
		write(&e.w, ") {\n")
		indent(&e.w)
		emit_wgsl_stmts(e, s.then_body[:])
		dedent(&e.w)
		for ei in s.elseif_clauses {
			for _ in 0 ..< e.w.indent do write(&e.w, "\t")
			write(&e.w, "} else if (")
			emit_wgsl_expr(e, ei.condition)
			write(&e.w, ") {\n")
			indent(&e.w)
			emit_wgsl_stmts(e, ei.body[:])
			dedent(&e.w)
		}
		if len(s.else_body) > 0 {
			write_line(&e.w, "} else {")
			indent(&e.w)
			emit_wgsl_stmts(e, s.else_body[:])
			dedent(&e.w)
		}
		write_line(&e.w, "}")

	case ^IR_For:
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		write(&e.w, "for (var ", s.var_name, " = ")
		emit_wgsl_expr(e, s.start)
		write(&e.w, "; ", s.var_name, " <= ")
		emit_wgsl_expr(e, s.stop)
		write(&e.w, "; ", s.var_name)
		if s.step != nil {
			write(&e.w, " += ")
			emit_wgsl_expr(e, s.step)
		} else {
			write(&e.w, "++")
		}
		write(&e.w, ") {\n")
		indent(&e.w)
		emit_wgsl_stmts(e, s.body[:])
		dedent(&e.w)
		write_line(&e.w, "}")

	case ^IR_While:
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		write(&e.w, "while (")
		emit_wgsl_expr(e, s.condition)
		write(&e.w, ") {\n")
		indent(&e.w)
		emit_wgsl_stmts(e, s.body[:])
		dedent(&e.w)
		write_line(&e.w, "}")

	case ^IR_Expr_Stmt:
		for _ in 0 ..< e.w.indent do write(&e.w, "\t")
		emit_wgsl_expr(e, s.expr)
		write(&e.w, ";\n")

	case ^IR_Barrier:
		write_line(&e.w, "workgroupBarrier();")

	case ^IR_Discard:
		write_line(&e.w, "discard;")

	case ^IR_Break:
		write_line(&e.w, "break;")

	case ^IR_Continue:
		write_line(&e.w, "continue;")
	}
}

// -- Expressions --

@(private = "file")
emit_wgsl_expr :: proc(e: ^Wgsl_Emitter, expr: ^IR_Expr) {
	if expr == nil {
		write(&e.w, "/* nil */")
		return
	}

	switch d in expr.derived {
	case ^IR_Literal:
		switch v in d.value {
		case i64:
			write(&e.w, v)
		case f64:
			s := fmt.aprintf("%v", v)
			if !strings.contains(s, ".") && !strings.contains(s, "e") {
				write(&e.w, s, ".0")
			} else {
				write(&e.w, s)
			}
		case bool:
			write(&e.w, v ? "true" : "false")
		}

	case ^IR_Var_Ref:
		write(&e.w, d.name)

	case ^IR_Binary:
		write(&e.w, "(")
		emit_wgsl_expr(e, d.left)
		write(&e.w, " ", ir_op_to_wgsl(d.op), " ")
		emit_wgsl_expr(e, d.right)
		write(&e.w, ")")

	case ^IR_Unary:
		if d.op == .Neg {
			write(&e.w, "(-")
		} else {
			write(&e.w, "(!")
		}
		emit_wgsl_expr(e, d.operand)
		write(&e.w, ")")

	case ^IR_Call:
		wgsl_name := d.is_builtin ? builtin_to_wgsl(d.name) : d.name
		// Handle texture sampling specially — WGSL separates texture and sampler args
		if d.is_builtin && (d.name == "sample" || d.name == "sample_level" || d.name == "sample_shadow") && len(d.args) >= 2 {
			fn_name := d.name == "sample" ? "textureSample" : (d.name == "sample_level" ? "textureSampleLevel" : "textureSampleCompare")
			write(&e.w, fn_name, "(")
			if lb, ok := d.args[0].derived.(^IR_Load_Binding); ok {
				tex_name, samp_name := find_split_bindings(e.module, lb.name)
				write(&e.w, tex_name, ", ", samp_name, ", ")
			} else {
				emit_wgsl_expr(e, d.args[0])
				write(&e.w, ", default_sampler, ")
			}
			emit_wgsl_expr(e, d.args[1])
			// Extra args after uv (lod for sample_level, ref for sample_shadow)
			for i := 2; i < len(d.args); i += 1 {
				write(&e.w, ", ")
				emit_wgsl_expr(e, d.args[i])
			}
			write(&e.w, ")")
			return
		}
		write(&e.w, wgsl_name, "(")
		for arg, i in d.args {
			if i > 0 do write(&e.w, ", ")
			emit_wgsl_expr(e, arg)
		}
		write(&e.w, ")")

	case ^IR_Field_Access:
		emit_wgsl_expr(e, d.object)
		write(&e.w, ".", d.field_name)

	case ^IR_Swizzle:
		emit_wgsl_expr(e, d.object)
		write(&e.w, ".", d.components)

	case ^IR_Composite_Extract:
		emit_wgsl_expr(e, d.object)
		write(&e.w, ".", d.field_name)

	case ^IR_Vector_Shuffle:
		emit_wgsl_expr(e, d.object)
		write(&e.w, ".", swizzle_indices_to_string(d.components))

	case ^IR_Index:
		emit_wgsl_expr(e, d.object)
		write(&e.w, "[")
		emit_wgsl_expr(e, d.index)
		write(&e.w, "]")

	case ^IR_Construct:
		wgsl_name := resolved_type_to_wgsl(expr.type)
		write(&e.w, wgsl_name, "(")
		for arg, i in d.args {
			if i > 0 do write(&e.w, ", ")
			emit_wgsl_expr(e, arg)
		}
		write(&e.w, ")")

	case ^IR_Type_Cast:
		write(&e.w, resolved_type_to_wgsl(expr.type), "(")
		emit_wgsl_expr(e, d.value)
		write(&e.w, ")")

	case ^IR_Load_Binding:
		write(&e.w, d.name)

	case ^IR_Input_Field:
		// In WGSL, entry point uses struct access: input.field
		write(&e.w, d.param_name, ".", d.field_name)

	case ^IR_Builtin_Var:
		// Find the field name from the current function's I/O
		fn := e.current_fn
		if fn != nil {
			io_list := d.is_input ? fn.inputs[:] : fn.outputs[:]
			for io in io_list {
				if io.builtin == d.name {
					// Compute shaders use inline params, not struct access
					if fn.stage == .Compute {
						write(&e.w, io.name)
					} else if d.is_input && len(fn.params) > 0 {
						write(&e.w, fn.params[0].name, ".", io.name)
					} else {
						write(&e.w, "output.", io.name)
					}
					return
				}
			}
		}
		write(&e.w, d.name)

	case ^IR_Shared_Ref:
		write(&e.w, d.name)

	case ^IR_Select:
		write(&e.w, "select(")
		emit_wgsl_expr(e, d.false_val)
		write(&e.w, ", ")
		emit_wgsl_expr(e, d.true_val)
		write(&e.w, ", ")
		emit_wgsl_expr(e, d.condition)
		write(&e.w, ")")
	}
}

// -- Helpers --

@(private = "file")
ir_op_to_wgsl :: proc(op: IR_Op) -> string {
	switch op {
	case .Add: return "+"
	case .Sub: return "-"
	case .Mul: return "*"
	case .Div: return "/"
	case .Mod: return "%"
	case .Eq:  return "=="
	case .Neq: return "!="
	case .Lt:  return "<"
	case .Gt:  return ">"
	case .Lte: return "<="
	case .Gte: return ">="
	case .And: return "&&"
	case .Or:  return "||"
	case .Neg: return "-"
	case .Not: return "!"
	}
	return "?"
}

@(private = "file")
builtin_to_wgsl :: proc(name: string) -> string {
	switch name {
	case "sample":          return "textureSample"
	case "sample_level":    return "textureSampleLevel"
	case "sample_grad":     return "textureSampleGrad"
	case "sample_compare":  return "textureSampleCompare"
	case "texel_fetch":     return "textureLoad"
	case "texture_size":    return "textureDimensions"
	case "atan2":           return "atan2"
	case "dfdx":            return "dpdx"
	case "dfdy":            return "dpdy"
	case "fwidth":          return "fwidth"
	case "inversesqrt":     return "inverseSqrt"
	case "mod":             return "fmod" // WGSL: use % operator for integers, or manually for floats
	}
	return name
}

@(private = "file")
wgsl_builtin_name :: proc(name: string) -> string {
	switch name {
	case "position":              return "position"
	case "vertex_id":             return "vertex_index"
	case "instance_id":           return "instance_index"
	case "frag_coord":            return "position" // fragment input
	case "front_facing":          return "front_facing"
	case "local_invocation_id":   return "local_invocation_id"
	case "local_invocation_index": return "local_invocation_index"
	case "global_invocation_id":  return "global_invocation_id"
	case "workgroup_id":          return "workgroup_id"
	}
	return name
}

resolved_type_to_wgsl :: proc(t: ^Resolved_Type) -> string {
	if t == nil do return "void"
	switch v in t^ {
	case Type_Scalar:
		switch v.kind {
		case .Bool:  return "bool"
		case .Int:   return "i32"
		case .Uint:  return "u32"
		case .Float: return "f32"
		case .Half:  return "f16"
		}
	case Type_Vector:
		elem: string
		switch v.elem {
		case .Float: elem = "f32"
		case .Int:   elem = "i32"
		case .Uint:  elem = "u32"
		case .Bool:  elem = "bool"
		case .Half:  elem = "f16"
		}
		return fmt.aprintf("vec%d<%s>", v.size, elem)
	case Type_Matrix:
		elem: string
		#partial switch v.elem {
		case .Float: elem = "f32"
		case .Half:  elem = "f16"
		case:        elem = "f32"
		}
		return fmt.aprintf("mat%dx%d<%s>", v.cols, v.rows, elem)
	case Type_Struct_Resolved:
		return v.name
	case Type_Array_Resolved:
		elem_str := resolved_type_to_wgsl(v.elem)
		if v.size == 0 {
			return fmt.aprintf("array<%s>", elem_str)
		}
		return fmt.aprintf("array<%s, %d>", elem_str, v.size)
	case Type_Sampler:
		switch v.kind {
		case .Sampler2D:       return "texture_2d<f32>"
		case .Sampler3D:       return "texture_3d<f32>"
		case .SamplerCube:     return "texture_cube<f32>"
		case .Sampler2DArray:  return "texture_2d_array<f32>"
		case .Sampler2DShadow: return "texture_depth_2d"
		}
	case Type_Void:
		return "void"
	}
	return "void"
}

@(private = "file")
wgsl_struct_name :: proc(t: ^Resolved_Type) -> string {
	if t == nil do return "Unknown"
	#partial switch v in t^ {
	case Type_Struct_Resolved:
		return v.name
	}
	return type_to_string(t)
}