compile builtin shaders in dependency layers

This commit is contained in:
Kovid Goyal 2026-07-28 00:15:57 +05:30
parent d12a9d3dd4
commit d7741b4996
No known key found for this signature in database
GPG key ID: 06BC317B515ACE7C
2 changed files with 68 additions and 3 deletions

View file

@ -403,6 +403,30 @@ def topological_sort(graph: dict[str, SlangFile]) -> list[str]:
return order
def topological_layers(graph: dict[str, SlangFile]) -> list[list[str]]:
layer_of: dict[str, int] = {}
def compute_layer(node: str) -> int:
if node in layer_of:
return layer_of[node]
if node not in graph:
return -1
layer = max((compute_layer(dep) + 1 for dep in graph[node].imports), default=0)
layer_of[node] = layer
return layer
for node in graph:
compute_layer(node)
if not layer_of:
return []
max_layer = max(layer_of.values())
layers: list[list[str]] = [[] for _ in range(max_layer + 1)]
for node, layer in layer_of.items():
layers[layer].append(node)
return layers
def get_ordered_sources_in_tree(dirpath: str) -> OrderedDict[str, SlangFile]:
g = build_import_graph(dirpath)
return OrderedDict({k: g[k] for k in topological_sort(g)})
@ -819,8 +843,10 @@ def compile_builtin_shaders(build_dir: str, dest_dir: str, parallel_run: Paralle
source_tree = get_ordered_sources_in_tree(src_dir)
serialize_source_metadata(source_tree, dest_dir)
# First ensure all IR is generated
parallel_run(commands_to_compile_dir_to_ir(source_tree, src_dir, build_dir))
# Compile IR layer by layer so each module's dependencies finish before it starts
for layer in topological_layers(source_tree):
layer_sources = {k: source_tree[k] for k in layer}
parallel_run(commands_to_compile_dir_to_ir(layer_sources, src_dir, build_dir))
# Create the specializations
parallel_run(create_specialisations(source_tree, build_dir))
# Now Vulkan shaders

View file

@ -4,7 +4,7 @@
import os
import tempfile
from kitty.shaders.slang import EntryPoint, SlangFile, Stage, build_import_graph, parse_slang_text, topological_sort
from kitty.shaders.slang import EntryPoint, SlangFile, Stage, build_import_graph, parse_slang_text, topological_layers, topological_sort
from .base import BaseTest
@ -189,3 +189,42 @@ void vsMain() {}
f.write('not a slang file\n')
graph3 = build_import_graph(tmpdir)
self.assertNotIn('ignored', graph3)
def test_topological_layers(self):
# Linear chain a <- b <- c produces three layers
graph: dict[str, SlangFile] = {
'a': SlangFile('', '', frozenset(), frozenset(), 'a'),
'b': SlangFile('', '', frozenset({'a'}), frozenset(), 'b'),
'c': SlangFile('', '', frozenset({'b'}), frozenset(), 'c'),
}
layers = topological_layers(graph)
self.assertEqual(len(layers), 3)
self.assertIn('a', layers[0])
self.assertIn('b', layers[1])
self.assertIn('c', layers[2])
# Diamond: base <- left, base <- right, left+right <- top
# base is layer 0, left and right are layer 1, top is layer 2
diamond: dict[str, SlangFile] = {
'base': SlangFile('', '', frozenset(), frozenset(), 'base'),
'left': SlangFile('', '', frozenset({'base'}), frozenset(), 'left'),
'right': SlangFile('', '', frozenset({'base'}), frozenset(), 'right'),
'top': SlangFile('', '', frozenset({'left', 'right'}), frozenset(), 'top'),
}
layers2 = topological_layers(diamond)
self.assertEqual(len(layers2), 3)
self.assertIn('base', layers2[0])
self.assertIn('left', layers2[1])
self.assertIn('right', layers2[1])
self.assertIn('top', layers2[2])
# Node with import not in graph is treated as layer 0
partial: dict[str, SlangFile] = {
'x': SlangFile('', '', frozenset({'missing'}), frozenset(), 'x'),
}
layers3 = topological_layers(partial)
self.assertEqual(len(layers3), 1)
self.assertIn('x', layers3[0])
# Empty graph
self.assertEqual(topological_layers({}), [])