Skip to content
Merged
Show file tree
Hide file tree
Changes from 9 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Original file line number Diff line number Diff line change
@@ -1,27 +1,23 @@
package io.computenode.cyfra.juliaset

import io.computenode.cyfra.dsl.{*, given}
import io.computenode.cyfra.*
import io.computenode.cyfra.dsl.GStruct.Empty
import io.computenode.cyfra.dsl.Pure.pure
import io.computenode.cyfra.runtime.{GContext, GFunction}
import org.apache.commons.io.IOUtils
import org.junit.runner.RunWith
import io.computenode.cyfra.dsl.{*, given}
import io.computenode.cyfra.runtime.SpirvOptimizer.{Enable, O}
import io.computenode.cyfra.runtime.mem.Vec4FloatMem
import io.computenode.cyfra.runtime.{GContext, GFunction}
import io.computenode.cyfra.utility.ImageUtility
import munit.FunSuite

import java.io.File
import java.nio.file.Files
import scala.concurrent.ExecutionContext
import scala.concurrent.ExecutionContext.Implicits
import scala.concurrent.duration.DurationInt
import scala.concurrent.{Await, ExecutionContext}

class JuliaSet extends FunSuite:
given GContext = new GContext()
given ExecutionContext = Implicits.global
test("Render julia set"):

def runJuliaSet(referenceImgName: String)(using GContext): Unit = {
val dim = 4096
val max = 1
val RECURSION_LIMIT = 1000
Expand Down Expand Up @@ -76,6 +72,14 @@ class JuliaSet extends FunSuite:
val r = Vec4FloatMem(dim * dim).map(function).asInstanceOf[Vec4FloatMem].toArray
val outputTemp = File.createTempFile("julia", ".png")
ImageUtility.renderToImage(r, dim, outputTemp.toPath)
val referenceImage = getClass.getResource("julia.png")
val referenceImage = getClass.getResource(referenceImgName)
ImageTests.assertImagesEquals(outputTemp, new File(referenceImage.getPath))

}

test("Render julia set"):
given GContext = new GContext()
runJuliaSet("julia.png")

test("Render julia set optimized"):
given GContext = new GContext(spirvOptimization = Enable(O))
runJuliaSet("julia_O_optimized.png")
Original file line number Diff line number Diff line change
Expand Up @@ -4,20 +4,21 @@ import io.computenode.cyfra
import io.computenode.cyfra.*
import io.computenode.cyfra.foton.animation.AnimatedFunctionRenderer.Parameters
import io.computenode.cyfra.foton.animation.{AnimatedFunction, AnimatedFunctionRenderer}
import io.computenode.cyfra.given
import io.computenode.cyfra.runtime.*
import io.computenode.cyfra.dsl.*
import io.computenode.cyfra.dsl.Color.{InterpolationThemes, interpolate}
import io.computenode.cyfra.dsl.Math3D.*
import io.computenode.cyfra.dsl.given
import io.computenode.cyfra.foton.animation.AnimationFunctions.*
import io.computenode.cyfra.runtime.SpirvOptimizer.{Enable, O}

import java.nio.file.Paths
import scala.concurrent.duration.DurationInt

object AnimatedJulia:
given GContext = new GContext(spirvOptimization = Enable(O))
@main
def julia() =
def julia(): Unit =

def julia(uv: Vec2[Float32])(using AnimationInstant): Int32 =
val p = smooth(from = 0.355f, to = 0.4f, duration = 3.seconds)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,10 +11,11 @@ import io.computenode.cyfra.foton.animation.AnimationRenderer
import io.computenode.cyfra.foton.rt.ImageRtRenderer.RaytracingIteration
import io.computenode.cyfra.foton.rt.animation.AnimationRtRenderer.RaytracingIteration
import io.computenode.cyfra.foton.rt.RtRenderer
import io.computenode.cyfra.runtime.{GFunction, GContext}
import io.computenode.cyfra.runtime.{GContext, GFunction, SpirvValidator}
import io.computenode.cyfra.utility.Units.Milliseconds
import io.computenode.cyfra.utility.Utility.timed
import io.computenode.cyfra.dsl.Algebra.{*, given}
import io.computenode.cyfra.runtime.SpirvOptimizer.{Enable, O}
import io.computenode.cyfra.runtime.mem.GMem.fRGBA
import io.computenode.cyfra.runtime.mem.Vec4FloatMem

Expand All @@ -24,9 +25,7 @@ import scala.concurrent.{Await, ExecutionContext}
import scala.concurrent.duration.DurationInt


class AnimatedFunctionRenderer(params: AnimatedFunctionRenderer.Parameters) extends AnimationRenderer[AnimatedFunction, AnimatedFunctionRenderer.RenderFn](params):

