diff --git a/Tools/CodeGen/CalculationFunction.m b/Tools/CodeGen/CalculationFunction.m index ee0ba6f..37b417a 100644 --- a/Tools/CodeGen/CalculationFunction.m +++ b/Tools/CodeGen/CalculationFunction.m @@ -126,7 +126,8 @@ replaceDerivatives[x_, derivRules_] := x /. replaceStandard]; (* Return a CodeGen block which assigns dest by evaluating expr *) -assignVariableFromExpression[dest_, expr_] := Module[{tSym, type, cleanExpr, code}, +assignVariableFromExpression[dest_, expr_, declare_] := + Module[{tSym, type, cleanExpr, code}, tSym = Unique[]; @@ -135,7 +136,7 @@ assignVariableFromExpression[dest_, expr_] := Module[{tSym, type, cleanExpr, cod cleanExpr = ReplacePowers[expr] /. sym`t -> tSym; If[SOURCELANGUAGE == "C", - code = type <> " const " <> + code = If[declare, type <> " const ", ""] <> ToString[dest == cleanExpr, CForm, PageWidth -> 120] <> ";\n", code = ToString@dest <> ".eq." <> ToString[cleanExpr, FortranForm, PageWidth -> 120] <> "\n" ]; @@ -294,6 +295,14 @@ simplifyEquationList[eqs_] := simplifyEquation[lhs_ -> rhs_] := lhs -> Simplify[rhs]; +(* Given an input list l, return a list L such that L[[i]] is True + only if l[[i]] is the first occurrence of l[[i]] in l and is not in + the list "already" *) +markFirst[l_List, already_List] := + If[l =!= {}, + {!MemberQ[already, First[l]]} ~Join~ markFirst[Rest[l], already ~Join~ {First[l]}], + {}]; + VerifyListContent[l_, type_, while_] := Module[{types}, If[!(Head[l] === List), @@ -656,7 +665,7 @@ equationLoop[eqs_, pddefs_, where_, addToStencilWidth_, useLoopControl_, useCSE_] := Module[{rhss, lhss, gfsInRHS, gfsInLHS, gfsOnlyInRHS, localGFs, localMap, eqs2, derivSwitch, actualSyncGroups, code, functionName, calcCode, - syncCode, loopFunction}, + syncCode, loopFunction, eqsReplaced, declare}, rhss = Map[#[[2]] &, eqs]; lhss = Map[#[[1]] &, eqs]; @@ -688,8 +697,24 @@ equationLoop[eqs_, eqs2 = ReplaceDerivatives[defsWithShorts, eqs2, False]; checkEquationAssignmentOrder[eqs2, shorts]; - code = {(*InitialiseGridLoopVariables[derivSwitch, addToStencilWidth], *) - functionName = ToString@lookup[cleancalc, Name]; + code = {(*InitialiseGridLoopVariables[derivSwitch, addToStencilWidth], *) + functionName = ToString@lookup[cleancalc, Name]; + + (* Replace grid functions with their local forms *) + eqsReplaced = eqs2 /. localMap; + + If[useCSE, + eqsReplaced = CSE[eqsReplaced]]; + + (* Construct a list, corresponding to the list of equations, + marking those which need their LHS variables declared. We + declare variables at the same time as assigning to them as it + gives a performance increase over declaring them separately at + the start of the loop. The local variables for the grid + functions which appear in the RHSs have been declared and set + already (DeclareMaybeAssignVariableInLoop below), so assignments + to these do not generate declarations here. *) + declare = markFirst[First /@ eqsReplaced, Map[localName, gfsInRHS]]; (* calcCode = @@ -697,11 +722,16 @@ equationLoop[eqs_, replaceDerivatives[replaceWithDerivativesHidden[eqs2, localMap], {}] ]; *) +(* calcCode = Map[{assignVariableFromExpression[#[[1]], #[[2]]], "\n"} &, If[useCSE, CSE, Identity][ replaceDerivatives[replaceWithDerivativesHidden[eqs2, localMap], {}]] ]; +*) + calcCode = + MapThread[{assignVariableFromExpression[#1[[1]], #1[[2]], #2], "\n"} &, + {eqsReplaced, declare}]; Join[ (* @@ -728,7 +758,7 @@ equationLoop[eqs_, "CCTK_REAL", localName[#], GridName[#], StringMatchQ[ToString[GridName[#]], "eT" ~~ _ ~~ _ ~~ "[" ~~ __ ~~ "]"], "*stress_energy_state"] &, - gfsOnlyInRHS]], + (* gfsOnlyInRHS *) gfsInRHS]], (* CommentedBlock["Check for nans", diff --git a/Tools/CodeGen/CodeGen.m b/Tools/CodeGen/CodeGen.m index 2945d2b..d2d3ccb 100644 --- a/Tools/CodeGen/CodeGen.m +++ b/Tools/CodeGen/CodeGen.m @@ -294,8 +294,8 @@ MaybeAssignVariableInLoop[dest_, src_, cond_] := DeclareMaybeAssignVariableInLoop[type_, dest_, src_, mmaCond_, codeCond_] := If [mmaCond, - {type, " const ", dest, " = (", codeCond, ") ? (", src, ") : 0.0", EOL[]}, - {type, " const ", dest, " = ", src, EOL[]}]; + {type, " ", dest, " = (", codeCond, ") ? (", src, ") : 0.0", EOL[]}, + {type, " ", dest, " = ", src, EOL[]}]; (* TODO: move these into OpenMP loop *) DeclareVariablesInLoopVectorised[dests_, temps_, srcs_] :=