diff --git a/src/main/java/com/aparapi/internal/writer/BlockWriter.java b/src/main/java/com/aparapi/internal/writer/BlockWriter.java index 013239d0..c01d8766 100644 --- a/src/main/java/com/aparapi/internal/writer/BlockWriter.java +++ b/src/main/java/com/aparapi/internal/writer/BlockWriter.java @@ -76,6 +76,8 @@ to national security controls as identified on the Commerce Control List (curren public abstract class BlockWriter{ + private boolean hoistInlineDeclarations; + public final static String arrayLengthMangleSuffix = "__javaArrayLength"; public final static String arrayDimMangleSuffix = "__javaArrayDimension"; @@ -320,6 +322,10 @@ protected void writeGetterBlock(FieldEntry accessorVariableFieldEntry) { public void writeBlock(Instruction _first, Instruction _last) throws CodeGenException { write("{"); in(); + if (hoistInlineDeclarations) { + writeHoistedInlineDeclarations(_first, _last); + hoistInlineDeclarations = false; + } writeSequence(_first, _last); out(); newLine(); @@ -327,6 +333,62 @@ public void writeBlock(Instruction _first, Instruction _last) throws CodeGenExce write("}"); } + private void writeHoistedInlineDeclarations(Instruction _first, Instruction _last) throws CodeGenException { + final Map declarations = new LinkedHashMap(); + collectInlineDeclarations(_first, _last, declarations); + + for (final LocalVariableInfo localVariableInfo : declarations.values()) { + newLine(); + final String descriptor = localVariableInfo.getVariableDescriptor(); + if (descriptor.startsWith("[")) { + write(" __global "); + } + write(convertType(descriptor, true, false)); + write(localVariableInfo.getVariableName()); + write(";"); + } + } + + private void collectInlineDeclarations(Instruction _first, Instruction _last, + Map _declarations) { + for (Instruction instruction = _first; instruction != _last && instruction != null; + instruction = instruction.getNextExpr()) { + collectInlineDeclarations(instruction, _declarations); + } + } + + private void collectInlineDeclarations(Instruction _instruction, + Map _declarations) { + if (_instruction == null) { + return; + } + if (_instruction instanceof InlineAssignInstruction) { + final InlineAssignInstruction inlineAssignInstruction = (InlineAssignInstruction) _instruction; + final AssignToLocalVariable assignToLocalVariable = inlineAssignInstruction.getAssignToLocalVariable(); + final LocalVariableInfo localVariableInfo = assignToLocalVariable.getLocalVariableInfo(); + if (assignToLocalVariable.isDeclaration() && localVariableInfo != null) { + final String key = localVariableInfo.getVariableIndex() + ":" + + localVariableInfo.getVariableName() + ":" + + localVariableInfo.getVariableDescriptor(); + _declarations.put(key, localVariableInfo); + } + collectInlineDeclarations(inlineAssignInstruction.getRhs(), _declarations); + } + + if (_instruction instanceof MultiAssignInstruction) { + final MultiAssignInstruction multiAssignInstruction = (MultiAssignInstruction) _instruction; + collectInlineDeclarations(multiAssignInstruction.getFrom(), _declarations); + collectInlineDeclarations(multiAssignInstruction.getTo(), _declarations); + collectInlineDeclarations(multiAssignInstruction.getCommon(), _declarations); + } + + final Instruction firstChild = _instruction.getFirstChild(); + final Instruction lastChild = _instruction.getLastChild(); + if (firstChild != null && lastChild != null) { + collectInlineDeclarations(firstChild, lastChild.getNextExpr(), _declarations); + } + } + public Instruction writeConditional(BranchSet _branchSet) throws CodeGenException { return (writeConditional(_branchSet, false)); } @@ -698,11 +760,7 @@ public void writeInstruction(Instruction _instruction) throws CodeGenException { final AssignToLocalVariable assignToLocalVariable = inlineAssignInstruction.getAssignToLocalVariable(); final LocalVariableInfo localVariableInfo = assignToLocalVariable.getLocalVariableInfo(); - if (assignToLocalVariable.isDeclaration()) { - // this is bad! we need a general way to hoist up a required declaration - throw new CodeGenException("/* we can't declare this " + convertType(localVariableInfo.getVariableDescriptor(), true, false) - + " here */"); - } + // Declarations used as method arguments are emitted at the start of the method body. write(localVariableInfo.getVariableName()); write("="); writeInstruction(inlineAssignInstruction.getRhs()); @@ -870,6 +928,7 @@ public void writeMethodBody(MethodModel _methodModel) throws CodeGenException { FieldEntry accessorVariableFieldEntry = _methodModel.getAccessorVariableFieldEntry(); writeGetterBlock(accessorVariableFieldEntry); } else { + hoistInlineDeclarations = true; writeBlock(_methodModel.getExprHead(), null); } } diff --git a/src/test/java/com/aparapi/codegen/test/AssignAndPassAsParameterSimpleTest.java b/src/test/java/com/aparapi/codegen/test/AssignAndPassAsParameterSimpleTest.java index efff6cd5..510a04a5 100644 --- a/src/test/java/com/aparapi/codegen/test/AssignAndPassAsParameterSimpleTest.java +++ b/src/test/java/com/aparapi/codegen/test/AssignAndPassAsParameterSimpleTest.java @@ -15,23 +15,35 @@ */ package com.aparapi.codegen.test; -import com.aparapi.internal.exception.ClassParseException; -import com.aparapi.internal.exception.CodeGenException; -import org.junit.Ignore; +import static org.junit.Assert.assertTrue; + +import com.aparapi.internal.model.ClassModel; +import com.aparapi.internal.model.Entrypoint; +import com.aparapi.internal.writer.KernelWriter; import org.junit.Test; public class AssignAndPassAsParameterSimpleTest extends com.aparapi.codegen.CodeGenJUnitBase { - private static final String[] expectedOpenCL = null; - private static final Class expectedException = CodeGenException.class; + private String generateOpenCL() throws Exception { + Class testClass = com.aparapi.codegen.test.AssignAndPassAsParameterSimple.class; + ClassModel classModel = ClassModel.createClassModel(testClass); + Object kernelInstance = testClass.getConstructor((Class[]) null).newInstance(); + Entrypoint entrypoint = classModel.getEntrypoint("run", null); + return KernelWriter.writeToString(entrypoint); + } @Test - public void AssignAndPassAsParameterSimpleTest() { - test(com.aparapi.codegen.test.AssignAndPassAsParameterSimple.class, expectedException, expectedOpenCL); + public void AssignAndPassAsParameterSimpleTest() throws Exception { + String actual = generateOpenCL(); + assertTrue(actual.contains("int z;")); + assertTrue(actual.contains("z=1")); + assertTrue(actual.indexOf("int z;") < actual.indexOf("z=1")); } @Test - public void AssignAndPassAsParameterSimpleTestWorksWithCaching() { - test(com.aparapi.codegen.test.AssignAndPassAsParameterSimple.class, expectedException, expectedOpenCL); + public void AssignAndPassAsParameterSimpleTestWorksWithCaching() throws Exception { + String actual = generateOpenCL(); + assertTrue(actual.contains("int z;")); + assertTrue(actual.contains("z=1")); } }