given GContext = new GContext()
class AnimatedFunctionRenderer(params: AnimatedFunctionRenderer.Parameters)(using GContext) extends AnimationRenderer[AnimatedFunction, AnimatedFunctionRenderer.RenderFn](params):

given ExecutionContext = Implicits.global

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,18 +2,20 @@ package io.computenode.cyfra.runtime

import io.computenode.cyfra.dsl.Algebra.FromExpr
import io.computenode.cyfra.dsl.{GArray, GStruct, GStructSchema, UniformContext, Value}
import GStruct.Empty
import Value.{Float32, Vec4}
import io.computenode.cyfra.vulkan.VulkanContext
import io.computenode.cyfra.vulkan.compute.{Binding, ComputePipeline, InputBufferSize, LayoutInfo, LayoutSet, Shader, UniformSize}
import io.computenode.cyfra.vulkan.executor.{BufferAction, SequenceExecutor}
import SequenceExecutor.*
import io.computenode.cyfra.runtime.SpirvOptimizer.{Optimization, O}
import io.computenode.cyfra.runtime.SpirvValidator.Validation
import io.computenode.cyfra.runtime.mem.GMem.totalStride
import io.computenode.cyfra.spirv.SpirvTypes.typeStride
import io.computenode.cyfra.spirv.compilers.DSLCompiler
import io.computenode.cyfra.spirv.compilers.ExpressionCompiler.{UniformStructRef, WorkerIndex}
import io.computenode.cyfra.utility.Logger.logger
import mem.{FloatMem, GMem, Vec4FloatMem}
import org.lwjgl.system.{Configuration, MemoryUtil}
import org.lwjgl.system.Configuration
import izumi.reflect.Tag

import java.io.FileOutputStream
Expand All @@ -23,7 +25,7 @@ import java.util.concurrent.Executors
import scala.concurrent.{ExecutionContext, ExecutionContextExecutor}


class GContext:
class GContext(spirvValidation: Validation = SpirvValidator.Enable(), spirvOptimization: Optimization = SpirvOptimizer.Disable):

Configuration.STACK_SIZE.set(1024) // fix lwjgl stack size

Expand All @@ -32,9 +34,9 @@ class GContext:
implicit val ec: ExecutionContextExecutor = ExecutionContext.fromExecutor(Executors.newFixedThreadPool(16))

def compile[
G <: GStruct[G] : Tag : GStructSchema,
H <: Value : Tag : FromExpr,
R <: Value : Tag : FromExpr
G <: GStruct[G] : {Tag, GStructSchema},
H <: Value : {Tag, FromExpr},
R <: Value : {Tag, FromExpr}
](function: GFunction[G, H, R]): ComputePipeline = {
val uniformStructSchema = summon[GStructSchema[G]]
val uniformStruct = uniformStructSchema.fromTree(UniformStructRef)
Expand All @@ -46,22 +48,43 @@ class GContext:
GArray[H](0)
)
val shaderCode = DSLCompiler.compile(tree, function.arrayInputs, function.arrayOutputs, uniformStructSchema)
SpirvValidator.validateSpirv(shaderCode, spirvValidation)

dumpSpvToFile(shaderCode, "program.spv") // TODO remove before release

val inOut = 0 to 1 map (Binding(_, InputBufferSize(typeStride(summon[Tag[H]]))))
val uniform = Option.when(uniformStructSchema.fields.nonEmpty)(Binding(2, UniformSize(totalStride(uniformStructSchema))))
val layoutInfo = LayoutInfo(Seq(LayoutSet(0, inOut ++ uniform)))
val shader = new Shader(shaderCode, new org.joml.Vector3i(256, 1, 1), layoutInfo, "main", vkContext.device)

val shader = SpirvOptimizer.getOptimizedSpirv(shaderCode, spirvOptimization) match {
case None =>
new Shader(shaderCode, new org.joml.Vector3i(256, 1, 1), layoutInfo, "main", vkContext.device)

case Some(optimizedShaderCode) =>
dumpSpvToFile(optimizedShaderCode, "optimized_program.spv") // TODO remove before release
SpirvValidator.validateSpirv(optimizedShaderCode, spirvValidation)
new Shader(optimizedShaderCode, new org.joml.Vector3i(256, 1, 1), layoutInfo, "main", vkContext.device)
}

new ComputePipeline(shader, vkContext)
}

def time[R](block: => R): R = {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think we have this somewhere, like in Util object?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I wasn't able to find it

val start = System.nanoTime()
val result = block
val end = System.nanoTime()
logger.debug(s"Elapsed time: ${(end - start) / 1e6} ms")
result
}

private def dumpSpvToFile(code: ByteBuffer, path: String): Unit =
val fc: FileChannel = new FileOutputStream("program.spv").getChannel
fc.write(code)
fc.close()
code.rewind()

def execute[
G <: GStruct[G] : Tag : GStructSchema,
G <: GStruct[G] : {Tag, GStructSchema},
H <: Value,
R <: Value
](mem: GMem[H], fn: GFunction[?, H, R])(using uniformContext: UniformContext[_]): GMem[R] =
Expand All @@ -70,12 +93,12 @@ class GContext:
LayoutLocation(0, 0) -> BufferAction.LoadTo,
LayoutLocation(0, 1) -> BufferAction.LoadFrom
) ++ (
if isUniformEmpty then Map.empty
if isUniformEmpty then Map.empty
else Map(LayoutLocation(0, 2) -> BufferAction.LoadTo)
)
)
val sequence = ComputationSequence(Seq(Compute(fn.pipeline, actions)), Seq.empty)
val executor = new SequenceExecutor(sequence, vkContext)

val data = mem.toReadOnlyBuffer
val inData =
if isUniformEmpty then Seq(data)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
package io.computenode.cyfra.runtime

import io.computenode.cyfra.utility.Logger.logger

import java.nio.ByteBuffer

object SpirvDisassembler extends SpirvTool {

override type SpirvError = SpirvDisassemblerError
protected override val toolName: SupportedSpirVTools = SupportedSpirVTools.Disassembler

def getDisassembledSpirv(shaderCode: ByteBuffer, options: Param*): Option[String] = {
getOS.flatMap { os =>
getToolExecutableFromPath(
toolName, os)
} match {
case None =>
logger.warn("Shader code will not be disassembled.")
None
case Some(executable) =>
val cmd = Seq(executable) ++ options.flatMap(_.asStringParam.split(" ")) ++ Seq("-")
val (outputStream, errorStream, exitCode) = executeSpirvCmd(shaderCode, cmd)

if (exitCode == 0) {
logger.debug("SPIRV-Tools Disassembler succeeded.")
Some(outputStream.toString)
} else {
throw SpirvDisassemblerError(s"SPIRV-Tools Disassembler failed with exit code $exitCode. ${errorStream.toString}")
}
}
}

override protected def createError(message: String): SpirvDisassemblerError =

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

let's jsut get rid of all that and just do some SpirvToolException

SpirvDisassemblerError(message)

case class SpirvDisassemblerError(msg: String) extends RuntimeException(msg)

case object NoIndent extends FlagParam("--no-indent")

case object NoHeader extends FlagParam("--no-header")

case object RawId extends FlagParam("--raw-id")

case object NestedIndent extends FlagParam("--nested-indent")

case object ReorderBlocks extends FlagParam("--reorder-blocks")

case object Offsets extends FlagParam("--offsets")

case object Comment extends FlagParam("--comment")
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
package io.computenode.cyfra.runtime

import io.computenode.cyfra.runtime.SpirvDisassembler.executeSpirvCmd
import io.computenode.cyfra.utility.Logger.logger

import java.nio.ByteBuffer

object SpirvOptimizer extends SpirvTool {

override type SpirvError = SpirvOptimizationError
protected override val toolName: SupportedSpirVTools = SupportedSpirVTools.Optimizer

def getOptimizedSpirv(shaderCode: ByteBuffer, optimization: Optimization): Option[ByteBuffer] = {
optimization match {
case Disable => None
case Enable(settings*) =>
getOS.flatMap { os =>
getToolExecutableFromPath(
toolName, os)
} match {
case None =>
logger.warn("Shader code will not be optimized.")
None
case Some(executable) =>
val cmd = Seq(executable) ++ settings.flatMap(_.asStringParam.split(" ")) ++ Seq("-", "-o", "-")
val (outputStream, errorStream, exitCode) = executeSpirvCmd(shaderCode, cmd)

if (exitCode == 0) {
logger.debug("SPIRV-Tools Optimizer succeeded.")
Some(toDirectBuffer(ByteBuffer.wrap(outputStream.toByteArray)))
} else {
throw SpirvOptimizationError(s"SPIRV-Tools Optimizer failed with exit code $exitCode.\n${errorStream.toString()}")
}
}
}
}

private def toDirectBuffer(buf: ByteBuffer): ByteBuffer = {
val direct = ByteBuffer.allocateDirect(buf.remaining())
direct.put(buf)
direct.flip()
direct
}

override protected def createError(message: String): SpirvOptimizationError =
SpirvOptimizationError(message)

sealed trait Optimization

case class SpirvOptimizationError(msg: String) extends RuntimeException(msg)

case class TargetEnv(version: String) extends ParamWithArgs("--target-env", version)

case class Enable(settings: Param*) extends Optimization

case object O extends FlagParam("-O")

case object Os extends FlagParam("-Os")

case object Disable extends Optimization
}
Loading