diff --git a/ocelot/CMakeLists.txt b/ocelot/CMakeLists.txt index 2662d58f0..10f247527 100644 --- a/ocelot/CMakeLists.txt +++ b/ocelot/CMakeLists.txt @@ -420,6 +420,10 @@ endif() add_library(${PROJECT_NAME}_executive STATIC ${${PROJECT_NAME}_executive_sources}) +if(CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang") +set_source_files_properties(src/executive/CooperativeThreadArray.cpp PROPERTIES COMPILE_OPTIONS -frounding-math) +endif() + set_property(TARGET ${PROJECT_NAME}_executive PROPERTY CXX_STANDARD 14) set_property(TARGET ${PROJECT_NAME}_executive PROPERTY POSITION_INDEPENDENT_CODE ON) target_compile_definitions(${PROJECT_NAME}_executive PRIVATE ${${PROJECT_NAME}_DEFINITIONS}) diff --git a/ocelot/include/ocelot/executive/CooperativeThreadArray.h b/ocelot/include/ocelot/executive/CooperativeThreadArray.h index 84e51a003..562ae7e0e 100644 --- a/ocelot/include/ocelot/executive/CooperativeThreadArray.h +++ b/ocelot/include/ocelot/executive/CooperativeThreadArray.h @@ -456,9 +456,12 @@ namespace executive { void eval_AddC(CTAContext &context, const ir::PTXInstruction &instr); void eval_And(CTAContext &context, const ir::PTXInstruction &instr); void eval_Atom(CTAContext &context, const ir::PTXInstruction &instr); + void evalAtomicRMW(CTAContext &context, const ir::PTXInstruction &instr, + ir::PTXInstruction::AtomicOperation operation, bool writeback); void eval_Bar(CTAContext &context, const ir::PTXInstruction &instr); void eval_Bfi(CTAContext &context, const ir::PTXInstruction &instr); void eval_Bfind(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Bmsk(CTAContext &context, const ir::PTXInstruction &instr); void eval_Bfe(CTAContext &context, const ir::PTXInstruction &instr); void eval_Bra(CTAContext &context, const ir::PTXInstruction &instr); void eval_Brev(CTAContext &context, const ir::PTXInstruction &instr); @@ -472,18 +475,23 @@ namespace executive { void eval_Cvt(CTAContext &context, const ir::PTXInstruction &instr); void eval_Cvta(CTAContext &context, const ir::PTXInstruction &instr); void eval_Div(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Dp(CTAContext &context, const ir::PTXInstruction &instr); void eval_Ex2(CTAContext &context, const ir::PTXInstruction &instr); void eval_Exit(CTAContext &context, const ir::PTXInstruction &instr); void eval_Fma(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Fns(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Mma(CTAContext &context, const ir::PTXInstruction &instr); void eval_Isspacep(CTAContext &context, const ir::PTXInstruction &instr); void eval_Ld(CTAContext &context, const ir::PTXInstruction &instr); void eval_Ldu(CTAContext &context, const ir::PTXInstruction &instr); void eval_Lg2(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Lop3(CTAContext &context, const ir::PTXInstruction &instr); void eval_Mad24(CTAContext &context, const ir::PTXInstruction &instr); void eval_Mad(CTAContext &context, const ir::PTXInstruction &instr); void eval_Max(CTAContext &context, const ir::PTXInstruction &instr); void eval_Membar(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Fence(CTAContext &context, const ir::PTXInstruction &instr); void eval_Min(CTAContext &context, const ir::PTXInstruction &instr); void eval_Mov(CTAContext &context, const ir::PTXInstruction &instr); void eval_Mul24(CTAContext &context, const ir::PTXInstruction &instr); @@ -520,6 +528,8 @@ namespace executive { void eval_Sured(CTAContext &context, const ir::PTXInstruction &instr); void eval_Sust(CTAContext &context, const ir::PTXInstruction &instr); void eval_Suq(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Szext(CTAContext &context, const ir::PTXInstruction &instr); + void eval_Tanh(CTAContext &context, const ir::PTXInstruction &instr); void eval_TestP(CTAContext &context, const ir::PTXInstruction &instr); void eval_Tex(CTAContext &context, const ir::PTXInstruction &instr); void eval_Trap(CTAContext &context, const ir::PTXInstruction &instr); @@ -562,4 +572,3 @@ namespace executive { } #endif - diff --git a/ocelot/include/ocelot/executive/EmulatedKernel.h b/ocelot/include/ocelot/executive/EmulatedKernel.h index 1d36a91f5..0ac7fee48 100644 --- a/ocelot/include/ocelot/executive/EmulatedKernel.h +++ b/ocelot/include/ocelot/executive/EmulatedKernel.h @@ -233,7 +233,7 @@ namespace executive { TextureVector textures; /*! A handle to the current scheduler, or 0 if none is executing */ - EmulatedKernelScheduler* scheduler; + EmulatedKernelScheduler* scheduler = nullptr; private: /*! Maps program counter to the kernel that begins there */ diff --git a/ocelot/include/ocelot/executive/EmulatedKernelScheduler.h b/ocelot/include/ocelot/executive/EmulatedKernelScheduler.h index 795e6e4c9..e27ae2158 100644 --- a/ocelot/include/ocelot/executive/EmulatedKernelScheduler.h +++ b/ocelot/include/ocelot/executive/EmulatedKernelScheduler.h @@ -49,6 +49,8 @@ class EmulatedKernelScheduler const ir::Dim3& ctaDim, ir::PTXU32 sharedMemory, ir::PTXU64 stream); /*! \brief Get the argument memory for the current context */ ir::PTXU64 argumentMemory() const; + /*! \brief Get the argument memory size for the current context */ + ir::PTXU64 argumentMemorySize() const; private: class Context @@ -123,4 +125,3 @@ class EmulatedKernelScheduler } - diff --git a/ocelot/include/ocelot/executive/EmulatorCallStack.h b/ocelot/include/ocelot/executive/EmulatorCallStack.h index 8ed384673..33d1d618d 100644 --- a/ocelot/include/ocelot/executive/EmulatorCallStack.h +++ b/ocelot/include/ocelot/executive/EmulatorCallStack.h @@ -89,6 +89,9 @@ class EmulatorCallStack const RegisterType* registerFilePointer(unsigned int thread) const; /*! \brief Get a pointer to local memory for a given thread */ void* localMemoryPointer(unsigned int thread); + /*! \brief Test whether an address belongs to any active local frame */ + bool isLocalMemoryAddress(unsigned long long address, + unsigned int thread, unsigned int addressBits) const; /*! \brief Get a pointer to shared memory */ void* sharedMemoryPointer(); /*! \brief Get a pointer to global local memory */ diff --git a/ocelot/include/ocelot/ir/PTXInstruction.h b/ocelot/include/ocelot/ir/PTXInstruction.h index efa813a84..f27d612f6 100644 --- a/ocelot/include/ocelot/ir/PTXInstruction.h +++ b/ocelot/include/ocelot/ir/PTXInstruction.h @@ -22,6 +22,16 @@ namespace ir { Level_Invalid }; + enum Semantics { + Sc, + AcqRel, + Acquire, + Release, + Relaxed, + Weak, + Semantics_Invalid + }; + /*! List of opcodes for PTX instructions */ enum Opcode { Abs = 0, @@ -33,6 +43,7 @@ namespace ir { Bfe, Bfi, Bfind, + Bmsk, Bra, Brev, Brkpt, @@ -44,18 +55,24 @@ namespace ir { Cvt, Cvta, Div, + Dp2a, + Dp4a, Ex2, Exit, Fma, + Fns, Isspacep, Ld, Ldu, Lg2, + Lop3, Mad24, Mad, MadC, + Mma, Max, Membar, + Fence, Min, Mov, Mul24, @@ -91,6 +108,8 @@ namespace ir { Sured, Sust, Suq, + Szext, + Tanh, TestP, Tex, Tld4, @@ -136,6 +155,12 @@ namespace ir { approx = 8192,//< identify an approximate instruction ftz = 16384, //< flush to zero full = 32768, //< full division + nan = 65536, //< return NaN if either input is NaN + xorsign = 131072, //< XOR input sign bits + abs = 262144, //< compare absolute input values + relu = 524288, //< clamp negative floating-point results to zero + rna = 1048576, //< round to nearest, ties away from zero + satfinite = 2097152, //< clamp integer mma result to representable range Modifier_invalid = 0 }; @@ -143,6 +168,21 @@ namespace ir { None = 0, CC = 1 }; + + enum MmaShape { + MmaM16N8K8, + MmaM16N8K16, + MmaM8N8K16, + MmaM16N8K32, + MmaM16N8K4, + MmaM8N8K4, + MmaM8N8K32, + MmaM16N8K64, + MmaM8N8K128, + MmaM16N8K128, + MmaM16N8K256, + MmaShape_Invalid + }; enum Volatility { Nonvolatile = 0, @@ -200,6 +240,7 @@ namespace ir { Cg = 2, Cs = 3, Nc = 4, + Lu = 5, Wb = 0, Wt = 1, CacheOperation_Invalid @@ -344,6 +385,7 @@ namespace ir { public: static std::string toString( Level ); + static std::string toString( Semantics ); static std::string toString( CacheLevel cache ); static std::string toStringLoad( CacheOperation op ); static std::string toStringStore( CacheOperation op ); @@ -425,6 +467,13 @@ namespace ir { /*! indicates data type of instruction */ PTXOperand::DataType type; + /*! Second input type for packed dot-product instructions */ + PTXOperand::DataType bType; + + /*! Shape for MMA instructions */ + MmaShape mmaShape; + bool mmaAColumnMajor; + bool mmaBColumnMajor; /*! Flag containing one or more floating-point modifiers */ unsigned int modifier; @@ -441,7 +490,7 @@ namespace ir { /*! For membar, the visibility level in the thread hierarchy */ Level level; - + /*! Shift amount flag for bfind instructions */ bool shiftAmount; @@ -468,8 +517,30 @@ namespace ir { }; + /*! For fence, the memory ordering semantics -- deliberately outside the + level/shuffleMode/etc. union above: fence needs both a semantics + AND a level (scope) simultaneously, unlike the other union members + which are each only ever used by instructions with no need for + the others at the same time. */ + Semantics semantics; + + /*! For atom/red, the optional memory-ordering scope (.cta/.gpu/.sys) + -- deliberately its own field rather than reusing the level/ + addressSpace union member above: atom/red already use addressSpace + for their .global/.shared qualifier, so writing scope through the + shared level union slot would alias and corrupt it. */ + Level scope; + /*! For call instructions, indicates a tail call */ bool tailCall; + + /*! For ld/st, indicates the .mmio memory-mapped-I/O form + (always .sem.sys{.global}) -- deliberately a standalone flag + rather than inferred from semantics/scope alone: a plain + ld.acquire.sys and ld.mmio.acquire.sys produce the same + semantics/scope pair but are syntactically distinct and must + round-trip through the printer differently. */ + bool mmio; /*! If the instruction is predicated, the guard */ PTXOperand pg; @@ -503,9 +574,6 @@ namespace ir { /*! Indicates whether the target address space is volatile */ Volatility volatility; - /*! Is this a divide full instruction? */ - bool divideFull; - /*! If cvta instruction, indicates whether destination is generic address or if source is generic address - true if segmented address space, false if generic */ @@ -554,6 +622,10 @@ namespace ir { /*! Source operand c */ PTXOperand c; + /*! Lookup table and predicate input for lop3 */ + PTXOperand immLut; + PTXOperand q; + /* Runtime annotations The following members are used to annotate the instruction @@ -602,4 +674,3 @@ namespace ir { } #endif - diff --git a/ocelot/include/ocelot/ir/PTXOperand.h b/ocelot/include/ocelot/ir/PTXOperand.h index 853b719c2..b7e04b672 100644 --- a/ocelot/include/ocelot/ir/PTXOperand.h +++ b/ocelot/include/ocelot/ir/PTXOperand.h @@ -25,6 +25,7 @@ namespace ir { typedef int32_t PTXS32; typedef int64_t PTXS64; + typedef _Float16 PTXF16; typedef float PTXF32; typedef double PTXF64; @@ -62,12 +63,17 @@ namespace ir { u64, f16, f32, + bf16, + tf32, f64, b8, b16, b32, b64, - pred + pred, + f16x2, + bf16x2, + s4, u4, b1 // MMA element types; stored in packed b32 registers. }; /*! Special register names */ @@ -139,7 +145,8 @@ namespace ir { enum Vec { v1 = 1, //< scalar v2 = 2, //< vector2 - v4 = 4 //< vector4 + v4 = 4, //< vector4 + v8 = 8 //< eight-register MMA fragment }; enum VectorIndex { @@ -285,4 +292,3 @@ namespace std { } #endif - diff --git a/ocelot/include/ocelot/parser/PTXParser.h b/ocelot/include/ocelot/parser/PTXParser.h index 1527b55c4..88f70c5ae 100644 --- a/ocelot/include/ocelot/parser/PTXParser.h +++ b/ocelot/include/ocelot/parser/PTXParser.h @@ -137,6 +137,7 @@ namespace parser private: void _setImmediateTypes(); + void _setMovVectorImmediateTypes(); std::string _nameInContext( const std::string& name ); OperandWrapper* _getOperand( const std::string& name ); @@ -236,6 +237,7 @@ namespace parser void constantOperand( double value ); void indexedOperand( const std::string& name, YYLTYPE& location, long long int value ); + void vectorOperand( unsigned int elements ); void addressableOperand( const std::string& name, long long int value, YYLTYPE& location, bool invert ); @@ -259,6 +261,10 @@ namespace parser void vote( int token ); void shuffle( int token ); void level( int token ); + void semantics( int token ); + void scope( int token ); + void mmio( bool condition ); + void finalizeMmioAddressSpace(); void permute( int token ); void floatingPointMode( int token ); void defaultPermute(); @@ -267,6 +273,9 @@ namespace parser void instruction(); void instruction( const std::string& opcode, int dataType ); void instruction( const std::string& opcode ); + void dotType( int token ); + void lop3(); + void mma( int shape, int accumulatorType, int aType, int bType, int cType, bool aColumnMajor = false, bool bColumnMajor = true ); void tex( int dataType ); void tld4( int dataType ); void callPrototypeName( const std::string& identifier ); @@ -345,6 +354,7 @@ namespace parser static ir::PTXInstruction::VoteMode tokenToVoteMode( int ); static ir::PTXInstruction::ShuffleMode tokenToShuffleMode( int ); static ir::PTXInstruction::Level tokenToLevel( int ); + static ir::PTXInstruction::Semantics tokenToSemantics( int ); static ir::PTXInstruction::PermuteMode tokenToPermuteMode( int ); static ir::PTXInstruction::FloatingPointMode tokenToFloatingPointMode( int); @@ -366,4 +376,3 @@ namespace parser } #endif - diff --git a/ocelot/src/cuda/test/TestPTXAssembly.cpp b/ocelot/src/cuda/test/TestPTXAssembly.cpp index a37e66fe2..0a9b5e1e2 100644 --- a/ocelot/src/cuda/test/TestPTXAssembly.cpp +++ b/ocelot/src/cuda/test/TestPTXAssembly.cpp @@ -1021,15 +1021,6 @@ std::string testShf_PTX( ptx << "\tld.global.u32 %r2, [%rIn + " << 2 * ir::PTXOperand::bytes(ir::PTXOperand::b32) << "]; \n"; - if (mode == ir::PTXInstruction::ShiftMode::Wrap) - { - ptx << "\tand.b32 %r2, %r2, 31; \n"; // Wrap mode: limit to 5 bits - } - else if (mode == ir::PTXInstruction::ShiftMode::Clamp) - { - ptx << "\tmin.u32 %r2, %r2, 32; \n"; // Clamp mode: cap at 32 - } - ptx << "\tshf" << directionString << modeString << typeString << " %r3, %r0, %r1, %r2; \n"; @@ -1066,14 +1057,15 @@ void testShf_REF(void* output, void* input) c = (c > 32) ? 32 : c; // Clamp mode: limit to 32 } + const uint64_t pair = (static_cast(b) << 32) | a; U32 result; if (direction == ir::PTXInstruction::ShiftLeft) { - result = (b << c) | (a >> (32 - c)); + result = (pair << c) >> 32; } else if (direction == ir::PTXInstruction::ShiftRight) { - result = (b << (32 - c)) | (a >> c); + result = pair >> c; } else { @@ -1145,8 +1137,6 @@ std::string testLops_PTX(ir::PTXInstruction::Opcode opcode, ptx << "\tld.global.u32 %r1, [%rIn + " << std::max((size_t)ir::PTXOperand::bytes(type), sizeof(uint32_t)) << "]; \n"; - ptx << "\trem.u32 %r1, %r1, " - << 8 * ir::PTXOperand::bytes(type) << ";\n"; } else if(opcode == ir::PTXInstruction::And || opcode == ir::PTXInstruction::Or @@ -1264,7 +1254,8 @@ void testLops_REF(void* output, void* input) r0 = r0 < 64; } - type d = r0 >> (r1 % (sizeof(type) * 8)); + type d = r1 >= sizeof(type) * 8 + ? (r0 < 0 ? -1 : 0) : r0 >> r1; setParameter(output, 0, d); break; @@ -1275,7 +1266,7 @@ void testLops_REF(void* output, void* input) uint32_t r1 = getParameter(input, std::max(sizeof(type), sizeof(uint32_t))); - type d = r0 << (r1 % (sizeof(type) * 8)); + type d = r1 >= sizeof(type) * 8 ? 0 : r0 << r1; setParameter(output, 0, d); break; @@ -6886,6 +6877,21 @@ namespace test testLops_OUT(I64), testLops_IN(ir::PTXInstruction::Shl, I64), uniformRandom, 1, 1); + add("TestShr-b16", + testLops_REF, + testLops_PTX(ir::PTXInstruction::Shr, ir::PTXOperand::b16), + testLops_OUT(I16), testLops_IN(ir::PTXInstruction::Shr, I16), + uniformRandom, 1, 1); + add("TestShr-b32", testLops_REF, + testLops_PTX(ir::PTXInstruction::Shr, ir::PTXOperand::b32), + testLops_OUT(I32), testLops_IN(ir::PTXInstruction::Shr, I32), + uniformRandom, 1, 1); + add("TestShr-b64", + testLops_REF, + testLops_PTX(ir::PTXInstruction::Shr, ir::PTXOperand::b64), + testLops_OUT(I64), testLops_IN(ir::PTXInstruction::Shr, I64), + uniformRandom, 1, 1); + add("TestShr-u16", testLops_REF, testLops_PTX(ir::PTXInstruction::Shr, ir::PTXOperand::u16), @@ -7824,4 +7830,3 @@ int main(int argc, char** argv) } #endif - diff --git a/ocelot/src/executive/CooperativeThreadArray.cpp b/ocelot/src/executive/CooperativeThreadArray.cpp index 81acd6dd9..9ba6c3068 100644 --- a/ocelot/src/executive/CooperativeThreadArray.cpp +++ b/ocelot/src/executive/CooperativeThreadArray.cpp @@ -27,10 +27,12 @@ // Standard Library Includes #include +#include #include #include #include #include +#include #include // Preprocessor Macros @@ -111,6 +113,20 @@ static T CTAAbs(T a) { return a; } +template +static bool addressInRange(T address, T base, T size) { + return address >= base && address - base < size; +} + +template +static T CTARemainder(T a, T b) { + if (b == 0) { + report("warning: rem by zero is unspecified by PTX; Ocelot returns 0"); + return 0; + } + return a % b; +} + template bool issubnormal_(T r0) { @@ -118,6 +134,85 @@ bool issubnormal_(T r0) && !hydrazine::isinf(r0) && r0 != (T)0; } +static int setRoundingMode(int modifier) +{ + const int previous = hydrazine::fegetround(); + int rounding = FE_TONEAREST; + if (modifier & ir::PTXInstruction::rz) rounding = FE_TOWARDZERO; + else if (modifier & ir::PTXInstruction::rm) rounding = FE_DOWNWARD; + else if (modifier & ir::PTXInstruction::rp) rounding = FE_UPWARD; + hydrazine::fesetround(rounding); + return previous; +} + +template +static T roundedAdd(T a, T b, int modifier) +{ + const int previous = setRoundingMode(modifier); + T d = a + b; + hydrazine::fesetround(previous); + return d; +} + +template +static T roundedSub(T a, T b, int modifier) +{ + const int previous = setRoundingMode(modifier); + T d = a - b; + hydrazine::fesetround(previous); + return d; +} + +template +static T roundedMul(T a, T b, int modifier) +{ + const int previous = setRoundingMode(modifier); + T d = a * b; + hydrazine::fesetround(previous); + return d; +} + +template +static T roundedDiv(T a, T b, int modifier) +{ + const int previous = setRoundingMode(modifier); + T d = a / b; + hydrazine::fesetround(previous); + return d; +} + +template +static T roundedSqrt(T a, int modifier) +{ + const int previous = setRoundingMode(modifier); + T d = std::sqrt(a); + hydrazine::fesetround(previous); + return d; +} + +template +static T roundedFma(T a, T b, T c, int modifier) +{ + const int previous = setRoundingMode(modifier); + T d = std::fma(a, b, c); + hydrazine::fesetround(previous); + return d; +} + +// Truncate a * b + c to f32 and force the last bit to 1 if any bits were lost +// (round-to-odd), so the later RNE to bf16 cannot see a false tie. +static ir::PTXF32 fmaF32OneRound(ir::PTXF32 a, ir::PTXF32 b, ir::PTXF32 c) +{ + std::feclearexcept(FE_INEXACT); + const int previous = setRoundingMode(ir::PTXInstruction::rz); + volatile ir::PTXF32 d = static_cast( + static_cast(a) * b + c); // bf16 product is exact in double + const ir::PTXU32 lost = std::fetestexcept(FE_INEXACT) ? 1U : 0U; + hydrazine::fesetround(previous); + return hydrazine::bit_cast( + hydrazine::bit_cast(static_cast(d)) | lost); +} + static executive::ReconvergenceMechanism* getReconvergenceMechanism(executive::CooperativeThreadArray* cta) { @@ -372,6 +467,154 @@ static ir::PTXF32 ftz(int modifier, ir::PTXF32 f) { return f; } +static ir::PTXF32 minMaxF32(int modifier, ir::PTXF32 a, + ir::PTXF32 b, bool maximum) { + const bool xorSign = std::signbit(a) != std::signbit(b); + if (modifier & ir::PTXInstruction::xorsign) { + a = std::fabs(a); + b = std::fabs(b); + } + ir::PTXF32 d; + if ((modifier & ir::PTXInstruction::nan) + && (hydrazine::isnan(a) || hydrazine::isnan(b))) { + d = hydrazine::bit_cast(0x7fffffffU); + } else if (hydrazine::isnan(a)) { + d = b; + } else if (hydrazine::isnan(b)) { + d = a; + } else { + d = maximum ? (a > b ? a : b) : (a < b ? a : b); + } + d = ftz(modifier, d); + return (modifier & ir::PTXInstruction::xorsign) && !hydrazine::isnan(d) + ? hydrazine::copysign(d, xorSign ? -1.0f : 1.0f) : d; +} + +static ir::PTXU16 ftzF16(int modifier, ir::PTXU16 bits) +{ + const bool subnormal = (bits & 0x7c00u) == 0 && (bits & 0x03ffu) != 0; + if ((modifier & ir::PTXInstruction::ftz) && subnormal) { + return bits & 0x8000u; // preserve sign: +0 or -0 + } + return bits; +} + +static ir::PTXF32 f16ToF32(ir::PTXU16 bits) { + ir::PTXF16 half; + std::memcpy(&half, &bits, sizeof(half)); + return static_cast(half); +} + +static ir::PTXF32 bf16ToF32(ir::PTXU16 bits) { + return hydrazine::bit_cast(static_cast(bits) << 16); +} + +// PTX half formats use sign/magnitude masks; CUDA defines 0x7fff as canonical NaN. +static const ir::PTXU16 halfSignMask = 0x8000u; +static const ir::PTXU16 halfMagnitudeMask = 0x7fffu; +static const ir::PTXU16 halfCanonicalNan = 0x7fffu; +static const ir::PTXU16 f16ExponentMask = 0x7c00u; +static const ir::PTXU16 f16MantissaMask = 0x03ffu; +static const ir::PTXU16 bf16ExponentMask = 0x7f80u; +static const ir::PTXU16 bf16MantissaMask = 0x007fu; + +static bool isNanHalf(ir::PTXU16 value, ir::PTXOperand::DataType type) { + const bool bfloat = type == ir::PTXOperand::bf16 || type == ir::PTXOperand::bf16x2; + const ir::PTXU16 exponentMask = bfloat ? bf16ExponentMask : f16ExponentMask; + const ir::PTXU16 mantissaMask = bfloat ? bf16MantissaMask : f16MantissaMask; + return (value & exponentMask) == exponentMask && (value & mantissaMask) != 0; +} + +static ir::PTXU16 minMaxHalf(ir::PTXOperand::DataType type, int modifier, + ir::PTXU16 a, ir::PTXU16 b, bool maximum) { + const bool isF16 = type == ir::PTXOperand::f16 || type == ir::PTXOperand::f16x2; + if (isF16) { + a = ftzF16(modifier, a); + b = ftzF16(modifier, b); + } + ir::PTXU16 xorSign = 0; + if (modifier & ir::PTXInstruction::xorsign) { + xorSign = (a ^ b) & halfSignMask; + a &= halfMagnitudeMask; + b &= halfMagnitudeMask; + } + const bool aNan = isNanHalf(a, type); + const bool bNan = isNanHalf(b, type); + const bool bothZero = ((a | b) & halfMagnitudeMask) == 0; + if (aNan && bNan) { + return halfCanonicalNan; + } + if ((modifier & ir::PTXInstruction::nan) && (aNan || bNan)) { + return halfCanonicalNan; + } + ir::PTXU16 d; + if (aNan) { + d = b; + } + else if (bNan) { + d = a; + } + else if (bothZero) { + d = maximum ? (a & b) : (a | b); + } + else { + const ir::PTXF32 aValue = isF16 ? f16ToF32(a) : bf16ToF32(a); + const ir::PTXF32 bValue = isF16 ? f16ToF32(b) : bf16ToF32(b); + d = maximum ? (aValue > bValue ? a : b) : (aValue < bValue ? a : b); + } + if (modifier & ir::PTXInstruction::xorsign) { + return (d & halfMagnitudeMask) | xorSign; + } + return d; +} + +static ir::PTXF32 tf32FromF32(ir::PTXF32 value) { + ir::PTXU32 bits = hydrazine::bit_cast(value); + if ((bits & 0x7f800000u) == 0x7f800000u) return value; + const ir::PTXU32 discarded = bits & 0x1fffu; + bits &= ~0x1fffu; + if (discarded > 0x1000u || + (discarded == 0x1000u && (bits & 0x2000u))) { + bits += 0x2000u; + } + return hydrazine::bit_cast(bits); +} + +static ir::PTXU32 tf32FromF32Rna(ir::PTXF32 value) { + // Model TF32 by retaining the upper 10 f32 fraction bits. + const ir::PTXU32 exponentMask = 0x7f800000u; + const ir::PTXU32 discardedMask = (1u << 13) - 1; + const ir::PTXU32 halfway = 1u << 12; + const ir::PTXU32 retainedUnit = 1u << 13; + ir::PTXU32 bits = hydrazine::bit_cast(value); + if ((bits & exponentMask) == exponentMask) return bits; + const ir::PTXU32 discarded = bits & discardedMask; + bits &= ~discardedMask; + if (discarded >= halfway) bits += retainedUnit; + return bits; +} + +template< typename Source > +static ir::PTXU16 toF16(Source value, int modifier); + +static ir::PTXU16 mmaHalfBits(executive::CooperativeThreadArray& cta, + int threadID, const ir::PTXOperand& operand, unsigned int half) +{ + if (operand.bytes() == 4) { + return static_cast( + (cta.operandAsB32(threadID, operand) >> (half * 16)) & 0xffffu); + } + return cta.operandAsU16(threadID, operand); +} + +static ir::PTXF32 mmaHalf(executive::CooperativeThreadArray& cta, + int threadID, const ir::PTXOperand& operand, unsigned int half, + ir::PTXOperand::DataType type) +{ + ir::PTXU16 bits = mmaHalfBits(cta, threadID, operand, half); + return type == ir::PTXOperand::bf16 ? bf16ToF32(bits) : f16ToF32(bits); +} + void executive::CooperativeThreadArray::trace() { if (traceEvents) { currentEvent.contextStackSize = @@ -479,6 +722,8 @@ void executive::CooperativeThreadArray::execute(int PC) { eval_Bfi(context, instr); break; case ir::PTXInstruction::Bfind: eval_Bfind(context, instr); break; + case ir::PTXInstruction::Bmsk: + eval_Bmsk(context, instr); break; case ir::PTXInstruction::Bfe: eval_Bfe(context, instr); break; case ir::PTXInstruction::Bra: @@ -503,18 +748,27 @@ void executive::CooperativeThreadArray::execute(int PC) { eval_Cvta(context, instr); break; case ir::PTXInstruction::Div: eval_Div(context, instr); break; + case ir::PTXInstruction::Dp2a: + case ir::PTXInstruction::Dp4a: + eval_Dp(context, instr); break; case ir::PTXInstruction::Ex2: eval_Ex2(context, instr); break; case ir::PTXInstruction::Exit: eval_Exit(context, instr); break; case ir::PTXInstruction::Fma: eval_Fma(context, instr); break; + case ir::PTXInstruction::Fns: + eval_Fns(context, instr); break; + case ir::PTXInstruction::Mma: + eval_Mma(context, instr); break; case ir::PTXInstruction::Isspacep: eval_Isspacep(context, instr); break; case ir::PTXInstruction::Ld: eval_Ld(context, instr); break; case ir::PTXInstruction::Lg2: eval_Lg2(context, instr); break; + case ir::PTXInstruction::Lop3: + eval_Lop3(context, instr); break; case ir::PTXInstruction::Ldu: eval_Ldu(context, instr); break; case ir::PTXInstruction::Mad24: @@ -525,6 +779,8 @@ void executive::CooperativeThreadArray::execute(int PC) { eval_Max(context, instr); break; case ir::PTXInstruction::Membar: eval_Membar(context, instr); break; + case ir::PTXInstruction::Fence: + eval_Fence(context, instr); break; case ir::PTXInstruction::Min: eval_Min(context, instr); break; case ir::PTXInstruction::Mov: @@ -585,6 +841,10 @@ void executive::CooperativeThreadArray::execute(int PC) { eval_Sub(context, instr); break; case ir::PTXInstruction::SubC: eval_SubC(context, instr); break; + case ir::PTXInstruction::Szext: + eval_Szext(context, instr); break; + case ir::PTXInstruction::Tanh: + eval_Tanh(context, instr); break; case ir::PTXInstruction::TestP: eval_TestP(context, instr); break; case ir::PTXInstruction::Tex: @@ -1657,7 +1917,25 @@ void executive::CooperativeThreadArray::setFunctionParameter(int threadID, void executive::CooperativeThreadArray::eval_Abs(CTAContext &context, const ir::PTXInstruction &instr) { trace(); - if (instr.type == ir::PTXOperand::f32) { + if (instr.type == ir::PTXOperand::f16 || instr.type == ir::PTXOperand::bf16) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + ir::PTXU16 a = operandAsU16(threadID, instr.a); + if (instr.type == ir::PTXOperand::f16) a = ftzF16(instr.modifier, a); + setRegAsB16(threadID, instr.d.reg, a & 0x7fff); + } + } + else if (instr.type == ir::PTXOperand::f16x2 || instr.type == ir::PTXOperand::bf16x2) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + ir::PTXU32 a = operandAsU32(threadID, instr.a); + if (instr.type == ir::PTXOperand::f16x2) + a = static_cast(ftzF16(instr.modifier, a)) | + (static_cast(ftzF16(instr.modifier, a >> 16)) << 16); + setRegAsU32(threadID, instr.d.reg, a & 0x7fff7fffu); + } + } + else if (instr.type == ir::PTXOperand::f32) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; @@ -1717,12 +1995,46 @@ void executive::CooperativeThreadArray::eval_Abs(CTAContext &context, void executive::CooperativeThreadArray::eval_Add(CTAContext &context, const ir::PTXInstruction &instr) { trace(); - if (instr.type == ir::PTXOperand::f32) { + if (instr.type == ir::PTXOperand::f16x2) { + const int effectiveModifier = instr.modifier | ir::PTXInstruction::rn; + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + ir::PTXU32 a = operandAsU32(threadID, instr.a); + ir::PTXU32 b = operandAsU32(threadID, instr.b); + ir::PTXU16 al = static_cast(a), ah = a >> 16; + ir::PTXU16 bl = static_cast(b), bh = b >> 16; + ir::PTXU16 dl = ftzF16(effectiveModifier, toF16( + roundedAdd(f16ToF32(ftzF16(effectiveModifier, al)), + f16ToF32(ftzF16(effectiveModifier, bl)), effectiveModifier), + effectiveModifier)); + ir::PTXU16 dh = ftzF16(effectiveModifier, toF16( + roundedAdd(f16ToF32(ftzF16(effectiveModifier, ah)), + f16ToF32(ftzF16(effectiveModifier, bh)), effectiveModifier), + effectiveModifier)); + setRegAsU32(threadID, instr.d.reg, static_cast(dl) | + (static_cast(dh) << 16)); + } + } + else if (instr.type == ir::PTXOperand::f16) { + const int effectiveModifier = instr.modifier | ir::PTXInstruction::rn; + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + + ir::PTXF32 a = f16ToF32(ftzF16(effectiveModifier, + operandAsU16(threadID, instr.a))); + ir::PTXF32 b = f16ToF32(ftzF16(effectiveModifier, + operandAsU16(threadID, instr.b))); + ir::PTXU16 d = toF16(roundedAdd(a, b, effectiveModifier), effectiveModifier); + setRegAsB16(threadID, instr.d.reg, ftzF16(effectiveModifier, d)); + } + } + else if (instr.type == ir::PTXOperand::f32) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; ir::PTXF32 d, a = ftz(instr.modifier, operandAsF32(threadID, instr.a)), b = ftz(instr.modifier, operandAsF32(threadID, instr.b)); - d = ftz(instr.modifier, sat(instr.modifier, a + b)); + d = ftz(instr.modifier, sat(instr.modifier, + roundedAdd(a, b, instr.modifier))); setRegAsF32(threadID, instr.d.reg, d); } } @@ -1731,7 +2043,7 @@ void executive::CooperativeThreadArray::eval_Add(CTAContext &context, if (!context.predicated(threadID, instr)) continue; ir::PTXF64 d, a = operandAsF64(threadID, instr.a), b = operandAsF64(threadID, instr.b); - d = a + b; + d = roundedAdd(a, b, instr.modifier); setRegAsF64(threadID, instr.d.reg, d); } } @@ -1917,6 +2229,17 @@ void executive::CooperativeThreadArray::eval_And(CTAContext &context, */ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir::PTXInstruction &instr) { + evalAtomicRMW(context, instr, instr.atomicOperation, true); +} + +// atom writes the pre-update value back to instr.d; red performs the same +// read-modify-write but discards it (writeback == false). red carries its +// operation in instr.reductionOperation (a distinct enum/op set from atom's +// instr.atomicOperation), so the caller resolves and passes the operation +// explicitly rather than this function reading instr.atomicOperation itself. +void executive::CooperativeThreadArray::evalAtomicRMW(CTAContext &context, + const ir::PTXInstruction &instr, ir::PTXInstruction::AtomicOperation operation, + bool writeback) { size_t elementSize = 0; switch (instr.type) { @@ -1929,7 +2252,9 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: } break; case ir::PTXOperand::b64: // fall through - case ir::PTXOperand::u64: + case ir::PTXOperand::u64: // fall through + case ir::PTXOperand::s64: // fall through + case ir::PTXOperand::f64: { elementSize = sizeof(ir::PTXU64); } @@ -2006,53 +2331,77 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: context.PC, instr); } - switch (instr.atomicOperation) { + switch (operation) { case ir::PTXInstruction::AtomicAnd: { - if(instr.type != ir::PTXOperand::b32 - && instr.type != ir::PTXOperand::s32 - && instr.type != ir::PTXOperand::u32) { + if (instr.type == ir::PTXOperand::b32) { + ir::PTXB32 d = *((ir::PTXB32*)source); + if (writeback) setRegAsB32(threadID, instr.d.reg, d); + ir::PTXB32 b = operandAsB32(threadID, instr.b); + *((ir::PTXB32*)source) = d & b; + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXB32*)source) ); + } + else if (instr.type == ir::PTXOperand::b64) { + ir::PTXB64 d = *((ir::PTXB64*)source); + if (writeback) setRegAsB64(threadID, instr.d.reg, d); + ir::PTXB64 b = operandAsB64(threadID, instr.b); + *((ir::PTXB64*)source) = d & b; + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXB64*)source) ); + } + else { throw RuntimeException("invalid data type", context.PC, instr); } - ir::PTXB32 d = *((ir::PTXB32*)source); - setRegAsB32(threadID, instr.d.reg, d); - ir::PTXB32 b = operandAsB32(threadID, instr.b); - *((ir::PTXB32*)source) = d & b; - reportE(REPORT_ATOM, "Atomically updated " << d << " to " - << *((ir::PTXB32*)source) ); } break; case ir::PTXInstruction::AtomicOr: { - if(instr.type != ir::PTXOperand::b32 - && instr.type != ir::PTXOperand::s32 - && instr.type != ir::PTXOperand::u32) { + if (instr.type == ir::PTXOperand::b32) { + ir::PTXB32 d = *((ir::PTXB32*)source); + if (writeback) setRegAsB32(threadID, instr.d.reg, d); + ir::PTXB32 b = operandAsB32(threadID, instr.b); + *((ir::PTXB32*)source) = d | b; + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXB32*)source) ); + } + else if (instr.type == ir::PTXOperand::b64) { + ir::PTXB64 d = *((ir::PTXB64*)source); + if (writeback) setRegAsB64(threadID, instr.d.reg, d); + ir::PTXB64 b = operandAsB64(threadID, instr.b); + *((ir::PTXB64*)source) = d | b; + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXB64*)source) ); + } + else { throw RuntimeException("invalid data type", context.PC, instr); } - ir::PTXB32 d = *((ir::PTXB32*)source); - setRegAsB32(threadID, instr.d.reg, d); - ir::PTXB32 b = operandAsB32(threadID, instr.b); - *((ir::PTXB32*)source) = d | b; - reportE(REPORT_ATOM, "Atomically updated " << d << " to " - << *((ir::PTXB32*)source) ); } break; case ir::PTXInstruction::AtomicXor: { - if(instr.type != ir::PTXOperand::b32 - && instr.type != ir::PTXOperand::s32 - && instr.type != ir::PTXOperand::u32) { + if (instr.type == ir::PTXOperand::b32) { + ir::PTXB32 d = *((ir::PTXB32*)source); + if (writeback) setRegAsB32(threadID, instr.d.reg, d); + ir::PTXB32 b = operandAsB32(threadID, instr.b); + *((ir::PTXB32*)source) = d ^ b; + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXB32*)source) ); + } + else if (instr.type == ir::PTXOperand::b64) { + ir::PTXB64 d = *((ir::PTXB64*)source); + if (writeback) setRegAsB64(threadID, instr.d.reg, d); + ir::PTXB64 b = operandAsB64(threadID, instr.b); + *((ir::PTXB64*)source) = d ^ b; + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXB64*)source) ); + } + else { throw RuntimeException("invalid data type", context.PC, instr); } - ir::PTXB32 d = *((ir::PTXB32*)source); - setRegAsB32(threadID, instr.d.reg, d); - ir::PTXB32 b = operandAsB32(threadID, instr.b); - *((ir::PTXB32*)source) = d ^ b; - reportE(REPORT_ATOM, "Atomically updated " << d << " to " - << *((ir::PTXB32*)source) ); } break; case ir::PTXInstruction::AtomicCas: @@ -2061,7 +2410,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: || instr.type == ir::PTXOperand::s32 || instr.type == ir::PTXOperand::u32) { ir::PTXB32 d = *((ir::PTXB32*)source); - setRegAsB32(threadID, instr.d.reg, d); + if (writeback) setRegAsB32(threadID, instr.d.reg, d); ir::PTXB32 b = operandAsB32(threadID, instr.b); ir::PTXB32 c = operandAsB32(threadID, instr.c); *((ir::PTXB32*)source) = (d==b) ? c : d; @@ -2072,7 +2421,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: || instr.type == ir::PTXOperand::s64 || instr.type == ir::PTXOperand::u64) { ir::PTXB64 d = *((ir::PTXB64*)source); - setRegAsB64(threadID, instr.d.reg, d); + if (writeback) setRegAsB64(threadID, instr.d.reg, d); ir::PTXB64 b = operandAsB64(threadID, instr.b); ir::PTXB64 c = operandAsB64(threadID, instr.c); *((ir::PTXB64*)source) = (d==b) ? c : d; @@ -2091,7 +2440,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: || instr.type == ir::PTXOperand::s32 || instr.type == ir::PTXOperand::u32) { ir::PTXB32 d = *((ir::PTXB32*)source); - setRegAsB32(threadID, instr.d.reg, d); + if (writeback) setRegAsB32(threadID, instr.d.reg, d); ir::PTXB32 b = operandAsB32(threadID, instr.b); *((ir::PTXB32*)source) = b; reportE(REPORT_ATOM, "Atomically updated " << d @@ -2101,7 +2450,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: || instr.type == ir::PTXOperand::s64 || instr.type == ir::PTXOperand::u64) { ir::PTXB64 d = *((ir::PTXB64*)source); - setRegAsB64(threadID, instr.d.reg, d); + if (writeback) setRegAsB64(threadID, instr.d.reg, d); ir::PTXB64 b = operandAsB64(threadID, instr.b); *((ir::PTXB64*)source) = b; reportE(REPORT_ATOM, "Atomically updated " << d @@ -2117,7 +2466,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: { if (instr.type == ir::PTXOperand::u32) { ir::PTXU32 d = *((ir::PTXU32*)source); - setRegAsU32(threadID, instr.d.reg, d); + if (writeback) setRegAsU32(threadID, instr.d.reg, d); ir::PTXU32 b = operandAsU32(threadID, instr.b); *((ir::PTXU32*)source) = b + d; reportE(REPORT_ATOM, "Atomically updated " << d @@ -2125,7 +2474,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: } else if (instr.type == ir::PTXOperand::s32) { ir::PTXS32 d = *((ir::PTXS32*)source); - setRegAsS32(threadID, instr.d.reg, d); + if (writeback) setRegAsS32(threadID, instr.d.reg, d); ir::PTXS32 b = operandAsS32(threadID, instr.b); *((ir::PTXS32*)source) = b + d; reportE(REPORT_ATOM, "Atomically updated " << d @@ -2133,7 +2482,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: } else if (instr.type == ir::PTXOperand::f32) { ir::PTXF32 d = *((ir::PTXF32*)source); - setRegAsF32(threadID, instr.d.reg, d); + if (writeback) setRegAsF32(threadID, instr.d.reg, d); ir::PTXF32 b = operandAsF32(threadID, instr.b); *((ir::PTXF32*)source) = b + d; reportE(REPORT_ATOM, "Atomically updated " << d @@ -2141,12 +2490,20 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: } else if (instr.type == ir::PTXOperand::u64) { ir::PTXU64 d = *((ir::PTXU64*)source); - setRegAsU64(threadID, instr.d.reg, d); + if (writeback) setRegAsU64(threadID, instr.d.reg, d); ir::PTXU64 b = operandAsU64(threadID, instr.b); *((ir::PTXU64*)source) = b + d; reportE(REPORT_ATOM, "Atomically updated " << d << " to " << *((ir::PTXU64*)source) ); } + else if (instr.type == ir::PTXOperand::f64) { + ir::PTXF64 d = *((ir::PTXF64*)source); + if (writeback) setRegAsF64(threadID, instr.d.reg, d); + ir::PTXF64 b = operandAsF64(threadID, instr.b); + *((ir::PTXF64*)source) = b + d; + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXF64*)source) ); + } else { throw RuntimeException("invalid data type", context.PC, instr); @@ -2160,7 +2517,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: context.PC, instr); } ir::PTXU32 d = *((ir::PTXU32*)source); - setRegAsU32(threadID, instr.d.reg, d); + if (writeback) setRegAsU32(threadID, instr.d.reg, d); ir::PTXU32 b = operandAsU32(threadID, instr.b); *((ir::PTXU32*)source) = (d >= b) ? 0 : d + 1; reportE(REPORT_ATOM, "Atomically updated " << d @@ -2174,7 +2531,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: context.PC, instr); } ir::PTXU32 d = *((ir::PTXU32*)source); - setRegAsU32(threadID, instr.d.reg, d); + if (writeback) setRegAsU32(threadID, instr.d.reg, d); ir::PTXU32 b = operandAsU32(threadID, instr.b); *((ir::PTXU32*)source) = ((d == 0) || (d > b)) ? b : d - 1; reportE(REPORT_ATOM, "Atomically updated " << d @@ -2185,7 +2542,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: { if (instr.type == ir::PTXOperand::u32) { ir::PTXU32 d = *((ir::PTXU32*)source); - setRegAsU32(threadID, instr.d.reg, d); + if (writeback) setRegAsU32(threadID, instr.d.reg, d); ir::PTXU32 b = operandAsU32(threadID, instr.b); *((ir::PTXU32*)source) = min(b, d); reportE(REPORT_ATOM, "Atomically updated " << d @@ -2193,19 +2550,27 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: } else if (instr.type == ir::PTXOperand::s32) { ir::PTXS32 d = *((ir::PTXS32*)source); - setRegAsS32(threadID, instr.d.reg, d); + if (writeback) setRegAsS32(threadID, instr.d.reg, d); ir::PTXS32 b = operandAsS32(threadID, instr.b); *((ir::PTXS32*)source) = min(b, d); reportE(REPORT_ATOM, "Atomically updated " << d << " to " << *((ir::PTXS32*)source) ); } - else if (instr.type == ir::PTXOperand::f32) { - ir::PTXF32 d = *((ir::PTXF32*)source); - setRegAsF32(threadID, instr.d.reg, d); - ir::PTXF32 b = operandAsF32(threadID, instr.b); - *((ir::PTXF32*)source) = min(b, d); + else if (instr.type == ir::PTXOperand::u64) { + ir::PTXU64 d = *((ir::PTXU64*)source); + if (writeback) setRegAsU64(threadID, instr.d.reg, d); + ir::PTXU64 b = operandAsU64(threadID, instr.b); + *((ir::PTXU64*)source) = min(b, d); reportE(REPORT_ATOM, "Atomically updated " << d - << " to " << *((ir::PTXF32*)source) ); + << " to " << *((ir::PTXU64*)source) ); + } + else if (instr.type == ir::PTXOperand::s64) { + ir::PTXS64 d = *((ir::PTXS64*)source); + if (writeback) setRegAsS64(threadID, instr.d.reg, d); + ir::PTXS64 b = operandAsS64(threadID, instr.b); + *((ir::PTXS64*)source) = min(b, d); + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXS64*)source) ); } else { throw RuntimeException("invalid data type", @@ -2217,7 +2582,7 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: { if (instr.type == ir::PTXOperand::u32) { ir::PTXU32 d = *((ir::PTXU32*)source); - setRegAsU32(threadID, instr.d.reg, d); + if (writeback) setRegAsU32(threadID, instr.d.reg, d); ir::PTXU32 b = operandAsU32(threadID, instr.b); *((ir::PTXU32*)source) = max(b, d); reportE(REPORT_ATOM, "Atomically updated " << d @@ -2225,19 +2590,27 @@ void executive::CooperativeThreadArray::eval_Atom(CTAContext &context, const ir: } else if (instr.type == ir::PTXOperand::s32) { ir::PTXS32 d = *((ir::PTXS32*)source); - setRegAsS32(threadID, instr.d.reg, d); + if (writeback) setRegAsS32(threadID, instr.d.reg, d); ir::PTXS32 b = operandAsS32(threadID, instr.b); *((ir::PTXS32*)source) = max(b, d); reportE(REPORT_ATOM, "Atomically updated " << d << " to " << *((ir::PTXS32*)source) ); } - else if (instr.type == ir::PTXOperand::f32) { - ir::PTXF32 d = *((ir::PTXF32*)source); - setRegAsF32(threadID, instr.d.reg, d); - ir::PTXF32 b = operandAsF32(threadID, instr.b); - *((ir::PTXF32*)source) = max(b, d); + else if (instr.type == ir::PTXOperand::u64) { + ir::PTXU64 d = *((ir::PTXU64*)source); + if (writeback) setRegAsU64(threadID, instr.d.reg, d); + ir::PTXU64 b = operandAsU64(threadID, instr.b); + *((ir::PTXU64*)source) = max(b, d); reportE(REPORT_ATOM, "Atomically updated " << d - << " to " << *((ir::PTXF32*)source) ); + << " to " << *((ir::PTXU64*)source) ); + } + else if (instr.type == ir::PTXOperand::s64) { + ir::PTXS64 d = *((ir::PTXS64*)source); + if (writeback) setRegAsS64(threadID, instr.d.reg, d); + ir::PTXS64 b = operandAsS64(threadID, instr.b); + *((ir::PTXS64*)source) = max(b, d); + reportE(REPORT_ATOM, "Atomically updated " << d + << " to " << *((ir::PTXS64*)source) ); } else { throw RuntimeException("invalid data type", @@ -2263,7 +2636,7 @@ void executive::CooperativeThreadArray::eval_Bfi(CTAContext &context, const ir::PTXB32 a = operandAsB32(threadID, instr.a); ir::PTXU32 b = operandAsU32(threadID, instr.b); ir::PTXU32 c = operandAsU32(threadID, instr.c); - ir::PTXB32 d = hydrazine::bitFieldInsert(pq, a, b, c); + ir::PTXB32 d = hydrazine::bitFieldInsert(pq, a, b & 0xff, c & 0xff); setRegAsB32(threadID, instr.d.reg, d); } break; @@ -2275,7 +2648,7 @@ void executive::CooperativeThreadArray::eval_Bfi(CTAContext &context, const ir::PTXB64 a = operandAsB64(threadID, instr.a); ir::PTXU32 b = operandAsU32(threadID, instr.b); ir::PTXU32 c = operandAsU32(threadID, instr.c); - ir::PTXB64 d = hydrazine::bitFieldInsert(pq, a, b, c); + ir::PTXB64 d = hydrazine::bitFieldInsert(pq, a, b & 0xff, c & 0xff); setRegAsB64(threadID, instr.d.reg, d); } break; @@ -2307,7 +2680,10 @@ void executive::CooperativeThreadArray::eval_Bfind(CTAContext &context, for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - ir::PTXS32 a = operandAsS32(threadID, instr.a); + const ir::PTXS32 signedA = operandAsS32(threadID, instr.a); + const ir::PTXU32 a = signedA < 0 + ? ~static_cast(signedA) + : static_cast(signedA); ir::PTXU32 d = hydrazine::bfind(a, instr.shiftAmount); setRegAsU32(threadID, instr.d.reg, d); } @@ -2317,7 +2693,10 @@ void executive::CooperativeThreadArray::eval_Bfind(CTAContext &context, for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - ir::PTXS64 a = operandAsS64(threadID, instr.a); + const ir::PTXS64 signedA = operandAsS64(threadID, instr.a); + const ir::PTXU64 a = signedA < 0 + ? ~static_cast(signedA) + : static_cast(signedA); ir::PTXU32 d = hydrazine::bfind(a, instr.shiftAmount); setRegAsU32(threadID, instr.d.reg, d); } @@ -2339,6 +2718,31 @@ void executive::CooperativeThreadArray::eval_Bfind(CTAContext &context, } } +void executive::CooperativeThreadArray::eval_Bmsk(CTAContext &context, + const ir::PTXInstruction &instr) { + trace(); + for (int threadID = 0; threadID < threadCount; ++threadID) { + if (!context.predicated(threadID, instr)) continue; + + const ir::PTXU32 a = operandAsU32(threadID, instr.a); + const ir::PTXU32 b = operandAsU32(threadID, instr.b); + const ir::PTXU32 a1 = a & 0x1f; + const ir::PTXU32 b1 = b & 0x1f; + const ir::PTXU32 sum = a1 + b1; + const bool positionOverflow = instr.shiftMode == + ir::PTXInstruction::ShiftMode::Clamp && a >= 32; + const bool widthOverflow = instr.shiftMode == + ir::PTXInstruction::ShiftMode::Clamp && b >= 32; + const ir::PTXU32 mask0 = positionOverflow ? 0 : (~0u << a1); + ir::PTXU32 mask1 = 0; + if (sum < 32 && !positionOverflow && !widthOverflow) { + mask1 = b1 == 0 ? ~0u : (~0u << sum); + } + + setRegAsU32(threadID, instr.d.reg, mask0 & ~mask1); + } +} + void executive::CooperativeThreadArray::eval_Brev(CTAContext& context, const ir::PTXInstruction& instr) { trace(); @@ -2378,19 +2782,11 @@ void executive::CooperativeThreadArray::eval_Bfe(CTAContext &context, || instr.type == ir::PTXOperand::s32); bool isSigned = (instr.type == ir::PTXOperand::s32 || instr.type == ir::PTXOperand::s64); - - ir::PTXU32 pos = operandAsU32(tid, instr.b); - ir::PTXU32 len = operandAsU32(tid, instr.c); - ir::PTXU64 a = operandAsU64(tid, instr.a); - ir::PTXU64 mask = ((1 << len) - 1); - ir::PTXU32 msb = min((pos+len-1), size32bit ? 31 : 63); - ir::PTXU32 sign = ((a>>(msb))&1); - ir::PTXU64 result = 0; - - if (isSigned) { - result = (sign ? -1 : 0) & (~mask); - } - result |= ((a >> pos) & mask); + ir::PTXU64 result = size32bit + ? hydrazine::bfe(operandAsU32(tid, instr.a), + operandAsU32(tid, instr.b), operandAsU32(tid, instr.c), isSigned) + : hydrazine::bfe(operandAsU64(tid, instr.a), + operandAsU32(tid, instr.b), operandAsU32(tid, instr.c), isSigned); if (size32bit) { setRegAsU32(tid, instr.d.reg, hydrazine::bit_cast(result)); @@ -2845,6 +3241,42 @@ void executive::CooperativeThreadArray::eval_CNot(CTAContext &context, } } +void executive::CooperativeThreadArray::eval_Fns(CTAContext &context, + const ir::PTXInstruction &instr) { + trace(); + for (int threadID = 0; threadID < threadCount; ++threadID) { + if (!context.predicated(threadID, instr)) continue; + + const ir::PTXU32 mask = operandAsB32(threadID, instr.a); + const ir::PTXU32 base = operandAsB32(threadID, instr.b); + const ir::PTXS32 offset = operandAsS32(threadID, instr.c); + ir::PTXU32 result = 0xffffffffu; + if (base > 31) { + report("warning: fns base outside 0..31 is undefined by PTX; " + "Ocelot returns 0xffffffff"); + } + else if (offset == 0) { + if ((mask >> base) & 1u) result = base; + } + else { + int pos = static_cast(base); + ir::PTXS64 count = std::abs(static_cast(offset)) - 1; + const int inc = offset > 0 ? 1 : -1; + while (pos >= 0 && pos < 32) { + if ((mask >> pos) & 1u) { + if (count == 0) { + result = pos; + break; + } + --count; + } + pos += inc; + } + } + setRegAsB32(threadID, instr.d.reg, result); + } +} + void executive::CooperativeThreadArray::eval_CopySign(CTAContext &context, const ir::PTXInstruction &instr) { trace(); @@ -2891,8 +3323,9 @@ void executive::CooperativeThreadArray::eval_Cos(CTAContext &context, for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - ir::PTXF32 d, a = operandAsF32(threadID, instr.a); - d = (ir::PTXF32)cos(a); + ir::PTXF32 d, + a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); + d = ftz(instr.modifier, (ir::PTXF32)cos(a)); setRegAsF32(threadID, instr.d.reg, d); } } @@ -2910,6 +3343,40 @@ void executive::CooperativeThreadArray::eval_Cos(CTAContext &context, } } +template< typename Destination, typename Source > +static Destination cvtInteger(Source value, int modifier) { + if( !(modifier & ir::PTXInstruction::sat) ) { + return static_cast(value); + } + + if( std::numeric_limits::is_signed ) { + const ir::PTXS64 signedValue = static_cast(value); + if( signedValue < 0 ) { + if( !std::numeric_limits::is_signed ) return 0; + + const ir::PTXS64 minimum = static_cast( + (std::numeric_limits::min)()); + return signedValue < minimum + ? (std::numeric_limits::min)() + : static_cast(value); + } + } + + const ir::PTXU64 unsignedValue = static_cast(value); + const ir::PTXU64 maximum = static_cast( + (std::numeric_limits::max)()); + return unsignedValue > maximum + ? (std::numeric_limits::max)() + : static_cast(value); +} + +template< typename Float > +static Float cvtSaturate(Float value, int modifier) { + if( !(modifier & ir::PTXInstruction::sat) ) return value; + if( hydrazine::isnan(value) || value <= 0 ) return 0; + return value >= 1 ? 1 : value; +} + template< typename Int > static ir::PTXF32 toF32(Int value, int modifier) { int mode = hydrazine::fegetround(); @@ -2922,7 +3389,7 @@ static ir::PTXF32 toF32(Int value, int modifier) { } else if (modifier & ir::PTXInstruction::rp) { hydrazine::fesetround(FE_UPWARD); } - ir::PTXF32 d = value; + ir::PTXF32 d = cvtSaturate(static_cast(value), modifier); hydrazine::fesetround(mode); return d; } @@ -2939,17 +3406,67 @@ static ir::PTXF64 toF64(Int value, int modifier) { } else if (modifier & ir::PTXInstruction::rp) { hydrazine::fesetround(FE_UPWARD); } - ir::PTXF64 d = value; + ir::PTXF64 d = cvtSaturate(static_cast(value), modifier); hydrazine::fesetround(mode); return d; } +template< typename Source > +static ir::PTXU16 toF16(Source value, int modifier) { + if( (modifier & ir::PTXInstruction::relu) && hydrazine::isnan(value) ) { + return 0x7fff; + } + if( (modifier & ir::PTXInstruction::relu) && value < 0 ) value = 0; + int mode = hydrazine::fegetround(); + if (modifier & ir::PTXInstruction::rn) { + hydrazine::fesetround(FE_TONEAREST); + } else if (modifier & ir::PTXInstruction::rz) { + hydrazine::fesetround(FE_TOWARDZERO); + } else if (modifier & ir::PTXInstruction::rm) { + hydrazine::fesetround(FE_DOWNWARD); + } else if (modifier & ir::PTXInstruction::rp) { + hydrazine::fesetround(FE_UPWARD); + } + ir::PTXF16 half = value; + if( modifier & ir::PTXInstruction::sat ) { + half = cvtSaturate(static_cast(half), modifier); + } + hydrazine::fesetround(mode); + ir::PTXU16 bits; + std::memcpy(&bits, &half, sizeof(bits)); + return bits; +} + +static ir::PTXU16 f32ToBF16(ir::PTXF32 value, int modifier) { + if( (modifier & ir::PTXInstruction::relu) && value < 0 ) value = 0; + ir::PTXU32 bits = hydrazine::bit_cast(value); + if ((bits & 0x7fffffffU) > 0x7f800000U) { + // NVIDIA canonical NaN + return 0x7fffU; + } + ir::PTXU16 upper = static_cast(bits >> 16); + if( modifier & ir::PTXInstruction::rz ) return upper; + ir::PTXU16 lower = static_cast(bits & 0xffffU); + if (lower > 0x8000U) { + return upper + 1U; + } else if (lower < 0x8000U) { + return upper; + } + return upper + (upper & 1U); +} + +static ir::PTXU16 f32ToBF16Rn(ir::PTXF32 value) { + return f32ToBF16(value, ir::PTXInstruction::rn); +} + template< typename Float > -static Float roundToInt(Float a, int modifier, executive::CTAContext &context, - const ir::PTXInstruction &instr) { +static Float roundToInt(Float a, int modifier) { Float fd = 0; if (modifier & ir::PTXInstruction::rni) { + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_TONEAREST); fd = hydrazine::nearbyintf(a); + hydrazine::fesetround(previous); } else if (modifier & ir::PTXInstruction::rzi) { fd = hydrazine::trunc(a); } else if (modifier & ir::PTXInstruction::rmi) { @@ -2963,12 +3480,42 @@ static Float roundToInt(Float a, int modifier, executive::CTAContext &context, return fd; } +template< typename Float > +static Float roundToInt(Float a, int modifier, executive::CTAContext &, + const ir::PTXInstruction &) { + return roundToInt(a, modifier); +} + +template< typename Destination, typename Float > +static Destination cvtFloatToInteger(Float value, int modifier) { + if (value != value) return 0; + const Float rounded = roundToInt(value, modifier); + const int bits = sizeof(Destination) * CHAR_BIT; + const int magnitudeBits = std::numeric_limits::is_signed + ? bits - 1 : bits; + const Float upper = std::ldexp(Float(1), magnitudeBits); + if (rounded >= upper) return (std::numeric_limits::max)(); + if (std::numeric_limits::is_signed && rounded + < -upper) return (std::numeric_limits::min)(); + if (!std::numeric_limits::is_signed && rounded < 0) return 0; + return static_cast(rounded); +} + /*! */ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, const ir::PTXInstruction &instr) { trace(); + auto setCvtB16 = [this](int threadID, ir::PTXOperand::RegisterType reg, + ir::PTXU16 value) { + setRegAsU64(threadID, reg, value); + }; + auto setCvtF32 = [this](int threadID, ir::PTXOperand::RegisterType reg, + ir::PTXF32 value) { + setRegAsU64(threadID, reg, + hydrazine::bit_cast(value)); + }; for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; @@ -2977,12 +3524,40 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, if (instr.a.relaxedType != ir::PTXOperand::TypeSpecifier_invalid) { sourceType = instr.a.relaxedType; } + if( instr.type == ir::PTXOperand::f16x2 ) { + const ir::PTXU32 d = + static_cast(toF16( + operandAsF32(threadID, instr.a), instr.modifier)) << 16 + | toF16(operandAsF32(threadID, instr.b), instr.modifier); + setRegAsU64(threadID, instr.d.reg, d); + continue; + } + if( instr.type == ir::PTXOperand::bf16x2 ) { + const ir::PTXU32 d = + static_cast(f32ToBF16( + operandAsF32(threadID, instr.a), instr.modifier)) << 16 + | f32ToBF16(operandAsF32(threadID, instr.b), instr.modifier); + setRegAsU64(threadID, instr.d.reg, d); + continue; + } + if( instr.type == ir::PTXOperand::tf32 ) { + setRegAsU64(threadID, instr.d.reg, + tf32FromF32Rna(operandAsF32(threadID, instr.a))); + continue; + } switch (sourceType) { case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: { switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsB8(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: // fall through @@ -3001,17 +3576,14 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::s8: { - ir::PTXU8 a = operandAsU8(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, CHAR_MAX); - } - ir::PTXS8 d = a; + ir::PTXS8 d = cvtInteger( + operandAsU8(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsB8(threadID, instr.a), instr.modifier)); } @@ -3033,8 +3605,15 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::s8: { switch (instr.type) { - case ir::PTXOperand::s8: // fall through - case ir::PTXOperand::s16: // fall through + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsS8(threadID, instr.a), + instr.modifier)); + } + break; + case ir::PTXOperand::s8: // fall through + case ir::PTXOperand::s16: // fall through case ir::PTXOperand::s32: // fall through case ir::PTXOperand::s64: { @@ -3044,24 +3623,34 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::u8: // fall through - case ir::PTXOperand::b8: // fall through + case ir::PTXOperand::b8: + { + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS8(threadID, instr.a), instr.modifier)); + } + break; case ir::PTXOperand::b16: // fall through - case ir::PTXOperand::u16: // fall through + case ir::PTXOperand::u16: + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS8(threadID, instr.a), instr.modifier)); + break; case ir::PTXOperand::b32: // fall through - case ir::PTXOperand::u32: // fall through + case ir::PTXOperand::u32: + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS8(threadID, instr.a), instr.modifier)); + break; case ir::PTXOperand::b64: // fall through case ir::PTXOperand::u64: - { - ir::PTXS8 a = operandAsS8(threadID, instr.a); - if (instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - setRegAsU64(threadID, instr.d.reg, a); - } + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS8(threadID, instr.a), instr.modifier)); break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsS8(threadID, instr.a), instr.modifier)); } @@ -3084,12 +3673,19 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::u16: { switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsB16(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: { - ir::PTXU16 a = operandAsU16(threadID, instr.a); - ir::PTXU8 d = a; + ir::PTXU8 d = cvtInteger( + operandAsU16(threadID, instr.a), instr.modifier); setRegAsU64(threadID, instr.d.reg, d); } break; @@ -3108,27 +3704,21 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::s8: { - ir::PTXU16 a = operandAsU16(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, CHAR_MAX); - } - ir::PTXS8 d = a; + ir::PTXS8 d = cvtInteger( + operandAsU16(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s16: { - ir::PTXU16 a = operandAsU16(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, SHRT_MAX); - } - ir::PTXS16 d = a; + ir::PTXS16 d = cvtInteger( + operandAsU16(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsB16(threadID, instr.a), instr.modifier)); } @@ -3151,10 +3741,17 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, { // s16 to one of the following switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsS16(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::s8: { - ir::PTXS16 a = operandAsS16(threadID, instr.a); - ir::PTXS8 d = a; + ir::PTXS8 d = cvtInteger( + operandAsS16(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; @@ -3170,31 +3767,32 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::u8: // fall through case ir::PTXOperand::b8: { - ir::PTXS16 a = operandAsS16(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - ir::PTXU8 d = a; + ir::PTXU8 d = cvtInteger( + operandAsS16(threadID, instr.a), instr.modifier); setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::b16: // fall through - case ir::PTXOperand::u16: // fall through + case ir::PTXOperand::u16: + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS16(threadID, instr.a), instr.modifier)); + break; case ir::PTXOperand::b32: // fall through - case ir::PTXOperand::u32: // fall through + case ir::PTXOperand::u32: + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS16(threadID, instr.a), instr.modifier)); + break; case ir::PTXOperand::b64: // fall through case ir::PTXOperand::u64: - { - ir::PTXS16 a = operandAsS16(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - setRegAsU64(threadID, instr.d.reg, a); - } + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS16(threadID, instr.a), instr.modifier)); break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsS16(threadID, instr.a), instr.modifier)); } @@ -3217,20 +3815,27 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::u32: { switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsU32(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: { - ir::PTXU32 a = operandAsU32(threadID, instr.a); - ir::PTXU8 d = a; + ir::PTXU8 d = cvtInteger( + operandAsU32(threadID, instr.a), instr.modifier); setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::u16: // fall through case ir::PTXOperand::b16: { - ir::PTXU32 a = operandAsU32(threadID, instr.a); - ir::PTXU16 d = a; + ir::PTXU16 d = cvtInteger( + operandAsU32(threadID, instr.a), instr.modifier); setRegAsU64(threadID, instr.d.reg, d); } break; @@ -3246,37 +3851,28 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::s8: { - ir::PTXU32 a = operandAsU32(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, CHAR_MAX); - } - ir::PTXS8 d = a; + ir::PTXS8 d = cvtInteger( + operandAsU32(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s16: { - ir::PTXU32 a = operandAsU32(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, SHRT_MAX); - } - ir::PTXS16 d = a; + ir::PTXS16 d = cvtInteger( + operandAsU32(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s32: { - ir::PTXU32 a = operandAsU32(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, INT_MAX); - } - ir::PTXS32 d = a; + ir::PTXS32 d = cvtInteger( + operandAsU32(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsU32(threadID, instr.a), instr.modifier)); } @@ -3298,52 +3894,53 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::s32: { switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsS32(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: { - ir::PTXS32 a = operandAsS32(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - ir::PTXU8 d = a; - setRegAsS64(threadID, instr.d.reg, d); + ir::PTXU8 d = cvtInteger( + operandAsS32(threadID, instr.a), instr.modifier); + setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::u16: // fall through case ir::PTXOperand::b16: { - ir::PTXS32 a = operandAsS32(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - ir::PTXU16 d = a; - setRegAsS64(threadID, instr.d.reg, d); + ir::PTXU16 d = cvtInteger( + operandAsS32(threadID, instr.a), instr.modifier); + setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::b32: // fall through - case ir::PTXOperand::u32: // fall through + case ir::PTXOperand::u32: + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS32(threadID, instr.a), instr.modifier)); + break; case ir::PTXOperand::b64: // fall through case ir::PTXOperand::u64: - { - ir::PTXS32 a = operandAsS32(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - setRegAsS64(threadID, instr.d.reg, a); - } + setRegAsU64(threadID, instr.d.reg, + cvtInteger( + operandAsS32(threadID, instr.a), instr.modifier)); break; case ir::PTXOperand::s8: { - ir::PTXS32 a = operandAsS32(threadID, instr.a); - ir::PTXS8 d = a; + ir::PTXS8 d = cvtInteger( + operandAsS32(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s16: { - ir::PTXS32 a = operandAsS32(threadID, instr.a); - ir::PTXS16 d = a; + ir::PTXS16 d = cvtInteger( + operandAsS32(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; @@ -3356,7 +3953,7 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsS32(threadID, instr.a), instr.modifier)); } @@ -3378,38 +3975,36 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::s64: { switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsS64(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: // fall through { - ir::PTXS64 a = operandAsS64(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - ir::PTXU8 d = a; - setRegAsS64(threadID, instr.d.reg, d); + ir::PTXU8 d = cvtInteger( + operandAsS64(threadID, instr.a), instr.modifier); + setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::u16: // fall through case ir::PTXOperand::b16: // fall through { - ir::PTXS64 a = operandAsS64(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - ir::PTXU16 d = a; - setRegAsS64(threadID, instr.d.reg, d); + ir::PTXU16 d = cvtInteger( + operandAsS64(threadID, instr.a), instr.modifier); + setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::b32: // fall through case ir::PTXOperand::u32: { - ir::PTXS64 a = operandAsS64(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = max(a, 0); - } - ir::PTXU32 d = a; - setRegAsS64(threadID, instr.d.reg, d); + ir::PTXU32 d = cvtInteger( + operandAsS64(threadID, instr.a), instr.modifier); + setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::b64: // fall through @@ -3424,22 +4019,22 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::s8: { - ir::PTXS64 a = operandAsS64(threadID, instr.a); - ir::PTXS8 d = a; + ir::PTXS8 d = cvtInteger( + operandAsS64(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s16: { - ir::PTXS64 a = operandAsS64(threadID, instr.a); - ir::PTXS16 d = a; + ir::PTXS16 d = cvtInteger( + operandAsS64(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s32: { - ir::PTXS64 a = operandAsS64(threadID, instr.a); - ir::PTXS32 d = a; + ir::PTXS32 d = cvtInteger( + operandAsS64(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; @@ -3451,7 +4046,7 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsS64(threadID, instr.a), instr.modifier)); } @@ -3474,28 +4069,35 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::u64: { switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsU64(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: { - ir::PTXU64 a = operandAsU64(threadID, instr.a); - ir::PTXU8 d = a; + ir::PTXU8 d = cvtInteger( + operandAsU64(threadID, instr.a), instr.modifier); setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::b16: // fall through case ir::PTXOperand::u16: { - ir::PTXU64 a = operandAsU64(threadID, instr.a); - ir::PTXU16 d = a; + ir::PTXU16 d = cvtInteger( + operandAsU64(threadID, instr.a), instr.modifier); setRegAsU64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::b32: // fall through case ir::PTXOperand::u32: { - ir::PTXU64 a = operandAsU64(threadID, instr.a); - ir::PTXU32 d = a; + ir::PTXU32 d = cvtInteger( + operandAsU64(threadID, instr.a), instr.modifier); setRegAsU64(threadID, instr.d.reg, d); } break; @@ -3508,47 +4110,35 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, break; case ir::PTXOperand::s8: { - ir::PTXU64 a = operandAsU64(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, CHAR_MAX); - } - ir::PTXS8 d = a; + ir::PTXS8 d = cvtInteger( + operandAsU64(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s16: { - ir::PTXU64 a = operandAsU64(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, SHRT_MAX); - } - ir::PTXS16 d = a; + ir::PTXS16 d = cvtInteger( + operandAsU64(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s32: { - ir::PTXU64 a = operandAsU64(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, INT_MAX); - } - ir::PTXS32 d = a; + ir::PTXS32 d = cvtInteger( + operandAsU64(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::s64: { - ir::PTXU64 a = operandAsU64(threadID, instr.a); - if(instr.modifier & ir::PTXInstruction::sat) { - a = min(a, LLONG_MAX); - } - ir::PTXS64 d = a; + ir::PTXS64 d = cvtInteger( + operandAsU64(threadID, instr.a), instr.modifier); setRegAsS64(threadID, instr.d.reg, d); } break; case ir::PTXOperand::f32: { - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, toF32(operandAsU64(threadID, instr.a), instr.modifier)); } @@ -3567,184 +4157,118 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, } } break; + case ir::PTXOperand::f16: // fall through case ir::PTXOperand::f32: { + ir::PTXF32 a = sourceType == ir::PTXOperand::f16 + ? f16ToF32(operandAsU16(threadID, instr.a)) + : ftz(instr.modifier, operandAsF32(threadID, instr.a)); switch (instr.type) { + case ir::PTXOperand::f16: + { + a = roundToInt(a, instr.modifier); + setCvtB16(threadID, instr.d.reg, + toF16(a, instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU8 d = 0; - if(fd > UCHAR_MAX) { - d = UCHAR_MAX; - } - else if(fd < 0) { - d = 0; - } - else { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::b16: // fall through case ir::PTXOperand::u16: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU16 d = 0; - if(fd > USHRT_MAX) { - d = USHRT_MAX; - } - else if(fd < 0) { - d = 0; - } - else { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::b32: // fall through case ir::PTXOperand::u32: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU32 d = 0; - if(fd > static_cast(UINT_MAX)) { - d = UINT_MAX; - } - else if(fd < 0) { - d = 0; - } - else { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::b64: // fall through case ir::PTXOperand::u64: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU64 d = 0; - if(fd > static_cast(ULLONG_MAX)) { - d = ULLONG_MAX; - } - else if(fd < 0) { - d = 0; - } - else { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s8: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS8 d = 0; - if(fd > CHAR_MAX) { - d = CHAR_MAX; - } - else if(fd < CHAR_MIN) { - d = CHAR_MIN; - } - else { - d = fd; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s16: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS16 d = 0; - if(fd > SHRT_MAX) { - d = SHRT_MAX; - } - else if(fd < SHRT_MIN) { - d = SHRT_MIN; - } - else { - d = fd; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s32: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS32 d = 0; - if(fd > static_cast(INT_MAX)) { - d = INT_MAX; - } - else if(fd < INT_MIN) { - d = INT_MIN; - } - else { - d = fd; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s64: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF32 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS64 d = 0; - if(fd > static_cast(LLONG_MAX)) { - d = LLONG_MAX; - } - else if(fd < LLONG_MIN) { - d = LLONG_MIN; - } - else { - d = fd; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::f32: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - a = roundToInt(a, instr.modifier, context, instr); - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, sat(instr.modifier, a)); } break; case ir::PTXOperand::f64: { - ir::PTXF32 a = operandAsF32(threadID, instr.a); - ir::PTXF64 d = toF64(a, instr.modifier); + ir::PTXF64 d = roundToInt(a, instr.modifier); + d = toF64(d, instr.modifier); setRegAsF64(threadID, instr.d.reg, d); } break; + case ir::PTXOperand::bf16: + { + ir::PTXU16 d = f32ToBF16(a, instr.modifier); + setCvtB16(threadID, instr.d.reg, d); + } + break; + default: + throw RuntimeException("conversion not implemented", + context.PC, instr); + break; + } + } + break; + case ir::PTXOperand::bf16: + { + switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(bf16ToF32(operandAsU16(threadID, + instr.a)), instr.modifier)); + } + break; + case ir::PTXOperand::f32: + { + setCvtF32(threadID, instr.d.reg, + bf16ToF32(operandAsU16(threadID, instr.a))); + } + break; default: throw RuntimeException("conversion not implemented", context.PC, instr); @@ -3755,183 +4279,93 @@ void executive::CooperativeThreadArray::eval_Cvt(CTAContext &context, case ir::PTXOperand::f64: { switch (instr.type) { + case ir::PTXOperand::f16: + { + setCvtB16(threadID, instr.d.reg, + toF16(operandAsF64(threadID, instr.a), + instr.modifier)); + } + break; case ir::PTXOperand::pred: // fall through case ir::PTXOperand::b8: // fall through case ir::PTXOperand::u8: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF64 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU8 d = 0; - if(fd > UCHAR_MAX) { - d = UCHAR_MAX; - } - else if(fd < 0) { - d = 0; - } - else { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::b16: // fall through case ir::PTXOperand::u16: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF64 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU16 d = 0; - if(fd > USHRT_MAX) { - d = USHRT_MAX; - } - else if(fd < 0) { - d = 0; - } - else { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::b32: // fall through case ir::PTXOperand::u32: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF64 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU32 d = 0; - if(fd > UINT_MAX) { - d = UINT_MAX; - } - else if(fd < 0) { - d = 0; - } - else - { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::b64: // fall through case ir::PTXOperand::u64: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0f; - ir::PTXF64 fd = roundToInt(a, instr.modifier, - context, instr); - ir::PTXU64 d = 0; - if(fd > static_cast(ULLONG_MAX)) { - d = ULLONG_MAX; - } - else if(fd < 0) { - d = 0; - } - else - { - d = fd; - } - setRegAsU64(threadID, instr.d.reg, d); + setRegAsU64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s8: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0; - a = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS8 d = 0; - if(a > CHAR_MAX) { - d = CHAR_MAX; - } - else if(a < CHAR_MIN) { - d = CHAR_MIN; - } - else { - d = a; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s16: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0; - a = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS16 d = 0; - if(a > SHRT_MAX) { - d = SHRT_MAX; - } - else if(a < SHRT_MIN) { - d = SHRT_MIN; - } - else { - d = a; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s32: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0; - a = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS32 d = 0; - if(a > INT_MAX) { - d = INT_MAX; - } - else if(a < INT_MIN) { - d = INT_MIN; - } - else { - d = a; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::s64: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - if (a != a) a = 0.0; - a = roundToInt(a, instr.modifier, - context, instr); - ir::PTXS64 d = 0; - if(a > static_cast(LLONG_MAX)) { - d = LLONG_MAX; - } - else if(a < LLONG_MIN) { - d = LLONG_MIN; - } - else { - d = a; - } - setRegAsS64(threadID, instr.d.reg, d); + setRegAsS64(threadID, instr.d.reg, + cvtFloatToInteger(a, instr.modifier)); } break; case ir::PTXOperand::f32: { ir::PTXF64 a = operandAsF64(threadID, instr.a); - a = toF32(a, instr.modifier); + a = ftz(instr.modifier, toF32(a, instr.modifier)); if(instr.modifier & ir::PTXInstruction::sat) { if (a != a) a = 0.0; a = min(1.0, a); a = max(a, 0.0); } - setRegAsF32(threadID, instr.d.reg, + setCvtF32(threadID, instr.d.reg, sat(instr.modifier, a)); } break; case ir::PTXOperand::f64: { ir::PTXF64 a = operandAsF64(threadID, instr.a); + a = roundToInt(a, instr.modifier); setRegAsF64(threadID, instr.d.reg, - sat(instr.modifier, a)); + sat(instr.modifier, a)); } break; default: @@ -3976,6 +4410,11 @@ void executive::CooperativeThreadArray::eval_Cvta(CTAContext &context, case ir::PTXInstruction::Global: // DO NOTHING case ir::PTXInstruction::Local: // DO NOTHING break; + case ir::PTXInstruction::Param: + addrSpaceBase = (ir::PTXU32)(kernel->scheduler + ? kernel->scheduler->argumentMemory() + : (ir::PTXU64)kernel->ArgumentMemory); + break; case ir::PTXInstruction::Shared: { hydrazine::bit_cast(addrSpaceBase, @@ -3999,7 +4438,8 @@ void executive::CooperativeThreadArray::eval_Cvta(CTAContext &context, if (instr.addressSpace == ir::PTXInstruction::Local) { ir::PTXU32 localMemPtr; - if (!instr.a.isGlobalLocal) { + if (instr.a.addressMode != ir::PTXOperand::Address + || !instr.a.isGlobalLocal) { hydrazine::bit_cast(localMemPtr, functionCallStack.localMemoryPointer(tid)); } @@ -4028,6 +4468,11 @@ void executive::CooperativeThreadArray::eval_Cvta(CTAContext &context, case ir::PTXInstruction::Global: // DO NOTHING case ir::PTXInstruction::Local: // DO NOTHING break; + case ir::PTXInstruction::Param: + addrSpaceBase = kernel->scheduler + ? kernel->scheduler->argumentMemory() + : (ir::PTXU64)kernel->ArgumentMemory; + break; case ir::PTXInstruction::Shared: { hydrazine::bit_cast(addrSpaceBase, @@ -4051,7 +4496,8 @@ void executive::CooperativeThreadArray::eval_Cvta(CTAContext &context, if (instr.addressSpace == ir::PTXInstruction::Local) { ir::PTXU64 localMemPtr; - if (!instr.a.isGlobalLocal) { + if (instr.a.addressMode != ir::PTXOperand::Address + || !instr.a.isGlobalLocal) { hydrazine::bit_cast(localMemPtr, functionCallStack.localMemoryPointer(tid)); } @@ -4102,6 +4548,12 @@ void executive::CooperativeThreadArray::eval_Cvta(CTAContext &context, } break; + case ir::PTXInstruction::Param: + addrSpaceBase = (ir::PTXU32)(kernel->scheduler + ? kernel->scheduler->argumentMemory() + : (ir::PTXU64)kernel->ArgumentMemory); + addrSpaceSize = kernel->argumentMemorySize(); + break; case ir::PTXInstruction::Const: { hydrazine::bit_cast(addrSpaceBase, kernel->ConstMemory); @@ -4179,6 +4631,12 @@ void executive::CooperativeThreadArray::eval_Cvta(CTAContext &context, } break; + case ir::PTXInstruction::Param: + addrSpaceBase = kernel->scheduler + ? kernel->scheduler->argumentMemory() + : (ir::PTXU64)kernel->ArgumentMemory; + addrSpaceSize = kernel->argumentMemorySize(); + break; case ir::PTXInstruction::Const: { hydrazine::bit_cast(addrSpaceBase, kernel->ConstMemory); @@ -4253,18 +4711,21 @@ void executive::CooperativeThreadArray::eval_Div(CTAContext &context, ir::PTXF32 d, a = ftz(instr.modifier, operandAsF32(threadID, instr.a)), b = ftz(instr.modifier, operandAsF32(threadID, instr.b)); if(ir::PTXInstruction::approx & instr.modifier) { - if(issubnormal_(a) || issubnormal_(b)) - { - d = a / b; - } - else - { - d = a * ( 1.0f / b ); + if(std::fabs(b) > std::ldexp(1.0f, 126) && !hydrazine::isinf(b)) { + if(hydrazine::isinf(a)) { + d = std::numeric_limits::quiet_NaN(); + } else { + const bool negative = std::signbit(a) != std::signbit(b); + d = hydrazine::copysign(0.0f, negative ? -1.0f : 1.0f); + } + } else { + d = roundedMul(a, roundedDiv(1.0f, b, + ir::PTXInstruction::rn), ir::PTXInstruction::rn); } + } else { + d = roundedDiv(a, b, instr.modifier); } - else { - d = ftz(instr.modifier, a / b); - } + d = ftz(instr.modifier, d); setRegAsF32(threadID, instr.d.reg, d); } } @@ -4274,7 +4735,7 @@ void executive::CooperativeThreadArray::eval_Div(CTAContext &context, ir::PTXF64 d, a = operandAsF64(threadID, instr.a), b = operandAsF64(threadID, instr.b); - d = a / b; + d = roundedDiv(a, b, instr.modifier); setRegAsF64(threadID, instr.d.reg, d); } } @@ -4367,6 +4828,36 @@ void executive::CooperativeThreadArray::eval_Div(CTAContext &context, } } +void executive::CooperativeThreadArray::eval_Dp(CTAContext &context, + const ir::PTXInstruction &instr) { + trace(); + const bool fourWay = instr.opcode == ir::PTXInstruction::Dp4a; + const unsigned int lanes = fourWay ? 4 : 2; + const unsigned int aWidth = fourWay ? 8 : 16; + const unsigned int bStart = !fourWay + && instr.modifier == ir::PTXInstruction::hi ? 2 : 0; + auto extract = [](ir::PTXU32 value, unsigned int bit, + unsigned int width, ir::PTXOperand::DataType type) { + const ir::PTXU32 field = (value >> bit) & ((1u << width) - 1); + if( type == ir::PTXOperand::s32 && (field & (1u << (width - 1))) ) { + return static_cast(field) - (1 << width); + } + return static_cast(field); + }; + for (int threadID = 0; threadID < threadCount; ++threadID) { + if (!context.predicated(threadID, instr)) continue; + const ir::PTXU32 a = operandAsU32(threadID, instr.a); + const ir::PTXU32 b = operandAsU32(threadID, instr.b); + ir::PTXU32 d = operandAsU32(threadID, instr.c); + for (unsigned int i = 0; i < lanes; ++i) { + const ir::PTXS32 va = extract(a, i * aWidth, aWidth, instr.type); + const ir::PTXS32 vb = extract(b, (bStart + i) * 8, 8, instr.bType); + d += static_cast(va * vb); + } + setRegAsU32(threadID, instr.d.reg, d); + } +} + /*! */ @@ -4377,11 +4868,22 @@ void executive::CooperativeThreadArray::eval_Ex2(CTAContext &context, for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - ir::PTXF32 d, a = operandAsF32(threadID, instr.a); + ir::PTXF32 d, + a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); d = ftz(instr.modifier, hydrazine::exp2f(a)); setRegAsF32(threadID, instr.d.reg, d); } } + else if (instr.type == ir::PTXOperand::f16) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + + ir::PTXF32 a = f16ToF32(ftzF16(instr.modifier, + operandAsU16(threadID, instr.a))); + setRegAsB16(threadID, instr.d.reg, + toF16(hydrazine::exp2f(a), instr.modifier)); + } + } else { throw RuntimeException("unsupported data type", context.PC, instr); } @@ -4403,7 +4905,26 @@ void executive::CooperativeThreadArray::eval_Exit(CTAContext &context, void executive::CooperativeThreadArray::eval_Fma(CTAContext &context, const ir::PTXInstruction &instr) { trace(); - if (instr.type == ir::PTXOperand::f32) { + if (instr.type == ir::PTXOperand::f16x2) { + for (int tid = 0; tid < threadCount; tid++) { + if (!context.predicated(tid, instr)) continue; + ir::PTXU32 a = operandAsU32(tid, instr.a), b = operandAsU32(tid, instr.b); + ir::PTXU32 c = operandAsU32(tid, instr.c); + ir::PTXU16 al = static_cast(a), ah = a >> 16; + ir::PTXU16 bl = static_cast(b), bh = b >> 16; + ir::PTXU16 cl = static_cast(c), ch = c >> 16; + auto fma = [&](ir::PTXU16 x, ir::PTXU16 y, ir::PTXU16 z) { + return ftzF16(instr.modifier, toF16(roundedFma( + f16ToF32(ftzF16(instr.modifier, x)), + f16ToF32(ftzF16(instr.modifier, y)), + f16ToF32(ftzF16(instr.modifier, z)), instr.modifier), instr.modifier)); + }; + ir::PTXU16 dl = fma(al, bl, cl), dh = fma(ah, bh, ch); + setRegAsU32(tid, instr.d.reg, static_cast(dl) | + (static_cast(dh) << 16)); + } + } + else if (instr.type == ir::PTXOperand::f32) { for (int tid = 0; tid < threadCount; tid++) { if (!context.predicated(tid, instr)) continue; ir::PTXF32 d = 0, @@ -4411,7 +4932,8 @@ void executive::CooperativeThreadArray::eval_Fma(CTAContext &context, b = ftz(instr.modifier, operandAsF32(tid, instr.b)), c = ftz(instr.modifier, operandAsF32(tid, instr.c)); - d = ftz(instr.modifier, sat(instr.modifier, a * b + c)); + d = ftz(instr.modifier, sat(instr.modifier, + roundedFma(a, b, c, instr.modifier))); setRegAsF32(tid, instr.d.reg, d); } @@ -4422,15 +4944,564 @@ void executive::CooperativeThreadArray::eval_Fma(CTAContext &context, ir::PTXF64 d, a = operandAsF64(tid, instr.a), b = operandAsF64(tid, instr.b), c = operandAsF64(tid, instr.c); - d = a * b + c; + d = roundedFma(a, b, c, instr.modifier); setRegAsF64(tid, instr.d.reg, d); } } + else if (instr.type == ir::PTXOperand::f16) { + for (int tid = 0; tid < threadCount; tid++) { + if (!context.predicated(tid, instr)) continue; + + ir::PTXF32 a = f16ToF32(ftzF16(instr.modifier, + operandAsU16(tid, instr.a))); + ir::PTXF32 b = f16ToF32(ftzF16(instr.modifier, + operandAsU16(tid, instr.b))); + ir::PTXF32 c = f16ToF32(ftzF16(instr.modifier, + operandAsU16(tid, instr.c))); + setRegAsB16(tid, instr.d.reg, + ftzF16(instr.modifier, toF16(roundedFma( + a, b, c, instr.modifier), instr.modifier))); + } + } + else if (instr.type == ir::PTXOperand::bf16) { + for (int tid = 0; tid < threadCount; tid++) { + if (!context.predicated(tid, instr)) continue; + + ir::PTXF32 a = bf16ToF32(operandAsU16(tid, instr.a)); + ir::PTXF32 b = bf16ToF32(operandAsU16(tid, instr.b)); + ir::PTXF32 c = bf16ToF32(operandAsU16(tid, instr.c)); + ir::PTXU16 d = f32ToBF16(fmaF32OneRound(a, b, c), instr.modifier); + setRegAsB16(tid, instr.d.reg, d); + } + } + else if (instr.type == ir::PTXOperand::bf16x2) { + for (int tid = 0; tid < threadCount; tid++) { + if (!context.predicated(tid, instr)) continue; + ir::PTXU32 a = operandAsU32(tid, instr.a), b = operandAsU32(tid, instr.b), + c = operandAsU32(tid, instr.c); + ir::PTXU16 al = static_cast(a), ah = a >> 16, + bl = static_cast(b), bh = b >> 16; + ir::PTXU16 cl = static_cast(c), ch = c >> 16; + auto fma = [&](ir::PTXU16 x, ir::PTXU16 y, ir::PTXU16 z) { + return f32ToBF16(fmaF32OneRound(bf16ToF32(x), bf16ToF32(y), + bf16ToF32(z)), instr.modifier); + }; + ir::PTXU16 dl = fma(al, bl, cl), dh = fma(ah, bh, ch); + setRegAsU32(tid, instr.d.reg, static_cast(dl) | + (static_cast(dh) << 16)); + } + } else { throw RuntimeException("unsupported data type", context.PC, instr); } } +void executive::CooperativeThreadArray::eval_Mma(CTAContext &context, + const ir::PTXInstruction &instr) { + trace(); + + const std::string error = instr.valid(); + if (!error.empty()) { + throw RuntimeException(error, context.PC, instr); + } + + const int completeThreads = (threadCount / 32) * 32; + for (int threadID = completeThreads; threadID < threadCount; ++threadID) { + if (context.predicated(threadID, instr)) { + throw RuntimeException("mma requires all warp lanes to participate", + context.PC, instr); + } + } + + const ir::PTXOperand::DataType inputType = instr.a.type; + if (inputType == ir::PTXOperand::b1) { + const int rows = instr.mmaShape == ir::PTXInstruction::MmaM8N8K128 ? 8 : 16; + const int kWords = instr.mmaShape == ir::PTXInstruction::MmaM16N8K256 ? 8 : 4; + for (int warpStart = 0; warpStart < completeThreads; warpStart += 32) { + int participants = 0; + for (int lane = 0; lane < 32; ++lane) participants += context.predicated(warpStart + lane, instr); + if (!participants) continue; + if (participants != 32) { + throw RuntimeException("mma requires all warp lanes to participate", context.PC, instr); + } + // Each word holds 32 consecutive K bits; A register pairs select rows and K halves. + ir::PTXU32 A[16][8], B[8][8], D[16][8]; + for (int lane = 0; lane < 32; ++lane) { + const int threadID = warpStart + lane; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for (unsigned i = 0; i < instr.a.array.size(); ++i) { + const int row = groupID + (i & 1) * 8, word = threadInGroup + (i / 2) * 4; + A[row][word] = getRegAsB32(threadID, instr.a.array[i].reg); + } + for (unsigned i = 0; i < instr.b.array.size(); ++i) { + const int word = threadInGroup + i * 4; + B[word][groupID] = getRegAsB32(threadID, instr.b.array[i].reg); + } + for (unsigned i = 0; i < instr.c.array.size(); ++i) { + const int row = groupID + (i / 2) * 8, col = threadInGroup * 2 + (i & 1); + D[row][col] = getRegAsB32(threadID, instr.c.array[i].reg); + } + } + for (int row = 0; row < rows; ++row) { + for (int col = 0; col < 8; ++col) { + for (int word = 0; word < kWords; ++word) { + const ir::PTXU32 bits = instr.booleanOperator == ir::PTXInstruction::BoolXor + ? A[row][word] ^ B[word][col] : A[row][word] & B[word][col]; + D[row][col] += hydrazine::popc(bits); + } + } + } + for (int lane = 0; lane < 32; ++lane) { + for (unsigned i = 0; i < instr.d.array.size(); ++i) { + const int row = (lane >> 2) + (i / 2) * 8, col = (lane & 3) * 2 + (i & 1); + setRegAsB32(warpStart + lane, instr.d.array[i].reg, D[row][col]); + } + } + } + return; + } + const bool subByteInput = inputType == ir::PTXOperand::s4 || + inputType == ir::PTXOperand::u4; + const bool intInput = subByteInput || inputType == ir::PTXOperand::s8 || + inputType == ir::PTXOperand::u8; + + if (instr.mmaShape == ir::PTXInstruction::MmaM8N8K4 && + inputType == ir::PTXOperand::f64) { + for (int warpStart = 0; warpStart < completeThreads; warpStart += 32) { + int participants = 0; + for (int lane = 0; lane < 32; ++lane) { + if (context.predicated(warpStart + lane, instr)) ++participants; + } + if (participants == 0) continue; + if (participants != 32) { + throw RuntimeException("mma requires all warp lanes to participate", + context.PC, instr); + } + + ir::PTXF64 A[8][4], B[4][8], D[8][8]; + for (int lane = 0; lane < 32; ++lane) { + const int threadID = warpStart + lane; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + A[groupID][threadInGroup] = getRegAsF64(threadID, instr.a.array[0].reg); + B[threadInGroup][groupID] = getRegAsF64(threadID, instr.b.array[0].reg); + for (int i = 0; i < 2; ++i) { + const int col = 2 * threadInGroup + i; + D[groupID][col] = getRegAsF64(threadID, instr.c.array[i].reg); + } + } + + for (int row = 0; row < 8; ++row) { + for (int col = 0; col < 8; ++col) { + for (int k = 0; k < 4; ++k) { + D[row][col] = roundedFma(A[row][k], B[k][col], + D[row][col], instr.modifier); + } + } + } + + for (int lane = 0; lane < 32; ++lane) { + const int threadID = warpStart + lane; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for (int i = 0; i < 2; ++i) { + const int col = 2 * threadInGroup + i; + setRegAsF64(threadID, instr.d.array[i].reg, D[groupID][col]); + } + } + } + return; + } + if (inputType == ir::PTXOperand::f64) { + throw RuntimeException("unsupported f64 mma shape", context.PC, instr); + } + if (intInput) { + // Packed 8-bit and 4-bit integer inputs accumulate into s32. + auto signExtend = [](ir::PTXU32 reg, unsigned int elementIndex, + ir::PTXOperand::DataType type) -> ir::PTXS32 { + const unsigned bits = (type == ir::PTXOperand::s4 || + type == ir::PTXOperand::u4) ? 4 : 8; + const ir::PTXU32 field = (reg >> (elementIndex * bits)) & ((1u << bits) - 1); + if ((type == ir::PTXOperand::s4 || type == ir::PTXOperand::s8) && + (field & (1u << (bits - 1)))) { + return static_cast(field) - (1 << bits); + } + return static_cast(field); + }; + + auto satfiniteClamp = [&instr](int64_t acc) -> ir::PTXS32 { + if (instr.modifier & ir::PTXInstruction::satfinite) { + const int64_t upper = (std::numeric_limits::max)(); + const int64_t lower = (std::numeric_limits::min)(); + if (acc > upper) acc = upper; + else if (acc < lower) acc = lower; + } + return static_cast(acc); + }; + + for (int warpStart = 0; warpStart < completeThreads; warpStart += 32) { + int participants = 0; + for (int lane = 0; lane < 32; ++lane) { + if (context.predicated(warpStart + lane, instr)) { + ++participants; + } + } + if (participants == 0) continue; + if (participants != 32) { + throw RuntimeException("mma requires all warp lanes to participate", + context.PC, instr); + } + + if (instr.mmaShape == ir::PTXInstruction::MmaM8N8K16 || + instr.mmaShape == ir::PTXInstruction::MmaM8N8K32) { + const int elementsPerRegister = subByteInput ? 8 : 4; + const int kCount = subByteInput ? 32 : 16; + ir::PTXS32 A[8][32] = {}; + ir::PTXS32 B[32][8] = {}; + ir::PTXS32 C[8][8] = {}; + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + + ir::PTXU32 aReg = operandAsU32(threadID, instr.a.array[0]); + for (int i = 0; i < elementsPerRegister; ++i) { + A[groupID][threadInGroup * elementsPerRegister + i] = + signExtend(aReg, i, instr.a.type); + } + + ir::PTXU32 bReg = operandAsU32(threadID, instr.b.array[0]); + for (int i = 0; i < elementsPerRegister; ++i) { + B[threadInGroup * elementsPerRegister + i][groupID] = + signExtend(bReg, i, instr.b.type); + } + + for (int i = 0; i < 2; ++i) { + C[groupID][threadInGroup * 2 + i] = + operandAsS32(threadID, instr.c.array[i]); + } + } + + ir::PTXS32 D[8][8]; + for (int row = 0; row < 8; ++row) { + for (int col = 0; col < 8; ++col) { + int64_t acc = C[row][col]; + for (int k = 0; k < kCount; ++k) { + acc += static_cast(A[row][k]) * + static_cast(B[k][col]); + } + D[row][col] = satfiniteClamp(acc); + } + } + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + for (int i = 0; i < 2; ++i) { + setRegAsS32(threadID, instr.d.array[i].reg, + D[groupID][threadInGroup * 2 + i]); + } + } + } + else if ((instr.mmaShape == ir::PTXInstruction::MmaM16N8K32 && !subByteInput) || + instr.mmaShape == ir::PTXInstruction::MmaM16N8K64) { + const int elementsPerRegister = subByteInput ? 8 : 4; + const int kCount = elementsPerRegister * 8; + ir::PTXS32 A[16][64] = {}; + ir::PTXS32 B[64][8] = {}; + ir::PTXS32 C[16][8] = {}; + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + + for (int i = 0; i < 4 * elementsPerRegister; ++i) { + int row = ((i / elementsPerRegister) & 1) ? groupID + 8 : groupID; + int col = threadInGroup * elementsPerRegister + i % elementsPerRegister + + (i >= 2 * elementsPerRegister ? kCount / 2 : 0); + ir::PTXU32 reg = operandAsU32(threadID, instr.a.array[i / elementsPerRegister]); + A[row][col] = signExtend(reg, i % elementsPerRegister, instr.a.type); + } + + for (int i = 0; i < 2 * elementsPerRegister; ++i) { + int row = threadInGroup * elementsPerRegister + i % elementsPerRegister + + (i >= elementsPerRegister ? kCount / 2 : 0); + int col = groupID; + ir::PTXU32 reg = operandAsU32(threadID, instr.b.array[i / elementsPerRegister]); + B[row][col] = signExtend(reg, i % elementsPerRegister, instr.b.type); + } + + for (int i = 0; i < 4; ++i) { + int row = groupID + (i >= 2 ? 8 : 0); + int col = threadInGroup * 2 + (i & 1); + C[row][col] = operandAsS32(threadID, instr.c.array[i]); + } + } + + ir::PTXS32 D[16][8]; + for (int row = 0; row < 16; ++row) { + for (int col = 0; col < 8; ++col) { + int64_t acc = C[row][col]; + for (int k = 0; k < kCount; ++k) { + acc += static_cast(A[row][k]) * + static_cast(B[k][col]); + } + D[row][col] = satfiniteClamp(acc); + } + } + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + for (int i = 0; i < 4; ++i) { + int row = groupID + (i >= 2 ? 8 : 0); + int col = threadInGroup * 2 + (i & 1); + setRegAsS32(threadID, instr.d.array[i].reg, D[row][col]); + } + } + } + else if (instr.mmaShape == ir::PTXInstruction::MmaM16N8K16 || + (instr.mmaShape == ir::PTXInstruction::MmaM16N8K32 && subByteInput)) { + const int elementsPerRegister = subByteInput ? 8 : 4; + const int kCount = elementsPerRegister * 4; + ir::PTXS32 A[16][32] = {}; + ir::PTXS32 B[32][8] = {}; + ir::PTXS32 C[16][8] = {}; + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + + for (int i = 0; i < 2 * elementsPerRegister; ++i) { + int row = (i < elementsPerRegister) ? groupID : groupID + 8; + int col = threadInGroup * elementsPerRegister + i % elementsPerRegister; + ir::PTXU32 reg = operandAsU32(threadID, instr.a.array[i / elementsPerRegister]); + A[row][col] = signExtend(reg, i % elementsPerRegister, instr.a.type); + } + + for (int i = 0; i < elementsPerRegister; ++i) { + int row = threadInGroup * elementsPerRegister + i; + int col = groupID; + ir::PTXU32 reg = operandAsU32(threadID, instr.b.array[0]); + B[row][col] = signExtend(reg, i, instr.b.type); + } + + for (int i = 0; i < 4; ++i) { + int row = groupID + (i >= 2 ? 8 : 0); + int col = threadInGroup * 2 + (i & 1); + C[row][col] = operandAsS32(threadID, instr.c.array[i]); + } + } + + ir::PTXS32 D[16][8]; + for (int row = 0; row < 16; ++row) { + for (int col = 0; col < 8; ++col) { + int64_t acc = C[row][col]; + for (int k = 0; k < kCount; ++k) { + acc += static_cast(A[row][k]) * + static_cast(B[k][col]); + } + D[row][col] = satfiniteClamp(acc); + } + } + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + for (int i = 0; i < 4; ++i) { + int row = groupID + (i >= 2 ? 8 : 0); + int col = threadInGroup * 2 + (i & 1); + setRegAsS32(threadID, instr.d.array[i].reg, D[row][col]); + } + } + } + else { + throw RuntimeException("unsupported integer mma shape", + context.PC, instr); + } + } + return; + } + + if (instr.mmaShape == ir::PTXInstruction::MmaM8N8K4 && + inputType == ir::PTXOperand::f16) { + for (int warp = 0; warp < completeThreads; warp += 32) { + int participants = 0; + for (int lane = 0; lane < 32; ++lane) participants += context.predicated(warp + lane, instr); + if (!participants) continue; + if (participants != 32) throw RuntimeException("mma requires all warp lanes to participate", context.PC, instr); + // Four independent products: lanes 4*g..4*g+3 and 16+4*g..19+4*g. + ir::PTXF32 A[4][8][4], B[4][4][8], D[4][8][8]; + for (int lane = 0; lane < 32; ++lane) { + const int group = (lane >> 2) & 3, row = (lane & 3) + (lane >= 16 ? 4 : 0); + for (int i = 0; i < 4; ++i) { + const int aRow = instr.mmaAColumnMajor ? i + (lane >= 16 ? 4 : 0) : row; + const int aCol = instr.mmaAColumnMajor ? lane & 3 : i; + const int bRow = instr.mmaBColumnMajor ? i : lane & 3; + const int bCol = instr.mmaBColumnMajor ? row : i + (lane >= 16 ? 4 : 0); + A[group][aRow][aCol] = mmaHalf(*this, warp + lane, instr.a.array[i / 2], i & 1, inputType); + B[group][bRow][bCol] = mmaHalf(*this, warp + lane, instr.b.array[i / 2], i & 1, inputType); + } + for (int i = 0; i < 8; ++i) { + if (instr.c.type == ir::PTXOperand::f32) { + const int cRow = (lane & 1) + (i & 2) + (lane >= 16 ? 4 : 0); + const int cCol = (i & 4) + (lane & 2) + (i & 1); + D[group][cRow][cCol] = operandAsF32(warp + lane, instr.c.array[i]); + } else D[group][row][i] = mmaHalf(*this, warp + lane, instr.c.array[i / 2], i & 1, inputType); + } + } + for (int group = 0; group < 4; ++group) + for (int row = 0; row < 8; ++row) + for (int col = 0; col < 8; ++col) + for (int k = 0; k < 4; ++k) { + D[group][row][col] = std::fma(A[group][row][k], B[group][k][col], D[group][row][col]); + if (instr.type == ir::PTXOperand::f16) D[group][row][col] = f16ToF32(toF16(D[group][row][col], 0)); + } + for (int lane = 0; lane < 32; ++lane) { + const int group = (lane >> 2) & 3, row = (lane & 3) + (lane >= 16 ? 4 : 0); + if (instr.type == ir::PTXOperand::f32) { + for (int i = 0; i < 8; ++i) { + const int dRow = (lane & 1) + (i & 2) + (lane >= 16 ? 4 : 0); + const int dCol = (i & 4) + (lane & 2) + (i & 1); + setRegAsF32(warp + lane, instr.d.array[i].reg, D[group][dRow][dCol]); + } + } else for (int i = 0; i < 4; ++i) setRegAsB32(warp + lane, instr.d.array[i].reg, + toF16(D[group][row][2 * i], 0) | (static_cast(toF16(D[group][row][2 * i + 1], 0)) << 16)); + } + } + return; + } + + if (instr.mmaShape == ir::PTXInstruction::MmaM8N8K4) { + throw RuntimeException("unsupported m8n8k4 input type", context.PC, instr); + } + + const bool halfAccumulator = instr.type == ir::PTXOperand::f16; + const bool tf32Input = inputType == ir::PTXOperand::tf32; + const bool m16n8k8 = instr.mmaShape == ir::PTXInstruction::MmaM16N8K8; + const bool m16n8k4 = instr.mmaShape == ir::PTXInstruction::MmaM16N8K4; + for (int warpStart = 0; warpStart < completeThreads; warpStart += 32) { + int participants = 0; + for (int lane = 0; lane < 32; ++lane) { + if (context.predicated(warpStart + lane, instr)) { + ++participants; + } + } + if (participants == 0) continue; + if (participants != 32) { + throw RuntimeException("mma requires all warp lanes to participate", + context.PC, instr); + } + + ir::PTXF32 A[16][16] = {}; + ir::PTXF32 B[16][8] = {}; + ir::PTXF32 C[16][8] = {}; + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + + if (tf32Input) { + for (int i = 0; i < (m16n8k4 ? 2 : 4); ++i) { + int row = (i & 1) ? groupID + 8 : groupID; + int col = threadInGroup + (i >= 2 ? 4 : 0); + A[row][col] = tf32FromF32(operandAsF32(threadID, + instr.a.array[i])); + } + + for (int i = 0; i < (m16n8k4 ? 1 : 2); ++i) { + int row = threadInGroup + (i >= 1 ? 4 : 0); + int col = groupID; + B[row][col] = tf32FromF32(operandAsF32(threadID, + instr.b.array[i])); + } + } + else if (m16n8k8) { + for (int i = 0; i < 4; ++i) { + int row = i < 2 ? groupID : groupID + 8; + int col = threadInGroup * 2 + (i & 1); + A[row][col] = mmaHalf(*this, threadID, + instr.a.array[i / 2], i & 1, inputType); + } + + for (int i = 0; i < 2; ++i) { + int row = threadInGroup * 2 + (i & 1); + B[row][groupID] = mmaHalf(*this, threadID, + instr.b.array[0], i, inputType); + } + } + else if (instr.mmaShape == ir::PTXInstruction::MmaM16N8K16) { + for (int i = 0; i < 8; ++i) { + int row = (i < 2 || (i >= 4 && i < 6)) ? groupID : groupID + 8; + int col = threadInGroup * 2 + (i & 1) + (i >= 4 ? 8 : 0); + A[row][col] = mmaHalf(*this, threadID, + instr.a.array[i / 2], i & 1, inputType); + } + + for (int i = 0; i < 4; ++i) { + int row = threadInGroup * 2 + (i & 1) + (i >= 2 ? 8 : 0); + int col = groupID; + B[row][col] = mmaHalf(*this, threadID, + instr.b.array[i / 2], i & 1, inputType); + } + } + + for (int i = 0; i < 4; ++i) { + int row = groupID + (i >= 2 ? 8 : 0); + int col = threadInGroup * 2 + (i & 1); + C[row][col] = halfAccumulator + ? mmaHalf(*this, threadID, instr.c.array[i / 2], i & 1, + ir::PTXOperand::f16) + : operandAsF32(threadID, instr.c.array[i]); + } + } + + ir::PTXF32 D[16][8]; + for (int row = 0; row < 16; ++row) { + for (int col = 0; col < 8; ++col) { + D[row][col] = C[row][col]; + for (int k = 0; k < (m16n8k4 ? 4 : m16n8k8 ? 8 : 16); ++k) { + D[row][col] = std::fma(A[row][k], B[k][col], D[row][col]); + if (halfAccumulator) { + D[row][col] = f16ToF32(toF16(D[row][col], instr.modifier)); + } + } + } + } + + for (int lane = 0; lane < 32; ++lane) { + int threadID = warpStart + lane; + int groupID = lane >> 2; + int threadInGroup = lane & 3; + if (halfAccumulator) { + ir::PTXU32 packed[2] = {0, 0}; + for (int i = 0; i < 4; ++i) { + int row = groupID + (i >= 2 ? 8 : 0); + int col = threadInGroup * 2 + (i & 1); + packed[i / 2] |= static_cast( + toF16(D[row][col], instr.modifier)) << (16 * (i & 1)); + } + setRegAsB32(threadID, instr.d.array[0].reg, packed[0]); + setRegAsB32(threadID, instr.d.array[1].reg, packed[1]); + } + else for (int i = 0; i < 4; ++i) { + int row = groupID + (i >= 2 ? 8 : 0); + int col = threadInGroup * 2 + (i & 1); + setRegAsF32(threadID, instr.d.array[i].reg, + D[row][col]); + } + } + } +} + /*! @@ -4440,24 +5511,74 @@ void executive::CooperativeThreadArray::eval_Isspacep(CTAContext &context, trace(); switch (instr.addressSpace) { + case ir::PTXInstruction::Const: + { + if (instr.a.type == ir::PTXOperand::u32) { + for (int tid = 0; tid < threadCount; tid++) { + if (!context.predicated(tid, instr)) continue; + ir::PTXU32 ptr = operandAsU32(tid, instr.a); + ir::PTXU32 base; + hydrazine::bit_cast(base, kernel->ConstMemory); + const ir::PTXU32 size = kernel->constMemorySize(); + setRegAsPredicate(tid, instr.d.reg, + addressInRange(ptr, base, size)); + } + } + else { + for (int tid = 0; tid < threadCount; tid++) { + if (!context.predicated(tid, instr)) continue; + ir::PTXU64 ptr = operandAsU64(tid, instr.a); + ir::PTXU64 base; + hydrazine::bit_cast(base, kernel->ConstMemory); + const ir::PTXU64 size = kernel->constMemorySize(); + setRegAsPredicate(tid, instr.d.reg, + addressInRange(ptr, base, size)); + } + } + } + break; + case ir::PTXInstruction::Param: + { + if (instr.a.type == ir::PTXOperand::u32) { + const ir::PTXU32 base = static_cast(kernel->scheduler + ? kernel->scheduler->argumentMemory() + : reinterpret_cast(kernel->ArgumentMemory)); + const ir::PTXU32 size = static_cast(kernel->scheduler + ? kernel->scheduler->argumentMemorySize() + : kernel->argumentMemorySize()); + for (int tid = 0; tid < threadCount; ++tid) { + if (!context.predicated(tid, instr)) continue; + const ir::PTXU32 ptr = operandAsU32(tid, instr.a); + setRegAsPredicate(tid, instr.d.reg, + addressInRange(ptr, base, size)); + } + } + else { + const ir::PTXU64 base = kernel->scheduler + ? kernel->scheduler->argumentMemory() + : reinterpret_cast(kernel->ArgumentMemory); + const ir::PTXU64 size = kernel->scheduler + ? kernel->scheduler->argumentMemorySize() + : kernel->argumentMemorySize(); + for (int tid = 0; tid < threadCount; ++tid) { + if (!context.predicated(tid, instr)) continue; + const ir::PTXU64 ptr = operandAsU64(tid, instr.a); + setRegAsPredicate(tid, instr.d.reg, + addressInRange(ptr, base, size)); + } + } + } + break; case ir::PTXInstruction::Local: { - if (sizeof(void *) == 4) { + if (instr.a.type == ir::PTXOperand::u32) { for (int tid = 0; tid < threadCount; tid++) { if (!context.predicated(tid, instr)) { continue; } - ir::PTXU32 ptr = operandAsU32(tid, instr.a); - ir::PTXU32 localMemPtr; - ir::PTXU32 localMemSize = functionCallStack.localMemorySize(); - hydrazine::bit_cast(localMemPtr, - functionCallStack.localMemoryPointer(tid)); - if (ptr >= localMemPtr && localMemPtr + localMemSize > ptr) { - setRegAsPredicate(tid, instr.d.reg, 1); - } - else { - setRegAsPredicate(tid, instr.d.reg, 0); - } + const ir::PTXU32 ptr = operandAsU32(tid, instr.a); + setRegAsPredicate(tid, instr.d.reg, + functionCallStack.isLocalMemoryAddress(ptr, tid, 32)); } } else { @@ -4465,24 +5586,16 @@ void executive::CooperativeThreadArray::eval_Isspacep(CTAContext &context, if (!context.predicated(tid, instr)) { continue; } - ir::PTXU64 ptr = operandAsU64(tid, instr.a); - ir::PTXU64 localMemPtr; - ir::PTXU64 localMemSize = functionCallStack.localMemorySize(); - hydrazine::bit_cast(localMemPtr, - functionCallStack.localMemoryPointer(tid)); - if (ptr >= localMemPtr && localMemPtr + localMemSize > ptr) { - setRegAsPredicate(tid, instr.d.reg, 1); - } - else { - setRegAsPredicate(tid, instr.d.reg, 0); - } + const ir::PTXU64 ptr = operandAsU64(tid, instr.a); + setRegAsPredicate(tid, instr.d.reg, + functionCallStack.isLocalMemoryAddress(ptr, tid, 64)); } } } break; case ir::PTXInstruction::Shared: { - if (sizeof(void *) == 4) { + if (instr.a.type == ir::PTXOperand::u32) { for (int tid = 0; tid < threadCount; tid++) { if (!context.predicated(tid, instr)) { continue; @@ -4492,7 +5605,7 @@ void executive::CooperativeThreadArray::eval_Isspacep(CTAContext &context, ir::PTXU32 sharedMemSize = functionCallStack.sharedMemorySize(); hydrazine::bit_cast(sharedMemPtr, functionCallStack.sharedMemoryPointer()); - if (ptr >= sharedMemPtr && sharedMemPtr + sharedMemSize > ptr) { + if (addressInRange(ptr, sharedMemPtr, sharedMemSize)) { setRegAsPredicate(tid, instr.d.reg, 1); } else { @@ -4505,12 +5618,12 @@ void executive::CooperativeThreadArray::eval_Isspacep(CTAContext &context, if (!context.predicated(tid, instr)) { continue; } - ir::PTXU64 ptr = operandAsU32(tid, instr.a); + ir::PTXU64 ptr = operandAsU64(tid, instr.a); ir::PTXU64 sharedMemPtr; ir::PTXU64 sharedMemSize = functionCallStack.sharedMemorySize(); hydrazine::bit_cast(sharedMemPtr, functionCallStack.sharedMemoryPointer()); - if (ptr >= sharedMemPtr && sharedMemPtr + sharedMemSize > ptr) { + if (addressInRange(ptr, sharedMemPtr, sharedMemSize)) { setRegAsPredicate(tid, instr.d.reg, 1); } else { @@ -4522,28 +5635,24 @@ void executive::CooperativeThreadArray::eval_Isspacep(CTAContext &context, break; case ir::PTXInstruction::Global: { - if (sizeof(void *) == 4) { + if (instr.a.type == ir::PTXOperand::u32) { for (int tid = 0; tid < threadCount; tid++) { if (!context.predicated(tid, instr)) { continue; } ir::PTXU32 ptr = operandAsU32(tid, instr.a); - ir::PTXU32 localMemPtr; - ir::PTXU32 localMemSize = functionCallStack.localMemorySize(); ir::PTXU32 sharedMemPtr; ir::PTXU32 sharedMemSize = functionCallStack.sharedMemorySize(); - hydrazine::bit_cast(localMemPtr, - functionCallStack.localMemoryPointer(tid)); + ir::PTXU32 constMemPtr; + const ir::PTXU32 constMemSize = kernel->constMemorySize(); hydrazine::bit_cast(sharedMemPtr, functionCallStack.sharedMemoryPointer()); - if ((ptr >= sharedMemPtr && sharedMemPtr + sharedMemSize > ptr) - || (ptr >= localMemPtr - && localMemPtr + localMemSize > ptr)) { - setRegAsPredicate(tid, instr.d.reg, 0); - } - else { - setRegAsPredicate(tid, instr.d.reg, 1); - } + hydrazine::bit_cast(constMemPtr, kernel->ConstMemory); + const bool inShared = addressInRange(ptr, sharedMemPtr, sharedMemSize); + const bool inLocal = + functionCallStack.isLocalMemoryAddress(ptr, tid, 32); + const bool inConst = addressInRange(ptr, constMemPtr, constMemSize); + setRegAsPredicate(tid, instr.d.reg, !(inShared || inLocal || inConst)); } } else { @@ -4552,22 +5661,18 @@ void executive::CooperativeThreadArray::eval_Isspacep(CTAContext &context, continue; } ir::PTXU64 ptr = operandAsU64(tid, instr.a); - ir::PTXU64 localMemPtr; - ir::PTXU64 localMemSize = functionCallStack.localMemorySize(); ir::PTXU64 sharedMemPtr; ir::PTXU64 sharedMemSize = functionCallStack.sharedMemorySize(); - hydrazine::bit_cast(localMemPtr, - functionCallStack.localMemoryPointer(tid)); + ir::PTXU64 constMemPtr; + const ir::PTXU64 constMemSize = kernel->constMemorySize(); hydrazine::bit_cast(sharedMemPtr, functionCallStack.sharedMemoryPointer()); - if ((ptr >= sharedMemPtr && sharedMemPtr + sharedMemSize > ptr) - || (ptr >= localMemPtr - && localMemPtr + localMemSize > ptr)) { - setRegAsPredicate(tid, instr.d.reg, 0); - } - else { - setRegAsPredicate(tid, instr.d.reg, 1); - } + hydrazine::bit_cast(constMemPtr, kernel->ConstMemory); + const bool inShared = addressInRange(ptr, sharedMemPtr, sharedMemSize); + const bool inLocal = + functionCallStack.isLocalMemoryAddress(ptr, tid, 64); + const bool inConst = addressInRange(ptr, constMemPtr, constMemSize); + setRegAsPredicate(tid, instr.d.reg, !(inShared || inLocal || inConst)); } } } @@ -5085,6 +6190,38 @@ void executive::CooperativeThreadArray::eval_Lg2(CTAContext &context, } } +void executive::CooperativeThreadArray::eval_Lop3(CTAContext &context, + const ir::PTXInstruction &instr) { + trace(); + for (int threadID = 0; threadID < threadCount; ++threadID) { + if (!context.predicated(threadID, instr)) continue; + + const ir::PTXU32 a = operandAsU32(threadID, instr.a); + const ir::PTXU32 b = operandAsU32(threadID, instr.b); + const ir::PTXU32 c = operandAsU32(threadID, instr.c); + const ir::PTXU32 immLut = operandAsU32(threadID, instr.immLut); + ir::PTXU32 d = 0; + for (unsigned int bit = 0; bit < 32; ++bit) { + const unsigned int aBit = (a >> bit) & 1u; + const unsigned int bBit = (b >> bit) & 1u; + const unsigned int cBit = (c >> bit) & 1u; + const unsigned int index = (aBit << 2) | (bBit << 1) | cBit; + const unsigned int resultBit = (immLut >> index) & 1u; + + d |= resultBit << bit; + } + if( instr.d.addressMode != ir::PTXOperand::BitBucket ) { + setRegAsU32(threadID, instr.d.reg, d); + } + if( instr.pq.addressMode != ir::PTXOperand::Invalid ) { + const bool q = operandAsPredicate(threadID, instr.q); + const bool p = instr.booleanOperator == ir::PTXInstruction::BoolAnd + ? d != 0 && q : d != 0 || q; + setRegAsPredicate(threadID, instr.pq.reg, p); + } + } +} + /*! */ @@ -5283,7 +6420,8 @@ void executive::CooperativeThreadArray::eval_Mad(CTAContext &context, b = ftz(instr.modifier, operandAsF32(threadID, instr.b)), c = ftz(instr.modifier, operandAsF32(threadID, instr.c)); - d = ftz(instr.modifier, sat(instr.modifier, a * b + c)); + d = ftz(instr.modifier, + sat(instr.modifier, roundedFma(a, b, c, instr.modifier))); setRegAsF32(threadID, instr.d.reg, d); } @@ -5297,10 +6435,7 @@ void executive::CooperativeThreadArray::eval_Mad(CTAContext &context, b = operandAsF64(threadID, instr.b), c = operandAsF64(threadID, instr.c); - d = a * b + c; - if (instr.modifier & ir::PTXInstruction::sat) { - if (d < 0) d = 0; else if (d > 1) d = 1; - } + d = roundedFma(a, b, c, instr.modifier); setRegAsF64(threadID, instr.d.reg, d); } } break; @@ -5319,24 +6454,10 @@ void executive::CooperativeThreadArray::eval_Max(CTAContext &context, if (instr.type == ir::PTXOperand::f32) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - - ir::PTXF32 d, a = operandAsF32(threadID, instr.a), + ir::PTXF32 a = operandAsF32(threadID, instr.a), b = operandAsF32(threadID, instr.b); - - if(hydrazine::isnan(a)) - { - d = ftz(instr.modifier, b); - } - else if(hydrazine::isnan(b)) - { - d = ftz(instr.modifier, a); - } - else - { - d = ftz(instr.modifier, a > b ? a : b); - } - - setRegAsF32(threadID, instr.d.reg, d); + setRegAsF32(threadID, instr.d.reg, + minMaxF32(instr.modifier, a, b, true)); } } else if (instr.type == ir::PTXOperand::f64) { @@ -5362,6 +6483,32 @@ void executive::CooperativeThreadArray::eval_Max(CTAContext &context, setRegAsF64(threadID, instr.d.reg, d); } } + else if (instr.type == ir::PTXOperand::f16 || instr.type == ir::PTXOperand::bf16) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + + ir::PTXU16 d = minMaxHalf(instr.type, instr.modifier, + operandAsU16(threadID, instr.a), operandAsU16(threadID, instr.b), true); + + setRegAsB16(threadID, instr.d.reg, d); + } + } + else if (instr.type == ir::PTXOperand::f16x2 || instr.type == ir::PTXOperand::bf16x2) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + const ir::PTXU32 a = operandAsU32(threadID, instr.a); + const ir::PTXU32 b = operandAsU32(threadID, instr.b); + const ir::PTXU16 aLow = static_cast(a); + const ir::PTXU16 aHigh = static_cast(a >> 16); + const ir::PTXU16 bLow = static_cast(b); + const ir::PTXU16 bHigh = static_cast(b >> 16); + const ir::PTXU32 low = static_cast( + minMaxHalf(instr.type, instr.modifier, aLow, bLow, true)); + const ir::PTXU32 high = static_cast( + minMaxHalf(instr.type, instr.modifier, aHigh, bHigh, true)); + setRegAsU32(threadID, instr.d.reg, low | (high << 16)); + } + } else if (instr.type == ir::PTXOperand::s16) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; @@ -5433,27 +6580,36 @@ void executive::CooperativeThreadArray::eval_Max(CTAContext &context, void executive::CooperativeThreadArray::eval_Min(CTAContext &context, const ir::PTXInstruction &instr) { trace(); - if (instr.type == ir::PTXOperand::f32) { + if (instr.type == ir::PTXOperand::f16 || instr.type == ir::PTXOperand::bf16) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - - ir::PTXF32 d, a = operandAsF32(threadID, instr.a), - b = operandAsF32(threadID, instr.b); - - if(hydrazine::isnan(a)) - { - d = ftz(instr.modifier, b); - } - else if(hydrazine::isnan(b)) - { - d = ftz(instr.modifier, a); - } - else - { - d = ftz(instr.modifier, (a < b ? a : b)); - } - - setRegAsF32(threadID, instr.d.reg, d); + setRegAsB16(threadID, instr.d.reg, minMaxHalf(instr.type, instr.modifier, + operandAsU16(threadID, instr.a), operandAsU16(threadID, instr.b), false)); + } + } + else if (instr.type == ir::PTXOperand::f16x2 || instr.type == ir::PTXOperand::bf16x2) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + const ir::PTXU32 a = operandAsU32(threadID, instr.a); + const ir::PTXU32 b = operandAsU32(threadID, instr.b); + const ir::PTXU16 aLow = static_cast(a); + const ir::PTXU16 aHigh = static_cast(a >> 16); + const ir::PTXU16 bLow = static_cast(b); + const ir::PTXU16 bHigh = static_cast(b >> 16); + const ir::PTXU32 low = static_cast( + minMaxHalf(instr.type, instr.modifier, aLow, bLow, false)); + const ir::PTXU32 high = static_cast( + minMaxHalf(instr.type, instr.modifier, aHigh, bHigh, false)); + setRegAsU32(threadID, instr.d.reg, low | (high << 16)); + } + } + else if (instr.type == ir::PTXOperand::f32) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + ir::PTXF32 a = operandAsF32(threadID, instr.a), + b = operandAsF32(threadID, instr.b); + setRegAsF32(threadID, instr.d.reg, + minMaxF32(instr.modifier, a, b, false)); } } else if (instr.type == ir::PTXOperand::f64) { @@ -5552,6 +6708,16 @@ void executive::CooperativeThreadArray::eval_Membar(CTAContext &context, const i /*! No need to do anything here. */ } +/*! + +*/ +void executive::CooperativeThreadArray::eval_Fence(CTAContext &context, const ir::PTXInstruction &instr) { + trace(); + /*! No need to do anything here -- see eval_Membar; this simulator executes + every instruction to completion for every thread before the next one + starts, so there is no reordering for fence to guard against. */ +} + //////////////////////////////////////////////////////////////////////////////// /*! @@ -5813,6 +6979,12 @@ void executive::CooperativeThreadArray::eval_Mov_imm(CTAContext &context, setRegAsU16(threadID, instr.d.reg, a); } break; + case ir::PTXOperand::f16: + { + ir::PTXU16 a = operandAsU16(threadID, instr.a); + setRegAsB16(threadID, instr.d.reg, a); + } + break; case ir::PTXOperand::u32: case ir::PTXOperand::s32: case ir::PTXOperand::b32: @@ -5950,13 +7122,47 @@ void executive::CooperativeThreadArray::eval_Mul24(CTAContext &context, const ir */ void executive::CooperativeThreadArray::eval_Mul(CTAContext &context, const ir::PTXInstruction &instr) { trace(); - if (instr.type == ir::PTXOperand::f32) { + if (instr.type == ir::PTXOperand::f16x2) { + const int effectiveModifier = instr.modifier | ir::PTXInstruction::rn; + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + ir::PTXU32 a = operandAsU32(threadID, instr.a); + ir::PTXU32 b = operandAsU32(threadID, instr.b); + ir::PTXU16 al = static_cast(a), ah = a >> 16; + ir::PTXU16 bl = static_cast(b), bh = b >> 16; + ir::PTXU16 dl = ftzF16(effectiveModifier, toF16( + roundedMul(f16ToF32(ftzF16(effectiveModifier, al)), + f16ToF32(ftzF16(effectiveModifier, bl)), effectiveModifier), + effectiveModifier)); + ir::PTXU16 dh = ftzF16(effectiveModifier, toF16( + roundedMul(f16ToF32(ftzF16(effectiveModifier, ah)), + f16ToF32(ftzF16(effectiveModifier, bh)), effectiveModifier), + effectiveModifier)); + setRegAsU32(threadID, instr.d.reg, static_cast(dl) | + (static_cast(dh) << 16)); + } + } + else if (instr.type == ir::PTXOperand::f16) { + const int effectiveModifier = instr.modifier | ir::PTXInstruction::rn; + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + + ir::PTXF32 a = f16ToF32(ftzF16(effectiveModifier, + operandAsU16(threadID, instr.a))); + ir::PTXF32 b = f16ToF32(ftzF16(effectiveModifier, + operandAsU16(threadID, instr.b))); + ir::PTXU16 d = toF16(roundedMul(a, b, effectiveModifier), effectiveModifier); + setRegAsB16(threadID, instr.d.reg, ftzF16(effectiveModifier, d)); + } + } + else if (instr.type == ir::PTXOperand::f32) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; ir::PTXF32 d, a = ftz(instr.modifier, operandAsF32(threadID, instr.a)), b = ftz(instr.modifier, operandAsF32(threadID, instr.b)); - d = ftz(instr.modifier, sat(instr.modifier, a * b)); + d = ftz(instr.modifier, sat(instr.modifier, + roundedMul(a, b, instr.modifier))); setRegAsF32(threadID, instr.d.reg, d); } } @@ -5965,7 +7171,7 @@ void executive::CooperativeThreadArray::eval_Mul(CTAContext &context, const ir:: if (!context.predicated(threadID, instr)) continue; ir::PTXF64 d, a = operandAsF64(threadID, instr.a), b = operandAsF64(threadID, instr.b); - d = a * b; + d = roundedMul(a, b, instr.modifier); setRegAsF64(threadID, instr.d.reg, d); } } @@ -6113,7 +7319,37 @@ void executive::CooperativeThreadArray::eval_Mul(CTAContext &context, const ir:: */ void executive::CooperativeThreadArray::eval_Neg(CTAContext &context, const ir::PTXInstruction &instr) { trace(); - if (instr.type == ir::PTXOperand::f32) { + if (instr.type == ir::PTXOperand::f16x2 || instr.type == ir::PTXOperand::bf16x2) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + ir::PTXU32 a = operandAsU32(threadID, instr.a); + a = static_cast(ftzF16(instr.modifier, + static_cast(a))) | + (static_cast(ftzF16(instr.modifier, a >> 16)) << 16); + setRegAsU32(threadID, instr.d.reg, a ^ 0x80008000u); + } + } + else if (instr.type == ir::PTXOperand::bf16) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + + ir::PTXF32 a = ftz(instr.modifier, + bf16ToF32(operandAsU16(threadID, instr.a))); + ir::PTXU16 d = f32ToBF16Rn(-a); + setRegAsB16(threadID, instr.d.reg, d); + } + } + else if (instr.type == ir::PTXOperand::f16) { + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + + ir::PTXF32 a = f16ToF32(ftzF16(instr.modifier, + operandAsU16(threadID, instr.a))); + ir::PTXU16 d = toF16(-a, instr.modifier); + setRegAsB16(threadID, instr.d.reg, ftzF16(instr.modifier, d)); + } + } + else if (instr.type == ir::PTXOperand::f32) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; @@ -6136,7 +7372,7 @@ void executive::CooperativeThreadArray::eval_Neg(CTAContext &context, const ir:: if (!context.predicated(threadID, instr)) continue; ir::PTXS16 d, a = operandAsS16(threadID, instr.a); - d = -a; + d = a == (std::numeric_limits::min)() ? a : -a; setRegAsS16(threadID, instr.d.reg, d); } } @@ -6145,7 +7381,7 @@ void executive::CooperativeThreadArray::eval_Neg(CTAContext &context, const ir:: if (!context.predicated(threadID, instr)) continue; ir::PTXS32 d, a = operandAsS32(threadID, instr.a); - d = -a; + d = a == (std::numeric_limits::min)() ? a : -a; setRegAsS32(threadID, instr.d.reg, d); } } @@ -6154,7 +7390,7 @@ void executive::CooperativeThreadArray::eval_Neg(CTAContext &context, const ir:: if (!context.predicated(threadID, instr)) continue; ir::PTXS64 d, a = operandAsS64(threadID, instr.a); - d = -a; + d = a == (std::numeric_limits::min)() ? a : -a; setRegAsS64(threadID, instr.d.reg, d); } } @@ -6479,8 +7715,14 @@ void executive::CooperativeThreadArray::eval_Rcp(CTAContext &context, for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; + // Without .ftz, .approx must still support subnormal inputs + // (e.g. rcp.approx.f32(2^-127) ~= 2^127); only .ftz flushes + // the input to signed zero, whose reciprocal is signed + // infinity through the division below. ir::PTXF32 d, a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); - d = ftz(instr.modifier, 1.0f/a); + d = ftz(instr.modifier, roundedDiv(1.0f, a, + instr.modifier & ir::PTXInstruction::approx + ? ir::PTXInstruction::rn : instr.modifier)); setRegAsF32(threadID, instr.d.reg, d); } } @@ -6489,7 +7731,22 @@ void executive::CooperativeThreadArray::eval_Rcp(CTAContext &context, if (!context.predicated(threadID, instr)) continue; ir::PTXF64 d, a = operandAsF64(threadID, instr.a); - d = 1.0/a; + if (instr.modifier & ir::PTXInstruction::approx) { + const ir::PTXU64 high = 0xffffffff00000000ull; + if (hydrazine::isnan(a)) { + d = hydrazine::bit_cast(0x7fffffff00000000ull); + } else { + a = hydrazine::bit_cast( + hydrazine::bit_cast(a) & high); + if (issubnormal_(a)) a = hydrazine::copysign(0.0, a); + d = roundedDiv(1.0, a, ir::PTXInstruction::rn); + if (issubnormal_(d)) d = hydrazine::copysign(0.0, d); + d = hydrazine::bit_cast( + hydrazine::bit_cast(d) & high); + } + } else { + d = roundedDiv(1.0, a, instr.modifier); + } setRegAsF64(threadID, instr.d.reg, d); } } @@ -6502,8 +7759,21 @@ void executive::CooperativeThreadArray::eval_Rcp(CTAContext &context, */ void executive::CooperativeThreadArray::eval_Red(CTAContext &context, const ir::PTXInstruction &instr) { - trace(); - throw RuntimeException("instruction not implemented", context.PC, instr); + ir::PTXInstruction::AtomicOperation operation; + switch (instr.reductionOperation) { + case ir::PTXInstruction::ReductionAnd: operation = ir::PTXInstruction::AtomicAnd; break; + case ir::PTXInstruction::ReductionXor: operation = ir::PTXInstruction::AtomicXor; break; + case ir::PTXInstruction::ReductionOr: operation = ir::PTXInstruction::AtomicOr; break; + case ir::PTXInstruction::ReductionAdd: operation = ir::PTXInstruction::AtomicAdd; break; + case ir::PTXInstruction::ReductionInc: operation = ir::PTXInstruction::AtomicInc; break; + case ir::PTXInstruction::ReductionDec: operation = ir::PTXInstruction::AtomicDec; break; + case ir::PTXInstruction::ReductionMin: operation = ir::PTXInstruction::AtomicMin; break; + case ir::PTXInstruction::ReductionMax: operation = ir::PTXInstruction::AtomicMax; break; + default: + throw RuntimeException("unsupported reduction operation", + context.PC, instr); + } + evalAtomicRMW(context, instr, operation, false); } /*! @@ -6518,11 +7788,7 @@ void executive::CooperativeThreadArray::eval_Rem(CTAContext &context, ir::PTXS16 d, a = operandAsS16(threadID, instr.a), b = operandAsS16(threadID, instr.b); - if(b == 0) { - throw RuntimeException("Modulus by zero at: " - + kernel->location(context.PC), context.PC, instr); - } - d = a % b; + d = CTARemainder(a, b); setRegAsS16(threadID, instr.d.reg, d); } } @@ -6532,11 +7798,7 @@ void executive::CooperativeThreadArray::eval_Rem(CTAContext &context, ir::PTXS32 d, a = operandAsS32(threadID, instr.a), b = operandAsS32(threadID, instr.b); - if(b == 0) { - throw RuntimeException("Modulus by zero at: " - + kernel->location(context.PC), context.PC, instr); - } - d = a % b; + d = CTARemainder(a, b); setRegAsS32(threadID, instr.d.reg, d); } } @@ -6546,11 +7808,7 @@ void executive::CooperativeThreadArray::eval_Rem(CTAContext &context, ir::PTXS64 d, a = operandAsS64(threadID, instr.a), b = operandAsS64(threadID, instr.b); - if(b == 0) { - throw RuntimeException("Modulus by zero at: " - + kernel->location(context.PC), context.PC, instr); - } - d = a % b; + d = CTARemainder(a, b); setRegAsS64(threadID, instr.d.reg, d); } } @@ -6560,11 +7818,7 @@ void executive::CooperativeThreadArray::eval_Rem(CTAContext &context, ir::PTXU16 d, a = operandAsU16(threadID, instr.a), b = operandAsU16(threadID, instr.b); - if(b == 0) { - throw RuntimeException("Modulus by zero at: " - + kernel->location(context.PC), context.PC, instr); - } - d = a % b; + d = CTARemainder(a, b); setRegAsU16(threadID, instr.d.reg, d); } } @@ -6574,11 +7828,7 @@ void executive::CooperativeThreadArray::eval_Rem(CTAContext &context, ir::PTXU32 d, a = operandAsU32(threadID, instr.a), b = operandAsU32(threadID, instr.b); - if(b == 0) { - throw RuntimeException("Modulus by zero at: " - + kernel->location(context.PC), context.PC, instr); - } - d = a % b; + d = CTARemainder(a, b); setRegAsU32(threadID, instr.d.reg, d); } } @@ -6588,11 +7838,7 @@ void executive::CooperativeThreadArray::eval_Rem(CTAContext &context, ir::PTXU64 d, a = operandAsU64(threadID, instr.a), b = operandAsU64(threadID, instr.b); - if(b == 0) { - throw RuntimeException("Modulus by zero at: " - + kernel->location(context.PC), context.PC, instr); - } - d = a % b; + d = CTARemainder(a, b); setRegAsU64(threadID, instr.d.reg, d); } } @@ -6689,7 +7935,8 @@ void executive::CooperativeThreadArray::eval_Rsqrt(CTAContext &context, if (!context.predicated(threadID, instr)) continue; ir::PTXF32 d, a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); - d = ftz(instr.modifier, 1.0f/(ir::PTXF32)std::sqrt(a)); + d = ftz(instr.modifier, roundedDiv(1.0f, + roundedSqrt(a, ir::PTXInstruction::rn), ir::PTXInstruction::rn)); setRegAsF32(threadID, instr.d.reg, d); } } @@ -6698,7 +7945,25 @@ void executive::CooperativeThreadArray::eval_Rsqrt(CTAContext &context, if (!context.predicated(threadID, instr)) continue; ir::PTXF64 d, a = operandAsF64(threadID, instr.a); - d = 1.0/sqrt(a); + if (instr.modifier & ir::PTXInstruction::ftz) { + const ir::PTXU64 high = 0xffffffff00000000ull; + if (hydrazine::isnan(a)) { + d = hydrazine::bit_cast(0x7fffffff00000000ull); + } else { + a = hydrazine::bit_cast( + hydrazine::bit_cast(a) & high); + if (issubnormal_(a)) a = hydrazine::copysign(0.0, a); + d = roundedDiv(1.0, + roundedSqrt(a, ir::PTXInstruction::rn), + ir::PTXInstruction::rn); + if (issubnormal_(d)) d = hydrazine::copysign(0.0, d); + d = hydrazine::bit_cast( + hydrazine::bit_cast(d) & high); + } + } else { + d = roundedDiv(1.0, roundedSqrt(a, + ir::PTXInstruction::rn), ir::PTXInstruction::rn); + } setRegAsF64(threadID, instr.d.reg, d); } } @@ -7021,8 +8286,10 @@ void executive::CooperativeThreadArray::eval_SetP(CTAContext &context, << " condition = " << t << ", input = " << c << " " << instr.d.identifier << " = " << p << ", q = " << q ); - setRegAsPredicate(threadID, instr.d.reg, p); - if (instr.pq.addressMode != ir::PTXOperand::Invalid) { + if (instr.d.addressMode != ir::PTXOperand::BitBucket) + setRegAsPredicate(threadID, instr.d.reg, p); + if (instr.pq.addressMode != ir::PTXOperand::Invalid + && instr.pq.addressMode != ir::PTXOperand::BitBucket) { setRegAsPredicate(threadID, instr.pq.reg, q); } } @@ -7124,120 +8391,54 @@ void executive::CooperativeThreadArray::eval_SetP(CTAContext &context, << " condition = " << t << ", input = " << c << " " << instr.d.identifier << " = " << p << ", q = " << q ); - setRegAsPredicate(threadID, instr.d.reg, p); - if (instr.pq.addressMode != ir::PTXOperand::Invalid) { + if (instr.d.addressMode != ir::PTXOperand::BitBucket) + setRegAsPredicate(threadID, instr.d.reg, p); + if (instr.pq.addressMode != ir::PTXOperand::Invalid + && instr.pq.addressMode != ir::PTXOperand::BitBucket) { setRegAsPredicate(threadID, instr.pq.reg, q); } } } break; - // single-precision float + // floating-point types [widened to double after type-specific FTZ] + case ir::PTXOperand::f16: case ir::PTXOperand::f32: + case ir::PTXOperand::f64: { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - ir::PTXF32 a = ftz(instr.modifier, operandAsF32(threadID, instr.a)), - b = ftz(instr.modifier, operandAsF32(threadID, instr.b)); - bool c = true; // read operator somehow - bool t = false; - - if (instr.c.addressMode == ir::PTXOperand::Register) { - c = operandAsPredicate(threadID, instr.c); - } - - // any branch predictor worth its salt will get this wrong twice or less - switch (instr.comparisonOperator) { - case ir::PTXInstruction::Equ: - case ir::PTXInstruction::Eq: - t = (a == b); - break; - case ir::PTXInstruction::Neu: - case ir::PTXInstruction::Ne: - t = (a != b); - break; - - case ir::PTXInstruction::Ltu: - case ir::PTXInstruction::Lo: // fall through - case ir::PTXInstruction::Lt: - t = (a < b); - break; - - case ir::PTXInstruction::Leu: - case ir::PTXInstruction::Ls: // fall through - case ir::PTXInstruction::Le: - t = (a <= b); - break; - - case ir::PTXInstruction::Gtu: - case ir::PTXInstruction::Hi: // fall through - case ir::PTXInstruction::Gt: - t = (a > b); - break; - - case ir::PTXInstruction::Geu: - case ir::PTXInstruction::Hs: // fall through - case ir::PTXInstruction::Ge: - t = (a >= b); - break; - - case ir::PTXInstruction::Num: - t = !hydrazine::isnan(a) && !hydrazine::isnan(b); - break; - case ir::PTXInstruction::Nan: - t = hydrazine::isnan(a) || hydrazine::isnan(b); - break; - - default: - throw RuntimeException("invalid comparison operator " - "for unsigned int type", context.PC, instr); - } + ir::PTXF64 a; + ir::PTXF64 b; - // now apply the bool op - bool p = false, q = false; - switch (instr.booleanOperator) { - case ir::PTXInstruction::BoolAnd: - p = (t && c); - q = (!t && c); - break; - case ir::PTXInstruction::BoolOr: - p = (t || c); - q = (!t || c); + switch (instr.type) { + case ir::PTXOperand::f16: + a = f16ToF32(ftzF16(instr.modifier, operandAsU16(threadID, instr.a))); + b = f16ToF32(ftzF16(instr.modifier, operandAsU16(threadID, instr.b))); break; - case ir::PTXInstruction::BoolXor: - p = (t && !c) || (!t && c); - q = (!t && !c) || (t && c); + case ir::PTXOperand::f32: + a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); + b = ftz(instr.modifier, operandAsF32(threadID, instr.b)); break; default: - p = t; - q = !t; + a = operandAsF64(threadID, instr.a); + b = operandAsF64(threadID, instr.b); break; } - reportE(REPORT_SETP, " " << instr.a.identifier << " = " << a - << ", " << instr.b.identifier << " = " << b - << " condition = " << t << ", input = " << c << " " - << instr.d.identifier << " = " << p << ", q = " << q ); - - setRegAsPredicate(threadID, instr.d.reg, p); - if (instr.pq.addressMode != ir::PTXOperand::Invalid) { - setRegAsPredicate(threadID, instr.pq.reg, q); - } - } - } - break; + bool c = true; // read operator somehow + bool t = false; - // double-precision float - case ir::PTXOperand::f64: - { - for (int threadID = 0; threadID < threadCount; threadID++) { - if (!context.predicated(threadID, instr)) continue; + const bool hasNaN = hydrazine::isnan(a) || hydrazine::isnan(b); - ir::PTXF64 a = operandAsF64(threadID, instr.a), - b = operandAsF64(threadID, instr.b); - bool c = true; - bool t = false; + const bool unorderedOp = + instr.comparisonOperator == ir::PTXInstruction::Equ || + instr.comparisonOperator == ir::PTXInstruction::Neu || + instr.comparisonOperator == ir::PTXInstruction::Ltu || + instr.comparisonOperator == ir::PTXInstruction::Leu || + instr.comparisonOperator == ir::PTXInstruction::Gtu || + instr.comparisonOperator == ir::PTXInstruction::Geu; if (instr.c.addressMode == ir::PTXOperand::Register) { c = operandAsPredicate(threadID, instr.c); @@ -7245,39 +8446,37 @@ void executive::CooperativeThreadArray::eval_SetP(CTAContext &context, // any branch predictor worth its salt will get this wrong twice or less switch (instr.comparisonOperator) { - case ir::PTXInstruction::Equ: case ir::PTXInstruction::Eq: - t = (a == b); + t = hasNaN ? unorderedOp : (a == b); break; - case ir::PTXInstruction::Neu: case ir::PTXInstruction::Ne: - t = (a != b); + t = hasNaN ? unorderedOp : (a != b); break; case ir::PTXInstruction::Ltu: case ir::PTXInstruction::Lo: // fall through case ir::PTXInstruction::Lt: - t = (a < b); + t = hasNaN ? unorderedOp : (a < b); break; case ir::PTXInstruction::Leu: case ir::PTXInstruction::Ls: // fall through case ir::PTXInstruction::Le: - t = (a <= b); + t = hasNaN ? unorderedOp : (a <= b); break; case ir::PTXInstruction::Gtu: case ir::PTXInstruction::Hi: // fall through case ir::PTXInstruction::Gt: - t = (a > b); + t = hasNaN ? unorderedOp : (a > b); break; case ir::PTXInstruction::Geu: case ir::PTXInstruction::Hs: // fall through case ir::PTXInstruction::Ge: - t = (a >= b); + t = hasNaN ? unorderedOp : (a >= b); break; case ir::PTXInstruction::Num: @@ -7292,7 +8491,6 @@ void executive::CooperativeThreadArray::eval_SetP(CTAContext &context, "for unsigned int type", context.PC, instr); } - // now apply the bool op bool p = false, q = false; switch (instr.booleanOperator) { @@ -7319,8 +8517,10 @@ void executive::CooperativeThreadArray::eval_SetP(CTAContext &context, << " condition = " << t << ", input = " << c << " " << instr.d.identifier << " = " << p << ", q = " << q ); - setRegAsPredicate(threadID, instr.d.reg, p); - if (instr.pq.addressMode != ir::PTXOperand::Invalid) { + if (instr.d.addressMode != ir::PTXOperand::BitBucket) + setRegAsPredicate(threadID, instr.d.reg, p); + if (instr.pq.addressMode != ir::PTXOperand::Invalid + && instr.pq.addressMode != ir::PTXOperand::BitBucket) { setRegAsPredicate(threadID, instr.pq.reg, q); } } @@ -7550,14 +8750,23 @@ void executive::CooperativeThreadArray::eval_Set(CTAContext &context, } break; - // single-precision float + // floating-point types, with f16 widened before comparison + case ir::PTXOperand::f16: case ir::PTXOperand::f32: { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - ir::PTXF32 a = ftz(instr.modifier, operandAsF32(threadID, instr.a)), + ir::PTXF32 a, b; + if (instr.a.type == ir::PTXOperand::f16) { + a = f16ToF32(ftzF16(instr.modifier, + operandAsU16(threadID, instr.a))); + b = f16ToF32(ftzF16(instr.modifier, + operandAsU16(threadID, instr.b))); + } else { + a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); b = ftz(instr.modifier, operandAsF32(threadID, instr.b)); + } bool c = true; // read operator somehow bool t = false; @@ -7642,12 +8851,15 @@ void executive::CooperativeThreadArray::eval_Set(CTAContext &context, break; } - switch (instr.type) { - case ir::PTXOperand::s32: - case ir::PTXOperand::u32: - setRegAsU32(threadID, instr.d.reg, (t ? 0xFFFFFFFF : 0x00)); - break; - case ir::PTXOperand::f32: + switch (instr.type) { + case ir::PTXOperand::s32: + case ir::PTXOperand::u32: + setRegAsU32(threadID, instr.d.reg, (t ? 0xFFFFFFFF : 0x00)); + break; + case ir::PTXOperand::f16: + setRegAsB16(threadID, instr.d.reg, t ? 0x3c00 : 0x0000); + break; + case ir::PTXOperand::f32: setRegAsF32(threadID, instr.d.reg, (t ? 1.0f : 0.0f)); break; default: @@ -7681,9 +8893,12 @@ void executive::CooperativeThreadArray::eval_Set(CTAContext &context, break; case ir::PTXInstruction::Neu: - case ir::PTXInstruction::Ne: t = (a != b); break; + case ir::PTXInstruction::Ne: + t = !hydrazine::isnan(a) && !hydrazine::isnan(b) + && (a != b); + break; case ir::PTXInstruction::Ltu: case ir::PTXInstruction::Lo: // fall through @@ -7728,7 +8943,6 @@ void executive::CooperativeThreadArray::eval_Set(CTAContext &context, case ir::PTXInstruction::Leu: case ir::PTXInstruction::Gtu: case ir::PTXInstruction::Geu: - case ir::PTXInstruction::Num: case ir::PTXInstruction::Nan: // if either is NaN, set t to true t = (hydrazine::isnan(a) || hydrazine::isnan(b) || t); @@ -7817,10 +9031,11 @@ void executive::CooperativeThreadArray::eval_Shf(CTAContext &context, const ir:: throw RuntimeException("unsupported shift mode", context.PC, instr); } + const ir::PTXB64 pair = (static_cast(b) << 32) | a; if (instr.shiftDirection == ir::PTXInstruction::ShiftLeft) { - d = (b << n) | (a >> (32 - n)); + d = (pair << n) >> 32; } else if (instr.shiftDirection == ir::PTXInstruction::ShiftRight) { - d = (b << (32 - n)) | (a >> n); + d = pair >> n; } else { throw RuntimeException("unsupported shift direction", context.PC, instr); } @@ -7861,11 +9076,7 @@ void executive::CooperativeThreadArray::eval_Shl(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 16 ) - { - b = 16; - } - d = a << b; + d = b >= 16 ? 0 : a << b; setRegAsB16(threadID, instr.d.reg, d); } } @@ -7894,11 +9105,7 @@ void executive::CooperativeThreadArray::eval_Shl(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 32 ) - { - b = 32; - } - d = a << b; + d = b >= 32 ? 0 : a << b; setRegAsB32(threadID, instr.d.reg, d); } } @@ -7927,11 +9134,7 @@ void executive::CooperativeThreadArray::eval_Shl(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 64 ) - { - b = 64; - } - d = a << b; + d = b >= 64 ? 0 : a << b; setRegAsB64(threadID, instr.d.reg, d); } } @@ -7972,11 +9175,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 16 ) - { - b = 16; - } - d = a >> b; + d = b >= 16 ? 0 : a >> b; setRegAsB16(threadID, instr.d.reg, d); } } @@ -8005,11 +9204,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 32 ) - { - b = 32; - } - d = a >> b; + d = b >= 32 ? 0 : a >> b; setRegAsB32(threadID, instr.d.reg, d); } } @@ -8038,11 +9233,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 64 ) - { - b = 64; - } - d = a >> b; + d = b >= 64 ? 0 : a >> b; setRegAsB64(threadID, instr.d.reg, d); } } @@ -8071,11 +9262,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 16 ) - { - b = 16; - } - d = a >> b; + d = b >= 16 ? (a < 0 ? -1 : 0) : a >> b; setRegAsS16(threadID, instr.d.reg, d); } } @@ -8104,11 +9291,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 32 ) - { - b = 32; - } - d = a >> b; + d = b >= 32 ? (a < 0 ? -1 : 0) : a >> b; setRegAsS32(threadID, instr.d.reg, d); } } @@ -8137,11 +9320,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 64 ) - { - b = 64; - } - d = a >> b; + d = b >= 64 ? (a < 0 ? -1 : 0) : a >> b; setRegAsS64(threadID, instr.d.reg, d); } } @@ -8170,11 +9349,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 16 ) - { - b = 16; - } - d = a >> b; + d = b >= 16 ? 0 : a >> b; setRegAsU16(threadID, instr.d.reg, d); } } @@ -8203,11 +9378,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 32 ) - { - b = 32; - } - d = a >> b; + d = b >= 32 ? 0 : a >> b; setRegAsU32(threadID, instr.d.reg, d); } } @@ -8236,11 +9407,7 @@ void executive::CooperativeThreadArray::eval_Shr(CTAContext &context, const ir:: throw RuntimeException("unsupported data type", context.PC, instr); } - if( b > 64 ) - { - b = 64; - } - d = a >> b; + d = b >= 64 ? 0 : a >> b; setRegAsU64(threadID, instr.d.reg, d); } } @@ -8259,7 +9426,8 @@ void executive::CooperativeThreadArray::eval_Sin(CTAContext &context, for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; - ir::PTXF32 d, a = operandAsF32(threadID, instr.a); + ir::PTXF32 d, + a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); d = ftz(instr.modifier, (ir::PTXF32)sin(a)); setRegAsF32(threadID, instr.d.reg, d); } @@ -8361,14 +9529,8 @@ void executive::CooperativeThreadArray::eval_Sqrt(CTAContext &context, ir::PTXF32 d, a = ftz(instr.modifier, operandAsF32(threadID, instr.a)); - if(a < 0.0f || hydrazine::isnan(a)) - { - d = std::numeric_limits::signaling_NaN(); - } - else - { - d = std::sqrt(a); - } + d = roundedSqrt(a, instr.modifier & ir::PTXInstruction::approx + ? ir::PTXInstruction::rn : instr.modifier); setRegAsF32(threadID, instr.d.reg, ftz(instr.modifier, d)); } @@ -8378,7 +9540,7 @@ void executive::CooperativeThreadArray::eval_Sqrt(CTAContext &context, if (!context.predicated(threadID, instr)) continue; ir::PTXF64 d, a = operandAsF64(threadID, instr.a); - d = sqrt(a); + d = roundedSqrt(a, instr.modifier); setRegAsF64(threadID, instr.d.reg, d); } } @@ -8750,13 +9912,47 @@ void executive::CooperativeThreadArray::eval_St(CTAContext &context, void executive::CooperativeThreadArray::eval_Sub(CTAContext &context, const ir::PTXInstruction &instr) { trace(); - if (instr.type == ir::PTXOperand::f32) { + if (instr.type == ir::PTXOperand::f16x2) { + const int effectiveModifier = instr.modifier | ir::PTXInstruction::rn; + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + ir::PTXU32 a = operandAsU32(threadID, instr.a); + ir::PTXU32 b = operandAsU32(threadID, instr.b); + ir::PTXU16 al = static_cast(a), ah = a >> 16; + ir::PTXU16 bl = static_cast(b), bh = b >> 16; + ir::PTXU16 dl = ftzF16(effectiveModifier, toF16( + roundedSub(f16ToF32(ftzF16(effectiveModifier, al)), + f16ToF32(ftzF16(effectiveModifier, bl)), effectiveModifier), + effectiveModifier)); + ir::PTXU16 dh = ftzF16(effectiveModifier, toF16( + roundedSub(f16ToF32(ftzF16(effectiveModifier, ah)), + f16ToF32(ftzF16(effectiveModifier, bh)), effectiveModifier), + effectiveModifier)); + setRegAsU32(threadID, instr.d.reg, static_cast(dl) | + (static_cast(dh) << 16)); + } + } + else if (instr.type == ir::PTXOperand::f16) { + const int effectiveModifier = instr.modifier | ir::PTXInstruction::rn; + for (int threadID = 0; threadID < threadCount; threadID++) { + if (!context.predicated(threadID, instr)) continue; + + ir::PTXF32 a = f16ToF32(ftzF16(effectiveModifier, + operandAsU16(threadID, instr.a))); + ir::PTXF32 b = f16ToF32(ftzF16(effectiveModifier, + operandAsU16(threadID, instr.b))); + ir::PTXU16 d = toF16(roundedSub(a, b, effectiveModifier), effectiveModifier); + setRegAsB16(threadID, instr.d.reg, ftzF16(effectiveModifier, d)); + } + } + else if (instr.type == ir::PTXOperand::f32) { for (int threadID = 0; threadID < threadCount; threadID++) { if (!context.predicated(threadID, instr)) continue; ir::PTXF32 d, a = ftz(instr.modifier, operandAsF32(threadID, instr.a)), b = ftz(instr.modifier, operandAsF32(threadID, instr.b)); - d = ftz(instr.modifier, sat(instr.modifier, a - b)); + d = ftz(instr.modifier, sat(instr.modifier, + roundedSub(a, b, instr.modifier))); setRegAsF32(threadID, instr.d.reg, d); } } @@ -8766,7 +9962,7 @@ void executive::CooperativeThreadArray::eval_Sub(CTAContext &context, ir::PTXF64 d, a = operandAsF64(threadID, instr.a), b = operandAsF64(threadID, instr.b); - d = a - b; + d = roundedSub(a, b, instr.modifier); setRegAsF64(threadID, instr.d.reg, d); } } @@ -8963,6 +10159,32 @@ void executive::CooperativeThreadArray::eval_Sust(CTAContext &context, /*! */ +void executive::CooperativeThreadArray::eval_Tanh(CTAContext &context, + const ir::PTXInstruction &instr) { + trace(); + auto tanhF16 = [](ir::PTXU16 a) { + return toF16(std::tanh(f16ToF32(a)), ir::PTXInstruction::rn); + }; + for (int threadID = 0; threadID < threadCount; ++threadID) { + if (!context.predicated(threadID, instr)) continue; + if (instr.type == ir::PTXOperand::f16) { + setRegAsB16(threadID, instr.d.reg, + tanhF16(operandAsU16(threadID, instr.a))); + } + else if (instr.type == ir::PTXOperand::f16x2) { + const ir::PTXU32 a = operandAsU32(threadID, instr.a); + setRegAsU32(threadID, instr.d.reg, tanhF16(a) + | (static_cast(tanhF16(a >> 16)) << 16)); + } + else { + const ir::PTXF32 a = operandAsF32(threadID, instr.a); + const ir::PTXF32 d = std::fpclassify(a) == FP_SUBNORMAL + ? a : std::tanh(a); + setRegAsF32(threadID, instr.d.reg, d); + } + } +} + void executive::CooperativeThreadArray::eval_TestP(CTAContext &context, const ir::PTXInstruction &instr) { trace(); @@ -9051,7 +10273,7 @@ void executive::CooperativeThreadArray::eval_TestP(CTAContext &context, break; case ir::PTXInstruction::SubNormal: { - d = !hydrazine::isnormal(a) && !hydrazine::isnan(a) && !hydrazine::isinf(a); + d = issubnormal_(a); } break; default: assertM(false, "Invalid floating point mode."); @@ -9634,6 +10856,28 @@ void executive::CooperativeThreadArray::eval_Suq(CTAContext &context, const ir:: eval_Txq(context, instr); } +void executive::CooperativeThreadArray::eval_Szext(CTAContext &context, + const ir::PTXInstruction &instr) { + trace(); + for (int threadID = 0; threadID < threadCount; ++threadID) { + if (!context.predicated(threadID, instr)) continue; + + const ir::PTXU32 a = operandAsU32(threadID, instr.a); + const ir::PTXU32 b = operandAsU32(threadID, instr.b); + const ir::PTXU32 b1 = b & 0x1f; + const bool tooLarge = b >= 32 + && instr.shiftMode == ir::PTXInstruction::ShiftMode::Clamp; + const ir::PTXU32 mask = tooLarge ? 0 : (~0u << b1); + const ir::PTXU32 signPos = (b1 - 1) & 0x1f; + const bool signBit = b1 != 0 && !tooLarge + && instr.type == ir::PTXOperand::s32 + && ((a >> signPos) & 1u); + const ir::PTXU32 d = (a & ~mask) | (signBit ? mask : 0); + + setRegAsU32(threadID, instr.d.reg, d); + } +} + void executive::CooperativeThreadArray::eval_Txq(CTAContext &context, const ir::PTXInstruction &instr) { trace(); const ir::Texture& texture = *context.kernel->textures[instr.a.reg]; @@ -9764,16 +11008,6 @@ void executive::CooperativeThreadArray::eval_Xor(CTAContext &context, setRegAsB64(threadID, instr.d.reg, d); } } - else if (instr.type == ir::PTXOperand::pred) { - for (int threadID = 0; threadID < threadCount; threadID++) { - if (!context.predicated(threadID, instr)) continue; - - bool d, a = operandAsPredicate(threadID, instr.a), - b = operandAsPredicate(threadID, instr.b); - d = a ^ b; - setRegAsPredicate(threadID, instr.d.reg, d); - } - } else { throw RuntimeException("unsupported data type", context.PC, instr); } diff --git a/ocelot/src/executive/EmulatedKernel.cpp b/ocelot/src/executive/EmulatedKernel.cpp index 625ffa52e..69c191e5f 100644 --- a/ocelot/src/executive/EmulatedKernel.cpp +++ b/ocelot/src/executive/EmulatedKernel.cpp @@ -265,7 +265,7 @@ void executive::EmulatedKernel::constructInstructionSequence() { After emitting the instruction sequence, visit each memory move operation and replace references to parameters with offsets into parameter memory. - Data movement instructions: ld, st + Address-bearing instructions: ld, st, cvta */ void executive::EmulatedKernel::updateParamReferences() { using namespace std; @@ -274,7 +274,9 @@ void executive::EmulatedKernel::updateParamReferences() { i_it != instructions.end(); ++i_it) { ir::PTXInstruction& instr = *i_it; if (instr.addressSpace == ir::PTXInstruction::Param) { - if (instr.opcode == ir::PTXInstruction::Ld + if ((instr.opcode == ir::PTXInstruction::Ld + || (instr.opcode == ir::PTXInstruction::Cvta + && !instr.toAddrSpace)) && instr.a.addressMode == ir::PTXOperand::Address) { ir::Parameter *pParam = getParameter(instr.a.identifier); diff --git a/ocelot/src/executive/EmulatedKernelScheduler.cpp b/ocelot/src/executive/EmulatedKernelScheduler.cpp index d61a6d9aa..507814126 100644 --- a/ocelot/src/executive/EmulatedKernelScheduler.cpp +++ b/ocelot/src/executive/EmulatedKernelScheduler.cpp @@ -81,6 +81,11 @@ ir::PTXU64 EmulatedKernelScheduler::argumentMemory() const return (ir::PTXU64)_getExecutingContext()->argumentMemory.data(); } +ir::PTXU64 EmulatedKernelScheduler::argumentMemorySize() const +{ + return _getExecutingContext()->argumentMemory.size(); +} + void EmulatedKernelScheduler::_scheduler() { while(!_executingContexts.empty()) @@ -271,4 +276,3 @@ void EmulatedKernelScheduler::Context::_yieldBarrier() } - diff --git a/ocelot/src/executive/EmulatorCallStack.cpp b/ocelot/src/executive/EmulatorCallStack.cpp index bcf36d49a..286a30a6c 100644 --- a/ocelot/src/executive/EmulatorCallStack.cpp +++ b/ocelot/src/executive/EmulatorCallStack.cpp @@ -142,6 +142,31 @@ namespace executive return (RegisterType*)_stackBase(_localMemoryBase + thread * localMemorySize()); } + + bool EmulatorCallStack::isLocalMemoryAddress(unsigned long long address, + unsigned int thread, unsigned int addressBits) const + { + assert(thread < _threadCount); + const unsigned long long mask = addressBits == 32 + ? 0xffffffffull : ~0ull; + address &= mask; + unsigned int stackPointer = 0; + for (unsigned int frame = 0; frame < _localMemorySizes.size(); ++frame) + { + const unsigned int localSize = _localMemorySizes[frame]; + const unsigned int localBase = stackPointer + + align(3 * sizeof(unsigned int)) + + _stackFrameSizes[frame + 1] * _threadCount + + thread * localSize; + const unsigned long long base = + reinterpret_cast(_stackBase(localBase)) & mask; + if (address >= base && address - base < localSize) return true; + stackPointer += (_stackFrameSizes[frame + 1] + + _registerFileSizes[frame] * sizeof(RegisterType) + localSize) + * _threadCount + align(3 * sizeof(unsigned int)); + } + return false; + } void* EmulatorCallStack::sharedMemoryPointer() { diff --git a/ocelot/src/executive/test/TestInstructions.cpp b/ocelot/src/executive/test/TestInstructions.cpp index ce779bb88..bce800131 100644 --- a/ocelot/src/executive/test/TestInstructions.cpp +++ b/ocelot/src/executive/test/TestInstructions.cpp @@ -9,16 +9,20 @@ #include #include +#include #include +#include #include #include #include +#include #include #include #include #include +#include using namespace std; using namespace ir; @@ -26,6 +30,17 @@ using namespace executive; namespace test { +class ConstMemoryTestKernel : public EmulatedKernel { +public: + ConstMemoryTestKernel(ir::IRKernel* kernel) + : EmulatedKernel(kernel, 0, false) {} + void setConstMemory(unsigned int size) { + delete[] ConstMemory; + ConstMemory = new char[size]; + _constMemorySize = size; + } +}; + class TestInstructions: public Test { public: int threadCount; @@ -42,7 +57,7 @@ class TestInstructions: public Test { status << "Test output:\n"; - threadCount = 16; + threadCount = 32; const std::string ptx = "TestInstructions_ptx"; @@ -313,6 +328,41 @@ class TestInstructions: public Test { } } + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_abs() { .reg .b16 d16, a16; " + << ".reg .b32 d32, a32; abs.f16 d16, a16; " + << "abs.ftz.f16x2 d32, a32; abs.bf16 d16, a16; " + << "abs.bf16x2 d32, a32; ret; }\n"; + try { Module parsed; parsed.load(ptx); } + catch (const std::exception& error) { + status << "half abs parse failed: " << error.what() << "\n"; + return false; + } + auto abs16 = [&](PTXOperand::DataType type, int modifier, + PTXU16 a, PTXU16 expected) { + ins.opcode = PTXInstruction::Abs; ins.type = type; ins.modifier = modifier; + ins.a = reg("a", PTXOperand::b16, 1); ins.d = reg("d", PTXOperand::b16, 0); + cta->setRegAsU16(0, 1, a); cta->eval_Abs(cta->getActiveContext(), ins); + return cta->getRegAsU16(0, 0) == expected; + }; + result = result && abs16(PTXOperand::f16, 0, 0xbe00, 0x3e00); + result = result && abs16(PTXOperand::f16, 0, 0x8001, 0x0001); + result = result && abs16(PTXOperand::f16, PTXInstruction::ftz, 0x8001, 0); + result = result && abs16(PTXOperand::bf16, 0, 0xbfc0, 0x3fc0); + result = result && abs16(PTXOperand::bf16, 0, 0x8001, 0x0001); + auto absPacked = [&](PTXOperand::DataType type, int modifier, + PTXU32 a, PTXU32 expected) { + ins.type = type; ins.modifier = modifier; + ins.a = reg("a", PTXOperand::b32, 1); ins.d = reg("d", PTXOperand::b32, 0); + cta->setRegAsU32(0, 1, a); cta->eval_Abs(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, 0) == expected; + }; + result = result && absPacked(PTXOperand::f16x2, 0, 0xbe003e00, 0x3e003e00); + result = result && absPacked(PTXOperand::f16x2, 0, 0x80018001, 0x00010001); + result = result && absPacked(PTXOperand::f16x2, PTXInstruction::ftz, 0x80018001, 0); + result = result && absPacked(PTXOperand::bf16x2, 0, 0xbfc03fc0, 0x3fc03fc0); + result = result && absPacked(PTXOperand::bf16x2, 0, 0x80010000, 0x00010000); status << "Abs test passed.\n"; return result; @@ -323,6 +373,121 @@ class TestInstructions: public Test { PTXInstruction ins; + // f16 + // + if (result) { + ins.opcode = PTXInstruction::Add; + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.b = reg("r2", PTXOperand::b16, 1); + ins.d = reg("r3", PTXOperand::b16, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, 0x3e00); // 1.5 + cta->setRegAsU16(i, 1, 0x4000); // 2.0 + cta->setRegAsU16(i, 2, 0); + } + if (!ins.valid().empty()) { + result = false; + status << "add.f16 rejected\n"; + } + else { + cta->eval_Add(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 2) != 0x4300) { // 3.5 + result = false; + status << "add.f16 incorrect\n"; + break; + } + } + } + } + if (result) { + ins.modifier = 0; // PTX defaults to round-to-nearest-even. + cta->setRegAsU16(0, 0, 0x3c00); // 1.0 + cta->setRegAsU16(0, 1, 0x1000); // half an f16 ULP at 1.0 + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + cta->eval_Add(cta->getActiveContext(), ins); + const bool roundedNearest = cta->getRegAsU16(0, 2) == 0x3c00; + const bool restored = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(previous); + if (!roundedNearest || !restored) { + status << "add.f16 default rounding failed\n"; + result = false; + } + } + ins.modifier = PTXInstruction::rz; + if (ins.valid().empty()) { + status << "add.rz.f16 accepted\n"; + result = false; + } + ins.modifier = PTXInstruction::ftz | PTXInstruction::sat; + if (!ins.valid().empty()) { + status << "add.ftz.sat.f16 rejected\n"; + result = false; + } + + // f16x2 + // + if (result) { + ins.type = PTXOperand::f16x2; + ins.modifier = 0; + ins.a = reg("r1", PTXOperand::b32, 0); + ins.b = reg("r2", PTXOperand::b32, 1); + ins.d = reg("r3", PTXOperand::b32, 2); + if (!ins.valid().empty()) result = false; + ins.modifier = PTXInstruction::rp; + if (ins.valid().empty()) { + status << "add.rp.f16x2 accepted\n"; + result = false; + } + cta->reset(); + auto packedAdd = [&](int modifier, PTXU32 a, PTXU32 b, + PTXU32 expected, bool alias) { + ins.modifier = modifier; + ins.d.reg = alias ? 0 : 2; + cta->setRegAsU32(0, 0, a); + cta->setRegAsU32(0, 1, b); + cta->setRegAsU32(0, 2, 0xdeadbeef); + cta->eval_Add(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, ins.d.reg) == expected; + }; + result = result && packedAdd(0, 0x40003c00, 0x3c004000, 0x42004200, false); + result = result && packedAdd(0, 0x3c013c00, 0x10001000, 0x3c023c00, false); + result = result && packedAdd(PTXInstruction::rn, 0x3c013c00, + 0x10001000, 0x3c023c00, false); + result = result && packedAdd(PTXInstruction::ftz, 0x80010001, + 0x80000000, 0x80000000, false); + result = result && packedAdd(PTXInstruction::sat, 0x7e004000, + 0x00000000, 0x00003c00, false); + result = result && packedAdd(0, 0x40003c00, 0x3c004000, + 0x42004200, true); + if (result) { + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + const bool defaultRn = packedAdd(0, 0x3c013c00, + 0x10001000, 0x3c023c00, false); + const bool restoredUpward = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(FE_DOWNWARD); + const bool cancellation = packedAdd(0, 0xbc003c00, + 0x3c00bc00, 0, false); + const bool restoredDownward = hydrazine::fegetround() == FE_DOWNWARD; + hydrazine::fesetround(previous); + if (!defaultRn || !cancellation || !restoredUpward || !restoredDownward) result = false; + } + ins.modifier = 0; + ins.d.reg = 2; + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsU32(0, 2, 0xcafebabe); + cta->eval_Add(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 2) != 0xcafebabe) result = false; + ins.pg.condition = PTXOperand::PT; + } + // u16 // if (result) { @@ -597,6 +762,129 @@ class TestInstructions: public Test { PTXInstruction ins; ins.opcode = PTXInstruction::Sub; + // f16 + // + if (result) { + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.b = reg("r2", PTXOperand::b16, 1); + ins.d = reg("r3", PTXOperand::b16, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, 0x4000); // 2.0 + cta->setRegAsU16(i, 1, 0x3e00); // 1.5 + cta->setRegAsU16(i, 2, 0); + } + if (!ins.valid().empty()) { + result = false; + status << "sub.f16 rejected\n"; + } + else { + cta->eval_Sub(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 2) != 0x3800) { // 0.5 + result = false; + status << "sub.f16 incorrect\n"; + break; + } + } + } + } + if (result) { + ins.type = PTXOperand::f16; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.b = reg("r2", PTXOperand::b16, 1); + ins.d = reg("r3", PTXOperand::b16, 2); + ins.modifier = 0; // PTX defaults to round-to-nearest-even. + cta->setRegAsU16(0, 0, 0x3c00); // 1.0 + cta->setRegAsU16(0, 1, 0x9000); // -(half an f16 ULP at 1.0) + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + cta->eval_Sub(cta->getActiveContext(), ins); + const bool roundedNearest = cta->getRegAsU16(0, 2) == 0x3c00; + const bool restored = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(previous); + if (!roundedNearest || !restored) { + status << "sub.f16 default rounding failed\n"; + result = false; + } + } + ins.modifier = PTXInstruction::rz; + if (ins.valid().empty()) { + status << "sub.rz.f16 accepted\n"; + result = false; + } + ins.modifier = PTXInstruction::ftz | PTXInstruction::sat; + if (!ins.valid().empty()) { + status << "sub.ftz.sat.f16 rejected\n"; + result = false; + } + + // f16x2 + // + if (result) { + ins.type = PTXOperand::f16x2; + ins.modifier = 0; + ins.a = reg("r1", PTXOperand::b32, 0); + ins.b = reg("r2", PTXOperand::b32, 1); + ins.d = reg("r3", PTXOperand::b32, 2); + if (!ins.valid().empty()) result = false; + ins.modifier = PTXInstruction::rp; + if (ins.valid().empty()) { + status << "sub.rp.f16x2 accepted\n"; + result = false; + } + cta->reset(); + auto packedSub = [&](int modifier, PTXU32 a, PTXU32 b, + PTXU32 expected, bool alias) { + ins.modifier = modifier; + ins.d.reg = alias ? 0 : 2; + cta->setRegAsU32(0, 0, a); + cta->setRegAsU32(0, 1, b); + cta->setRegAsU32(0, 2, 0xdeadbeef); + cta->eval_Sub(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, ins.d.reg) == expected; + }; + result = result && packedSub(0, 0x3e004000, 0x40003e00, 0xb8003800, false); + result = result && packedSub(0, 0x3c013c00, 0x90009000, 0x3c023c00, false); + result = result && packedSub(PTXInstruction::rn, 0x3c013c00, + 0x90009000, 0x3c023c00, false); + result = result && packedSub(PTXInstruction::ftz, 0x80010001, + 0, 0x80000000, false); + result = result && packedSub(PTXInstruction::ftz, 0x84010401, + 0x84000400, 0x80000000, false); + result = result && packedSub(PTXInstruction::sat, 0x40000000, + 0x003c00, 0x3c000000, false); + result = result && packedSub(PTXInstruction::sat, 0x00007e00, + 0, 0, false); + result = result && packedSub(0, 0x3e004000, 0x40003e00, + 0xb8003800, true); + if (result) { + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + const bool defaultRn = packedSub(0, 0x3c013c00, + 0x90009000, 0x3c023c00, false); + const bool restoredUpward = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(FE_DOWNWARD); + const bool cancellation = packedSub(0, 0xbc003c00, + 0xbc003c00, 0, false); + const bool restoredDownward = hydrazine::fegetround() == FE_DOWNWARD; + hydrazine::fesetround(previous); + if (!defaultRn || !cancellation || !restoredUpward || !restoredDownward) + result = false; + } + ins.modifier = 0; + ins.d.reg = 2; + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsU32(0, 2, 0xcafebabe); + cta->eval_Sub(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 2) != 0xcafebabe) result = false; + ins.pg.condition = PTXOperand::PT; + } + // u16 // if (result) { @@ -786,6 +1074,85 @@ class TestInstructions: public Test { return result; } + bool test_AddSubRounding() { + PTXInstruction ins; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.b = reg("r2", PTXOperand::f32, 1); + ins.d = reg("r3", PTXOperand::f32, 2); + const int modes[] = {PTXInstruction::rn, PTXInstruction::rz, PTXInstruction::rm, PTXInstruction::rp}; + const PTXU32 add32[][2] = {{0x3f800000, 0xbf800000}, + {0x3f800000, 0xbf800000}, {0x3f800000, 0xbf800001}, + {0x3f800001, 0xbf800000}}; + const PTXU32 sub32[] = {0x3f800000, 0x3f7fffff, 0x3f7fffff, 0x3f800000}; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + ins.opcode = PTXInstruction::Add; + cta->setRegAsF32(0, 0, 1.0f); + cta->setRegAsF32(0, 1, std::ldexp(1.0f, -24)); + cta->setRegAsF32(1, 0, -1.0f); + cta->setRegAsF32(1, 1, -std::ldexp(1.0f, -24)); + cta->eval_Add(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != add32[i][0] + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != add32[i][1]) { + status << "add.f32 rounding failed\n"; + return false; + } + ins.opcode = PTXInstruction::Sub; + cta->setRegAsF32(0, 0, 1.0f); + cta->setRegAsF32(0, 1, std::ldexp(1.0f, -25)); + cta->eval_Sub(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != sub32[i]) { + status << "sub.f32 rounding failed\n"; + return false; + } + } + ins.type = PTXOperand::f64; + ins.a.type = ins.b.type = ins.d.type = PTXOperand::f64; + const PTXU64 expected64[][3] = { + {0x3ff0000000000000ull, 0xbff0000000000000ull, 0x3ff0000000000000ull}, + {0x3ff0000000000000ull, 0xbff0000000000000ull, 0x3fefffffffffffffull}, + {0x3ff0000000000000ull, 0xbff0000000000001ull, 0x3fefffffffffffffull}, + {0x3ff0000000000001ull, 0xbff0000000000000ull, 0x3ff0000000000000ull}}; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + ins.opcode = PTXInstruction::Add; + cta->setRegAsF64(0, 0, 1.0); + cta->setRegAsF64(0, 1, std::ldexp(1.0, -53)); + cta->setRegAsF64(1, 0, -1.0); + cta->setRegAsF64(1, 1, -std::ldexp(1.0, -53)); + cta->eval_Add(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) != expected64[i][0] + || hydrazine::bit_cast(cta->getRegAsF64(1, 2)) != expected64[i][1]) { + status << "add.f64 rounding failed\n"; + return false; + } + ins.opcode = PTXInstruction::Sub; + cta->setRegAsF64(0, 1, std::ldexp(1.0, -54)); + cta->eval_Sub(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) != expected64[i][2]) { + status << "sub.f64 rounding failed\n"; + return false; + } + } + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + ins.modifier = 0; + ins.opcode = PTXInstruction::Add; + cta->setRegAsF64(0, 0, 1.0); + cta->setRegAsF64(0, 1, std::ldexp(1.0, -53)); + cta->eval_Add(cta->getActiveContext(), ins); + const bool defaultRn = hydrazine::bit_cast( + cta->getRegAsF64(0, 2)) == 0x3ff0000000000000ull; + const bool restored = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(previous); + if (!defaultRn || !restored) { + status << "default rounding or rounding-mode restoration failed\n"; + return false; + } + return true; + } + bool test_SubC() { bool result = true; @@ -889,7 +1256,7 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { PTXU16 a = (i * 2), b = (4 + i), c = 2; PTXU16 expected = c + ((a < b) ? b-a : a-b); - if (cta->getRegAsU16(i, 2) != expected) { + if (cta->getRegAsU16(i, 3) != expected) { result = false; status << "sad.u16 incorrect\n"; break; @@ -916,7 +1283,7 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { PTXU32 a = (i * 2), b = (4 + i), c = 2; PTXU32 expected = c + ((a < b) ? b-a : a-b); - if (cta->getRegAsU32(i, 2) != expected) { + if (cta->getRegAsU32(i, 3) != expected) { result = false; status << "sad.u32 incorrect\n"; break; @@ -943,7 +1310,7 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { PTXU64 a = (i * 2), b = (4 + i), c = 2; PTXU64 expected = c + ((a < b) ? b-a : a-b); - if (cta->getRegAsU64(i, 2) != expected) { + if (cta->getRegAsU64(i, 3) != expected) { result = false; status << "sad.u64 incorrect\n"; break; @@ -961,16 +1328,16 @@ class TestInstructions: public Test { ins.d = reg("r4", PTXOperand::s16, 3); for (int i = 0; i < threadCount; i++) { - cta->setRegAsS16(i, 0, (PTXS16)(i * 2)); + cta->setRegAsS16(i, 0, (PTXS16)(-i * 2)); cta->setRegAsS16(i, 1, (PTXS16)(4 + i)); cta->setRegAsS16(i, 2, 2); cta->setRegAsS16(i, 3, 0); } cta->eval_Sad(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS16 a = (i * 2), b = (4 + i), c = 2; + PTXS16 a = (-i * 2), b = (4 + i), c = 2; PTXS16 expected = c + ((a < b) ? b-a : a-b); - if (cta->getRegAsS16(i, 2) != expected) { + if (cta->getRegAsS16(i, 3) != expected) { result = false; status << "sad.s16 incorrect\n"; break; @@ -988,16 +1355,16 @@ class TestInstructions: public Test { ins.d = reg("r4", PTXOperand::s32, 3); for (int i = 0; i < threadCount; i++) { - cta->setRegAsS32(i, 0, (PTXS32)(i * 2)); + cta->setRegAsS32(i, 0, (PTXS32)(-i * 2)); cta->setRegAsS32(i, 1, (PTXS32)(4 + i)); cta->setRegAsS32(i, 2, 2); cta->setRegAsS32(i, 3, 0); } cta->eval_Sad(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS32 a = (i * 2), b = (4 + i), c = 2; + PTXS32 a = (-i * 2), b = (4 + i), c = 2; PTXS32 expected = c + ((a < b) ? b-a : a-b); - if (cta->getRegAsS32(i, 2) != expected) { + if (cta->getRegAsS32(i, 3) != expected) { result = false; status << "sad.s32 incorrect\n"; break; @@ -1009,22 +1376,22 @@ class TestInstructions: public Test { // if (result) { ins.type = PTXOperand::s64; - ins.a = reg("r1", PTXOperand::u64, 0); - ins.b = reg("r2", PTXOperand::u64, 1); - ins.c = reg("r3", PTXOperand::u64, 2); - ins.d = reg("r4", PTXOperand::u64, 3); + ins.a = reg("r1", PTXOperand::s64, 0); + ins.b = reg("r2", PTXOperand::s64, 1); + ins.c = reg("r3", PTXOperand::s64, 2); + ins.d = reg("r4", PTXOperand::s64, 3); for (int i = 0; i < threadCount; i++) { - cta->setRegAsS64(i, 0, (PTXS64)(i * 2)); + cta->setRegAsS64(i, 0, (PTXS64)(-i * 2)); cta->setRegAsS64(i, 1, (PTXS64)(4 + i)); cta->setRegAsS64(i, 2, 2); cta->setRegAsS64(i, 3, 0); } cta->eval_Sad(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS64 a = (i * 2), b = (4 + i), c = 2; + PTXS64 a = (-i * 2), b = (4 + i), c = 2; PTXS64 expected = c + ((a < b) ? b-a : a-b); - if (cta->getRegAsS64(i, 2) != expected) { + if (cta->getRegAsS64(i, 3) != expected) { result = false; status << "sad.s64 incorrect\n"; break; @@ -1035,219 +1402,980 @@ class TestInstructions: public Test { return result; } -#define argmin(a, b) ((a) > (b) ? (b) : (a)) -#define argmax(a, b) ((b) > (a) ? (b) : (a)) - - bool test_Min() { + bool test_LdLu() { bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_ld_lu() {\n" + << " .reg .f32 d;\n" + << " .reg .u64 a;\n" + << " ld.lu.f32 d, [a];\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse ld.lu example: " << error.what() << "\n"; + return false; + } + bool foundLu = false; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (!ptxIns || ptxIns->opcode != PTXInstruction::Ld) continue; + if (ptxIns->cacheOperation == PTXInstruction::Lu) foundLu = true; + } + if (!foundLu) { + result = false; + status << "ld.lu did not parse with cacheOperation == Lu\n"; + } PTXInstruction ins; - ins.opcode = PTXInstruction::Min; + ins.opcode = PTXInstruction::Ld; + ins.type = PTXOperand::f32; + ins.addressSpace = PTXInstruction::Global; + ins.cacheOperation = PTXInstruction::Lu; + ins.volatility = PTXInstruction::Volatile; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.d = reg("d", PTXOperand::f32, 1); + if (ins.valid().empty()) { + result = false; + status << "ld.lu combined with .volatile was incorrectly accepted\n"; + } - // u16 - // - if (result) { - ins.type = PTXOperand::u16; - ins.a = reg("r1", PTXOperand::u16, 0); - ins.b = reg("r2", PTXOperand::u16, 1); - ins.d = reg("r3", PTXOperand::u16, 2); + return result; + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsU16(i, 0, (PTXU16)(i * 2)); - cta->setRegAsU16(i, 1, (PTXU16)(4 + i)); - cta->setRegAsU16(i, 2, 0); - } - cta->eval_Min(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - PTXU16 expected = argmin(i*2, 4+i); - if (cta->getRegAsU16(i, 2) != expected) { - result = false; - status << "min.u16 incorrect\n"; - break; - } - } + // suld/sust only support .cop = {.ca, .cg, .cs, .cv} per PTX ISA; the + // grammar's cacheOperation production is shared with ld, which also + // accepts .nc and .lu, so opcode-specific validation must reject those + // two on surface ops. + bool test_SuldSustCacheOperator() { + bool result = true; + PTXInstruction ins; + ins.opcode = PTXInstruction::Suld; + ins.formatMode = PTXInstruction::Unformatted; + ins.type = PTXOperand::b32; + if (!ins.valid().empty()) { + result = false; + status << "suld with default cache operator was incorrectly rejected\n"; + } + ins.cacheOperation = PTXInstruction::Cv; + if (!ins.valid().empty()) { + result = false; + status << "suld.cv was incorrectly rejected\n"; + } + ins.cacheOperation = PTXInstruction::Nc; + if (ins.valid().empty()) { + result = false; + status << "suld.nc was incorrectly accepted\n"; + } + ins.cacheOperation = PTXInstruction::Lu; + if (ins.valid().empty()) { + result = false; + status << "suld.lu was incorrectly accepted\n"; } - // u32 - // - if (result) { - ins.type = PTXOperand::u32; - ins.a = reg("r1", PTXOperand::u32, 0); - ins.b = reg("r2", PTXOperand::u32, 1); - ins.d = reg("r3", PTXOperand::u32, 2); + ins.opcode = PTXInstruction::Sust; + ins.cacheOperation = PTXInstruction::Cs; + if (!ins.valid().empty()) { + result = false; + status << "sust.cs was incorrectly rejected\n"; + } + ins.cacheOperation = PTXInstruction::Nc; + if (ins.valid().empty()) { + result = false; + status << "sust.nc was incorrectly accepted\n"; + } + ins.cacheOperation = PTXInstruction::Lu; + if (ins.valid().empty()) { + result = false; + status << "sust.lu was incorrectly accepted\n"; + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsU32(i, 0, (PTXU32)(i * 2)); - cta->setRegAsU32(i, 1, (PTXU32)(4 + i)); - cta->setRegAsU32(i, 2, 0); - } - cta->eval_Min(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - PTXU32 expected = argmin(i*2, 4+i); - if (cta->getRegAsU32(i, 2) != expected) { - result = false; - status << "min.u32 incorrect\n"; - break; - } - } + return result; + } + + // st only supports .cop = {.wb, .cg, .cs, .wt}; it previously shared + // ld's grammar production, which wrongly accepted .ca/.cv/.nc/.lu and + // couldn't parse .wb/.wt at all (tokens declared but never wired into + // any rule). st now has its own storeCacheOperation production. + bool test_StCacheOperator() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_st_cop() {\n" + << " .reg .f32 d;\n" + << " .reg .u64 a;\n" + << " st.wb.f32 [a], d;\n" + << " st.wt.f32 [a], d;\n" + << " st.cg.f32 [a], d;\n" + << " st.cs.f32 [a], d;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse st.wb/.wt/.cg/.cs example: " << error.what() << "\n"; + return false; + } + unsigned int found = 0; + const PTXInstruction::CacheOperation expected[] = { + PTXInstruction::Wb, PTXInstruction::Wt, + PTXInstruction::Cg, PTXInstruction::Cs}; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (!ptxIns || ptxIns->opcode != PTXInstruction::St) continue; + if (found < 4 && ptxIns->cacheOperation == expected[found]) ++found; + } + if (found != 4) { + status << "st.wb/.wt/.cg/.cs did not all parse with the expected cache operator\n"; + return false; } - // u64 - // - if (result) { - ins.type = PTXOperand::u64; - ins.a = reg("r1", PTXOperand::u64, 0); - ins.b = reg("r2", PTXOperand::u64, 1); - ins.d = reg("r3", PTXOperand::u64, 2); + std::stringstream illegal; + illegal << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_st_cop_illegal() {\n" + << " .reg .f32 d;\n" + << " .reg .u64 a;\n" + << " st.ca.f32 [a], d;\n" + << " ret;\n}\n"; + Module illegalModule; + try { + illegalModule.load(illegal); + status << "st.ca was incorrectly accepted by the parser\n"; + return false; + } + catch (const std::exception&) { + // expected: .ca is not a legal st cache operator + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsU64(i, 0, (PTXU64)(i * 2)); - cta->setRegAsU64(i, 1, (PTXU64)(4 + i)); - cta->setRegAsU64(i, 2, 0); - } - cta->eval_Min(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - PTXU64 expected = argmin(i*2, 4+i); - if (cta->getRegAsU64(i, 2) != expected) { - result = false; - status << "min.u64 incorrect\n"; - break; - } - } + return true; + } + + // Bison only attaches a rule's trailing { action } to its LAST + // alternative, not all of them (rule : A | B | C { action } only runs + // action for C). Several grammar productions shared one action across + // multiple alternatives, so only the last-listed modifier actually got + // recorded; every other modifier silently kept the instruction's + // previous/default field value. This checks the first-listed + // (previously-broken) alternative of each affected production parses + // with the correct field, not just the last one. + bool test_GrammarSharedActionFix() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_shared_action() {\n" + << " .reg .f32 d;\n" + << " .reg .u64 a;\n" + << " .reg .pred p, q;\n" + << " ld.cg.f32 d, [a];\n" + << " lop3.and.b32 d|p, d, d, d, 0x3f, q;\n" + << " lop3.or.b32 d|p, d, d, d, 0x3f, q;\n" + << " bar.arrive 0;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse shared-action regression example: " + << error.what() << "\n"; + return false; + } + bool ldCg = false, lop3And = false, lop3Or = false, barArrive = false; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (!ptxIns) continue; + if (ptxIns->opcode == PTXInstruction::Ld + && ptxIns->cacheOperation == PTXInstruction::Cg) ldCg = true; + if (ptxIns->opcode == PTXInstruction::Lop3 + && ptxIns->booleanOperator == PTXInstruction::BoolAnd) lop3And = true; + if (ptxIns->opcode == PTXInstruction::Lop3 + && ptxIns->booleanOperator == PTXInstruction::BoolOr) lop3Or = true; + if (ptxIns->opcode == PTXInstruction::Bar + && ptxIns->barrierOperation == PTXInstruction::BarArrive) barArrive = true; + } + // require BOTH lop3 forms to be seen with their distinct operator, + // not just one, so a coincidental match against uninitialized + // memory can't make this pass for the wrong reason + if (!ldCg || !lop3And || !lop3Or || !barArrive) { + status << "shared-action regression: ld.cg=" << ldCg + << " lop3.and=" << lop3And << " lop3.or=" << lop3Or + << " bar.arrive=" << barArrive << "\n"; + return false; } - // s16 - // - if (result) { - ins.type = PTXOperand::s16; - ins.a = reg("r1", PTXOperand::s16, 0); - ins.b = reg("r2", PTXOperand::s16, 1); - ins.d = reg("r3", PTXOperand::s16, 2); + return true; + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsS16(i, 0, (PTXS16)(i * 2)); - cta->setRegAsS16(i, 1, (PTXS16)(4 + i)); - cta->setRegAsS16(i, 2, 0); - } - cta->eval_Min(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - PTXS16 expected = argmin(i*2, 4+i); - if (cta->getRegAsS16(i, 2) != expected) { - result = false; - status << "min.s16 incorrect\n"; - break; - } - } + // stringToOpcode() had no entries for "prefetch"/"prefetchu" at all, + // so both silently fell through to Nop and failed to parse with + // "NOP is not a valid instruction" -- unrelated to the shared-action + // grammar bug, found while regression-testing it. + bool test_PrefetchOpcode() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_prefetch() {\n" + << " .reg .u64 a;\n" + << " prefetch.global.L1 [a];\n" + << " prefetchu.L1 [a];\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse prefetch/prefetchu example: " + << error.what() << "\n"; + return false; + } + bool prefetch = false, prefetchu = false; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (!ptxIns) continue; + if (ptxIns->opcode == PTXInstruction::Prefetch + && ptxIns->cacheLevel == PTXInstruction::L1 + && ptxIns->addressSpace == PTXInstruction::Global) prefetch = true; + if (ptxIns->opcode == PTXInstruction::Prefetchu + && ptxIns->cacheLevel == PTXInstruction::L1) prefetchu = true; + } + if (!prefetch || !prefetchu) { + status << "prefetch=" << prefetch << " prefetchu=" << prefetchu << "\n"; + return false; } - // s32 - // - if (result) { - ins.type = PTXOperand::s32; - ins.a = reg("r1", PTXOperand::s32, 0); - ins.b = reg("r2", PTXOperand::s32, 1); - ins.d = reg("r3", PTXOperand::s32, 2); + return true; + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsS32(i, 0, (PTXS32)(i * 2)); - cta->setRegAsS32(i, 1, (PTXS32)(4 + i)); - cta->setRegAsS32(i, 2, 0); - } - cta->eval_Min(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - PTXS32 expected = argmin(i*2, 4+i); - if (cta->getRegAsS32(i, 2) != expected) { - result = false; - status << "min.s32 incorrect\n"; - break; + // eval_Red used to unconditionally throw "instruction not implemented". + // red is atom minus the writeback of the pre-update value; it shares + // evalAtomicRMW with eval_Atom but resolves its operation from the + // distinct reductionOperation field/enum rather than atomicOperation. + bool test_Red() { + PTXU32 mem = 10; + PTXInstruction ins; + ins.opcode = PTXInstruction::Red; + ins.type = PTXOperand::u32; + ins.addressSpace = PTXInstruction::Global; + ins.reductionOperation = PTXInstruction::ReductionAdd; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_uint("b", PTXOperand::u32, 5); + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU64(i, 0, (PTXU64)&mem); + } + cta->eval_Red(cta->getActiveContext(), ins); + if (mem != 10 + 5u * threadCount) { + status << "red.global.add.u32 gave " << mem << ", expected " + << (10 + 5u * threadCount) << "\n"; + return false; + } + + return true; + } + + // fence was entirely unimplemented -- no lexer token, no grammar rule, + // no Opcode value. Verifies parsing of the basic thread-fence form + // (sem optional, defaulting to acq_rel) and that execution is a no-op, + // mirroring the already-implemented membar. + bool test_Fence() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_fence() {\n" + << " fence.sc.gpu;\n" + << " fence.acquire.cta;\n" + << " fence.sys;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse fence example: " << error.what() << "\n"; + return false; + } + std::vector fences; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (ptxIns && ptxIns->opcode == PTXInstruction::Fence) fences.push_back(ptxIns); } - } + if (fences.size() != 3) { + status << "expected 3 fence instructions, found " << fences.size() << "\n"; + return false; + } + if (fences[0]->semantics != PTXInstruction::Sc + || fences[0]->level != PTXInstruction::GlobalLevel) { + status << "fence.sc.gpu: semantics=" << fences[0]->semantics + << " level=" << fences[0]->level << "\n"; + return false; + } + if (fences[1]->semantics != PTXInstruction::Acquire + || fences[1]->level != PTXInstruction::CtaLevel) { + status << "fence.acquire.cta: semantics=" << fences[1]->semantics + << " level=" << fences[1]->level << "\n"; + return false; + } + if (fences[2]->semantics != PTXInstruction::AcqRel + || fences[2]->level != PTXInstruction::SystemLevel) { + status << "fence.sys (no .sem, should default to acq_rel): semantics=" + << fences[2]->semantics << " level=" << fences[2]->level << "\n"; + return false; } - // s64 - // - if (result) { - ins.type = PTXOperand::s64; - ins.a = reg("r1", PTXOperand::s64, 0); - ins.b = reg("r2", PTXOperand::s64, 1); - ins.d = reg("r3", PTXOperand::s64, 2); + // trivial no-op check + PTXInstruction ins; + ins.opcode = PTXInstruction::Fence; + try { cta->eval_Fence(cta->getActiveContext(), ins); } + catch (const std::exception& error) { + status << "eval_Fence threw: " << error.what() << "\n"; + return false; + } + + return true; + } + + // atom/red previously could not parse the optional {.sem}{.scope} + // memory-ordering prefix at all. Verifies parsing of all legal atom + // .sem forms (relaxed/acquire/release/acq_rel, no .sc), all legal red + // .sem forms (relaxed/release only), optional .scope, and that .sem + // defaults to .relaxed when omitted, per PTX ISA 9.7.14.5/9.7.14.6. + bool test_AtomicRedSemantics() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_atomic_red_sem() {\n" + << " .reg .u64 a;\n" + << " .reg .u32 b, d;\n" + << " atom.relaxed.gpu.global.add.u32 d, [a], b;\n" + << " atom.acquire.cta.global.add.u32 d, [a], b;\n" + << " atom.global.add.u32 d, [a], b;\n" + << " red.relaxed.global.add.u32 a, b;\n" + << " red.release.gpu.global.add.u32 a, b;\n" + << " red.global.add.u32 a, b;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse atom/red .sem/.scope example: " + << error.what() << "\n"; + return false; + } + std::vector atoms, reds; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (!ptxIns) continue; + if (ptxIns->opcode == PTXInstruction::Atom) atoms.push_back(ptxIns); + if (ptxIns->opcode == PTXInstruction::Red) reds.push_back(ptxIns); + } + if (atoms.size() != 3 || reds.size() != 3) { + status << "expected 3 atom + 3 red instructions, found " + << atoms.size() << " atom, " << reds.size() << " red\n"; + return false; + } + if (atoms[0]->semantics != PTXInstruction::Relaxed + || atoms[0]->scope != PTXInstruction::GlobalLevel) { + status << "atom.relaxed.gpu: semantics=" << atoms[0]->semantics + << " scope=" << atoms[0]->scope << "\n"; + return false; + } + if (atoms[1]->semantics != PTXInstruction::Acquire + || atoms[1]->scope != PTXInstruction::CtaLevel) { + status << "atom.acquire.cta: semantics=" << atoms[1]->semantics + << " scope=" << atoms[1]->scope << "\n"; + return false; + } + if (atoms[2]->semantics != PTXInstruction::Relaxed + || atoms[2]->scope != PTXInstruction::Level_Invalid) { + status << "atom (no .sem/.scope, should default to relaxed): semantics=" + << atoms[2]->semantics << " scope=" << atoms[2]->scope << "\n"; + return false; + } + if (atoms[2]->addressSpace != PTXInstruction::Global) { + status << "atom (no .sem/.scope): addressSpace corrupted, got " + << atoms[2]->addressSpace << "\n"; + return false; + } + if (reds[0]->semantics != PTXInstruction::Relaxed + || reds[0]->scope != PTXInstruction::Level_Invalid) { + status << "red.relaxed: semantics=" << reds[0]->semantics + << " scope=" << reds[0]->scope << "\n"; + return false; + } + if (reds[1]->semantics != PTXInstruction::Release + || reds[1]->scope != PTXInstruction::GlobalLevel) { + status << "red.release.gpu: semantics=" << reds[1]->semantics + << " scope=" << reds[1]->scope << "\n"; + return false; + } + if (reds[2]->semantics != PTXInstruction::Relaxed + || reds[2]->scope != PTXInstruction::Level_Invalid) { + status << "red (no .sem/.scope, should default to relaxed): semantics=" + << reds[2]->semantics << " scope=" << reds[2]->scope << "\n"; + return false; + } + if (reds[2]->addressSpace != PTXInstruction::Global) { + status << "red (no .sem/.scope): addressSpace corrupted, got " + << reds[2]->addressSpace << "\n"; + return false; + } + + return true; + } + + // ld/st .weak / .relaxed.scope / .acquire.scope / .release.scope parsing. + bool test_LdStOrdering() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_ldst_ordering() {\n" + << " .reg .f32 b, d;\n" + << " .reg .u64 a;\n" + << " ld.weak.global.f32 d, [a];\n" + << " ld.relaxed.gpu.global.f32 d, [a];\n" + << " ld.acquire.cta.global.f32 d, [a];\n" + << " ld.global.f32 d, [a];\n" + << " st.weak.global.f32 [a], b;\n" + << " st.relaxed.gpu.global.f32 [a], b;\n" + << " st.release.cta.global.f32 [a], b;\n" + << " st.global.f32 [a], b;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse ld/st ordering example: " + << error.what() << "\n"; + return false; + } + std::vector lds, sts; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (!ptxIns) continue; + if (ptxIns->opcode == PTXInstruction::Ld) lds.push_back(ptxIns); + if (ptxIns->opcode == PTXInstruction::St) sts.push_back(ptxIns); + } + if (lds.size() != 4 || sts.size() != 4) { + status << "expected 4 ld + 4 st instructions, found " + << lds.size() << " ld, " << sts.size() << " st\n"; + return false; + } + if (lds[0]->semantics != PTXInstruction::Weak) { + status << "ld.weak: semantics=" << lds[0]->semantics << "\n"; + return false; + } + if (lds[1]->semantics != PTXInstruction::Relaxed + || lds[1]->scope != PTXInstruction::GlobalLevel) { + status << "ld.relaxed.gpu: semantics=" << lds[1]->semantics + << " scope=" << lds[1]->scope << "\n"; + return false; + } + if (lds[2]->semantics != PTXInstruction::Acquire + || lds[2]->scope != PTXInstruction::CtaLevel) { + status << "ld.acquire.cta: semantics=" << lds[2]->semantics + << " scope=" << lds[2]->scope << "\n"; + return false; + } + if (lds[3]->semantics != PTXInstruction::Weak + || lds[3]->scope != PTXInstruction::Level_Invalid) { + status << "ld (no qualifier, should default to weak): semantics=" + << lds[3]->semantics << " scope=" << lds[3]->scope << "\n"; + return false; + } + if (lds[3]->addressSpace != PTXInstruction::Global) { + status << "ld (no qualifier): addressSpace corrupted, got " + << lds[3]->addressSpace << "\n"; + return false; + } + if (sts[0]->semantics != PTXInstruction::Weak) { + status << "st.weak: semantics=" << sts[0]->semantics << "\n"; + return false; + } + if (sts[1]->semantics != PTXInstruction::Relaxed + || sts[1]->scope != PTXInstruction::GlobalLevel) { + status << "st.relaxed.gpu: semantics=" << sts[1]->semantics + << " scope=" << sts[1]->scope << "\n"; + return false; + } + if (sts[2]->semantics != PTXInstruction::Release + || sts[2]->scope != PTXInstruction::CtaLevel) { + status << "st.release.cta: semantics=" << sts[2]->semantics + << " scope=" << sts[2]->scope << "\n"; + return false; + } + if (sts[3]->semantics != PTXInstruction::Weak + || sts[3]->scope != PTXInstruction::Level_Invalid) { + status << "st (no qualifier, should default to weak): semantics=" + << sts[3]->semantics << " scope=" << sts[3]->scope << "\n"; + return false; + } + if (sts[3]->addressSpace != PTXInstruction::Global) { + status << "st (no qualifier): addressSpace corrupted, got " + << sts[3]->addressSpace << "\n"; + return false; + } + + return true; + } + + bool test_LdStMmio() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_ldst_mmio() {\n" + << " .reg .f32 b, d;\n" + << " .reg .u64 a;\n" + << " ld.mmio.acquire.sys.global.f32 d, [a];\n" + << " ld.mmio.relaxed.sys.f32 d, [a];\n" + << " st.mmio.relaxed.sys.global.f32 [a], b;\n" + << " st.mmio.release.sys.f32 [a], b;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse ld/st mmio example: " + << error.what() << "\n"; + return false; + } + std::vector lds, sts; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptxIns = dynamic_cast(*instruction); + if (!ptxIns) continue; + if (ptxIns->opcode == PTXInstruction::Ld) lds.push_back(ptxIns); + if (ptxIns->opcode == PTXInstruction::St) sts.push_back(ptxIns); + } + if (lds.size() != 2 || sts.size() != 2) { + status << "expected 2 ld + 2 st instructions, found " + << lds.size() << " ld, " << sts.size() << " st\n"; + return false; + } + if (!lds[0]->mmio || lds[0]->semantics != PTXInstruction::Acquire + || lds[0]->scope != PTXInstruction::SystemLevel + || lds[0]->addressSpace != PTXInstruction::Global) { + status << "ld.mmio.acquire.sys.global: mmio=" << lds[0]->mmio + << " semantics=" << lds[0]->semantics + << " scope=" << lds[0]->scope + << " addressSpace=" << lds[0]->addressSpace << "\n"; + return false; + } + if (!lds[1]->mmio || lds[1]->semantics != PTXInstruction::Relaxed + || lds[1]->scope != PTXInstruction::SystemLevel + || lds[1]->addressSpace != PTXInstruction::Global) { + status << "ld.mmio.relaxed.sys (no .global): mmio=" << lds[1]->mmio + << " semantics=" << lds[1]->semantics + << " scope=" << lds[1]->scope + << " addressSpace=" << lds[1]->addressSpace << "\n"; + return false; + } + if (!sts[0]->mmio || sts[0]->semantics != PTXInstruction::Relaxed + || sts[0]->scope != PTXInstruction::SystemLevel + || sts[0]->addressSpace != PTXInstruction::Global) { + status << "st.mmio.relaxed.sys.global: mmio=" << sts[0]->mmio + << " semantics=" << sts[0]->semantics + << " scope=" << sts[0]->scope + << " addressSpace=" << sts[0]->addressSpace << "\n"; + return false; + } + if (!sts[1]->mmio || sts[1]->semantics != PTXInstruction::Release + || sts[1]->scope != PTXInstruction::SystemLevel + || sts[1]->addressSpace != PTXInstruction::Global) { + status << "st.mmio.release.sys (no .global): mmio=" << sts[1]->mmio + << " semantics=" << sts[1]->semantics + << " scope=" << sts[1]->scope + << " addressSpace=" << sts[1]->addressSpace << "\n"; + return false; + } + + // Regression guard: .mmio must round-trip through the printer, not + // collapse into the already-supported ld.acquire.sys form. + std::string printed = lds[0]->toString(); + if (printed.find("mmio") == std::string::npos) { + status << "ld.mmio toString() lost the mmio marker: " << printed << "\n"; + return false; + } + return true; + } + + // Covers the atom/red type/operation table (PTX ISA Table 35/36) after + // correcting f32 min/max, and/or/xor bitness, and inc/dec typing. + bool test_AtomicRedTypeTable() { + bool result = true; + + // atom.and.b64: verify bitwise AND and old-value writeback. + { + PTXU64 mem = 0xFF00FF00FF00FF00ULL; + PTXInstruction ins; + ins.opcode = PTXInstruction::Atom; + ins.type = PTXOperand::b64; + ins.addressSpace = PTXInstruction::Global; + ins.atomicOperation = PTXInstruction::AtomicAnd; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_uint("b", PTXOperand::b64, 0x0F0F0F0F0F0F0F0FULL); + ins.d = reg("d", PTXOperand::b64, 1); for (int i = 0; i < threadCount; i++) { - cta->setRegAsS64(i, 0, (PTXS64)(i * 2)); - cta->setRegAsS64(i, 1, (PTXS64)(4 + i)); - cta->setRegAsS64(i, 2, 0); + cta->setRegAsU64(i, 0, (PTXU64)&mem); } - cta->eval_Min(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - PTXS64 expected = argmin(i*2, 4+i); - if (cta->getRegAsS64(i, 2) != expected) { - result = false; - status << "min.s64 incorrect\n"; - break; - } + cta->eval_Atom(cta->getActiveContext(), ins); + if (mem != 0x0F000F000F000F00ULL) { + status << "atom.and.b64 gave " << std::hex << mem + << ", expected 0x0f000f000f000f00\n" << std::dec; + result = false; + } + if (cta->getRegAsB64(0, 1) != 0xFF00FF00FF00FF00ULL) { + status << "atom.and.b64 old-value writeback incorrect\n"; + result = false; } } - // f32 - // - if (result) { - ins.type = PTXOperand::f32; - ins.a = reg("r1", PTXOperand::f32, 0); - ins.b = reg("r2", PTXOperand::f32, 1); - ins.d = reg("r3", PTXOperand::f32, 2); + // atom.add.f64: verify float64 add and old-value writeback. + { + PTXF64 mem = 10.5; + PTXInstruction ins; + ins.opcode = PTXInstruction::Atom; + ins.type = PTXOperand::f64; + ins.addressSpace = PTXInstruction::Global; + ins.atomicOperation = PTXInstruction::AtomicAdd; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_float("b", PTXOperand::f64, 2.25); + ins.d = reg("d", PTXOperand::f64, 1); + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU64(i, 0, (PTXU64)&mem); + } + cta->eval_Atom(cta->getActiveContext(), ins); + if (mem != 10.5 + 2.25 * threadCount) { + status << "atom.add.f64 gave " << mem << ", expected " + << (10.5 + 2.25 * threadCount) << "\n"; + result = false; + } + if (cta->getRegAsF64(0, 1) != 10.5) { + status << "atom.add.f64 old-value writeback incorrect\n"; + result = false; + } + } + // atom.min.s64: verify signed 64-bit min and old-value writeback. + { + PTXS64 mem = 100; + PTXInstruction ins; + ins.opcode = PTXInstruction::Atom; + ins.type = PTXOperand::s64; + ins.addressSpace = PTXInstruction::Global; + ins.atomicOperation = PTXInstruction::AtomicMin; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_int("b", PTXOperand::s64, -50); + ins.d = reg("d", PTXOperand::s64, 1); for (int i = 0; i < threadCount; i++) { - cta->setRegAsF32(i, 0, (PTXF32)(i * 2)); - cta->setRegAsF32(i, 1, (PTXF32)(4 + i)); - cta->setRegAsF32(i, 2, 0); + cta->setRegAsU64(i, 0, (PTXU64)&mem); } - cta->eval_Min(cta->getActiveContext(), ins); + cta->eval_Atom(cta->getActiveContext(), ins); + if (mem != -50) { + status << "atom.min.s64 gave " << mem << ", expected -50\n"; + result = false; + } + if (cta->getRegAsS64(0, 1) != 100) { + status << "atom.min.s64 old-value writeback incorrect\n"; + result = false; + } + } + + // red.and.b64: verify bitwise AND on a 64-bit value. + { + PTXU64 mem = 0xAAAAAAAAAAAAAAAAULL; + PTXInstruction ins; + ins.opcode = PTXInstruction::Red; + ins.type = PTXOperand::b64; + ins.addressSpace = PTXInstruction::Global; + ins.reductionOperation = PTXInstruction::ReductionAnd; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_uint("b", PTXOperand::b64, 0x0F0F0F0F0F0F0F0FULL); for (int i = 0; i < threadCount; i++) { - PTXF32 expected = argmin(i*2, 4+i); - if (cta->getRegAsF32(i, 2) != expected) { - result = false; - status << "min.f32 incorrect [" << i << "] - expected: " << (float)(i*2+4+i) - << ", got " << cta->getRegAsF32(i, 2) << "\n"; - break; - } + cta->setRegAsU64(i, 0, (PTXU64)&mem); + } + cta->eval_Red(cta->getActiveContext(), ins); + if (mem != 0x0A0A0A0A0A0A0A0AULL) { + status << "red.and.b64 gave " << std::hex << mem + << ", expected 0x0a0a0a0a0a0a0a0a\n" << std::dec; + result = false; } } - // f64 - // - if (result) { + // red.add.f64: verify float64 add. + { + PTXF64 mem = 1.0; + PTXInstruction ins; + ins.opcode = PTXInstruction::Red; ins.type = PTXOperand::f64; - ins.a = reg("r1", PTXOperand::f64, 0); - ins.b = reg("r2", PTXOperand::f64, 1); - ins.d = reg("r3", PTXOperand::f64, 2); - + ins.addressSpace = PTXInstruction::Global; + ins.reductionOperation = PTXInstruction::ReductionAdd; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_float("b", PTXOperand::f64, 0.5); for (int i = 0; i < threadCount; i++) { - cta->setRegAsF64(i, 0, (PTXF64)(i * 2)); - cta->setRegAsF64(i, 1, (PTXF64)(4 + i)); - cta->setRegAsF64(i, 2, 0.0); + cta->setRegAsU64(i, 0, (PTXU64)&mem); } - cta->eval_Min(cta->getActiveContext(), ins); + cta->eval_Red(cta->getActiveContext(), ins); + if (mem != 1.0 + 0.5 * threadCount) { + status << "red.add.f64 gave " << mem << ", expected " + << (1.0 + 0.5 * threadCount) << "\n"; + result = false; + } + } + + // red.max.u64: verify unsigned 64-bit max. + { + PTXU64 mem = 5; + PTXInstruction ins; + ins.opcode = PTXInstruction::Red; + ins.type = PTXOperand::u64; + ins.addressSpace = PTXInstruction::Global; + ins.reductionOperation = PTXInstruction::ReductionMax; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_uint("b", PTXOperand::u64, 100); for (int i = 0; i < threadCount; i++) { - PTXF64 expected = argmin(i*2, 4+i); - if (std::fabs(cta->getRegAsF64(i, 2) - expected) > 0.1) { - result = false; - status << "min.f64 incorrect [" << i << "] - expected: " << expected - << ", got " << cta->getRegAsF64(i, 2) << "\n"; - break; - } + cta->setRegAsU64(i, 0, (PTXU64)&mem); + } + cta->eval_Red(cta->getActiveContext(), ins); + if (mem != 100) { + status << "red.max.u64 gave " << mem << ", expected 100\n"; + result = false; + } + } + + // Negative validation cases: now-illegal type/operation combinations + // must be rejected by PTXInstruction::valid(). + { + PTXInstruction ins; + ins.opcode = PTXInstruction::Atom; + ins.type = PTXOperand::f32; + ins.addressSpace = PTXInstruction::Global; + ins.atomicOperation = PTXInstruction::AtomicMin; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_float("b", PTXOperand::f32, 1.0); + ins.d = reg("d", PTXOperand::f32, 1); + if (ins.valid().empty()) { + status << "atom.min.f32 was wrongly accepted as valid\n"; + result = false; + } + } + { + PTXInstruction ins; + ins.opcode = PTXInstruction::Red; + ins.type = PTXOperand::f32; + ins.addressSpace = PTXInstruction::Global; + ins.reductionOperation = PTXInstruction::ReductionMin; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_float("b", PTXOperand::f32, 1.0); + if (ins.valid().empty()) { + status << "red.min.f32 was wrongly accepted as valid\n"; + result = false; + } + } + { + PTXInstruction ins; + ins.opcode = PTXInstruction::Atom; + ins.type = PTXOperand::s32; + ins.addressSpace = PTXInstruction::Global; + ins.atomicOperation = PTXInstruction::AtomicAnd; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_uint("b", PTXOperand::s32, 1); + ins.d = reg("d", PTXOperand::s32, 1); + if (ins.valid().empty()) { + status << "atom.and.s32 was wrongly accepted as valid\n"; + result = false; + } + } + { + PTXInstruction ins; + ins.opcode = PTXInstruction::Atom; + ins.type = PTXOperand::u64; + ins.addressSpace = PTXInstruction::Global; + ins.atomicOperation = PTXInstruction::AtomicInc; + ins.a = reg("a", PTXOperand::u64, 0); + ins.a.addressMode = PTXOperand::Indirect; + ins.b = imm_uint("b", PTXOperand::u64, 1); + ins.d = reg("d", PTXOperand::u64, 1); + if (ins.valid().empty()) { + status << "atom.inc.u64 was wrongly accepted as valid\n"; + result = false; } } return result; } +#define argmin(a, b) ((a) > (b) ? (b) : (a)) +#define argmax(a, b) ((b) > (a) ? (b) : (a)) - bool test_Max() { + bool test_Min() { bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_minmax() {\n" + << " .reg .f32 d, a, b;\n" + << " .reg .b16 hd, ha, hb;\n" + << " .reg .b32 pd, pa, pb;\n" + << " min.NaN.f32 d, a, b;\n" + << " min.xorsign.abs.f32 d, a, b;\n" + << " min.NaN.f16 hd, ha, hb; min.xorsign.abs.f16 hd, ha, hb; max.NaN.f16 hd, ha, hb; max.xorsign.abs.f16 hd, ha, hb;\n" + << " min.bf16 hd, ha, hb; min.NaN.bf16 hd, ha, hb; min.xorsign.abs.bf16 hd, ha, hb;\n" + << " max.bf16 hd, ha, hb; max.NaN.bf16 hd, ha, hb; max.xorsign.abs.bf16 hd, ha, hb;\n" + << " min.f16x2 pd, pa, pb; min.ftz.f16x2 pd, pa, pb; min.NaN.f16x2 pd, pa, pb; min.xorsign.abs.f16x2 pd, pa, pb;\n" + << " max.f16x2 pd, pa, pb; max.ftz.f16x2 pd, pa, pb; max.NaN.f16x2 pd, pa, pb; max.xorsign.abs.f16x2 pd, pa, pb;\n" + << " min.bf16x2 pd, pa, pb; min.NaN.bf16x2 pd, pa, pb; min.xorsign.abs.bf16x2 pd, pa, pb;\n" + << " max.bf16x2 pd, pa, pb; max.NaN.bf16x2 pd, pa, pb; max.xorsign.abs.bf16x2 pd, pa, pb;\n" + << " max.NaN.f32 d, a, b;\n" + << " max.xorsign.abs.f32 d, a, b;\n ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse min/max modifier examples: " + << error.what() << "\n"; + return false; + } + const int expectedBf16x2[] = {0, PTXInstruction::nan, + PTXInstruction::xorsign | PTXInstruction::abs, 0, PTXInstruction::nan, + PTXInstruction::xorsign | PTXInstruction::abs}; + unsigned int parsedBf16x2 = 0; + for (auto kernel = parsed.kernels().begin(); kernel != parsed.kernels().end(); ++kernel) + for (auto block = kernel->second->cfg()->begin(); block != kernel->second->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* ptx = dynamic_cast(*instruction); + if (!ptx || ptx->type != PTXOperand::bf16x2) continue; + if (parsedBf16x2 >= 6 || ptx->modifier != expectedBf16x2[parsedBf16x2]) result = false; + ++parsedBf16x2; + } + if (parsedBf16x2 != 6) result = false; PTXInstruction ins; - ins.opcode = PTXInstruction::Max; - + ins.opcode = PTXInstruction::Min; + ins.type = PTXOperand::b16; + ins.a = reg("a", PTXOperand::b16, 0); + ins.b = reg("b", PTXOperand::b16, 1); + ins.d = reg("d", PTXOperand::b16, 2); + if (ins.valid().empty()) result = false; + ins.type = PTXOperand::b32; + ins.a = reg("a", PTXOperand::b32, 0); + ins.b = reg("b", PTXOperand::b32, 1); + ins.d = reg("d", PTXOperand::b32, 2); + if (ins.valid().empty()) result = false; + ins.type = PTXOperand::b64; + ins.a = reg("a", PTXOperand::b64, 0); + ins.b = reg("b", PTXOperand::b64, 1); + ins.d = reg("d", PTXOperand::b64, 2); + if (ins.valid().empty()) result = false; + ins.type = PTXOperand::u16; + ins.a = reg("a", PTXOperand::b16, 0); + ins.b = reg("b", PTXOperand::b16, 1); + ins.d = reg("d", PTXOperand::b16, 2); + ins.modifier = PTXInstruction::rn; + if (ins.valid().empty()) result = false; + ins.modifier = PTXInstruction::sat | PTXInstruction::relu; + if (ins.valid().empty()) result = false; + ins.modifier = 0; + ins.carry = PTXInstruction::CC; + if (ins.valid().empty()) result = false; + ins.carry = PTXInstruction::None; + ins.d.type = PTXOperand::b32; + if (ins.valid().empty()) result = false; + const PTXOperand::DataType halfTypes[] = {PTXOperand::f16, PTXOperand::f16x2, + PTXOperand::bf16, PTXOperand::bf16x2}; + const int halfModifiers[] = {PTXInstruction::nan, PTXInstruction::ftz | PTXInstruction::nan, + PTXInstruction::nan, PTXInstruction::xorsign | PTXInstruction::abs}; + for (unsigned int i = 0; i < 4; ++i) { + ins.type = halfTypes[i]; ins.modifier = halfModifiers[i]; + const PTXOperand::DataType container = i % 2 ? PTXOperand::b32 : PTXOperand::b16; + ins.a = reg("a", container, 0); ins.b = reg("b", container, 1); ins.d = reg("d", container, 2); + if (!ins.valid().empty()) result = false; + } + ins.type = PTXOperand::bf16; ins.modifier = PTXInstruction::ftz; + ins.a = reg("a", PTXOperand::b16, 0); ins.b = reg("b", PTXOperand::b16, 1); + ins.d = reg("d", PTXOperand::b16, 2); + if (ins.valid().empty()) result = false; + ins.type = PTXOperand::f16; ins.modifier = 0; + ins.a = reg("a", PTXOperand::b16, 0); ins.b = reg("b", PTXOperand::b16, 1); + ins.d = reg("d", PTXOperand::b16, 2); + cta->reset(); + struct HalfCase { int modifier; PTXU16 a, b, minimum, maximum; }; + const HalfCase cases[] = { + {0,0x3c00,0x4000,0x3c00,0x4000}, {0,0x7c00,0x7c00,0x7c00,0x7c00}, + {0,0,0x8000,0x8000,0}, {0,0x8000,0,0x8000,0}, + {0,0x7e00,0x3c00,0x3c00,0x3c00}, {0,0x3c00,0x7e00,0x3c00,0x3c00}, + {0,0x7e00,0xfe00,0x7fff,0x7fff}, {PTXInstruction::nan,0x7e00,0x3c00,0x7fff,0x7fff}, + {0,0x8001,2,0x8001,2}, {PTXInstruction::ftz,0x8001,2,0x8000,0}, + {PTXInstruction::xorsign|PTXInstruction::abs,0xbc00,0x3e00,0xbc00,0xbe00}, + {PTXInstruction::xorsign|PTXInstruction::abs,0xfe00,0x3c00,0xbc00,0xbc00}, + {PTXInstruction::ftz|PTXInstruction::nan|PTXInstruction::xorsign|PTXInstruction::abs, + 0x8001,0x7e00,0x7fff,0x7fff}}; + for (const HalfCase& c : cases) for (int maximum = 0; maximum < 2; ++maximum) { + ins.opcode = maximum ? PTXInstruction::Max : PTXInstruction::Min; + ins.modifier = c.modifier; cta->setRegAsU16(0, 0, c.a); cta->setRegAsU16(0, 1, c.b); + if (maximum) cta->eval_Max(cta->getActiveContext(), ins); + else cta->eval_Min(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 2) != (maximum ? c.maximum : c.minimum)) result = false; + } + ins.type = PTXOperand::bf16; ins.modifier = 0; + struct Bf16Case { int modifier; PTXU16 a, b, minimum, maximum; }; + const Bf16Case bf16Cases[] = { + {0,0x3f80,0x4000,0x3f80,0x4000}, {0,0x7c01,0x3f80,0x3f80,0x7c01}, + {0,0,0x8000,0x8000,0}, {0,0x8001,1,0x8001,1}, + {0,0x7fc1,0x3f80,0x3f80,0x3f80}, {0,0x7fc1,0xffc2,0x7fff,0x7fff}, + {PTXInstruction::nan,0x7fc1,0x3f80,0x7fff,0x7fff}, + {PTXInstruction::xorsign|PTXInstruction::abs,0xbf80,0x3fc0,0xbf80,0xbfc0}, + {PTXInstruction::xorsign|PTXInstruction::abs,0xffc1,0x3f80,0xbf80,0xbf80}}; + ins.a = reg("a", PTXOperand::b16, 0); ins.b = reg("b", PTXOperand::b16, 1); + ins.d = reg("d", PTXOperand::b16, 2); + for (const Bf16Case& c : bf16Cases) for (int maximum = 0; maximum < 2; ++maximum) { + ins.opcode = maximum ? PTXInstruction::Max : PTXInstruction::Min; + ins.modifier = c.modifier; cta->setRegAsU16(0, 0, c.a); cta->setRegAsU16(0, 1, c.b); + if (maximum) cta->eval_Max(cta->getActiveContext(), ins); + else cta->eval_Min(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 2) != (maximum ? c.maximum : c.minimum)) result = false; + } + struct PackedHalfCase { PTXOperand::DataType type; int modifier; PTXU32 a, b, minimum, maximum; }; + const PackedHalfCase packedCases[] = { + {PTXOperand::f16x2,0,0x40003c00,0x3c004000,0x3c003c00,0x40004000}, + {PTXOperand::f16x2,0,0x7e003c00,0x3c007e00,0x3c003c00,0x3c003c00}, + {PTXOperand::f16x2,PTXInstruction::nan,0x7e003c00,0x3c007e00,0x7fff7fff,0x7fff7fff}, + {PTXOperand::f16x2,0,0x00008000,0x80000000,0x80008000,0x00000000}, + {PTXOperand::f16x2,0,0x80010001,0x00020002,0x80010001,0x00020002}, + {PTXOperand::f16x2,PTXInstruction::ftz,0x80010001,0x00020002,0x80000000,0x00000000}, + {PTXOperand::f16x2,PTXInstruction::xorsign|PTXInstruction::abs,0xbc003c00,0x3e00be00,0xbc00bc00,0xbe00be00}, + {PTXOperand::bf16x2,0,0x40003f80,0x3f804000,0x3f803f80,0x40004000}, + {PTXOperand::bf16x2,0,0x7c013f80,0x3f807c01,0x3f803f80,0x7c017c01}, + {PTXOperand::bf16x2,0,0x7fc13f80,0x3f807fc1,0x3f803f80,0x3f803f80}, + {PTXOperand::bf16x2,PTXInstruction::nan,0x7fc13f80,0x3f807fc1,0x7fff7fff,0x7fff7fff}, + {PTXOperand::bf16x2,0,0x80010000,0x00010001,0x80010000,0x00010001}, + {PTXOperand::bf16x2,PTXInstruction::xorsign|PTXInstruction::abs,0xbf803f80,0x3fc0bfc0,0xbf80bf80,0xbfc0bfc0}}; + ins.a = reg("pa", PTXOperand::b32, 0); ins.b = reg("pb", PTXOperand::b32, 1); + ins.d = reg("pd", PTXOperand::b32, 2); + for (const PackedHalfCase& c : packedCases) for (int maximum = 0; maximum < 2; ++maximum) { + ins.type = c.type; ins.opcode = maximum ? PTXInstruction::Max : PTXInstruction::Min; + ins.modifier = c.modifier; cta->setRegAsU32(0, 0, c.a); cta->setRegAsU32(0, 1, c.b); + if (maximum) cta->eval_Max(cta->getActiveContext(), ins); + else cta->eval_Min(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 2) != (maximum ? c.maximum : c.minimum)) result = false; + } + ins.type = PTXOperand::f16x2; ins.opcode = PTXInstruction::Min; ins.modifier = 0; + ins.d = reg("pd", PTXOperand::b32, 0); cta->setRegAsU32(0, 0, 0x40003c00); cta->setRegAsU32(0, 1, 0x3c004000); + cta->eval_Min(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0x3c003c00) result = false; + ins.d = reg("pd", PTXOperand::b32, 2); ins.opcode = PTXInstruction::Max; + ins.pg.condition = PTXOperand::Pred; ins.pg.reg = 3; cta->setRegAsPredicate(0, 3, false); + cta->setRegAsU32(0, 2, 0xa5a5a5a5); cta->eval_Max(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 2) != 0xa5a5a5a5) result = false; + ins.opcode = PTXInstruction::Min; ins.modifier = 0; ins.carry = PTXInstruction::None; + ins.pg.condition = PTXOperand::PT; ins.pg.reg = 0; ins.d = reg("d", PTXOperand::b16, 2); // u16 // if (result) { @@ -1256,17 +2384,19 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::u16, 1); ins.d = reg("r3", PTXOperand::u16, 2); + const PTXU16 aValues[] = {0, 1, 0, (std::numeric_limits::max)(), 0x8000}; + const PTXU16 bValues[] = {0, 0, 1, 0x8000, 0}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsU16(i, 0, (PTXU16)(i * 2)); - cta->setRegAsU16(i, 1, (PTXU16)(4 + i)); + cta->setRegAsU16(i, 0, aValues[i % 5]); + cta->setRegAsU16(i, 1, bValues[i % 5]); cta->setRegAsU16(i, 2, 0); } - cta->eval_Max(cta->getActiveContext(), ins); + cta->eval_Min(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXU16 expected = argmax(i*2, 4+i); + PTXU16 expected = argmin(aValues[i % 5], bValues[i % 5]); if (cta->getRegAsU16(i, 2) != expected) { result = false; - status << "max.u16 incorrect\n"; + status << "min.u16 incorrect\n"; break; } } @@ -1280,17 +2410,19 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::u32, 1); ins.d = reg("r3", PTXOperand::u32, 2); + const PTXU32 aValues[] = {0, 1, 0, (std::numeric_limits::max)(), 0x80000000U}; + const PTXU32 bValues[] = {0, 0, 1, 0x80000000U, 0}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsU32(i, 0, (PTXU32)(i * 2)); - cta->setRegAsU32(i, 1, (PTXU32)(4 + i)); + cta->setRegAsU32(i, 0, aValues[i % 5]); + cta->setRegAsU32(i, 1, bValues[i % 5]); cta->setRegAsU32(i, 2, 0); } - cta->eval_Max(cta->getActiveContext(), ins); + cta->eval_Min(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXU32 expected = argmax(i*2, 4+i); + PTXU32 expected = argmin(aValues[i % 5], bValues[i % 5]); if (cta->getRegAsU32(i, 2) != expected) { result = false; - status << "max.u32 incorrect\n"; + status << "min.u32 incorrect\n"; break; } } @@ -1304,17 +2436,19 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::u64, 1); ins.d = reg("r3", PTXOperand::u64, 2); + const PTXU64 aValues[] = {0, 1, 0, (std::numeric_limits::max)(), 0x8000000000000000ULL}; + const PTXU64 bValues[] = {0, 0, 1, 0x8000000000000000ULL, 0}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsU64(i, 0, (PTXU64)(i * 2)); - cta->setRegAsU64(i, 1, (PTXU64)(4 + i)); + cta->setRegAsU64(i, 0, aValues[i % 5]); + cta->setRegAsU64(i, 1, bValues[i % 5]); cta->setRegAsU64(i, 2, 0); } - cta->eval_Max(cta->getActiveContext(), ins); + cta->eval_Min(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXU64 expected = argmax(i*2, 4+i); + PTXU64 expected = argmin(aValues[i % 5], bValues[i % 5]); if (cta->getRegAsU64(i, 2) != expected) { result = false; - status << "max.u64 incorrect\n"; + status << "min.u64 incorrect\n"; break; } } @@ -1328,17 +2462,20 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::s16, 1); ins.d = reg("r3", PTXOperand::s16, 2); + const PTXS16 aValues[] = {0, -1, 0, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; + const PTXS16 bValues[] = {0, 0, -1, (std::numeric_limits::min)(), 0}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsS16(i, 0, (PTXS16)(i * 2)); - cta->setRegAsS16(i, 1, (PTXS16)(4 + i)); + cta->setRegAsS16(i, 0, aValues[i % 5]); + cta->setRegAsS16(i, 1, bValues[i % 5]); cta->setRegAsS16(i, 2, 0); } - cta->eval_Max(cta->getActiveContext(), ins); + cta->eval_Min(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS16 expected = argmax(i*2, 4+i); + PTXS16 expected = argmin(aValues[i % 5], bValues[i % 5]); if (cta->getRegAsS16(i, 2) != expected) { result = false; - status << "max.s16 incorrect\n"; + status << "min.s16 incorrect\n"; break; } } @@ -1352,17 +2489,20 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::s32, 1); ins.d = reg("r3", PTXOperand::s32, 2); + const PTXS32 aValues[] = {0, -1, 0, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; + const PTXS32 bValues[] = {0, 0, -1, (std::numeric_limits::min)(), 0}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsS32(i, 0, (PTXS32)(i * 2)); - cta->setRegAsS32(i, 1, (PTXS32)(4 + i)); + cta->setRegAsS32(i, 0, aValues[i % 5]); + cta->setRegAsS32(i, 1, bValues[i % 5]); cta->setRegAsS32(i, 2, 0); } - cta->eval_Max(cta->getActiveContext(), ins); + cta->eval_Min(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS32 expected = argmax(i*2, 4+i); + PTXS32 expected = argmin(aValues[i % 5], bValues[i % 5]); if (cta->getRegAsS32(i, 2) != expected) { result = false; - status << "max.s32 incorrect\n"; + status << "min.s32 incorrect\n"; break; } } @@ -1376,20 +2516,312 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::s64, 1); ins.d = reg("r3", PTXOperand::s64, 2); + const PTXS64 aValues[] = {0, -1, 0, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; + const PTXS64 bValues[] = {0, 0, -1, (std::numeric_limits::min)(), 0}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsS64(i, 0, (PTXS64)(i * 2)); - cta->setRegAsS64(i, 1, (PTXS64)(4 + i)); + cta->setRegAsS64(i, 0, aValues[i % 5]); + cta->setRegAsS64(i, 1, bValues[i % 5]); cta->setRegAsS64(i, 2, 0); } - cta->eval_Max(cta->getActiveContext(), ins); + cta->eval_Min(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS64 expected = argmax(i*2, 4+i); + PTXS64 expected = argmin(aValues[i % 5], bValues[i % 5]); if (cta->getRegAsS64(i, 2) != expected) { result = false; - status << "max.s64 incorrect\n"; + status << "min.s64 incorrect\n"; + break; + } + } + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsS64(0, 2, 0x123456789LL); + cta->eval_Min(cta->getActiveContext(), ins); + if (cta->getRegAsS64(0, 2) != 0x123456789LL) result = false; + ins.pg.condition = PTXOperand::PT; + } + + // f32 + // + if (result) { + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.b = reg("r2", PTXOperand::f32, 1); + ins.d = reg("r3", PTXOperand::f32, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsF32(i, 0, (PTXF32)(i * 2)); + cta->setRegAsF32(i, 1, (PTXF32)(4 + i)); + cta->setRegAsF32(i, 2, 0); + } + cta->eval_Min(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXF32 expected = argmin(i*2, 4+i); + if (cta->getRegAsF32(i, 2) != expected) { + result = false; + status << "min.f32 incorrect [" << i << "] - expected: " << (float)(i*2+4+i) + << ", got " << cta->getRegAsF32(i, 2) << "\n"; + break; + } + } + } + + // f64 + // + if (result) { + ins.type = PTXOperand::f64; + ins.a = reg("r1", PTXOperand::f64, 0); + ins.b = reg("r2", PTXOperand::f64, 1); + ins.d = reg("r3", PTXOperand::f64, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsF64(i, 0, (PTXF64)(i * 2)); + cta->setRegAsF64(i, 1, (PTXF64)(4 + i)); + cta->setRegAsF64(i, 2, 0.0); + } + cta->eval_Min(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXF64 expected = argmin(i*2, 4+i); + if (std::fabs(cta->getRegAsF64(i, 2) - expected) > 0.1) { + result = false; + status << "min.f64 incorrect [" << i << "] - expected: " << expected + << ", got " << cta->getRegAsF64(i, 2) << "\n"; + break; + } + } + } + + if (result) { + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::nan | PTXInstruction::xorsign + | PTXInstruction::abs; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.b = reg("r2", PTXOperand::f32, 1); + ins.d = reg("r3", PTXOperand::f32, 2); + cta->setRegAsF32(0, 0, + hydrazine::bit_cast(0xffc00000U)); + cta->setRegAsF32(0, 1, 2.0f); + cta->setRegAsF32(1, 0, -4.0f); + cta->setRegAsF32(1, 1, 2.0f); + cta->setRegAsF32(2, 0, 0.0f); + cta->setRegAsF32(2, 1, -0.0f); + cta->eval_Min(cta->getActiveContext(), ins); + result = hydrazine::bit_cast(cta->getRegAsF32(0, 2)) + == 0x7fffffffU && cta->getRegAsF32(1, 2) == -2.0f + && std::signbit(cta->getRegAsF32(2, 2)); + ins.opcode = PTXInstruction::Max; + cta->eval_Max(cta->getActiveContext(), ins); + result = result && hydrazine::bit_cast( + cta->getRegAsF32(0, 2)) == 0x7fffffffU + && cta->getRegAsF32(1, 2) == -4.0f + && std::signbit(cta->getRegAsF32(2, 2)); + if (!result) status << "min/max modifiers incorrect\n"; + } + + return result; + } + + + bool test_Max() { + bool result = true; + + PTXInstruction ins; + ins.opcode = PTXInstruction::Max; + ins.type = PTXOperand::b16; + ins.a = reg("a", PTXOperand::b16, 0); + ins.b = reg("b", PTXOperand::b16, 1); + ins.d = reg("d", PTXOperand::b16, 2); + if (ins.valid().empty()) result = false; + ins.type = PTXOperand::b32; + ins.a = reg("a", PTXOperand::b32, 0); + ins.b = reg("b", PTXOperand::b32, 1); + ins.d = reg("d", PTXOperand::b32, 2); + if (ins.valid().empty()) result = false; + ins.type = PTXOperand::b64; + ins.a = reg("a", PTXOperand::b64, 0); + ins.b = reg("b", PTXOperand::b64, 1); + ins.d = reg("d", PTXOperand::b64, 2); + if (ins.valid().empty()) result = false; + ins.type = PTXOperand::u16; + ins.a = reg("a", PTXOperand::b16, 0); + ins.b = reg("b", PTXOperand::b16, 1); + ins.d = reg("d", PTXOperand::b16, 2); + ins.modifier = PTXInstruction::rn; + if (ins.valid().empty()) result = false; + ins.modifier = PTXInstruction::sat | PTXInstruction::relu; + if (ins.valid().empty()) result = false; + ins.modifier = 0; + ins.carry = PTXInstruction::CC; + if (ins.valid().empty()) result = false; + ins.carry = PTXInstruction::None; + ins.d.type = PTXOperand::b32; + if (ins.valid().empty()) result = false; + + // u16 + // + if (result) { + ins.type = PTXOperand::u16; + ins.a = reg("r1", PTXOperand::u16, 0); + ins.b = reg("r2", PTXOperand::u16, 1); + ins.d = reg("r3", PTXOperand::u16, 2); + + const PTXU16 aValues[] = {0, 1, 0, (std::numeric_limits::max)(), 0x8000}; + const PTXU16 bValues[] = {0, 0, 1, 0x8000, 0}; + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, aValues[i % 5]); + cta->setRegAsU16(i, 1, bValues[i % 5]); + cta->setRegAsU16(i, 2, 0); + } + cta->eval_Max(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXU16 expected = argmax(aValues[i % 5], bValues[i % 5]); + if (cta->getRegAsU16(i, 2) != expected) { + result = false; + status << "max.u16 incorrect\n"; + break; + } + } + } + + // u32 + // + if (result) { + ins.type = PTXOperand::u32; + ins.a = reg("r1", PTXOperand::u32, 0); + ins.b = reg("r2", PTXOperand::u32, 1); + ins.d = reg("r3", PTXOperand::u32, 2); + + const PTXU32 aValues[] = {0, 1, 0, (std::numeric_limits::max)(), 0x80000000U}; + const PTXU32 bValues[] = {0, 0, 1, 0x80000000U, 0}; + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU32(i, 0, aValues[i % 5]); + cta->setRegAsU32(i, 1, bValues[i % 5]); + cta->setRegAsU32(i, 2, 0); + } + cta->eval_Max(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXU32 expected = argmax(aValues[i % 5], bValues[i % 5]); + if (cta->getRegAsU32(i, 2) != expected) { + result = false; + status << "max.u32 incorrect\n"; + break; + } + } + } + + // u64 + // + if (result) { + ins.type = PTXOperand::u64; + ins.a = reg("r1", PTXOperand::u64, 0); + ins.b = reg("r2", PTXOperand::u64, 1); + ins.d = reg("r3", PTXOperand::u64, 2); + + const PTXU64 aValues[] = {0, 1, 0, (std::numeric_limits::max)(), 0x8000000000000000ULL}; + const PTXU64 bValues[] = {0, 0, 1, 0x8000000000000000ULL, 0}; + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU64(i, 0, aValues[i % 5]); + cta->setRegAsU64(i, 1, bValues[i % 5]); + cta->setRegAsU64(i, 2, 0); + } + cta->eval_Max(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXU64 expected = argmax(aValues[i % 5], bValues[i % 5]); + if (cta->getRegAsU64(i, 2) != expected) { + result = false; + status << "max.u64 incorrect\n"; + break; + } + } + } + + // s16 + // + if (result) { + ins.type = PTXOperand::s16; + ins.a = reg("r1", PTXOperand::s16, 0); + ins.b = reg("r2", PTXOperand::s16, 1); + ins.d = reg("r3", PTXOperand::s16, 2); + + const PTXS16 aValues[] = {0, -1, 0, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; + const PTXS16 bValues[] = {0, 0, -1, (std::numeric_limits::min)(), 0}; + for (int i = 0; i < threadCount; i++) { + cta->setRegAsS16(i, 0, aValues[i % 5]); + cta->setRegAsS16(i, 1, bValues[i % 5]); + cta->setRegAsS16(i, 2, 0); + } + cta->eval_Max(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXS16 expected = argmax(aValues[i % 5], bValues[i % 5]); + if (cta->getRegAsS16(i, 2) != expected) { + result = false; + status << "max.s16 incorrect\n"; + break; + } + } + } + + // s32 + // + if (result) { + ins.type = PTXOperand::s32; + ins.a = reg("r1", PTXOperand::s32, 0); + ins.b = reg("r2", PTXOperand::s32, 1); + ins.d = reg("r3", PTXOperand::s32, 2); + + const PTXS32 aValues[] = {0, -1, 0, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; + const PTXS32 bValues[] = {0, 0, -1, (std::numeric_limits::min)(), 0}; + for (int i = 0; i < threadCount; i++) { + cta->setRegAsS32(i, 0, aValues[i % 5]); + cta->setRegAsS32(i, 1, bValues[i % 5]); + cta->setRegAsS32(i, 2, 0); + } + cta->eval_Max(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXS32 expected = argmax(aValues[i % 5], bValues[i % 5]); + if (cta->getRegAsS32(i, 2) != expected) { + result = false; + status << "max.s32 incorrect\n"; + break; + } + } + } + + // s64 + // + if (result) { + ins.type = PTXOperand::s64; + ins.a = reg("r1", PTXOperand::s64, 0); + ins.b = reg("r2", PTXOperand::s64, 1); + ins.d = reg("r3", PTXOperand::s64, 2); + + const PTXS64 aValues[] = {0, -1, 0, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; + const PTXS64 bValues[] = {0, 0, -1, (std::numeric_limits::min)(), 0}; + for (int i = 0; i < threadCount; i++) { + cta->setRegAsS64(i, 0, aValues[i % 5]); + cta->setRegAsS64(i, 1, bValues[i % 5]); + cta->setRegAsS64(i, 2, 0); + } + cta->eval_Max(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + PTXS64 expected = argmax(aValues[i % 5], bValues[i % 5]); + if (cta->getRegAsS64(i, 2) != expected) { + result = false; + status << "max.s64 incorrect\n"; break; } } + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsS64(0, 2, 0x123456789LL); + cta->eval_Max(cta->getActiveContext(), ins); + if (cta->getRegAsS64(0, 2) != 0x123456789LL) result = false; + ins.pg.condition = PTXOperand::PT; } // f32 @@ -1451,20 +2883,126 @@ class TestInstructions: public Test { PTXInstruction ins; ins.opcode = PTXInstruction::Neg; + // bf16 + // + if (result) { + ins.type = PTXOperand::bf16; + ins.modifier = PTXInstruction::rn; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.d = reg("r3", PTXOperand::b16, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, 0x3fc0); // 1.5 + cta->setRegAsU16(i, 2, 0); + } + if (!ins.valid().empty()) { + result = false; + status << "neg.bf16 rejected\n"; + } + else { + cta->eval_Neg(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 2) != 0xbfc0) { // -1.5 + result = false; + status << "neg.bf16 incorrect\n"; + break; + } + } + } + } + + // f16 + // + if (result) { + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.d = reg("r3", PTXOperand::b16, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, 0x3e00); // 1.5 + cta->setRegAsU16(i, 2, 0); + } + if (!ins.valid().empty()) { + result = false; + status << "neg.f16 rejected\n"; + } + else { + cta->eval_Neg(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 2) != 0xbe00) { // -1.5 + result = false; + status << "neg.f16 incorrect\n"; + break; + } + } + } + } + + // f16x2 + // + if (result) { + ins.type = PTXOperand::f16x2; + ins.modifier = 0; + ins.a = reg("r1", PTXOperand::b32, 0); + ins.d = reg("r3", PTXOperand::b32, 2); + if (!ins.valid().empty()) { + status << "neg.f16x2 rejected\n"; + return false; + } + cta->reset(); + auto packedNeg = [&](int modifier, PTXU32 a, PTXU32 expected) { + ins.modifier = modifier; + cta->setRegAsU32(0, 0, a); + cta->setRegAsU32(0, 2, 0xdeadbeef); + cta->eval_Neg(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, 2) == expected; + }; + result = result && packedNeg(0, 0x3e00be00, 0xbe003e00); + result = result && packedNeg(PTXInstruction::ftz, 0x80010001, + 0x00008000); + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_neg() { .reg .b32 d, a; " + << "neg.bf16x2 d, a; ret; }\n"; + try { Module parsed; parsed.load(ptx); } + catch (const std::exception& error) { + status << "neg.bf16x2 parse failed: " << error.what() << "\n"; + return false; + } + ins.type = PTXOperand::bf16x2; + ins.modifier = 0; + if (!ins.valid().empty()) return false; + result = result && packedNeg(0, 0x3fc03f80, 0xbfc0bf80); + result = result && packedNeg(0, 0x80010000, 0x00018000); + } + // s16 // if (result) { ins.type = PTXOperand::s16; + ins.modifier = PTXInstruction::rn; ins.a = reg("r1", PTXOperand::s16, 0); ins.d = reg("r3", PTXOperand::s16, 2); - - for (int i = 0; i < threadCount; i++) { - cta->setRegAsS16(i, 0, (PTXS16)(i * 2)); + if (ins.valid().empty()) result = false; + ins.modifier = PTXInstruction::sat; + if (ins.valid().empty()) result = false; + ins.modifier = PTXInstruction::ftz; + if (ins.valid().empty()) result = false; + ins.modifier = 0; + ins.carry = PTXInstruction::CC; + if (ins.valid().empty()) result = false; + ins.carry = PTXInstruction::None; + + const PTXS16 values[] = {0, 1, -1, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; + for (int i = 0; i < threadCount; i++) { + cta->setRegAsS16(i, 0, values[i % 5]); cta->setRegAsS16(i, 2, 0); } cta->eval_Neg(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS16 expected = -(i*2); + PTXS16 expected = values[i % 5] == values[4] ? values[4] : -values[i % 5]; if (cta->getRegAsS16(i, 2) != expected) { result = false; status << "neg.s16 incorrect\n"; @@ -1481,13 +3019,15 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::s32, 1); ins.d = reg("r3", PTXOperand::s32, 2); + const PTXS32 values[] = {0, 1, -1, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsS32(i, 0, (PTXS32)(i * 2)); + cta->setRegAsS32(i, 0, values[i % 5]); cta->setRegAsS32(i, 2, 0); } cta->eval_Neg(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS32 expected = -(i*2); + PTXS32 expected = values[i % 5] == values[4] ? values[4] : -values[i % 5]; if (cta->getRegAsS32(i, 2) != expected) { result = false; status << "neg.s32 incorrect\n"; @@ -1503,19 +3043,29 @@ class TestInstructions: public Test { ins.a = reg("r1", PTXOperand::s64, 0); ins.d = reg("r3", PTXOperand::s64, 2); + const PTXS64 values[] = {0, 1, -1, (std::numeric_limits::max)(), + (std::numeric_limits::min)()}; for (int i = 0; i < threadCount; i++) { - cta->setRegAsS64(i, 0, (PTXS64)(i * 2)); + cta->setRegAsS64(i, 0, values[i % 5]); cta->setRegAsS64(i, 2, 0); } cta->eval_Neg(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS64 expected = -(i*2); + PTXS64 expected = values[i % 5] == values[4] ? values[4] : -values[i % 5]; if (cta->getRegAsS64(i, 2) != expected) { result = false; status << "neg.s64 incorrect\n"; break; } } + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsS64(0, 0, 1); + cta->setRegAsS64(0, 2, 0x123456789LL); + cta->eval_Neg(cta->getActiveContext(), ins); + if (cta->getRegAsS64(0, 2) != 0x123456789LL) result = false; + ins.pg.condition = PTXOperand::PT; } // f32 @@ -1584,12 +3134,12 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { cta->setRegAsU16(i, 0, (PTXU16)(i * 8 + 8)); - cta->setRegAsU16(i, 1, (PTXU16)(4 + i)); + cta->setRegAsU16(i, 1, i == 0 ? 0 : (PTXU16)(4 + i)); cta->setRegAsU16(i, 2, 0); } cta->eval_Rem(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXU16 expected = ((i * 8 + 8) % (4 + i)); + PTXU16 expected = i == 0 ? 0 : ((i * 8 + 8) % (4 + i)); if (cta->getRegAsU16(i, 2) != expected) { result = false; status << "rem.u16 incorrect\n"; @@ -1608,12 +3158,12 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { cta->setRegAsU32(i, 0, (PTXU32)(i * 8 + 8)); - cta->setRegAsU32(i, 1, (PTXU32)(4 + i)); + cta->setRegAsU32(i, 1, i == 0 ? 0 : (PTXU32)(4 + i)); cta->setRegAsU32(i, 2, 0); } cta->eval_Rem(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXU32 expected = ((i * 8 + 8) % (4 + i)); + PTXU32 expected = i == 0 ? 0 : ((i * 8 + 8) % (4 + i)); if (cta->getRegAsU32(i, 2) != expected) { result = false; status << "rem.u32 incorrect\n"; @@ -1632,12 +3182,12 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { cta->setRegAsU64(i, 0, (PTXU64)(i * 8 + 8)); - cta->setRegAsU64(i, 1, (PTXU64)(4 + i)); + cta->setRegAsU64(i, 1, i == 0 ? 0 : (PTXU64)(4 + i)); cta->setRegAsU64(i, 2, 0); } cta->eval_Rem(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXU64 expected = ((i * 8 + 8) % (4 + i)); + PTXU64 expected = i == 0 ? 0 : ((i * 8 + 8) % (4 + i)); if (cta->getRegAsU64(i, 2) != expected) { result = false; status << "rem.u64 incorrect\n"; @@ -1656,12 +3206,12 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { cta->setRegAsS16(i, 0, (PTXS16)(i * 8 + 8)); - cta->setRegAsS16(i, 1, (PTXS16)(4 + i)); + cta->setRegAsS16(i, 1, i == 0 ? 0 : (PTXS16)(4 + i)); cta->setRegAsS16(i, 2, 0); } cta->eval_Rem(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS16 expected = ((i * 8 + 8) % (4 + i)); + PTXS16 expected = i == 0 ? 0 : ((i * 8 + 8) % (4 + i)); if (cta->getRegAsS16(i, 2) != expected) { result = false; status << "rem.s16 incorrect\n"; @@ -1680,12 +3230,12 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { cta->setRegAsS32(i, 0, (PTXS32)(i * 8 + 8)); - cta->setRegAsS32(i, 1, (PTXS32)(4 + i)); + cta->setRegAsS32(i, 1, i == 0 ? 0 : (PTXS32)(4 + i)); cta->setRegAsS32(i, 2, 0); } cta->eval_Rem(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS32 expected = ((i * 8 + 8) % (4 + i)); + PTXS32 expected = i == 0 ? 0 : ((i * 8 + 8) % (4 + i)); if (cta->getRegAsS32(i, 2) != expected) { result = false; status << "rem.s32 incorrect\n"; @@ -1704,12 +3254,12 @@ class TestInstructions: public Test { for (int i = 0; i < threadCount; i++) { cta->setRegAsS64(i, 0, (PTXS64)(i * 8 + 8)); - cta->setRegAsS64(i, 1, (PTXS64)(4 + i)); + cta->setRegAsS64(i, 1, i == 0 ? 0 : (PTXS64)(4 + i)); cta->setRegAsS64(i, 2, 0); } cta->eval_Rem(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXS64 expected = ((i * 8 + 8) % (4 + i)); + PTXS64 expected = i == 0 ? 0 : ((i * 8 + 8) % (4 + i)); if (cta->getRegAsS64(i, 2) != expected) { result = false; status << "rem.s64 incorrect\n"; @@ -1876,8 +3426,17 @@ class TestInstructions: public Test { // if (result) { ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn; ins.a = reg("r1", PTXOperand::f32, 0); + ins.b = reg("r2", PTXOperand::f32, 1); ins.d = reg("r3", PTXOperand::f32, 2); + if (!ins.valid().empty()) { status << ins.valid() << "\n"; return false; } + PTXInstruction invalid = ins; + invalid.modifier = 0; + if (invalid.valid().empty()) { + status << "div.f32 accepted without a mode\n"; + return false; + } for (int i = 0; i < threadCount; i++) { cta->setRegAsF32(i, 0, (PTXF32)(i * 8 + 8)); @@ -1901,7 +3460,14 @@ class TestInstructions: public Test { if (result) { ins.type = PTXOperand::f64; ins.a = reg("r1", PTXOperand::f64, 0); + ins.b = reg("r2", PTXOperand::f64, 1); ins.d = reg("r3", PTXOperand::f64, 2); + PTXInstruction invalid = ins; + invalid.modifier |= PTXInstruction::ftz; + if (invalid.valid().empty()) { + status << "div.ftz.f64 accepted\n"; + return false; + } for (int i = 0; i < threadCount; i++) { cta->setRegAsF64(i, 0, (PTXF64)(i * 8 + 8)); @@ -1910,7 +3476,7 @@ class TestInstructions: public Test { } cta->eval_Div(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXF32 expected = ((PTXF64)(i * 8 + 8) / (PTXF64)(4 + i)); + PTXF64 expected = ((PTXF64)(i * 8 + 8) / (PTXF64)(4 + i)); if (std::fabs(cta->getRegAsF64(i, 2) - expected) > 0.1) { result = false; status << "div.f64 incorrect [" << i << "] - expected: " << expected @@ -1920,11 +3486,86 @@ class TestInstructions: public Test { } } - return result; - } - - bool test_Mad() { - bool result = true; + const int modes[] = {PTXInstruction::rn, PTXInstruction::rz, + PTXInstruction::rm, PTXInstruction::rp}; + const PTXU32 expected[][2] = {{0x3dcccccd, 0xbdcccccd}, + {0x3dcccccc, 0xbdcccccc}, {0x3dcccccc, 0xbdcccccd}, + {0x3dcccccd, 0xbdcccccc}}; + ins.type = PTXOperand::f32; + ins.a.type = ins.b.type = ins.d.type = PTXOperand::f32; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + cta->setRegAsF32(0, 0, 1.0f); + cta->setRegAsF32(0, 1, 10.0f); + cta->setRegAsF32(1, 0, -1.0f); + cta->setRegAsF32(1, 1, 10.0f); + cta->eval_Div(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != expected[i][0] + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != expected[i][1]) { + status << "div.f32 rounding failed\n"; + return false; + } + } + ins.type = PTXOperand::f64; + ins.a.type = ins.b.type = ins.d.type = PTXOperand::f64; + ins.modifier = PTXInstruction::rz; + cta->setRegAsF64(0, 0, 1.0); + cta->setRegAsF64(0, 1, 10.0); + cta->eval_Div(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) + != 0x3fb9999999999999ull) return false; + + ins.type = PTXOperand::f32; + ins.a.type = ins.b.type = ins.d.type = PTXOperand::f32; + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + cta->setRegAsF32(0, 0, 1.0f); + cta->setRegAsF32(1, 0, -1.0f); + cta->setRegAsF32(2, 0, std::numeric_limits::infinity()); + cta->setRegAsF32(3, 0, std::numeric_limits::quiet_NaN()); + for (int i = 0; i < 4; ++i) { + cta->setRegAsF32(i, 1, std::numeric_limits::max()); + } + cta->eval_Div(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0 + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != 0x80000000 + || !hydrazine::isnan(cta->getRegAsF32(2, 2)) + || hydrazine::bit_cast(cta->getRegAsF32(3, 2)) != 0) { + status << "div.approx.f32 large divisor behavior failed\n"; + return false; + } + cta->setRegAsF32(0, 0, std::numeric_limits::denorm_min()); + cta->setRegAsF32(0, 1, 1.0f); + cta->eval_Div(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 1) { + status << "div.approx.f32 flushed a subnormal without ftz\n"; + return false; + } + + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->setRegAsF32(0, 0, -std::numeric_limits::denorm_min()); + cta->setRegAsF32(0, 1, 1.0f); + cta->eval_Div(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0x80000000) { + status << "div.approx.ftz.f32 input flushing failed\n"; + return false; + } + + ins.modifier = PTXInstruction::full | PTXInstruction::ftz; + if (!ins.valid().empty()) return false; + cta->setRegAsF32(0, 0, std::numeric_limits::min()); + cta->setRegAsF32(0, 1, 2.0f); + cta->eval_Div(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0) { + status << "div.full.ftz.f32 result flushing failed\n"; + return false; + } + + return result; + } + + bool test_Mad() { + bool result = true; PTXInstruction ins; ins.opcode = PTXInstruction::Mad; @@ -2091,10 +3732,21 @@ class TestInstructions: public Test { // if (result) { ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn; ins.a = reg("r1", PTXOperand::f32, 0); ins.b = reg("r2", PTXOperand::f32, 1); ins.c = reg("r3", PTXOperand::f32, 2); ins.d = reg("r4", PTXOperand::f32, 3); + if (!ins.valid().empty()) { + status << "mad.rn.f32 rejected\n"; + return false; + } + PTXInstruction invalid = ins; + invalid.modifier = 0; + if (invalid.valid().empty()) { + status << "mad.f32 accepted without a rounding modifier\n"; + return false; + } for (int i = 0; i < threadCount; i++) { cta->setRegAsF32(i, 0, (PTXF32)(i - 1)); @@ -2122,6 +3774,12 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::f64, 1); ins.c = reg("r3", PTXOperand::f64, 2); ins.d = reg("r4", PTXOperand::f64, 3); + PTXInstruction invalid = ins; + invalid.modifier |= PTXInstruction::sat; + if (invalid.valid().empty()) { + status << "mad.sat.f64 accepted\n"; + return false; + } for (int i = 0; i < threadCount; i++) { cta->setRegAsF64(i, 0, (PTXF64)(i - 1)); @@ -2131,7 +3789,7 @@ class TestInstructions: public Test { } cta->eval_Mad(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXF32 expected = (PTXF64)(i - 1) * (PTXF64)(4 + 2*i) + (PTXF64)(i); + PTXF64 expected = (PTXF64)(i - 1) * (PTXF64)(4 + 2*i) + (PTXF64)(i); if (std::fabs(cta->getRegAsF64(i, 3) - expected) > 0.1) { result = false; status << "mad.f64 incorrect [" << i << "] - expected: " << expected @@ -2141,6 +3799,49 @@ class TestInstructions: public Test { } } + ins.type = PTXOperand::f32; + ins.a.type = ins.b.type = ins.c.type = ins.d.type = PTXOperand::f32; + cta->setRegAsF32(0, 0, 1.0f + std::ldexp(1.0f, -23)); + cta->setRegAsF32(0, 1, 1.0f - std::ldexp(1.0f, -23)); + cta->setRegAsF32(0, 2, -1.0f); + cta->eval_Mad(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 3)) != 0xa8800000) { + status << "mad.f32 was not fused\n"; + return false; + } + + ins.type = PTXOperand::f64; + ins.a.type = ins.b.type = ins.c.type = ins.d.type = PTXOperand::f64; + cta->setRegAsF64(0, 0, 1.0 + std::ldexp(1.0, -52)); + cta->setRegAsF64(0, 1, 1.0 - std::ldexp(1.0, -52)); + cta->setRegAsF64(0, 2, -1.0); + cta->eval_Mad(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 3)) + != 0xb970000000000000ull) { + status << "mad.f64 was not fused\n"; + return false; + } + + ins.type = PTXOperand::f32; + ins.a.type = ins.b.type = ins.c.type = ins.d.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn | PTXInstruction::ftz; + cta->setRegAsF32(0, 0, std::numeric_limits::min()); + cta->setRegAsF32(0, 1, 0.5f); + cta->setRegAsF32(0, 2, 0.0f); + cta->eval_Mad(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 3)) != 0) { + status << "mad.ftz.f32 failed\n"; + return false; + } + ins.modifier = PTXInstruction::rn | PTXInstruction::sat; + cta->setRegAsF32(0, 0, std::numeric_limits::quiet_NaN()); + cta->setRegAsF32(0, 1, 1.0f); + cta->eval_Mad(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 3) != 0.0f) { + status << "mad.sat.f32 failed\n"; + return false; + } + return result; } @@ -2150,6 +3851,137 @@ class TestInstructions: public Test { PTXInstruction ins; ins.opcode = PTXInstruction::Mul; + // f16 + // + if (result) { + ins.type = PTXOperand::f16; + ins.modifier = 0; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.b = reg("r2", PTXOperand::b16, 1); + ins.c = reg("r3", PTXOperand::b16, 2); + ins.d = reg("r4", PTXOperand::b16, 3); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, 0x3e00); // 1.5 + cta->setRegAsU16(i, 1, 0x4000); // 2.0 + cta->setRegAsU16(i, 2, 0); + cta->setRegAsU16(i, 3, 0); + } + std::string error = ins.valid(); + if (!error.empty()) { + result = false; + status << "mul.f16 rejected: " << error << "\n"; + } + else { + cta->eval_Mul(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 3) != 0x4200) { // 3.0 + result = false; + status << "mul.f16 incorrect\n"; + break; + } + } + } + } + if (result) { + ins.type = PTXOperand::f16; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.b = reg("r2", PTXOperand::b16, 1); + ins.d = reg("r4", PTXOperand::b16, 3); + ins.modifier = 0; // PTX defaults to round-to-nearest-even. + cta->setRegAsU16(0, 0, 0x3c01); // 1 + 2^-10 + cta->setRegAsU16(0, 1, 0x3c01); // 1 + 2^-10 + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + cta->eval_Mul(cta->getActiveContext(), ins); + const bool roundedNearest = cta->getRegAsU16(0, 3) == 0x3c02; + const bool restored = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(previous); + if (!roundedNearest || !restored) { + status << "mul.f16 default rounding failed\n"; + result = false; + } + } + ins.modifier = PTXInstruction::rz; + if (ins.valid().empty()) { + status << "mul.rz.f16 accepted\n"; + result = false; + } + ins.modifier = PTXInstruction::ftz | PTXInstruction::sat; + if (!ins.valid().empty()) { + status << "mul.ftz.sat.f16 rejected\n"; + result = false; + } + + // f16x2 + // + if (result) { + ins.type = PTXOperand::f16x2; + ins.modifier = 0; + ins.a = reg("r1", PTXOperand::b32, 0); + ins.b = reg("r2", PTXOperand::b32, 1); + ins.d = reg("r3", PTXOperand::b32, 2); + if (!ins.valid().empty()) result = false; + ins.modifier = PTXInstruction::rp; + if (ins.valid().empty()) { + status << "mul.rp.f16x2 accepted\n"; + result = false; + } + cta->reset(); + auto packedMul = [&](int modifier, PTXU32 a, PTXU32 b, + PTXU32 expected, bool alias) { + ins.modifier = modifier; + ins.d.reg = alias ? 0 : 2; + cta->setRegAsU32(0, 0, a); + cta->setRegAsU32(0, 1, b); + cta->setRegAsU32(0, 2, 0xdeadbeef); + cta->eval_Mul(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, ins.d.reg) == expected; + }; + result = result && packedMul(0, 0xc000be00, 0x3e00c000, + 0xc2004200, false); + result = result && packedMul(0, 0x3c033c01, 0x3e003e00, + 0x3e043e02, false); + result = result && packedMul(PTXInstruction::rn, 0x3c033c01, + 0x3e003e00, 0x3e043e02, false); + result = result && packedMul(0, 0x00010001, 0x3c003c00, + 0x00010001, false); + result = result && packedMul(PTXInstruction::ftz, 0x00010001, + 0x3c003c00, 0, false); + result = result && packedMul(0, 0x84000400, 0x38003800, + 0x82000200, false); + result = result && packedMul(PTXInstruction::ftz, 0x84000400, + 0x38003800, 0x80000000, false); + result = result && packedMul(PTXInstruction::sat, 0xbc004000, + 0x3c003c00, 0x00003c00, false); + result = result && packedMul(PTXInstruction::sat, 0x3c003800, + 0x3c003c00, 0x3c003800, false); + result = result && packedMul(PTXInstruction::sat, 0x7e007e00, + 0x3c003c00, 0, false); + result = result && packedMul(0, 0x00008000, 0xc0004000, + 0x80008000, false); + result = result && packedMul(0, 0xc000be00, 0x3e00c000, + 0xc2004200, true); + if (result) { + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + const bool defaultRn = packedMul(0, 0x3c033c01, + 0x3e003e00, 0x3e043e02, false); + const bool restored = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(previous); + if (!defaultRn || !restored) result = false; + } + ins.modifier = 0; + ins.d.reg = 2; + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsU32(0, 2, 0xcafebabe); + cta->eval_Mul(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 2) != 0xcafebabe) result = false; + ins.pg.condition = PTXOperand::PT; + } + // u16 // if (result) { @@ -2365,6 +4197,53 @@ class TestInstructions: public Test { return result; } + bool test_MulRounding() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Mul; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.b = reg("r2", PTXOperand::f32, 1); + ins.d = reg("r3", PTXOperand::f32, 2); + const int modes[] = {PTXInstruction::rn, PTXInstruction::rz, + PTXInstruction::rm, PTXInstruction::rp}; + const PTXU32 expected32[][2] = {{0x7f800000, 0xff800000}, + {0x7f7fffff, 0xff7fffff}, {0x7f7fffff, 0xff800000}, + {0x7f800000, 0xff7fffff}}; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + cta->setRegAsF32(0, 0, std::numeric_limits::max()); + cta->setRegAsF32(0, 1, 2.0f); + cta->setRegAsF32(1, 0, -std::numeric_limits::max()); + cta->setRegAsF32(1, 1, 2.0f); + cta->eval_Mul(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != expected32[i][0] + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != expected32[i][1]) { + status << "mul.f32 rounding failed\n"; + return false; + } + } + ins.type = PTXOperand::f64; + ins.a.type = ins.b.type = ins.d.type = PTXOperand::f64; + const PTXU64 expected64[][2] = {{0x7ff0000000000000ull, 0xfff0000000000000ull}, + {0x7fefffffffffffffull, 0xffefffffffffffffull}, + {0x7fefffffffffffffull, 0xfff0000000000000ull}, + {0x7ff0000000000000ull, 0xffefffffffffffffull}}; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + cta->setRegAsF64(0, 0, std::numeric_limits::max()); + cta->setRegAsF64(0, 1, 2.0); + cta->setRegAsF64(1, 0, -std::numeric_limits::max()); + cta->setRegAsF64(1, 1, 2.0); + cta->eval_Mul(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) != expected64[i][0] + || hydrazine::bit_cast(cta->getRegAsF64(1, 2)) != expected64[i][1]) { + status << "mul.f64 rounding failed\n"; + return false; + } + } + return true; + } + ///////////////////////////////////////////////////////////////////////////////////////////////// // // @@ -2375,9 +4254,43 @@ class TestInstructions: public Test { bool test_Rcp() { bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_rcp() {\n" + << " .reg .f32 f0, f1;\n" + << " .reg .f64 d0, d1;\n" + << " rcp.approx.f32 f0, f1;\n" + << " rcp.rz.ftz.f32 f0, f1;\n" + << " rcp.rp.f64 d0, d1;\n" + << " rcp.approx.ftz.f64 d0, d1;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse PTX 8.0 rcp forms: " << error.what() << "\n"; + return false; + } PTXInstruction ins; ins.opcode = PTXInstruction::Rcp; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + ins.modifier = 0; + if (ins.valid().empty()) { + status << "rcp.f32 accepted without .approx or rounding mode\n"; + return false; + } + ins.modifier = PTXInstruction::approx | PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.type = PTXOperand::f64; + ins.a.type = ins.d.type = PTXOperand::f64; + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + if (!ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx; + if (ins.valid().empty()) return false; double freq = 2.0f / (double)threadCount; @@ -2385,6 +4298,7 @@ class TestInstructions: public Test { // if (result) { ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn; ins.a = reg("r1", PTXOperand::f32, 0); ins.d = reg("r3", PTXOperand::f32, 2); @@ -2408,6 +4322,7 @@ class TestInstructions: public Test { // if (result) { ins.type = PTXOperand::f64; + ins.modifier = PTXInstruction::rn; ins.a = reg("r1", PTXOperand::f64, 0); ins.d = reg("r3", PTXOperand::f64, 2); @@ -2426,24 +4341,120 @@ class TestInstructions: public Test { } } + const int modes[] = {PTXInstruction::rn, PTXInstruction::rz, + PTXInstruction::rm, PTXInstruction::rp}; + const PTXU32 expected32[][2] = { + {0x3dcccccdu, 0xbdcccccdu}, {0x3dccccccu, 0xbdccccccu}, + {0x3dccccccu, 0xbdcccccdu}, {0x3dcccccdu, 0xbdccccccu} + }; + const PTXU64 expected64[][2] = { + {0x3fb999999999999aull, 0xbfb999999999999aull}, + {0x3fb9999999999999ull, 0xbfb9999999999999ull}, + {0x3fb9999999999999ull, 0xbfb999999999999aull}, + {0x3fb999999999999aull, 0xbfb9999999999999ull} + }; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + ins.type = PTXOperand::f32; + ins.a.type = ins.d.type = PTXOperand::f32; + cta->setRegAsF32(0, 0, 10.0f); + cta->setRegAsF32(1, 0, -10.0f); + cta->eval_Rcp(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != expected32[i][0] + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != expected32[i][1]) { + status << "rcp.f32 rounding failed\n"; + return false; + } + + ins.type = PTXOperand::f64; + ins.a.type = ins.d.type = PTXOperand::f64; + cta->setRegAsF64(0, 0, 10.0); + cta->setRegAsF64(1, 0, -10.0); + cta->eval_Rcp(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) != expected64[i][0] + || hydrazine::bit_cast(cta->getRegAsF64(1, 2)) != expected64[i][1]) { + status << "rcp.f64 rounding failed\n"; + return false; + } + } + + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->setRegAsF64(0, 0, 3.0); + cta->setRegAsF64(1, 0, std::numeric_limits::quiet_NaN()); + cta->eval_Rcp(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) + != 0x3fd5555500000000ull + || hydrazine::bit_cast(cta->getRegAsF64(1, 2)) + != 0x7fffffff00000000ull) { + status << "rcp.approx.ftz.f64 failed\n"; + return false; + } + + // Without .ftz, .approx must still support subnormal inputs and + // return a finite approximate reciprocal, not flush to infinity. + ins.type = PTXOperand::f32; + ins.a.type = ins.d.type = PTXOperand::f32; + ins.modifier = PTXInstruction::approx; + const PTXF32 subnormal = std::numeric_limits::min() + - std::numeric_limits::denorm_min(); + cta->setRegAsF32(0, 0, subnormal); + cta->setRegAsF32(1, 0, -subnormal); + cta->eval_Rcp(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0x7e800001u + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != 0xfe800001u) { + status << "rcp.approx.f32 subnormal input failed\n"; + return false; + } + + // With .ftz, the subnormal input flushes to signed zero first, so + // the reciprocal is signed infinity. + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->setRegAsF32(0, 0, subnormal); + cta->setRegAsF32(1, 0, -subnormal); + cta->eval_Rcp(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 2) != std::numeric_limits::infinity() + || cta->getRegAsF32(1, 2) != -std::numeric_limits::infinity()) { + status << "rcp.approx.ftz.f32 subnormal input should flush to infinity\n"; + return false; + } + return result; } bool test_Cos() { bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_cos() {\n" + << " .reg .f32 f0, f1;\n" + << " cos.approx.f32 f0, f1;\n" + << " cos.approx.ftz.f32 f0, f1;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse PTX 8.0 cos forms: " << error.what() << "\n"; + return false; + } PTXInstruction ins; + ins.opcode = PTXInstruction::Cos; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + ins.modifier = 0; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx; // f32 // if (result) { float freq = 2 * 3.14159f / (float)threadCount; - ins.opcode = PTXInstruction::Cos; - ins.type = PTXOperand::f32; - ins.a = reg("r1", PTXOperand::f32, 0); - ins.d = reg("r3", PTXOperand::f32, 2); - for (int i = 0; i < threadCount; i++) { cta->setRegAsF32(i, 0, (PTXF32)((float)i * freq)); cta->setRegAsF32(i, 2, 0); @@ -2459,14 +4470,45 @@ class TestInstructions: public Test { } } } + + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->setRegAsF32(0, 0, std::numeric_limits::denorm_min()); + cta->setRegAsF32(1, 0, -std::numeric_limits::denorm_min()); + cta->eval_Cos(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 2) != 1.0f + || cta->getRegAsF32(1, 2) != 1.0f) return false; + return result; } bool test_Sin() { bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_sin() {\n" + << " .reg .f32 f0, f1;\n" + << " sin.approx.f32 f0, f1;\n" + << " sin.approx.ftz.f32 f0, f1;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse PTX 8.0 sin forms: " << error.what() << "\n"; + return false; + } PTXInstruction ins; ins.opcode = PTXInstruction::Sin; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + ins.modifier = 0; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx; // f32 // @@ -2493,14 +4535,104 @@ class TestInstructions: public Test { } } + const PTXF32 subnormal = std::numeric_limits::denorm_min(); + cta->setRegAsF32(0, 0, subnormal); + cta->setRegAsF32(1, 0, -subnormal); + cta->eval_Sin(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 1 + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) + != 0x80000001u) return false; + + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->eval_Sin(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0 + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) + != 0x80000000u) return false; + return result; } + + bool test_Tanh() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_tanh() {\n" + << " .reg .f32 d, a;\n" + << " .reg .b16 hd, ha; .reg .b32 pd, pa;\n" + << " tanh.approx.f32 d, a;\n" + << " tanh.approx.f16 hd, ha; tanh.approx.f16x2 pd, pa;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse tanh example: " << error.what() << "\n"; + return false; + } + + PTXInstruction ins; + ins.opcode = PTXInstruction::Tanh; + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::approx; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::f32, 1); + for (int thread = 0; thread < threadCount; ++thread) { + const PTXF32 input = (thread - threadCount / 2) / 4.0f; + cta->setRegAsF32(thread, 1, input); + } + cta->eval_Tanh(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + const PTXF32 input = (thread - threadCount / 2) / 4.0f; + if (std::fabs(cta->getRegAsF32(thread, 0) - std::tanh(input)) > 1e-6f) + return false; + } + const PTXF32 subnormal = std::numeric_limits::denorm_min(); + cta->setRegAsF32(0, 1, subnormal); + cta->setRegAsF32(1, 1, -subnormal); + cta->eval_Tanh(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 0) != subnormal + || cta->getRegAsF32(1, 0) != -subnormal) return false; + + ins.type = PTXOperand::f16; + ins.d = reg("hd", PTXOperand::b16, 0); + ins.a = reg("ha", PTXOperand::b16, 1); + const PTXU16 inputs[] = {0, 0x8000, 0x3c00, 0xbc00, + 0x7c00, 0xfc00, 1, 0x8001}; + const PTXU16 expected[] = {0, 0x8000, 0x3a18, 0xba18, + 0x3c00, 0xbc00, 1, 0x8001}; + for (unsigned int i = 0; i < 8; ++i) { + cta->setRegAsU16(0, 1, inputs[i]); + cta->eval_Tanh(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 0) != expected[i]) return false; + } + cta->setRegAsU16(0, 1, 0x7e00); + cta->eval_Tanh(cta->getActiveContext(), ins); + const PTXU16 nan = cta->getRegAsU16(0, 0); + if ((nan & 0x7c00) != 0x7c00 || !(nan & 0x03ff)) return false; + + ins.type = PTXOperand::f16x2; + ins.d = reg("pd", PTXOperand::b32, 0); + ins.a = reg("pa", PTXOperand::b32, 1); + cta->setRegAsU32(0, 1, 0x3c00bc00); + cta->eval_Tanh(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0x3a18ba18) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + if (ins.valid().empty()) return false; + + ins.modifier = PTXInstruction::approx; + ins.type = PTXOperand::bf16; + ins.d = reg("bd", PTXOperand::b16, 0); + ins.a = reg("ba", PTXOperand::b16, 1); + if (ins.valid().empty()) return false; + ins.type = PTXOperand::bf16x2; + ins.d = reg("pd", PTXOperand::b32, 0); + ins.a = reg("pa", PTXOperand::b32, 1); + return !ins.valid().empty(); + } bool test_CopySign() { bool result = true; PTXInstruction ins; - ins.opcode = PTXInstruction::Fma; + ins.opcode = PTXInstruction::CopySign; // f32 // @@ -2517,24 +4649,18 @@ class TestInstructions: public Test { cta->setRegAsF32(i, 1, (PTXF32)((float)(bs * i) / (float)threadCount * 2.7f)); cta->setRegAsF32(i, 2, 0); } - cta->eval_Fma(cta->getActiveContext(), ins); + cta->eval_CopySign(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { PTXF32 got = cta->getRegAsF32(i, 2); PTXF32 a = cta->getRegAsF32(i, 0); PTXF32 b = cta->getRegAsF32(i, 1); - PTXF32 exp = b; - if (a < 0) { - exp = -std::fabs(b); - } - else { - exp = std::fabs(b); - } + PTXF32 exp = std::copysign(b, a); - if (std::fabs(got - exp) > 0.1f) { + if (got != exp) { result = false; - status << "fma.f32 incorrect [" << i << "] - expected: " + status << "copysign.f32 incorrect [" << i << "] - expected: " << (PTXF32)exp << ", got " << got << "\n"; break; @@ -2556,24 +4682,18 @@ class TestInstructions: public Test { cta->setRegAsF64(i, 1, (PTXF64)((double)(bs * i) / (double)threadCount * 7.7)); cta->setRegAsF64(i, 2, 0); } - cta->eval_Fma(cta->getActiveContext(), ins); + cta->eval_CopySign(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { PTXF64 got = cta->getRegAsF64(i, 2); PTXF64 a = cta->getRegAsF64(i, 0); PTXF64 b = cta->getRegAsF64(i, 1); - PTXF64 exp = b; - if (a < 0) { - exp = -std::fabs(b); - } - else { - exp = std::fabs(b); - } + PTXF64 exp = std::copysign(b, a); - if (std::fabs(got - exp) > 0.1) { + if (got != exp) { result = false; - status << "fma.f64 incorrect [" << i << "] - expected: " + status << "copysign.f64 incorrect [" << i << "] - expected: " << exp << ", got " << got << "\n"; break; @@ -2586,33 +4706,80 @@ class TestInstructions: public Test { bool test_Ex2() { bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_ex2() {\n" + << " .reg .f32 f0, f1;\n" + << " ex2.approx.f32 f0, f1;\n" + << " ex2.approx.ftz.f32 f0, f1;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse PTX 8.0 ex2 forms: " << error.what() << "\n"; + return false; + } PTXInstruction ins; ins.opcode = PTXInstruction::Ex2; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + ins.modifier = 0; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx; // f32 // if (result) { - ins.type = PTXOperand::f32; - ins.a = reg("r1", PTXOperand::f32, 0); - ins.d = reg("r3", PTXOperand::f32, 2); - for (int i = 0; i < threadCount; i++) { - cta->setRegAsF32(i, 0, (PTXF32)((float)i / (float)threadCount * 4.0f)); + cta->setRegAsF32(i, 0, + -4.0f + (PTXF32)i / (PTXF32)threadCount * 8.0f); cta->setRegAsF32(i, 2, 0); } cta->eval_Ex2(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - if (std::fabs(cta->getRegAsF32(i, 2) - (PTXF32)exp2((float)i / (float)threadCount * 4.0f)) > 0.1f) { - result = false; - status << "ex2.f32 incorrect [" << i << "] - expected: " - << (PTXF32)exp2((float)i / (float)threadCount * 4.0f) + const PTXF32 value = -4.0f + + (PTXF32)i / (PTXF32)threadCount * 8.0f; + const PTXF32 expected = (PTXF32)std::exp2((double)value); + const PTXU32 actualBits = hydrazine::bit_cast( + cta->getRegAsF32(i, 2)); + const PTXU32 expectedBits = hydrazine::bit_cast(expected); + const PTXU32 ulps = actualBits > expectedBits + ? actualBits - expectedBits : expectedBits - actualBits; + if (ulps > 2) { + result = false; + status << "ex2.f32 incorrect [" << i << "] - expected: " + << expected << ", got " << cta->getRegAsF32(i, 2) << "\n"; break; } } } + cta->setRegAsF32(0, 0, -std::numeric_limits::infinity()); + cta->setRegAsF32(1, 0, -0.0f); + cta->setRegAsF32(2, 0, 0.0f); + cta->setRegAsF32(3, 0, std::numeric_limits::infinity()); + cta->setRegAsF32(4, 0, std::numeric_limits::quiet_NaN()); + cta->eval_Ex2(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0 + || cta->getRegAsF32(1, 2) != 1.0f + || cta->getRegAsF32(2, 2) != 1.0f + || cta->getRegAsF32(3, 2) != std::numeric_limits::infinity() + || !hydrazine::isnan(cta->getRegAsF32(4, 2))) return false; + + cta->setRegAsF32(0, 0, -149.0f); + cta->eval_Ex2(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 1) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->eval_Ex2(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0) return false; + return result; } @@ -2621,6 +4788,7 @@ class TestInstructions: public Test { PTXInstruction ins; ins.opcode = PTXInstruction::Fma; + ins.modifier = PTXInstruction::rn; // f32 // @@ -2630,6 +4798,28 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::f32, 1); ins.c = reg("r4", PTXOperand::f32, 3); ins.d = reg("r3", PTXOperand::f32, 2); + if (!ins.valid().empty()) { + status << "fma.rn.f32 rejected\n"; + return false; + } + PTXInstruction invalid = ins; + invalid.modifier = 0; + if (invalid.valid().empty()) { + status << "fma.f32 accepted without a rounding modifier\n"; + return false; + } + invalid = ins; + invalid.modifier |= PTXInstruction::rz; + if (invalid.valid().empty()) { + status << "fma.f32 accepted multiple rounding modifiers\n"; + return false; + } + invalid = ins; + invalid.c.type = PTXOperand::s32; + if (invalid.valid().empty()) { + status << "fma.f32 accepted a mismatched C operand\n"; + return false; + } for (int i = 0; i < threadCount; i++) { cta->setRegAsF32(i, 0, (PTXF32)((float)i / (float)threadCount * 4.0f)); @@ -2660,17 +4850,29 @@ class TestInstructions: public Test { ins.b = reg("r2", PTXOperand::f64, 1); ins.c = reg("r4", PTXOperand::f64, 3); ins.d = reg("r3", PTXOperand::f64, 2); + PTXInstruction invalid = ins; + invalid.modifier |= PTXInstruction::ftz; + if (invalid.valid().empty()) { + status << "fma.ftz.f64 accepted\n"; + return false; + } + invalid = ins; + invalid.modifier |= PTXInstruction::sat; + if (invalid.valid().empty()) { + status << "fma.sat.f64 accepted\n"; + return false; + } for (int i = 0; i < threadCount; i++) { - cta->setRegAsF64(i, 0, (PTXF32)((double)i / (double)threadCount * 4.5)); - cta->setRegAsF64(i, 1, (PTXF32)((double)i / (double)threadCount * 2.25)); - cta->setRegAsF64(i, 3, (PTXF32)((double)i / (double)threadCount * 0.55)); + cta->setRegAsF64(i, 0, (PTXF64)((double)i / (double)threadCount * 4.5)); + cta->setRegAsF64(i, 1, (PTXF64)((double)i / (double)threadCount * 2.25)); + cta->setRegAsF64(i, 3, (PTXF64)((double)i / (double)threadCount * 0.55)); cta->setRegAsF64(i, 2, 0); } cta->eval_Fma(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - PTXF32 got = cta->getRegAsF32(i, 2); - PTXF32 exp = (double)i / (double)threadCount * 4.5 * (double)i / (double)threadCount * 2.25 + + PTXF64 got = cta->getRegAsF64(i, 2); + PTXF64 exp = (double)i / (double)threadCount * 4.5 * (double)i / (double)threadCount * 2.25 + (double)i / (double)threadCount * 0.55; if (std::fabs(got - exp) > 0.1) { @@ -2683,117 +4885,1624 @@ class TestInstructions: public Test { } } - return result; - } - - bool test_Lg2() { - bool result = true; - - PTXInstruction ins; - ins.opcode = PTXInstruction::Lg2; + const int modes[] = {PTXInstruction::rn, PTXInstruction::rz, + PTXInstruction::rm, PTXInstruction::rp}; + const PTXU32 expected32[][2] = {{0x7f800000, 0xff800000}, + {0x7f7fffff, 0xff7fffff}, {0x7f7fffff, 0xff800000}, + {0x7f800000, 0xff7fffff}}; + ins.type = PTXOperand::f32; + ins.a.type = ins.b.type = ins.c.type = ins.d.type = PTXOperand::f32; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + cta->setRegAsF32(0, 0, std::numeric_limits::max()); + cta->setRegAsF32(0, 1, 2.0f); + cta->setRegAsF32(0, 3, 0.0f); + cta->setRegAsF32(1, 0, -std::numeric_limits::max()); + cta->setRegAsF32(1, 1, 2.0f); + cta->setRegAsF32(1, 3, 0.0f); + cta->setRegAsF32(2, 0, 1.0f + std::ldexp(1.0f, -23)); + cta->setRegAsF32(2, 1, 1.0f - std::ldexp(1.0f, -23)); + cta->setRegAsF32(2, 3, -1.0f); + cta->eval_Fma(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != expected32[i][0] + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != expected32[i][1] + || hydrazine::bit_cast(cta->getRegAsF32(2, 2)) != 0xa8800000) { + status << "fma.f32 fused rounding failed\n"; + return false; + } + } + + const PTXU64 expected64[][2] = {{0x7ff0000000000000ull, 0xfff0000000000000ull}, + {0x7fefffffffffffffull, 0xffefffffffffffffull}, + {0x7fefffffffffffffull, 0xfff0000000000000ull}, + {0x7ff0000000000000ull, 0xffefffffffffffffull}}; + ins.type = PTXOperand::f64; + ins.a.type = ins.b.type = ins.c.type = ins.d.type = PTXOperand::f64; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + cta->setRegAsF64(0, 0, std::numeric_limits::max()); + cta->setRegAsF64(0, 1, 2.0); + cta->setRegAsF64(0, 3, 0.0); + cta->setRegAsF64(1, 0, -std::numeric_limits::max()); + cta->setRegAsF64(1, 1, 2.0); + cta->setRegAsF64(1, 3, 0.0); + cta->setRegAsF64(2, 0, 1.0 + std::ldexp(1.0, -52)); + cta->setRegAsF64(2, 1, 1.0 - std::ldexp(1.0, -52)); + cta->setRegAsF64(2, 3, -1.0); + cta->eval_Fma(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) != expected64[i][0] + || hydrazine::bit_cast(cta->getRegAsF64(1, 2)) != expected64[i][1] + || hydrazine::bit_cast(cta->getRegAsF64(2, 2)) != 0xb970000000000000ull) { + status << "fma.f64 fused rounding failed\n"; + return false; + } + } - // f32 - // - if (result) { - ins.type = PTXOperand::f32; - ins.a = reg("r1", PTXOperand::f32, 0); - ins.d = reg("r3", PTXOperand::f32, 2); + ins.type = PTXOperand::f32; + ins.a.type = ins.b.type = ins.c.type = ins.d.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn | PTXInstruction::ftz; + if (!ins.valid().empty()) { + status << "fma.rn.ftz.f32 rejected\n"; + return false; + } + cta->setRegAsF32(0, 0, std::numeric_limits::min()); + cta->setRegAsF32(0, 1, 0.5f); + cta->setRegAsF32(0, 3, 0.0f); + cta->setRegAsF32(1, 0, std::numeric_limits::denorm_min()); + cta->setRegAsF32(1, 1, std::numeric_limits::max()); + cta->setRegAsF32(1, 3, 0.0f); + cta->setRegAsF32(2, 0, -std::numeric_limits::min()); + cta->setRegAsF32(2, 1, 0.5f); + cta->setRegAsF32(2, 3, -0.0f); + cta->eval_Fma(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0 + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != 0 + || hydrazine::bit_cast(cta->getRegAsF32(2, 2)) != 0x80000000) { + status << "fma.rn.ftz.f32 failed\n"; + return false; + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsF32(i, 0, (PTXF32)(0.5f + (float)i / (float)threadCount * 4.0f)); - cta->setRegAsF32(i, 2, 0); - } - cta->eval_Lg2(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - if (std::fabs(cta->getRegAsF32(i, 2) - (PTXF32)log2(0.5f + (float)i / (float)threadCount * 4.0f)) > 0.1f) { - result = false; - status << "lg2.f32 incorrect [" << i - << "] - log2(" << (0.5f + (float)i / (float)threadCount * 4.0f) << ") - expected: " - << (PTXF32)log2(0.5f + (float)i / (float)threadCount * 4.0f) - << ", got " << cta->getRegAsF32(i, 2) << "\n"; - break; - } - } + ins.modifier = PTXInstruction::rn | PTXInstruction::sat; + cta->setRegAsF32(0, 0, 2.0f); + cta->setRegAsF32(0, 1, 1.0f); + cta->setRegAsF32(0, 3, 0.0f); + cta->setRegAsF32(1, 0, -1.0f); + cta->setRegAsF32(1, 1, 1.0f); + cta->setRegAsF32(1, 3, 0.0f); + cta->setRegAsF32(2, 0, std::numeric_limits::infinity()); + cta->setRegAsF32(2, 1, 0.0f); + cta->setRegAsF32(2, 3, 0.0f); + cta->eval_Fma(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 2) != 1.0f + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) != 0 + || hydrazine::bit_cast(cta->getRegAsF32(2, 2)) != 0) { + status << "fma.rn.sat.f32 failed\n"; + return false; } + return result; } - bool test_Sqrt() { - bool result = true; - + bool test_Bf16Fma() { PTXInstruction ins; - ins.opcode = PTXInstruction::Sqrt; - - double freq = 2.0f / (double)threadCount; - - // f32 - // - if (result) { - ins.type = PTXOperand::f32; - ins.a = reg("r1", PTXOperand::f32, 0); - ins.d = reg("r3", PTXOperand::f32, 2); + ins.opcode = PTXInstruction::Fma; + ins.type = PTXOperand::bf16; + ins.modifier = PTXInstruction::rn; + ins.a = reg("r1", PTXOperand::b16, 0); + ins.b = reg("r2", PTXOperand::b16, 1); + ins.c = reg("r4", PTXOperand::b16, 3); + ins.d = reg("r3", PTXOperand::b16, 2); + + // 1.0 * 1.0078125 + 0.00390625 is halfway between + // 0x3f81 and 0x3f82, so round to the even result 0x3f82. + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, 0x3f80); + cta->setRegAsU16(i, 1, 0x3f81); + cta->setRegAsU16(i, 3, 0x3b80); + cta->setRegAsU16(i, 2, 0); + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsF32(i, 0, (PTXF32)(0.1f + (float)i * freq)); - cta->setRegAsF32(i, 2, 0); - } - cta->eval_Sqrt(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - if (std::fabs(cta->getRegAsF32(i, 2) - (PTXF32)sqrt(0.1f + (float)i * freq)) > 0.1f) { - result = false; - status << "sqrt.f32 incorrect [" << i << "] - expected: " - << (PTXF32)sqrt(0.1f + (float)i * freq) - << ", got " << cta->getRegAsF32(i, 2) << "\n"; - break; - } + cta->eval_Fma(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 2) != 0x3f82) { + status << "fma.rn.bf16 incorrect [" << i << "]\n"; + return false; } } - // f64 - // - if (result) { - ins.type = PTXOperand::f64; - ins.a = reg("r1", PTXOperand::f64, 0); - ins.d = reg("r3", PTXOperand::f64, 2); + auto scalarFma = [&](int modifier, PTXU16 a, PTXU16 b, PTXU16 c, + PTXU16 expected) { + ins.modifier = modifier; + cta->setRegAsU16(0, 0, a); cta->setRegAsU16(0, 1, b); + cta->setRegAsU16(0, 3, c); cta->setRegAsU16(0, 2, 0xbeef); + cta->eval_Fma(cta->getActiveContext(), ins); + return cta->getRegAsU16(0, 2) == expected; + }; + if (!scalarFma(PTXInstruction::rn, 0x3fe0, 0x3f14, 0xab80, 0x3f81) + || !scalarFma(PTXInstruction::rn, 0xbfe0, 0x3f14, 0x2b80, 0xbf81) + || !scalarFma(PTXInstruction::rn | PTXInstruction::relu, + 0xbf80, 0x3f80, 0, 0x0000) + || !scalarFma(PTXInstruction::rn | PTXInstruction::relu, + 0x7fc0, 0x3f80, 0, 0x7fff) + || !scalarFma(PTXInstruction::rn, 0x7f7f, 0x4000, 0, 0x7f80)) { + status << "fma.bf16 scalar case failed\n"; + return false; + } + auto checkValid = [&](int modifier, bool shouldAccept) { + ins.modifier = modifier; + return ins.valid().empty() == shouldAccept; + }; + if (!checkValid(PTXInstruction::rn, true) + || !checkValid(PTXInstruction::rz, false) + || !checkValid(PTXInstruction::rn | PTXInstruction::ftz, false) + || !checkValid(PTXInstruction::rn | PTXInstruction::sat, false) + || !checkValid(PTXInstruction::rn | PTXInstruction::relu, true)) { + status << "fma.bf16 validation failed\n"; + return false; + } + ins.modifier = PTXInstruction::rn; + ins.type = PTXOperand::bf16x2; + ins.a = reg("r1", PTXOperand::b32, 0); ins.b = reg("r2", PTXOperand::b32, 1); + ins.c = reg("r4", PTXOperand::b32, 3); ins.d = reg("r3", PTXOperand::b32, 2); + if (!checkValid(PTXInstruction::rn, true)) { + status << "fma.rn.bf16x2 rejected\n"; + return false; + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsF64(i, 0, (PTXF64)(0.1f + (double)i * freq)); - cta->setRegAsF64(i, 2, 0); - } - cta->eval_Sqrt(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - if (std::fabs(cta->getRegAsF64(i, 2) - sqrt(0.1 + (double)i * freq)) > 0.1f) { - result = false; - status << "sqrt.f64 incorrect [" << i << "] - expected: " << sqrt(0.1 + (double)i * freq) - << ", got " << cta->getRegAsF64(i, 2) << "\n"; - break; - } - } + auto packedFma = [&](int modifier, PTXU32 a, PTXU32 b, PTXU32 c, + PTXU32 expected) { + ins.modifier = modifier; + cta->setRegAsU32(0, 0, a); cta->setRegAsU32(0, 1, b); + cta->setRegAsU32(0, 3, c); cta->setRegAsU32(0, 2, 0xdeadbeef); + cta->eval_Fma(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, 2) == expected; + }; + if (!packedFma(PTXInstruction::rn, 0xbfe03fe0, 0x3f143f14, + 0x2b80ab80, 0xbf813f81) + || !packedFma(PTXInstruction::rn | PTXInstruction::relu, + 0x3f80bf80, 0x3f803f80, 0, 0x3f800000)) { + status << "fma.bf16x2 packed case failed\n"; + return false; } - return result; + return true; } - bool test_Rsqrt() { - bool result = true; - + bool test_F16Fma() { PTXInstruction ins; - ins.opcode = PTXInstruction::Rsqrt; - - double freq = 2.0f / (double)threadCount; - - // f32 - // - if (result) { - ins.type = PTXOperand::f32; - ins.a = reg("r1", PTXOperand::f32, 0); - ins.d = reg("r3", PTXOperand::f32, 2); + ins.opcode = PTXInstruction::Fma; + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn; + ins.a = reg("a", PTXOperand::b16, 0); + ins.b = reg("b", PTXOperand::b16, 1); + ins.c = reg("c", PTXOperand::b16, 3); + ins.d = reg("d", PTXOperand::b16, 2); + + // 1.0 * (1.0 + 2^-10) - 2^-11 is halfway between + // 1.0 and the next half value, so round to the even result 1.0. + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU16(i, 0, 0x3c00); + cta->setRegAsU16(i, 1, 0x3c01); + cta->setRegAsU16(i, 3, 0x9000); + cta->setRegAsU16(i, 2, 0); + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsF32(i, 0, (PTXF32)(0.1f + (float)i * freq)); - cta->setRegAsF32(i, 2, 0); + cta->eval_Fma(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 2) != 0x3c00) { + status << "fma.rn.f16 incorrect [" << i << "]\n"; + return false; } - cta->eval_Rsqrt(cta->getActiveContext(), ins); + } + + auto scalarFma = [&](int modifier, PTXU16 a, PTXU16 b, PTXU16 c, + PTXU16 expected) { + ins.modifier = modifier; + cta->setRegAsU16(0, 0, a); cta->setRegAsU16(0, 1, b); + cta->setRegAsU16(0, 3, c); cta->setRegAsU16(0, 2, 0xbeef); + cta->eval_Fma(cta->getActiveContext(), ins); + return cta->getRegAsU16(0, 2) == expected; + }; + if (!scalarFma(PTXInstruction::rn, 0x3bfe, 0x9001, 0x3c01, 0x3c01) + || !scalarFma(PTXInstruction::rn, 0x0001, 0x3c00, 0, 0x0001) + || !scalarFma(PTXInstruction::rn | PTXInstruction::ftz, + 0x0001, 0x3c00, 0, 0x0000) + || !scalarFma(PTXInstruction::rn | PTXInstruction::ftz, + 0x0400, 0x3800, 0, 0x0000) + || !scalarFma(PTXInstruction::rn | PTXInstruction::ftz, + 0x8400, 0x3800, 0, 0x8000) + || !scalarFma(PTXInstruction::rn | PTXInstruction::sat, + 0x4000, 0x4000, 0, 0x3c00) + || !scalarFma(PTXInstruction::rn | PTXInstruction::relu, + 0xbc00, 0x3c00, 0, 0x0000) + || !scalarFma(PTXInstruction::rn | PTXInstruction::relu, + 0x7e00, 0x3c00, 0, 0x7fff)) { + status << "fma.f16 scalar case failed\n"; + return false; + } + ins.modifier = PTXInstruction::rn; + if (!ins.valid().empty()) { + status << "fma.rn.f16 rejected\n"; + return false; + } + ins.modifier = PTXInstruction::rz; + if (ins.valid().empty()) { + status << "fma.rz.f16 accepted\n"; + return false; + } + ins.modifier = PTXInstruction::rn | PTXInstruction::sat + | PTXInstruction::relu; + if (ins.valid().empty()) { + status << "fma.rn.sat.relu.f16 accepted\n"; + return false; + } + ins.modifier = PTXInstruction::rn | PTXInstruction::ftz + | PTXInstruction::relu; + if (!ins.valid().empty()) { + status << "fma.rn.ftz.relu.f16 rejected\n"; + return false; + } + ins.modifier = PTXInstruction::rn; + + ins.type = PTXOperand::f16x2; + ins.a = reg("a", PTXOperand::b32, 0); + ins.b = reg("b", PTXOperand::b32, 1); + ins.c = reg("c", PTXOperand::b32, 3); + ins.d = reg("d", PTXOperand::b32, 2); + if (!ins.valid().empty()) { + status << "fma.rn.f16x2 rejected\n"; + return false; + } + ins.modifier = PTXInstruction::rm; + if (ins.valid().empty()) { + status << "fma.rm.f16x2 accepted\n"; + return false; + } + ins.modifier = PTXInstruction::rn; + cta->reset(); + auto packedFma = [&](int modifier, PTXU32 a, PTXU32 b, PTXU32 c, + PTXU32 expected, bool alias) { + ins.modifier = modifier; + ins.d.reg = alias ? 0 : 2; + cta->setRegAsU32(0, 0, a); cta->setRegAsU32(0, 1, b); + cta->setRegAsU32(0, 3, c); cta->setRegAsU32(0, 2, 0xdeadbeef); + cta->eval_Fma(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, ins.d.reg) == expected; + }; + if (!packedFma(PTXInstruction::rn, 0x40003c01, 0x42003c01, + 0x3c00bc02, 0x47000010, false) + || !packedFma(PTXInstruction::rn, 0x3c033c01, 0x3e003e00, + 0x00018001, 0x3e053e01, false) + || !packedFma(PTXInstruction::rn | PTXInstruction::ftz, + 0x40003c01, 0x42003c01, 0x3c00bc02, 0x47000000, false) + || !packedFma(PTXInstruction::rn | PTXInstruction::ftz, + 0x3c033c01, 0x3e003e00, 0x00018001, 0x3e043e02, false) + || !packedFma(PTXInstruction::rn, 0x00010001, 0x3c003c00, + 0, 0x00010001, false) + || !packedFma(PTXInstruction::rn | PTXInstruction::ftz, + 0x00010001, 0x3c003c00, 0, 0, false) + || !packedFma(PTXInstruction::rn | PTXInstruction::sat, + 0x7e004000, 0x3c003c00, 0, 0x00003c00, false) + || !packedFma(PTXInstruction::rn | PTXInstruction::relu, + 0xbc007e00, 0x3c003c00, 0, 0x00007fff, false) + || !packedFma(PTXInstruction::rn, 0x40003c01, 0x42003c01, + 0x3c00bc02, 0x47000010, true)) return false; + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + const bool hostRn = packedFma(PTXInstruction::rn, 0x3c003c00, + 0x3c013c01, 0x90009000, 0x3c003c00, false); + const bool restored = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(previous); + if (!hostRn || !restored) return false; + ins.modifier = PTXInstruction::rn; + ins.d.reg = 2; + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 4; + cta->setRegAsPredicate(0, 4, false); + cta->setRegAsU32(0, 2, 0xcafebabe); + cta->eval_Fma(cta->getActiveContext(), ins); + ins.pg.condition = PTXOperand::PT; + if (cta->getRegAsU32(0, 2) != 0xcafebabe) return false; + + return true; + } + + bool test_MmaWarpParticipation() { + for (const char* form : { + "m8n8k128.row.col.s32.b1.b1.s32.xor.popc {%r0,%r1},{%r4},{%r5},{%r2,%r3}", + "m8n8k4.row.col.f64.f64.f64.f64 {%r0,%r1},{%r4},{%r5},{%r2,%r3}", + "m8n8k16.row.col.s32.s8.s8.s32 {%r0,%r1},{%r4},{%r5},{%r2,%r3}", + "m8n8k4.row.col.f16.f16.f16.f16 {%r0,%r1,%r2,%r3},{%r8,%r9},{%r10,%r11},{%r4,%r5,%r6,%r7}", + "m16n8k8.row.col.f32.f16.f16.f32 {%r0,%r1,%r2,%r3},{%r8,%r9},{%r10},{%r4,%r5,%r6,%r7}"}) { + const bool fp64 = std::string(form).find("f64") != std::string::npos; + std::stringstream source(std::string(".version 8.0\n.target sm_86\n.entry warp_mma() { .reg .pred %p; .reg .") + + (fp64 ? "b64" : "b32") + " %r<12>; @%p mma.sync.aligned." + form + "; ret; }"); + Module parsed; parsed.load(source); + EmulatedKernel localKernel(parsed.getKernel("warp_mma"), nullptr); + for (int size : {33, 65, 1, 31}) { + localKernel.setKernelShape(size, 1, 1); + CooperativeThreadArray localCTA(&localKernel, Dim3(), false); + for (const auto& ins : localKernel.instructions) if (ins.opcode == PTXInstruction::Mma) { + const int completeThreads = (size / 32) * 32; + for (int scenario : {1, 0, 2, 3, 4}) { + if (scenario == 3 && completeThreads == 0) continue; + localCTA.reset(); + auto& context = localCTA.getActiveContext(); + for (int thread = 0; thread < size; ++thread) { + for (unsigned regID = 0; regID < localKernel.registerCount(); ++regID) localCTA.setRegAsB64(thread, regID, 0); + for (const auto& dest : ins.d.array) localCTA.setRegAsB64(thread, dest.reg, 0xdeadbeef); + const bool predicate = scenario == 2 || scenario == 4 || + (scenario != 0 && thread < completeThreads && (scenario != 3 || thread != 0)); + localCTA.setRegAsPredicate(thread, ins.pg.reg, predicate); + if (scenario == 4 && thread >= completeThreads) context.active.reset(thread); + } + bool rejected = false; + try { localCTA.eval_Mma(context, ins); } catch (const RuntimeException&) { rejected = true; } + if (rejected != (scenario == 2 || scenario == 3)) { + status << "mma participation mismatch: " << form << " CTA " << size << " scenario " << scenario << "\n"; + return false; + } + if (!rejected) for (int thread = 0; thread < size; ++thread) for (const auto& dest : ins.d.array) { + const PTXU64 expected = scenario != 0 && thread < completeThreads ? 0 : 0xdeadbeef; + if (localCTA.getRegAsB64(thread, dest.reg) != expected) { status << "mma inactive lane changed or active result incorrect\n"; return false; } + } + } + } + } + } + return true; + } + + bool test_MmaFloatShapes() { + for (const char* shape : {"m8n8k128", "m16n8k128", "m16n8k256"}) for (const char* op : {"xor", "and"}) { + const bool small = std::string(shape) == "m8n8k128", k256 = std::string(shape) == "m16n8k256"; + const std::string d = small ? "{%r0,%r1}" : "{%r0,%r1,%r2,%r3}"; + const std::string a = small ? "{%r4}" : k256 ? "{%r4,%r5,%r6,%r7}" : "{%r4,%r5}"; + const std::string b = k256 ? "{%r8,%r9}" : "{%r8}"; + const std::string c = small ? "{%r10,%r11}" : "{%r10,%r11,%r12,%r13}"; + const std::string mnemonic = std::string("mma.sync.aligned.") + shape + ".row.col.s32.b1.b1.s32." + op + ".popc"; + const std::string source = ".version 8.0\n.target sm_86\n.entry binary_mma() { .reg .b32 %r<14>; " + mnemonic + " " + d + ", " + a + ", " + b + ", " + c + "; ret; }"; + std::stringstream input(source); Module parsed; parsed.load(input); + std::stringstream printed; parsed.write(printed); Module roundtrip; roundtrip.load(printed); + if (printed.str().find(mnemonic) == std::string::npos) return false; + for (const auto& statement : parsed.statements()) if (statement.instruction.opcode == PTXInstruction::Mma) { + const auto& ins = statement.instruction; + PTXInstruction invalid = ins; invalid.booleanOperator = PTXInstruction::BoolOr; + if (invalid.valid().empty()) return false; + if (ins.a.type != PTXOperand::b1 || ins.booleanOperator != (std::string(op) == "xor" ? PTXInstruction::BoolXor : PTXInstruction::BoolAnd)) return false; + PTXInstruction run = ins; + for (auto* fragment : {&run.a, &run.b, &run.c, &run.d}) + for (auto& element : fragment->array) element.reg = std::stoi(element.identifier.substr(2)); + run.c = run.d; + auto aBit = [](int row, int k) { return (row * 13 + k * 7 + (k / 32) * 3) % 11 < 5; }; + auto bBit = [](int k, int col) { return (k * 5 + col * 3 + k / 64) % 13 < 6; }; + for (PTXU32 base : {0x7fffffc0u, 0xffffffe0u}) { + cta->reset(); + for (int lane = 0; lane < 32; ++lane) { + const int group = lane / 4, t = lane % 4; + for (unsigned i = 0; i < run.a.array.size(); ++i) { + PTXU32 packed = 0; + for (int bit = 0; bit < 32; ++bit) packed |= PTXU32(aBit(group + (i & 1) * 8, (t + (i / 2) * 4) * 32 + bit)) << bit; + cta->setRegAsB32(lane, run.a.array[i].reg, packed); + } + for (unsigned i = 0; i < run.b.array.size(); ++i) { + PTXU32 packed = 0; + for (int bit = 0; bit < 32; ++bit) packed |= PTXU32(bBit((t + i * 4) * 32 + bit, group)) << bit; + cta->setRegAsB32(lane, run.b.array[i].reg, packed); + } + for (unsigned i = 0; i < run.c.array.size(); ++i) + cta->setRegAsB32(lane, run.c.array[i].reg, base + (group + (i / 2) * 8) * 8 + 2 * t + (i & 1)); + } + cta->eval_Mma(cta->getActiveContext(), run); + for (int lane = 0; lane < 32; ++lane) for (unsigned i = 0; i < run.d.array.size(); ++i) { + const int row = lane / 4 + (i / 2) * 8, col = (lane % 4) * 2 + (i & 1); + PTXU32 expected = base + row * 8 + col; + for (int k = 0; k < (k256 ? 256 : 128); ++k) + expected += run.booleanOperator == PTXInstruction::BoolXor ? aBit(row, k) != bBit(k, col) : aBit(row, k) && bBit(k, col); + if (cta->getRegAsB32(lane, run.d.array[i].reg) != expected) { status << "binary mma mismatch " << shape << " " << op << " lane " << lane << " register " << i << "\n"; return false; } + } + } + } + for (const auto& replacement : {std::make_pair(std::string(".popc"), std::string("")), + std::make_pair(std::string(".") + op + ".popc", std::string(".or.popc")), + std::make_pair(std::string("row.col"), std::string("col.col")), + std::make_pair(std::string("row.col"), std::string("row.col.satfinite")), + std::make_pair(std::string(shape), std::string("m8n8k16")), + std::make_pair(std::string(".reg .b32"), std::string(".reg .b1")), + std::make_pair(d, std::string("{%r0}"))}) { + std::string bad = source; bad.replace(bad.find(replacement.first), replacement.first.size(), replacement.second); + std::stringstream badInput(bad); bool rejected = false; + try { Module module; module.load(badInput); } catch (const std::exception&) { rejected = true; } + if (!rejected) { status << "invalid binary mma accepted: " << bad << "\n"; return false; } + } + } + // Every fragment must retain its exact register count and register mode. + for (const char* operation : { + "m8n8k4.row.col.f64.f64.f64.f64", + "m8n8k4.row.col.f16.f16.f16.f16", + "m16n8k16.row.col.f32.f16.f16.f32", + "m16n8k4.row.col.f32.tf32.tf32.f32", + "m8n8k16.row.col.s32.s8.u8.s32", + "m8n8k32.row.col.s32.s4.u4.s32", + "m8n8k128.row.col.s32.b1.b1.s32.xor.popc"}) { + const std::string op(operation); + const bool fp64 = op.find("f64") != std::string::npos; + const bool legacy = op.find("f16.f16.f16.f16") != std::string::npos; + const bool modern = op.find("m16n8") == 0; + const std::string a = op.find("m16n8k16") == 0 ? "{%r4,%r5,%r11,%r12}" : + modern || legacy ? "{%r4,%r5}" : "{%r4}"; + const std::string d = modern || legacy ? "{%r0,%r1,%r2,%r3}" : "{%r0,%r1}"; + const std::string b = op.find("m16n8k16") == 0 || legacy ? "{%r6,%r13}" : "{%r6}"; + const std::string c = modern || legacy ? "{%r7,%r8,%r9,%r10}" : "{%r7,%r8}"; + const std::string instruction = "mma.sync.aligned." + op + " " + d + ", " + a + ", " + b + ", " + c + ";"; + const std::string prefix = ".version 8.0\n.target sm_86\n.global ." + std::string(fp64 ? "b64" : "b32") + + " ga;\n.entry fragments() { .reg ." + (fp64 ? "b64" : "b32") + " %r<14>; " + + ".reg .v2 ." + (fp64 ? "b64" : "b32") + " va; " + + ".reg .v4 ." + (fp64 ? "b64" : "b32") + " vb; "; + std::stringstream validSource(prefix + instruction + " ret; }"); + Module validModule; validModule.load(validSource); + for (const std::string& fragment : {d, a, b, c}) { + for (int malformed = 0; malformed < 4; ++malformed) { + std::string bad = instruction; + const auto position = bad.find(fragment); + if (position == std::string::npos) continue; + if (malformed == 0) bad.insert(position + fragment.size() - 1, ",%r13"); + else bad.replace(position + 1, 3, malformed == 1 ? "ga" : malformed == 2 ? "va" : "vb"); + std::stringstream source(prefix + bad + " ret; }"); + bool rejected = false; + try { Module module; module.load(source); } + catch (const std::exception&) { rejected = true; } + if (!rejected) { status << "invalid mma fragment accepted: " << bad << "\n"; return false; } + } + } + } + // PTX 8.0: m8n8k16/m16n8k32/m16n8k64 are integer-only shapes. + for (const char* shape : {"m16n8k4", "m16n8k8", "m16n8k16", "m8n8k16", "m16n8k32", "m16n8k64", "m8n8k128", "m16n8k128", "m16n8k256"}) { + for (const char* input : {"f16", "bf16", "tf32"}) { + for (const char* accumulator : {"f16", "f32"}) { + const bool small = std::string(shape) == "m16n8k8"; + const bool k4 = std::string(shape) == "m16n8k4"; + const bool tf32 = std::string(input) == "tf32"; + const bool half = std::string(accumulator) == "f16"; + const bool legalShape = k4 ? tf32 : small || std::string(shape) == "m16n8k16"; + const bool legal = legalShape && (!tf32 || small || k4) + && (!half || std::string(input) == "f16"); + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry mma_shapes() {\n.reg ." + << (half ? "f16x2" : "b32") << " %d<4>;\n.reg ." + << (tf32 ? "b32" : "f16x2") << " %r<10>;\n" + << ".reg ." << (half ? "f16x2" : "b32") << " %c<4>;\n" + << "mma.sync.aligned." << shape << ".row.col." + << accumulator << "." << input << "." << input << "." + << accumulator << (half ? " {%d0,%d1}, " : " {%d0,%d1,%d2,%d3}, ") + << (k4 || (small && !tf32) ? "{%r4,%r5}, {%r8}, " + : "{%r4,%r5,%r6,%r7}, {%r8,%r9}, ") + << (half ? "{%c0,%c1};\n" : "{%c0,%c1,%c2,%c3};\n") + << "ret;\n}\n"; + Module parsed; + bool accepted = true; + try { + parsed.load(ptx); + std::stringstream printed; parsed.write(printed); + Module roundtrip; roundtrip.load(printed); + } + catch (const std::exception&) { accepted = false; } + if (accepted != legal) { + status << "mma " << shape << " " << input << " -> " + << accumulator << (accepted ? " incorrectly accepted\n" + : " incorrectly rejected\n"); + return false; + } + } + } + } + for (const char* a : {"s4", "u4"}) for (const char* b : {"s4", "u4"}) { + for (const char* saturation : {"", "satfinite."}) { + const std::string instruction = std::string("mma.sync.aligned.m8n8k32.row.col.") + saturation + + "s32." + a + "." + b + ".s32 {%r0,%r1}, {%r2}, {%r3}, {%r4,%r5};"; + const std::string source = ".version 8.0\n.target sm_86\n.entry int4() { .reg .b32 %r<6>; " + instruction + " ret; }"; + std::stringstream stream(source); Module parsed; parsed.load(stream); + std::stringstream printed; parsed.write(printed); Module roundtrip; roundtrip.load(printed); + if (printed.str().find(instruction.substr(0, instruction.find(' '))) == std::string::npos) return false; + const std::string wideInstruction = std::string("mma.sync.aligned.m16n8k32.row.col.") + saturation + + "s32." + a + "." + b + ".s32 {%r0,%r1,%r2,%r3}, {%r4,%r5}, {%r6}, {%r7,%r8,%r9,%r10};"; + std::stringstream wideSource(".version 8.0\n.target sm_86\n.entry int4_wide() { .reg .b32 %r<11>; " + wideInstruction + " ret; }"); + Module wideParsed; wideParsed.load(wideSource); + std::stringstream widePrinted; wideParsed.write(widePrinted); Module wideRoundtrip; wideRoundtrip.load(widePrinted); + if (widePrinted.str().find(wideInstruction.substr(0, wideInstruction.find(' '))) == std::string::npos) return false; + const std::string k64Instruction = std::string("mma.sync.aligned.m16n8k64.row.col.") + saturation + + "s32." + a + "." + b + ".s32 {%r0,%r1,%r2,%r3}, {%r4,%r5,%r6,%r7}, {%r8,%r9}, {%r10,%r11,%r12,%r13};"; + std::stringstream k64Source(".version 8.0\n.target sm_86\n.entry int4_k64() { .reg .b32 %r<14>; " + k64Instruction + " ret; }"); + Module k64Parsed; k64Parsed.load(k64Source); + std::stringstream k64Printed; k64Parsed.write(k64Printed); Module k64Roundtrip; k64Roundtrip.load(k64Printed); + if (k64Printed.str().find(k64Instruction.substr(0, k64Instruction.find(' '))) == std::string::npos) return false; + for (const char* invalid : {"m8n8k16", "m16n8k32"}) { + std::string bad = source; + bad.replace(bad.find("m8n8k32"), 7, invalid); + std::stringstream input(bad); bool rejected = false; + try { Module module; module.load(input); } + catch (const std::exception&) { rejected = true; } + if (!rejected) { status << "invalid sub-byte mma shape/fragments accepted\n"; return false; } + } + } + } + return true; + } + + bool test_Mma() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Mma; + ins.mmaShape = PTXInstruction::MmaM16N8K16; + ins.type = PTXOperand::f32; + ins.modifier = 0; + + auto vector = [this](PTXOperand::DataType type, + PTXOperand::DataType elementType, PTXOperand::Vec vec, + int firstRegister, int count) { + PTXOperand operand; + operand.addressMode = PTXOperand::Register; + operand.type = type; + operand.vec = vec; + for(int i = 0; i < count; ++i) { + operand.array.push_back(reg("mma", elementType, + (PTXOperand::RegisterType)(firstRegister + i))); + } + return operand; + }; + + ins.d = vector(PTXOperand::f32, PTXOperand::f32, + PTXOperand::v4, 0, 4); + ins.c = vector(PTXOperand::f32, PTXOperand::f32, + PTXOperand::v4, 6, 4); + ins.a = vector(PTXOperand::f16, PTXOperand::b32, + PTXOperand::v4, 0, 4); + ins.b = vector(PTXOperand::f16, PTXOperand::b32, + PTXOperand::v2, 4, 2); + for (auto& operand : ins.a.array) operand.type = PTXOperand::f16x2; + for (auto& operand : ins.b.array) operand.type = PTXOperand::f16x2; + if (!ins.valid().empty()) { + status << "mma rejected f16x2 input fragments: " << ins.valid() << "\n"; + return false; + } + auto rejects = [&](const PTXInstruction& invalid) { + bool rejected = false; + try { cta->eval_Mma(cta->getActiveContext(), invalid); } + catch (const RuntimeException&) { rejected = true; } + if (!rejected || invalid.valid().empty()) { + status << "invalid mma accepted (shape " << invalid.mmaShape + << "): " << invalid.toString() << "\n"; + return false; + } + return true; + }; + for (int fragment = 0; fragment < 4; ++fragment) { + PTXInstruction invalid = ins; + PTXOperand* operands[] = {&invalid.a, &invalid.b, &invalid.c, &invalid.d}; + operands[fragment]->addressMode = PTXOperand::Address; + if (!rejects(invalid)) return false; + for (int malformed = 0; malformed < 2; ++malformed) { + invalid = ins; + auto& element = operands[fragment]->array[0]; + if (malformed == 0) element.vec = PTXOperand::v2; + else element.array.push_back(element); + if (!rejects(invalid)) return false; + } + } + for (auto shape : {PTXInstruction::MmaShape_Invalid, + PTXInstruction::MmaM8N8K16, PTXInstruction::MmaM16N8K32}) { + PTXInstruction invalid = ins; + invalid.mmaShape = shape; + if (!rejects(invalid)) return false; + } + for (int bad = 0; bad < 4; ++bad) { + PTXInstruction invalid = ins; + if (bad == 0) invalid.a.type = PTXOperand::s32; + if (bad == 1) invalid.b.type = PTXOperand::bf16; + if (bad == 2) invalid.modifier = PTXInstruction::rn; + if (bad == 3) invalid.modifier = PTXInstruction::satfinite; + if (!rejects(invalid)) return false; + } + PTXInstruction tf32 = ins; + tf32.mmaShape = PTXInstruction::MmaM16N8K8; + tf32.a.type = tf32.b.type = PTXOperand::tf32; + if (tf32.valid().empty()) { + status << "tf32 mma incorrectly accepted f16x2 fragments\n"; + return false; + } + + const PTXU16 f16Values[17] = { + 0x0000, 0x3c00, 0x4000, 0x4200, 0x4400, 0x4500, + 0x4600, 0x4700, 0x4800, 0x4880, 0x4900, 0x4980, + 0x4a00, 0x4a80, 0x4b00, 0x4b80, 0x4c00 + }; + cta->reset(); + for(int thread = 0; thread < threadCount; ++thread) { + const int lane = thread & 31; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for(int regIndex = 0; regIndex < 4; ++regIndex) { + PTXU32 packed = 0; + for(int half = 0; half < 2; ++half) { + const int i = regIndex * 2 + half; + const int row = (i < 2 || (i >= 4 && i < 6)) + ? groupID : groupID + 8; + const int value = row + 1; + packed |= (PTXU32)f16Values[value] << (16 * half); + } + cta->setRegAsB32(thread, regIndex, packed); + } + for(int regIndex = 0; regIndex < 2; ++regIndex) { + PTXU32 packed = 0; + for(int half = 0; half < 2; ++half) { + const int col = groupID; + packed |= (PTXU32)f16Values[col + 1] << (16 * half); + } + cta->setRegAsB32(thread, 4 + regIndex, packed); + } + for(int i = 0; i < 4; ++i) { + const int row = groupID + (i >= 2 ? 8 : 0); + const int col = threadInGroup * 2 + (i & 1); + cta->setRegAsF32(thread, 6 + i, (PTXF32)(100 * row + col)); + } + } + + cta->eval_Mma(cta->getActiveContext(), ins); + for(int thread = 0; thread < threadCount; ++thread) { + const int lane = thread & 31; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for(int regIndex = 0; regIndex < 4; ++regIndex) { + const int row = groupID + (regIndex >= 2 ? 8 : 0); + const int col = threadInGroup * 2 + (regIndex & 1); + const PTXF32 expected = 16.0f * (row + 1) * (col + 1) + + (PTXF32)(100 * row + col); + if(std::fabs(cta->getRegAsF32(thread, regIndex) - expected) > 0.001f) { + status << "mma.m16n8k16.f16 incorrect [" + << thread << "]\n"; + return false; + } + } + } + + ins.a.type = PTXOperand::bf16; + ins.b.type = PTXOperand::bf16; + const PTXU32 bf16One = 0x3f803f80u; + const PTXU32 bf16Two = 0x40004000u; + cta->reset(); + for(int thread = 0; thread < threadCount; ++thread) { + for(int regIndex = 0; regIndex < 4; ++regIndex) { + cta->setRegAsB32(thread, regIndex, bf16One); + } + for(int regIndex = 4; regIndex < 6; ++regIndex) { + cta->setRegAsB32(thread, regIndex, bf16Two); + } + for(int regIndex = 6; regIndex < 10; ++regIndex) { + cta->setRegAsF32(thread, regIndex, 3.0f); + } + } + + cta->eval_Mma(cta->getActiveContext(), ins); + for(int thread = 0; thread < threadCount; ++thread) { + for(int regIndex = 0; regIndex < 4; ++regIndex) { + if(std::fabs(cta->getRegAsF32(thread, regIndex) - 35.0f) > 0.001f) { + status << "mma.m16n8k16.bf16 incorrect [" + << thread << "]\n"; + return false; + } + } + } + + // MMA is warp-collective: a partially active warp must not execute it. + cta->getActiveContext().active[0] = false; + bool rejectedPartialWarp = false; + try { + cta->eval_Mma(cta->getActiveContext(), ins); + } + catch (RuntimeException &) { + rejectedPartialWarp = true; + } + cta->getActiveContext().active[0] = true; + if (!rejectedPartialWarp) { + status << "mma.m16n8k16 accepted a partial warp\n"; + return false; + } + + // TF32 m16n8k4: A(r,k)=r-2k+1, B(k,c)=3k+c-2. + ins.mmaShape = PTXInstruction::MmaM16N8K4; + ins.a = vector(PTXOperand::tf32, PTXOperand::b32, PTXOperand::v2, 0, 2); + ins.b = vector(PTXOperand::tf32, PTXOperand::b32, PTXOperand::v1, 4, 1); + if (!ins.valid().empty()) { status << ins.valid() << "\n"; return false; } + cta->reset(); + for (int lane = 0; lane < 32; ++lane) { + const int group = lane >> 2, t = lane & 3; + cta->setRegAsF32(lane, 0, group - 2 * t + 1); + cta->setRegAsF32(lane, 1, group + 8 - 2 * t + 1); + cta->setRegAsF32(lane, 4, 3 * t + group - 2); + for (int i = 0; i < 4; ++i) + cta->setRegAsF32(lane, 6 + i, (group + (i >= 2 ? 8 : 0)) * 8 + t * 2 + (i & 1)); + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) { + for (int i = 0; i < 4; ++i) { + const int row = (lane >> 2) + (i >= 2 ? 8 : 0), col = (lane & 3) * 2 + (i & 1); + int expected = row * 8 + col; + for (int k = 0; k < 4; ++k) expected += (row - 2 * k + 1) * (3 * k + col - 2); + if (cta->getRegAsF32(lane, i) != expected) { + status << "mma.m16n8k4.tf32 incorrect lane " << lane << " register " << i << "\n"; + return false; + } + } + } + ins.mmaShape = PTXInstruction::MmaM8N8K4; + ins.type = PTXOperand::f64; + ins.a = vector(PTXOperand::f64, PTXOperand::f64, PTXOperand::v1, 0, 1); + ins.b = vector(PTXOperand::f64, PTXOperand::f64, PTXOperand::v1, 4, 1); + ins.c = vector(PTXOperand::f64, PTXOperand::f64, PTXOperand::v2, 6, 2); + ins.d = vector(PTXOperand::f64, PTXOperand::f64, PTXOperand::v2, 0, 2); + for (auto* operand : {&ins.a, &ins.b, &ins.c, &ins.d}) for (auto& element : operand->array) element.identifier = "%f" + std::to_string(element.reg); + PTXInstruction conflictingRounding = ins; + conflictingRounding.modifier = PTXInstruction::rn | PTXInstruction::rz; + if (!rejects(conflictingRounding)) return false; + for (int mode : {0, (int)PTXInstruction::rn, (int)PTXInstruction::rz, (int)PTXInstruction::rm, (int)PTXInstruction::rp}) { + ins.modifier = mode; + if (!ins.valid().empty()) { status << ins.valid() << "\n"; return false; } + std::stringstream source; + source << ".version 8.0\n.target sm_86\n.entry fp64() { .reg .f64 %f<8>; " << ins.toString() << "; ret; }"; + Module parsed; parsed.load(source); + std::stringstream printed; parsed.write(printed); Module roundtrip; roundtrip.load(printed); + for (int scenario : {0, 1, -1}) { + cta->reset(); + for (int lane = 0; lane < 32; ++lane) { + const int row = lane >> 2, k = lane & 3; + cta->setRegAsF64(lane, 0, scenario ? (k == 0 ? scenario * 0x1p-53 : 0) : row - 2 * k + 1); + cta->setRegAsF64(lane, 4, scenario ? 1 : 3 * k + row - 2); + for (int i = 0; i < 2; ++i) cta->setRegAsF64(lane, 6 + i, scenario ? scenario : row * 8 + 2 * k + i); + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) for (int i = 0; i < 2; ++i) { + const int row = lane >> 2, col = 2 * (lane & 3) + i; + double expected = scenario ? scenario : row * 8 + col; + if (!scenario) for (int k = 0; k < 4; ++k) expected += (row - 2 * k + 1) * (3 * k + col - 2); + else if ((scenario > 0 && mode == PTXInstruction::rp) || (scenario < 0 && mode == PTXInstruction::rm)) expected = std::nextafter(expected, 2 * expected); + if (cta->getRegAsF64(lane, i) != expected) { status << "f64 mma mismatch lane " << lane << " mode " << mode << "\n"; return false; } + } + } + } + + PTXInstruction wrongLayout = ins; + wrongLayout.mmaAColumnMajor = true; + if (wrongLayout.valid().empty()) { status << "f64 mma accepted col.col\n"; return false; } + ins.mmaShape = PTXInstruction::MmaM16N8K8; + bool rejectedF64Shape = false; + try { cta->eval_Mma(cta->getActiveContext(), ins); } + catch (RuntimeException&) { rejectedF64Shape = true; } + if (!rejectedF64Shape) { status << "unsupported f64 mma shape executed\n"; return false; } + ins.mmaShape = PTXInstruction::MmaM8N8K4; + + // Four different products, nonuniform rows/columns/K, and A/D aliasing. + ins.modifier = 0; + ins.type = PTXOperand::f16; + ins.a = vector(PTXOperand::f16, PTXOperand::f16x2, PTXOperand::v2, 0, 2); + ins.b = vector(PTXOperand::f16, PTXOperand::b32, PTXOperand::v2, 4, 2); + ins.c = vector(PTXOperand::f16, PTXOperand::f16x2, PTXOperand::v4, 6, 4); + ins.d = vector(PTXOperand::f16, PTXOperand::b32, PTXOperand::v4, 0, 4); + if (!ins.valid().empty()) { status << ins.valid() << "\n"; return false; } + for (auto* operand : {&ins.a, &ins.b, &ins.c, &ins.d}) for (auto& element : operand->array) element.identifier = "%h" + std::to_string(element.reg); + for (int bad : {0, 1, 2}) { + PTXInstruction invalid = ins; + if (bad == 0) invalid.a.type = PTXOperand::bf16; + if (bad == 1) invalid.c.array.pop_back(); + if (bad == 2) invalid.modifier = PTXInstruction::satfinite; + if (invalid.valid().empty()) { status << "invalid legacy f16 mma accepted\n"; return false; } + bool rejected = false; + try { cta->eval_Mma(cta->getActiveContext(), invalid); } + catch (RuntimeException&) { rejected = true; } + if (!rejected) { status << "invalid legacy f16 mma executed\n"; return false; } + } + std::stringstream source; + source << ".version 8.0\n.target sm_86\n.entry fp16_884() { .reg .b32 %h<10>; " << ins.toString() << "; ret; }"; + Module parsed; parsed.load(source); + std::stringstream printed; parsed.write(printed); Module roundtrip; roundtrip.load(printed); + cta->reset(); + for (int lane = 0; lane < 32; ++lane) { + const int group = (lane / 4) % 4, row = lane % 4 + (lane / 16) * 4; + for (int i = 0; i < 2; ++i) { + const PTXU32 packed = f16Values[group + row + 2 * i + 1] | ((PTXU32)f16Values[group + row + 2 * i + 2] << 16); + cta->setRegAsB32(lane, i, packed); + cta->setRegAsB32(lane, 4 + i, f16Values[group + row + 4 * i] | ((PTXU32)f16Values[group + row + 4 * i + 2] << 16)); + } + for (int i = 0; i < 4; ++i) cta->setRegAsB32(lane, 6 + i, f16Values[row + 2 * i] | ((PTXU32)f16Values[row + 2 * i + 1] << 16)); + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) for (int col = 0; col < 8; ++col) { + const int group = (lane / 4) % 4, row = lane % 4 + (lane / 16) * 4; + int expected = row + col; + for (int k = 0; k < 4; ++k) expected += (group + row + k + 1) * (group + col + 2 * k); + const PTXU16 bits = (cta->getRegAsB32(lane, col / 2) >> (16 * (col & 1))) & 0xffff; + if (std::ldexp(1024 + (bits & 1023), (bits >> 10) - 25) != expected) { + status << "mma.m8n8k4.f16 mismatch lane " << lane << " column " << col << "\n"; return false; + } + } + + for (int layout = 0; layout < 4; ++layout) { + ins.mmaAColumnMajor = layout & 1; + ins.mmaBColumnMajor = layout & 2; + std::stringstream source; + source << ".version 8.0\n.target sm_86\n.entry layouts() { .reg .b32 %h<10>; " << ins.toString() << "; ret; }"; + Module parsed; parsed.load(source); + std::stringstream printed; parsed.write(printed); Module roundtrip; roundtrip.load(printed); + if (printed.str().find(ins.toString()) == std::string::npos) { status << "mma layout lost in parsing\n"; return false; } + cta->reset(); + for (int lane = 0; lane < 32; ++lane) { + const int group = lane / 4 % 4, t = lane % 4, high = lane / 16 * 4; + PTXU32 packedA[2] = {}, packedB[2] = {}; + for (int i = 0; i < 4; ++i) { + const int aRow = ins.mmaAColumnMajor ? i + high : t + high; + const int aK = ins.mmaAColumnMajor ? t : i; + const int bK = ins.mmaBColumnMajor ? i : t; + const int bCol = ins.mmaBColumnMajor ? t + high : i + high; + packedA[i / 2] |= (PTXU32)f16Values[group + aRow + aK + 1] << (16 * (i % 2)); + packedB[i / 2] |= (PTXU32)f16Values[group + bCol + 2 * bK] << (16 * (i % 2)); + } + for (int i = 0; i < 2; ++i) { cta->setRegAsB32(lane, i, packedA[i]); cta->setRegAsB32(lane, 4 + i, packedB[i]); } + for (int i = 0; i < 4; ++i) cta->setRegAsB32(lane, 6 + i, f16Values[t + high + 2 * i] | ((PTXU32)f16Values[t + high + 2 * i + 1] << 16)); + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) for (int col = 0; col < 8; ++col) { + const int group = lane / 4 % 4, row = lane % 4 + lane / 16 * 4; + int expected = row + col; + for (int k = 0; k < 4; ++k) expected += (group + row + k + 1) * (group + col + 2 * k); + const PTXU16 bits = cta->getRegAsB32(lane, col / 2) >> (16 * (col % 2)); + if (std::ldexp(1024 + (bits & 1023), (bits >> 10) - 25) != expected) { status << "mma layout mismatch " << layout << "\n"; return false; } + } + } + + ins.type = PTXOperand::f32; + ins.mmaAColumnMajor = false; ins.mmaBColumnMajor = true; + ins.a = vector(PTXOperand::f16, PTXOperand::b32, PTXOperand::v2, 8, 2); + ins.b = vector(PTXOperand::f16, PTXOperand::b32, PTXOperand::v2, 10, 2); + ins.c = ins.d = vector(PTXOperand::f32, PTXOperand::f32, PTXOperand::v8, 0, 8); + for (auto* operand : {&ins.a, &ins.b, &ins.c, &ins.d}) for (auto& e : operand->array) e.identifier = "%f" + std::to_string(e.reg); + std::stringstream fp32Source; + fp32Source << ".version 8.0\n.target sm_86\n.entry fp32_884() { .reg .b32 %f<12>; " << ins.toString() << "; ret; }"; + Module fp32Parsed; fp32Parsed.load(fp32Source); + std::stringstream fp32Printed; fp32Parsed.write(fp32Printed); Module fp32Roundtrip; fp32Roundtrip.load(fp32Printed); + cta->reset(); + cta->functionCallStack.pushFrame(0, 12, 0, 0, 0, 0, 0); + for (int lane = 0; lane < 32; ++lane) { + const int group = lane / 4 % 4, t = lane % 4, high = lane / 16 * 4; + for (int i = 0; i < 2; ++i) { + cta->setRegAsB32(lane, 8 + i, f16Values[group + 2 * i + high + t + 1] | ((PTXU32)f16Values[group + 2 * i + high + t + 2] << 16)); + cta->setRegAsB32(lane, 10 + i, f16Values[group + t + high + 4 * i] | ((PTXU32)f16Values[group + t + high + 4 * i + 2] << 16)); + } + for (int i = 0; i < 8; ++i) cta->setRegAsF32(lane, i, 100 * (lane % 2 + (i & 2) + high) + (i & 4) + (lane & 2) + (i & 1) + 0.25f); + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) for (int i = 0; i < 8; ++i) { + const int group = lane / 4 % 4, row = lane % 2 + (i & 2) + lane / 16 * 4, col = (i & 4) + (lane & 2) + i % 2; + float expected = 100 * row + col + 0.25f; + for (int k = 0; k < 4; ++k) expected += (group + row + k + 1) * (group + col + 2 * k); + if (cta->getRegAsF32(lane, i) != expected) { status << "legacy mma f32 mismatch lane " << lane << " register " << i << "\n"; return false; } + } + PTXInstruction illegalMixed = ins; + illegalMixed.type = illegalMixed.d.type = PTXOperand::f16; + illegalMixed.d = vector(PTXOperand::f16, PTXOperand::b32, PTXOperand::v4, 0, 4); + if (illegalMixed.valid().empty()) { status << "mma accepted f32 C with f16 D\n"; return false; } + ins.c = vector(PTXOperand::f16, PTXOperand::f16x2, PTXOperand::v4, 0, 4); + for (auto& element : ins.c.array) element.identifier = "%f" + std::to_string(element.reg); + for (int layout = 0; layout < 4; ++layout) { + ins.mmaAColumnMajor = layout & 1; + ins.mmaBColumnMajor = layout & 2; + std::stringstream source; + source << ".version 8.0\n.target sm_86\n.entry mixed() { .reg .b32 %f<12>; " << ins.toString() << "; ret; }"; + Module parsed; parsed.load(source); + std::stringstream printed; parsed.write(printed); + Module roundtrip; roundtrip.load(printed); + for (int lane = 0; lane < 32; ++lane) { + const int group = lane / 4 % 4, t = lane % 4, high = lane / 16 * 4; + PTXU32 packedA[2] = {}, packedB[2] = {}; + for (int i = 0; i < 4; ++i) { + const int aRow = ins.mmaAColumnMajor ? i + high : t + high; + const int aK = ins.mmaAColumnMajor ? t : i; + const int bK = ins.mmaBColumnMajor ? i : t; + const int bCol = ins.mmaBColumnMajor ? t + high : i + high; + packedA[i / 2] |= (PTXU32)f16Values[group + aRow + aK + 1] << (16 * (i % 2)); + packedB[i / 2] |= (PTXU32)f16Values[group + bCol + 2 * bK] << (16 * (i % 2)); + } + for (int i = 0; i < 2; ++i) { + cta->setRegAsB32(lane, 8 + i, packedA[i]); + cta->setRegAsB32(lane, 10 + i, packedB[i]); + } + for (int i = 0; i < 4; ++i) { + const PTXU32 packedC = f16Values[t + high + 2 * i] | ((PTXU32)f16Values[t + high + 2 * i + 1] << 16); + cta->setRegAsB32(lane, i, packedC); + } + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) { + for (int i = 0; i < 8; ++i) { + const int group = lane / 4 % 4; + const int row = lane % 2 + (i & 2) + lane / 16 * 4; + const int col = (i & 4) + (lane & 2) + i % 2; + float expected = row + col; + for (int k = 0; k < 4; ++k) expected += (group + row + k + 1) * (group + col + 2 * k); + if (cta->getRegAsF32(lane, i) != expected) { + status << "mixed mma mismatch layout " << layout << " lane " << lane << " register " << i << "\n"; + return false; + } + } + } + } + // Keep a quarter-unit from FP16 C even when half-precision accumulation would lose it. + for (int lane = 0; lane < 32; ++lane) { + for (int i = 0; i < 4; ++i) cta->setRegAsB32(lane, i, 0x34003400u); + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) { + for (int i = 0; i < 8; ++i) { + const int group = lane / 4 % 4; + const int row = lane % 2 + (i & 2) + lane / 16 * 4; + const int col = (i & 4) + (lane & 2) + i % 2; + float expected = 0.25f; + for (int k = 0; k < 4; ++k) expected += (group + row + k + 1) * (group + col + 2 * k); + if (cta->getRegAsF32(lane, i) != expected) { status << "mixed mma lost FP32 accumulation precision\n"; return false; } + } + } + cta->functionCallStack.popFrame(); + + return true; + } + + // mma.m16n8k16.s32.s8/u8.s8/u8.s32: integer dot-product mma, PTX ISA 8.0 + // sec 9.7.15.5.9. Each case uses a constant-valued A/B fragment so the + // expected accumulation (C + 16 * aVal * bVal) is a closed form -- still + // exercises the real 16-term k-loop and the per-lane fragment layout. + bool test_MmaInt8() { + auto vector = [this](PTXOperand::DataType type, + PTXOperand::DataType elementType, PTXOperand::Vec vec, + int firstRegister, int count) { + PTXOperand operand; + operand.addressMode = PTXOperand::Register; + operand.type = type; + operand.vec = vec; + for(int i = 0; i < count; ++i) { + operand.array.push_back(reg("mma", elementType, + (PTXOperand::RegisterType)(firstRegister + i))); + } + return operand; + }; + + auto runCase = [&](PTXOperand::DataType aType, PTXOperand::DataType bType, + int aVal, int bVal, PTXS32 cVal, bool satfinite, + PTXS32 expected, const char *label, bool varyC = true) -> bool { + PTXInstruction ins; + ins.opcode = PTXInstruction::Mma; + ins.mmaShape = PTXInstruction::MmaM16N8K16; + ins.type = PTXOperand::s32; + ins.modifier = satfinite ? PTXInstruction::satfinite + : PTXInstruction::Modifier_invalid; + + ins.d = vector(PTXOperand::s32, PTXOperand::s32, PTXOperand::v4, 0, 4); + ins.c = vector(PTXOperand::s32, PTXOperand::s32, PTXOperand::v4, 6, 4); + ins.a = vector(aType, PTXOperand::b32, PTXOperand::v2, 0, 2); + ins.b = vector(bType, PTXOperand::b32, PTXOperand::v1, 4, 1); + + if (ins.valid() != "") { + status << label << ": instruction rejected as invalid: " + << ins.valid() << "\n"; + return false; + } + + PTXInstruction invalid = ins; + invalid.modifier |= PTXInstruction::rn; + if (invalid.valid().empty()) { + status << "integer mma accepted a rounding modifier\n"; + return false; + } + const PTXU32 aByte = static_cast(aVal) & 0xffu; + const PTXU32 bByte = static_cast(bVal) & 0xffu; + const PTXU32 aPacked = aByte | (aByte << 8) | (aByte << 16) | (aByte << 24); + const PTXU32 bPacked = bByte | (bByte << 8) | (bByte << 16) | (bByte << 24); + + cta->reset(); + for (int thread = 0; thread < threadCount; ++thread) { + cta->setRegAsB32(thread, 0, aPacked); + cta->setRegAsB32(thread, 1, aPacked); + cta->setRegAsB32(thread, 4, bPacked); + const int lane = thread & 31; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for (int i = 0; i < 4; ++i) { + const int row = groupID + (i >= 2 ? 8 : 0); + const int col = threadInGroup * 2 + (i & 1); + cta->setRegAsS32(thread, 6 + i, + varyC ? cVal + row * 8 + col : cVal); + } + } + + cta->eval_Mma(cta->getActiveContext(), ins); + + for (int thread = 0; thread < threadCount; ++thread) { + const int lane = thread & 31; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for (int i = 0; i < 4; ++i) { + const int row = groupID + (i >= 2 ? 8 : 0); + const int col = threadInGroup * 2 + (i & 1); + const PTXS32 got = cta->getRegAsS32(thread, i); + const PTXS32 want = varyC ? expected + row * 8 + col + : expected; + if (got != want) { + status << label << " incorrect [" << thread + << "]: got " << got << " want " << want << "\n"; + return false; + } + } + } + return true; + }; + + // s8 x s8: -3 * 5 * 16 terms = -240. + if (!runCase(PTXOperand::s8, PTXOperand::s8, -3, 5, 0, false, + -240, "mma s8xs8")) return false; + + // u8 x u8: byte 200 means 200 unsigned, -56 if wrongly sign-extended. + // 200 * 200 * 16 = 640000. + if (!runCase(PTXOperand::u8, PTXOperand::u8, 200, 200, 0, false, + 640000, "mma u8xu8")) return false; + + // mixed s8 x u8: -3 (signed) * 200 (unsigned) * 16 = -9600. + if (!runCase(PTXOperand::s8, PTXOperand::u8, -3, 200, 0, false, + -9600, "mma s8xu8")) return false; + + // satfinite: large positive C plus a large positive product overflows + // s32 and must clamp to INT32_MAX rather than wrap. + if (!runCase(PTXOperand::s8, PTXOperand::s8, 127, 127, + (std::numeric_limits::max)() - 100, true, + (std::numeric_limits::max)(), "mma satfinite positive", + false)) + return false; + + // satfinite: large negative C plus a large negative product + // underflows s32 and must clamp to INT32_MIN. + if (!runCase(PTXOperand::s8, PTXOperand::s8, 127, -128, + (std::numeric_limits::min)() + 100, true, + (std::numeric_limits::min)(), "mma satfinite negative", + false)) + return false; + + // Validation: integer mma requires an s32 accumulator, not f32. + { + PTXInstruction ins; + ins.opcode = PTXInstruction::Mma; + ins.mmaShape = PTXInstruction::MmaM16N8K16; + ins.type = PTXOperand::f32; + ins.d = vector(PTXOperand::f32, PTXOperand::f32, PTXOperand::v4, 0, 4); + ins.c = vector(PTXOperand::f32, PTXOperand::f32, PTXOperand::v4, 6, 4); + ins.a = vector(PTXOperand::s8, PTXOperand::b32, PTXOperand::v2, 0, 2); + ins.b = vector(PTXOperand::s8, PTXOperand::b32, PTXOperand::v1, 4, 1); + if (ins.valid() == "") { + status << "mma with s8 inputs and f32 accumulator " + "unexpectedly valid\n"; + return false; + } + } + + // Direct execution must not interpret unsupported shapes as m16n8k16. + for (auto shape : {PTXInstruction::MmaShape_Invalid, PTXInstruction::MmaM16N8K8}) { + PTXInstruction ins; + ins.opcode = PTXInstruction::Mma; + ins.mmaShape = shape; + ins.type = PTXOperand::s32; + ins.d = vector(PTXOperand::s32, PTXOperand::s32, PTXOperand::v4, 0, 4); + ins.c = vector(PTXOperand::s32, PTXOperand::s32, PTXOperand::v4, 6, 4); + ins.a = vector(PTXOperand::s8, PTXOperand::b32, PTXOperand::v2, 0, 2); + ins.b = vector(PTXOperand::s8, PTXOperand::b32, PTXOperand::v1, 4, 1); + cta->reset(); + bool rejected = false; + try { cta->eval_Mma(cta->getActiveContext(), ins); } + catch (const RuntimeException&) { rejected = true; } + if (!rejected) { + status << "integer mma executed an unsupported shape\n"; + return false; + } + } + + for (auto aType : {PTXOperand::s4, PTXOperand::u4}) for (auto bType : {PTXOperand::s4, PTXOperand::u4}) { + for (bool saturate : {false, true}) for (int sign : {-1, 1}) + for (auto shape : {PTXInstruction::MmaM8N8K32, PTXInstruction::MmaM16N8K32, PTXInstruction::MmaM16N8K64}) { + const bool wide = shape != PTXInstruction::MmaM8N8K32; + const bool k64 = shape == PTXInstruction::MmaM16N8K64; + PTXInstruction ins; + ins.opcode = PTXInstruction::Mma; ins.mmaShape = shape; + ins.type = PTXOperand::s32; ins.modifier = saturate ? PTXInstruction::satfinite : 0; + ins.a = vector(aType, PTXOperand::b32, k64 ? PTXOperand::v4 : wide ? PTXOperand::v2 : PTXOperand::v1, 4, k64 ? 4 : wide ? 2 : 1); + ins.b = vector(bType, PTXOperand::b32, k64 ? PTXOperand::v2 : PTXOperand::v1, 8, k64 ? 2 : 1); + ins.c = ins.d = vector(PTXOperand::s32, PTXOperand::s32, wide ? PTXOperand::v4 : PTXOperand::v2, 0, wide ? 4 : 2); + if (k64) for (auto invalidType : {PTXOperand::s8, PTXOperand::u8, PTXOperand::f16, PTXOperand::tf32}) { + PTXInstruction invalid = ins; + invalid.a.type = invalid.b.type = invalidType; + if (invalid.valid().empty()) { status << "m16n8k64 accepted non-int4 inputs\n"; return false; } + } + const PTXS32 base = sign > 0 ? INT32_MAX - 200 : INT32_MIN + 200; + auto decode = [](int value, PTXOperand::DataType type) { return type == PTXOperand::s4 && value >= 8 ? value - 16 : value; }; + cta->reset(); + for (int lane = 0; lane < 32; ++lane) { + const int group = lane / 4, t = lane % 4; + PTXU32 a = 0, b = 0; + PTXU32 upperA = 0, upperB = 0; + for (int i = 0; i < 8; ++i) { + a |= (PTXU32)((group + 8 * t + i + (t >= 2 ? 3 : 0)) % 16) << (4 * i); + b |= (PTXU32)((3 * (8 * t + i) + group + (t >= 2 ? 5 : 0)) % 16) << (4 * i); + upperA |= (PTXU32)((group + 8 * t + i + (t >= 2 ? 3 : 0) + 6) % 16) << (4 * i); + upperB |= (PTXU32)((3 * (8 * t + i) + group + (t >= 2 ? 5 : 0) + 10) % 16) << (4 * i); + } + cta->setRegAsB32(lane, 4, a); cta->setRegAsB32(lane, 8, b); + if (wide) cta->setRegAsB32(lane, 5, a ^ 0x88888888u); + if (k64) { + cta->setRegAsB32(lane, 6, upperA); cta->setRegAsB32(lane, 7, upperA ^ 0x88888888u); + cta->setRegAsB32(lane, 9, upperB); + } + for (int i = 0; i < (wide ? 4 : 2); ++i) cta->setRegAsS32(lane, i, + base + (group + (i >= 2 ? 8 : 0)) * 8 + 2 * t + (i & 1)); + } + cta->eval_Mma(cta->getActiveContext(), ins); + for (int lane = 0; lane < 32; ++lane) for (int i = 0; i < (wide ? 4 : 2); ++i) { + const int row = lane / 4 + (i >= 2 ? 8 : 0), col = 2 * (lane % 4) + (i & 1); + int64_t expected = (int64_t)base + row * 8 + col; + for (int k = 0; k < (k64 ? 64 : 32); ++k) expected += + decode((row + k + (k / 16) * 3) % 16, aType) * + decode((3 * k + col + (k / 16) * 5) % 16, bType); + if (saturate) expected = std::max(INT32_MIN, std::min(INT32_MAX, expected)); + const PTXS32 want = hydrazine::bit_cast((PTXU32)expected); + if (cta->getRegAsS32(lane, i) != want) { status << "int4 mma shape " << shape << " mismatch lane " << lane << " register " << i << "\n"; return false; } + } + } + } + return true; + } + + // mma.m8n8k16/m16n8k32.s32.s8/u8.s8/u8.s32: the two remaining integer mma + // shapes beyond m16n8k16 (PTX ISA 8.0 sec 9.7.15.5.3 / 9.7.15.5.10). + // Same closed-form approach as test_MmaInt8: constant-valued A/B + // fragments make the expected accumulation (C + k * aVal * bVal) a + // closed form while a row/col-varying C still exercises the real + // per-lane C/D fragment layout and the shape's k-loop bound. + bool test_MmaInt8WideShapes() { + auto vector = [this](PTXOperand::DataType type, + PTXOperand::DataType elementType, PTXOperand::Vec vec, + int firstRegister, int count) { + PTXOperand operand; + operand.addressMode = PTXOperand::Register; + operand.type = type; + operand.vec = vec; + for(int i = 0; i < count; ++i) { + operand.array.push_back(reg("mma", elementType, + (PTXOperand::RegisterType)(firstRegister + i))); + } + return operand; + }; + + // cdRowCol: m8n8k16 has a 8x8 C/D with 2 registers per lane; + // m16n8k16/m16n8k32 share the same 16x8 C/D with 4 registers per + // lane -- reuse that mapping for m16n8k32 directly. + auto runCase = [&](PTXInstruction::MmaShape shape, int k, + PTXOperand::Vec aVec, int aCount, PTXOperand::Vec bVec, int bCount, + PTXOperand::Vec cdVec, int cdCount, + auto cdRowCol, + PTXOperand::DataType aType, PTXOperand::DataType bType, + int aVal, int bVal, PTXS32 cVal, bool satfinite, + PTXS32 expected, const char *label, bool varyC = true) -> bool { + PTXInstruction ins; + ins.opcode = PTXInstruction::Mma; + ins.mmaShape = shape; + ins.type = PTXOperand::s32; + ins.modifier = satfinite ? PTXInstruction::satfinite + : PTXInstruction::Modifier_invalid; + + // Register numbering mirrors test_MmaInt8: D and A share the + // low registers (A is fully read into local matrices before D + // is written, so reuse is safe), keeping everything within the + // small register file the embedded test kernel declares. + ins.d = vector(PTXOperand::s32, PTXOperand::s32, cdVec, 0, cdCount); + ins.a = vector(aType, PTXOperand::b32, aVec, 0, aCount); + ins.b = vector(bType, PTXOperand::b32, bVec, 4, bCount); + ins.c = vector(PTXOperand::s32, PTXOperand::s32, cdVec, 6, cdCount); + + if (ins.valid() != "") { + status << label << ": instruction rejected as invalid: " + << ins.valid() << "\n"; + return false; + } + + const PTXU32 aByte = static_cast(aVal) & 0xffu; + const PTXU32 bByte = static_cast(bVal) & 0xffu; + const PTXU32 aPacked = aByte | (aByte << 8) | (aByte << 16) | (aByte << 24); + const PTXU32 bPacked = bByte | (bByte << 8) | (bByte << 16) | (bByte << 24); + + cta->reset(); + for (int thread = 0; thread < threadCount; ++thread) { + for (int i = 0; i < aCount; ++i) { + cta->setRegAsB32(thread, i, aPacked); + } + for (int i = 0; i < bCount; ++i) { + cta->setRegAsB32(thread, 4 + i, bPacked); + } + const int lane = thread & 31; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for (int i = 0; i < cdCount; ++i) { + int row, col; + cdRowCol(groupID, threadInGroup, i, row, col); + cta->setRegAsS32(thread, 6 + i, + varyC ? cVal + row * 8 + col : cVal); + } + } + + cta->eval_Mma(cta->getActiveContext(), ins); + + for (int thread = 0; thread < threadCount; ++thread) { + const int lane = thread & 31; + const int groupID = lane >> 2; + const int threadInGroup = lane & 3; + for (int i = 0; i < cdCount; ++i) { + int row, col; + cdRowCol(groupID, threadInGroup, i, row, col); + const PTXS32 got = cta->getRegAsS32(thread, i); + const PTXS32 want = varyC ? expected + row * 8 + col + : expected; + if (got != want) { + status << label << " incorrect [" << thread + << "]: got " << got << " want " << want << "\n"; + return false; + } + } + } + return true; + }; + + auto cdRowCol8x8 = [](int groupID, int threadInGroup, int i, + int &row, int &col) { + row = groupID; + col = threadInGroup * 2 + i; + }; + auto cdRowCol16x8 = [](int groupID, int threadInGroup, int i, + int &row, int &col) { + row = groupID + (i >= 2 ? 8 : 0); + col = threadInGroup * 2 + (i & 1); + }; + + // m8n8k16: 1 A register, 1 B register, 2 C/D registers, k = 16. + if (!runCase(PTXInstruction::MmaM8N8K16, 16, + PTXOperand::v1, 1, PTXOperand::v1, 1, PTXOperand::v2, 2, cdRowCol8x8, + PTXOperand::s8, PTXOperand::s8, -3, 5, 0, false, + -240, "mma.m8n8k16 s8xs8")) return false; + if (!runCase(PTXInstruction::MmaM8N8K16, 16, + PTXOperand::v1, 1, PTXOperand::v1, 1, PTXOperand::v2, 2, cdRowCol8x8, + PTXOperand::u8, PTXOperand::u8, 200, 200, 0, false, + 640000, "mma.m8n8k16 u8xu8")) return false; + if (!runCase(PTXInstruction::MmaM8N8K16, 16, + PTXOperand::v1, 1, PTXOperand::v1, 1, PTXOperand::v2, 2, cdRowCol8x8, + PTXOperand::s8, PTXOperand::u8, -3, 200, 0, false, + -9600, "mma.m8n8k16 s8xu8")) return false; + if (!runCase(PTXInstruction::MmaM8N8K16, 16, + PTXOperand::v1, 1, PTXOperand::v1, 1, PTXOperand::v2, 2, cdRowCol8x8, + PTXOperand::s8, PTXOperand::s8, 127, 127, + (std::numeric_limits::max)() - 100, true, + (std::numeric_limits::max)(), "mma.m8n8k16 satfinite positive", + false)) return false; + if (!runCase(PTXInstruction::MmaM8N8K16, 16, + PTXOperand::v1, 1, PTXOperand::v1, 1, PTXOperand::v2, 2, cdRowCol8x8, + PTXOperand::s8, PTXOperand::s8, 127, -128, + (std::numeric_limits::min)() + 100, true, + (std::numeric_limits::min)(), "mma.m8n8k16 satfinite negative", + false)) return false; + + // m16n8k32: 4 A registers, 2 B registers, 4 C/D registers, k = 32. + if (!runCase(PTXInstruction::MmaM16N8K32, 32, + PTXOperand::v4, 4, PTXOperand::v2, 2, PTXOperand::v4, 4, cdRowCol16x8, + PTXOperand::s8, PTXOperand::s8, -3, 5, 0, false, + -480, "mma.m16n8k32 s8xs8")) return false; + if (!runCase(PTXInstruction::MmaM16N8K32, 32, + PTXOperand::v4, 4, PTXOperand::v2, 2, PTXOperand::v4, 4, cdRowCol16x8, + PTXOperand::u8, PTXOperand::u8, 200, 200, 0, false, + 1280000, "mma.m16n8k32 u8xu8")) return false; + if (!runCase(PTXInstruction::MmaM16N8K32, 32, + PTXOperand::v4, 4, PTXOperand::v2, 2, PTXOperand::v4, 4, cdRowCol16x8, + PTXOperand::s8, PTXOperand::u8, -3, 200, 0, false, + -19200, "mma.m16n8k32 s8xu8")) return false; + if (!runCase(PTXInstruction::MmaM16N8K32, 32, + PTXOperand::v4, 4, PTXOperand::v2, 2, PTXOperand::v4, 4, cdRowCol16x8, + PTXOperand::s8, PTXOperand::s8, 127, 127, + (std::numeric_limits::max)() - 100, true, + (std::numeric_limits::max)(), "mma.m16n8k32 satfinite positive", + false)) return false; + if (!runCase(PTXInstruction::MmaM16N8K32, 32, + PTXOperand::v4, 4, PTXOperand::v2, 2, PTXOperand::v4, 4, cdRowCol16x8, + PTXOperand::s8, PTXOperand::s8, 127, -128, + (std::numeric_limits::min)() + 100, true, + (std::numeric_limits::min)(), "mma.m16n8k32 satfinite negative", + false)) return false; + + // Validation: m8n8k16 rejects m16n8k16-sized fragments (wrong vec). + { + PTXInstruction ins; + ins.opcode = PTXInstruction::Mma; + ins.mmaShape = PTXInstruction::MmaM8N8K16; + ins.type = PTXOperand::s32; + ins.d = vector(PTXOperand::s32, PTXOperand::s32, PTXOperand::v4, 0, 4); + ins.c = vector(PTXOperand::s32, PTXOperand::s32, PTXOperand::v4, 6, 4); + ins.a = vector(PTXOperand::s8, PTXOperand::b32, PTXOperand::v2, 0, 2); + ins.b = vector(PTXOperand::s8, PTXOperand::b32, PTXOperand::v1, 4, 1); + if (ins.valid() == "") { + status << "mma.m8n8k16 with m16n8k16-sized fragments " + "unexpectedly valid\n"; + return false; + } + } + + return true; + } + + bool test_Lg2() { + bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_lg2() {\n" + << " .reg .f32 f0, f1;\n" + << " lg2.approx.f32 f0, f1;\n" + << " lg2.approx.ftz.f32 f0, f1;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse PTX 8.0 lg2 forms: " << error.what() << "\n"; + return false; + } + + PTXInstruction ins; + ins.opcode = PTXInstruction::Lg2; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + ins.modifier = 0; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx; + + // f32 + // + if (result) { + for (int i = 0; i < threadCount; i++) { + const PTXF32 value = 0.51f + + (PTXF32)i / (PTXF32)threadCount * 1.48f; + cta->setRegAsF32(i, 0, value); + cta->setRegAsF32(i, 2, 0); + } + cta->eval_Lg2(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + const PTXF32 value = 0.51f + + (PTXF32)i / (PTXF32)threadCount * 1.48f; + const double expected = std::log2((double)value); + if (std::fabs(cta->getRegAsF32(i, 2) - expected) + > std::ldexp(1.0, -22)) { + result = false; + status << "lg2.f32 incorrect [" << i + << "] - expected: " << expected + << ", got " << cta->getRegAsF32(i, 2) << "\n"; + break; + } + } + } + + cta->setRegAsF32(0, 0, -std::numeric_limits::infinity()); + cta->setRegAsF32(1, 0, -1.0f); + cta->setRegAsF32(2, 0, -0.0f); + cta->setRegAsF32(3, 0, 0.0f); + cta->setRegAsF32(4, 0, std::numeric_limits::infinity()); + cta->setRegAsF32(5, 0, std::numeric_limits::quiet_NaN()); + cta->eval_Lg2(cta->getActiveContext(), ins); + if (!hydrazine::isnan(cta->getRegAsF32(0, 2)) + || !hydrazine::isnan(cta->getRegAsF32(1, 2)) + || cta->getRegAsF32(2, 2) != -std::numeric_limits::infinity() + || cta->getRegAsF32(3, 2) != -std::numeric_limits::infinity() + || cta->getRegAsF32(4, 2) != std::numeric_limits::infinity() + || !hydrazine::isnan(cta->getRegAsF32(5, 2))) + return false; + + const PTXF32 subnormal = std::numeric_limits::denorm_min(); + cta->setRegAsF32(0, 0, subnormal); + cta->eval_Lg2(cta->getActiveContext(), ins); + if (std::fabs((cta->getRegAsF32(0, 2) + 149.0) / 149.0) + > std::ldexp(1.0, -22)) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->eval_Lg2(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 2) + != -std::numeric_limits::infinity()) return false; + + return result; + } + + bool test_Sqrt() { + bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_sqrt() {\n" + << " .reg .f32 f0, f1;\n" + << " .reg .f64 d0, d1;\n" + << " sqrt.approx.f32 f0, f1;\n" + << " sqrt.rz.ftz.f32 f0, f1;\n" + << " sqrt.rp.f64 d0, d1;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse PTX 8.0 sqrt forms: " << error.what() << "\n"; + return false; + } + + PTXInstruction ins; + ins.opcode = PTXInstruction::Sqrt; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + ins.modifier = 0; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.type = PTXOperand::f64; + ins.a.type = ins.d.type = PTXOperand::f64; + ins.modifier = PTXInstruction::rp; + if (!ins.valid().empty()) return false; + ins.modifier = PTXInstruction::rn | PTXInstruction::ftz; + if (ins.valid().empty()) return false; + + double freq = 2.0f / (double)threadCount; + + // f32 + // + if (result) { + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsF32(i, 0, (PTXF32)(0.1f + (float)i * freq)); + cta->setRegAsF32(i, 2, 0); + } + cta->eval_Sqrt(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (std::fabs(cta->getRegAsF32(i, 2) - (PTXF32)sqrt(0.1f + (float)i * freq)) > 0.1f) { + result = false; + status << "sqrt.f32 incorrect [" << i << "] - expected: " + << (PTXF32)sqrt(0.1f + (float)i * freq) + << ", got " << cta->getRegAsF32(i, 2) << "\n"; + break; + } + } + } + + // f64 + // + if (result) { + ins.type = PTXOperand::f64; + ins.a = reg("r1", PTXOperand::f64, 0); + ins.d = reg("r3", PTXOperand::f64, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsF64(i, 0, (PTXF64)(0.1f + (double)i * freq)); + cta->setRegAsF64(i, 2, 0); + } + cta->eval_Sqrt(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (std::fabs(cta->getRegAsF64(i, 2) - sqrt(0.1 + (double)i * freq)) > 0.1f) { + result = false; + status << "sqrt.f64 incorrect [" << i << "] - expected: " << sqrt(0.1 + (double)i * freq) + << ", got " << cta->getRegAsF64(i, 2) << "\n"; + break; + } + } + } + + const int modes[] = {PTXInstruction::rn, PTXInstruction::rz, + PTXInstruction::rm, PTXInstruction::rp}; + const PTXU32 expected32[] = { + 0x3fb504f3u, 0x3fb504f3u, 0x3fb504f3u, 0x3fb504f4u}; + const PTXU64 expected64[] = {0x3ff6a09e667f3bcdull, + 0x3ff6a09e667f3bccull, 0x3ff6a09e667f3bccull, + 0x3ff6a09e667f3bcdull}; + for (unsigned int i = 0; i < 4; ++i) { + ins.modifier = modes[i]; + ins.type = PTXOperand::f32; + ins.a.type = ins.d.type = PTXOperand::f32; + cta->setRegAsF32(0, 0, 2.0f); + cta->eval_Sqrt(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) + != expected32[i]) return false; + + ins.type = PTXOperand::f64; + ins.a.type = ins.d.type = PTXOperand::f64; + cta->setRegAsF64(0, 0, 2.0); + cta->eval_Sqrt(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) + != expected64[i]) return false; + } + + ins.type = PTXOperand::f32; + ins.a.type = ins.d.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn | PTXInstruction::ftz; + const PTXF32 subnormal = std::numeric_limits::min() + - std::numeric_limits::denorm_min(); + cta->setRegAsF32(0, 0, subnormal); + cta->setRegAsF32(1, 0, -subnormal); + cta->eval_Sqrt(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF32(0, 2)) != 0 + || hydrazine::bit_cast(cta->getRegAsF32(1, 2)) + != 0x80000000u) return false; + + ins.modifier = PTXInstruction::approx; + cta->setRegAsF32(0, 0, subnormal); + cta->eval_Sqrt(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 2) == 0.0f) return false; + + return result; + } + + bool test_Rsqrt() { + bool result = true; + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_rsqrt() {\n" + << " .reg .f32 f0, f1;\n" + << " .reg .f64 d0, d1;\n" + << " rsqrt.approx.f32 f0, f1;\n" + << " rsqrt.approx.ftz.f32 f0, f1;\n" + << " rsqrt.approx.f64 d0, d1;\n" + << " rsqrt.approx.ftz.f64 d0, d1;\n" + << " ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const std::exception& error) { + status << "failed to parse PTX 8.0 rsqrt forms: " << error.what() << "\n"; + return false; + } + + PTXInstruction ins; + ins.opcode = PTXInstruction::Rsqrt; + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + ins.modifier = PTXInstruction::approx; + if (!ins.valid().empty()) return false; + ins.modifier = 0; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::approx | PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.type = PTXOperand::f64; + ins.a.type = ins.d.type = PTXOperand::f64; + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + if (!ins.valid().empty()) return false; + + double freq = 2.0f / (double)threadCount; + + // f32 + // + if (result) { + ins.type = PTXOperand::f32; + ins.a = reg("r1", PTXOperand::f32, 0); + ins.d = reg("r3", PTXOperand::f32, 2); + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsF32(i, 0, (PTXF32)(0.1f + (float)i * freq)); + cta->setRegAsF32(i, 2, 0); + } + cta->eval_Rsqrt(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { if (std::fabs(cta->getRegAsF32(i, 2) - 1.0f/(PTXF32)sqrt(0.1f + (float)i * freq)) > 0.1f) { result = false; @@ -2827,9 +6536,89 @@ class TestInstructions: public Test { } } + ins.type = PTXOperand::f64; + ins.a.type = ins.d.type = PTXOperand::f64; + ins.modifier = PTXInstruction::approx | PTXInstruction::ftz; + cta->setRegAsF64(0, 0, 3.0); + cta->setRegAsF64(1, 0, std::numeric_limits::quiet_NaN()); + cta->setRegAsF64(2, 0, -0.0); + cta->setRegAsF64(3, 0, std::numeric_limits::denorm_min()); + cta->eval_Rsqrt(cta->getActiveContext(), ins); + if (hydrazine::bit_cast(cta->getRegAsF64(0, 2)) + != 0x3fe279a700000000ull + || hydrazine::bit_cast(cta->getRegAsF64(1, 2)) + != 0x7fffffff00000000ull + || cta->getRegAsF64(2, 2) + != -std::numeric_limits::infinity() + || cta->getRegAsF64(3, 2) + != std::numeric_limits::infinity()) return false; + + ins.type = PTXOperand::f32; + ins.a.type = ins.d.type = PTXOperand::f32; + const PTXF32 subnormal = std::numeric_limits::min() + - std::numeric_limits::denorm_min(); + cta->setRegAsF32(0, 0, subnormal); + ins.modifier = PTXInstruction::approx; + cta->eval_Rsqrt(cta->getActiveContext(), ins); + if (!std::isfinite(cta->getRegAsF32(0, 2))) return false; + ins.modifier |= PTXInstruction::ftz; + cta->eval_Rsqrt(cta->getActiveContext(), ins); + if (cta->getRegAsF32(0, 2) + != std::numeric_limits::infinity()) return false; + return result; } + bool test_Dp() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_dp() {\n" + << " .reg .b32 d0, a0, b0, c0;\n" + << " dp4a.u32.u32 d0, a0, b0, c0;\n" + << " dp4a.u32.s32 d0, a0, b0, c0;\n" + << " dp2a.lo.u32.u32 d0, a0, b0, c0;\n" + << " dp2a.hi.u32.s32 d0, a0, b0, c0;\n ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse dp2a/dp4a examples: " << error.what() << "\n"; + return false; + } + + struct Case { + PTXInstruction::Opcode opcode; + PTXOperand::DataType aType, bType; + unsigned int mode; + PTXU32 a, b, c, expected; + }; + const Case cases[] = { + {PTXInstruction::Dp4a, PTXOperand::u32, PTXOperand::u32, 0, + 0x04030201u, 0x08070605u, 10, 80}, + {PTXInstruction::Dp4a, PTXOperand::u32, PTXOperand::s32, 0, + 0x04030201u, 0xfcfdfeffu, 10, 0xffffffecu}, + {PTXInstruction::Dp2a, PTXOperand::u32, PTXOperand::u32, PTXInstruction::lo, + 0x00030002u, 0x64640504u, 1, 24}, + {PTXInstruction::Dp2a, PTXOperand::s32, PTXOperand::u32, PTXInstruction::hi, + 0x0003fffeu, 0x05040000u, 1, 8} + }; + PTXInstruction ins; + ins.d = reg("d", PTXOperand::b32, 0); + for (const Case& test : cases) { + ins.opcode = test.opcode; + ins.type = test.aType; + ins.bType = test.bType; + ins.modifier = test.mode; + ins.a = imm_uint("a", test.aType, test.a); + ins.b = imm_uint("b", test.bType, test.b); + ins.c = imm_uint("c", PTXOperand::b32, test.c); + cta->eval_Dp(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU32(thread, 0) != test.expected) return false; + } + } + return true; + } + ///////////////////////////////////////////////////////////////////////////////////////////////// // // @@ -2837,6 +6626,510 @@ class TestInstructions: public Test { // ///////////////////////////////////////////////////////////////////////////////////////////////// + bool test_Fns() { + struct Case { + PTXU32 base; + PTXS32 offset; + PTXU32 expected; + }; + const Case cases[] = { + {3, 1, 3}, + {3, -1, 3}, + {2, 1, 3}, + {2, -1, 1} + }; + + PTXInstruction ins; + ins.opcode = PTXInstruction::Fns; + ins.type = PTXOperand::b32; + ins.d = reg("d", PTXOperand::b32, 0); + ins.a = imm_uint("mask", PTXOperand::b32, 0xaaaaaaaau); + + for (const Case& test : cases) { + ins.b = imm_uint("base", PTXOperand::b32, test.base); + ins.c = imm_int("offset", PTXOperand::s32, test.offset); + cta->eval_Fns(cta->getActiveContext(), ins); + + for (int thread = 0; thread < threadCount; ++thread) { + PTXU32 actual = cta->getRegAsB32(thread, ins.d.reg); + if (actual != test.expected) { + status << "fns.b32 failed for base " << test.base + << ", offset " << test.offset << ": expected " + << test.expected << ", got " << actual << "\n"; + return false; + } + } + } + return true; + } + + bool test_Szext() { + std::stringstream ptx; + ptx << ".version 8.0\n" + << ".target sm_86\n" + << ".address_size 64\n" + << ".visible .entry test_szext() {\n" + << " .reg .b32 rd;\n" + << " .reg .s32 ra;\n" + << " .reg .u32 rb;\n" + << " szext.clamp.s32 rd, ra, rb;\n" + << " szext.wrap.u32 rd, 0xffffffff, 0;\n" + << " ret;\n" + << "}\n"; + Module parsed; + try { + parsed.load(ptx); + } + catch (const hydrazine::Exception& error) { + status << "failed to parse PTX 8.0 szext examples: " + << error.what() << "\n"; + return false; + } + + PTXInstruction ins; + ins.opcode = PTXInstruction::Szext; + ins.type = PTXOperand::s32; + ins.shiftMode = PTXInstruction::ShiftMode::Clamp; + ins.d = reg("rd", PTXOperand::b32, 0); + ins.a = reg("ra", PTXOperand::s32, 1); + ins.b = reg("rb", PTXOperand::u32, 2); + for (int thread = 0; thread < threadCount; ++thread) { + cta->setRegAsU32(thread, 1, 0x80); + cta->setRegAsU32(thread, 2, 8); + } + cta->eval_Szext(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU32(thread, 0) != 0xffffff80u) return false; + } + + ins.type = PTXOperand::u32; + ins.shiftMode = PTXInstruction::ShiftMode::Wrap; + ins.a = imm_uint("a", PTXOperand::u32, 0xffffffffu); + ins.b = imm_uint("b", PTXOperand::u32, 0); + cta->eval_Szext(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU32(thread, 0) != 0) return false; + } + + return true; + } + + bool test_Bmsk() { + std::stringstream ptx; + ptx << ".version 8.0\n" + << ".target sm_86\n" + << ".address_size 64\n" + << ".visible .entry test_bmsk() {\n" + << " .reg .b32 rd, ra, rb;\n" + << " bmsk.clamp.b32 rd, ra, rb;\n" + << " bmsk.wrap.b32 rd, 1, 2;\n" + << " ret;\n" + << "}\n"; + Module parsed; + try { + parsed.load(ptx); + } + catch (const hydrazine::Exception& error) { + status << "failed to parse PTX 8.0 bmsk examples: " + << error.what() << "\n"; + return false; + } + + PTXInstruction ins; + ins.opcode = PTXInstruction::Bmsk; + ins.type = PTXOperand::b32; + ins.shiftMode = PTXInstruction::ShiftMode::Clamp; + ins.d = reg("rd", PTXOperand::b32, 0); + ins.a = reg("ra", PTXOperand::b32, 1); + ins.b = reg("rb", PTXOperand::b32, 2); + for (int thread = 0; thread < threadCount; ++thread) { + cta->setRegAsU32(thread, 1, 1); + cta->setRegAsU32(thread, 2, 2); + } + cta->eval_Bmsk(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU32(thread, 0) != 0x00000006u) return false; + } + + ins.shiftMode = PTXInstruction::ShiftMode::Wrap; + ins.a = imm_uint("a", PTXOperand::b32, 1); + ins.b = imm_uint("b", PTXOperand::b32, 2); + cta->eval_Bmsk(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU32(thread, 0) != 0x00000006u) return false; + } + + return true; + } + + bool test_Bfe() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Bfe; + ins.b = imm_uint("pos", PTXOperand::u32, 0); + ins.c = imm_uint("len", PTXOperand::u32, 0); + ins.d = reg("d", PTXOperand::u32, 0); + ins.a = reg("a", PTXOperand::u32, 1); + ins.type = PTXOperand::u32; + if (!ins.valid().empty()) return false; + PTXInstruction invalid = ins; + invalid.a.addressMode = PTXOperand::Immediate; + invalid.b = reg("pos", PTXOperand::f32, 2); + if (invalid.valid().empty()) return false; + + auto check32 = [&](PTXOperand::DataType type, PTXU32 value, + PTXU32 pos, PTXU32 len, PTXU32 expected) { + ins.type = type; + ins.d.type = type; + ins.a.type = type; + ins.b.imm_uint = pos; + ins.c.imm_uint = len; + cta->reset(); + for (int t = 0; t < threadCount; ++t) cta->setRegAsU32(t, 1, value); + cta->eval_Bfe(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; ++t) + if (cta->getRegAsU32(t, 0) != expected) return false; + return true; + }; + if (!check32(PTXOperand::u32, 0xD2, 1, 4, 0x9)) return false; + if (!check32(PTXOperand::u32, 0x12345678, 0, 32, 0x12345678)) return false; + if (!check32(PTXOperand::u32, 0xffffffff, 0, 0, 0)) return false; + if (!check32(PTXOperand::u32, 0xD2, 0x101, 0x104, 0x9)) return false; + if (!check32(PTXOperand::s32, 0x000000f0, 4, 4, 0xffffffffu)) return false; + if (!check32(PTXOperand::s32, 0x80000000, 30, 4, 0xfffffffeu)) return false; + if (!check32(PTXOperand::s32, 0x80000000, 40, 3, 0xffffffffu)) return false; + + ins.type = PTXOperand::u64; + ins.d.type = PTXOperand::u64; + ins.a.type = PTXOperand::u64; + ins.b.imm_uint = 4; + ins.c.imm_uint = 8; + cta->reset(); + for (int t = 0; t < threadCount; ++t) cta->setRegAsU64(t, 1, 0x123456789abcdef0ULL); + cta->eval_Bfe(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0xef) return false; + ins.b.imm_uint = 0; + ins.c.imm_uint = 64; + cta->eval_Bfe(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0x123456789abcdef0ULL) return false; + ins.type = PTXOperand::s64; + ins.d.type = PTXOperand::s64; + ins.a.type = PTXOperand::s64; + ins.b.imm_uint = 62; + ins.c.imm_uint = 4; + cta->reset(); + for (int t = 0; t < threadCount; ++t) cta->setRegAsU64(t, 1, 0x8000000000000000ULL); + cta->eval_Bfe(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0xfffffffffffffffeULL) return false; + + ins.type = PTXOperand::u32; + ins.d.type = PTXOperand::u32; + ins.a.type = PTXOperand::u32; + ins.b.imm_uint = 0; + ins.c.imm_uint = 4; + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->reset(); + for (int t = 0; t < threadCount; ++t) { + cta->setRegAsU32(t, 1, 0xf0); + cta->setRegAsU32(t, 0, 0xdeadbeef); + cta->setRegAsPredicate(t, 3, false); + } + cta->eval_Bfe(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0xdeadbeef) return false; + return true; + } + + bool test_Bfi() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Bfi; + ins.type = PTXOperand::b32; + ins.d = reg("d", PTXOperand::b32, 0); + ins.pq = imm_uint("pq", PTXOperand::b32, 0xf); + ins.a = imm_uint("a", PTXOperand::b32, 0xa5); + ins.b = imm_uint("pos", PTXOperand::u32, 0x104); + ins.c = imm_uint("len", PTXOperand::u32, 4); + if (!ins.valid().empty()) return false; + + auto check32 = [&](PTXU32 pq, PTXU32 a, PTXU32 pos, + PTXU32 len, PTXU32 expected) { + ins.pq.imm_uint = pq; + ins.a.imm_uint = a; + ins.b.imm_uint = pos; + ins.c.imm_uint = len; + cta->reset(); + cta->eval_Bfi(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; ++t) + if (cta->getRegAsU32(t, 0) != expected) return false; + return true; + }; + if (!check32(0xf, 0xa5, 0x104, 4, 0xf5)) return false; + if (!check32(0xf, 0x12345678, 4, 0x104, 0x123456f8)) return false; + if (!check32(0xf, 0x12345678, 4, 0x100, 0x12345678)) return false; + if (!check32(0xffffffff, 0, 30, 4, 0xc0000000)) return false; + if (!check32(0xffffffff, 0x12345678, 32, 4, 0x12345678)) return false; + + PTXInstruction invalid = ins; + invalid.c = reg("len", PTXOperand::u64, 1); + if (invalid.valid().empty()) return false; + invalid.c = reg("len", PTXOperand::f32, 1); + if (invalid.valid().empty()) return false; + + ins.type = PTXOperand::b64; + ins.d.type = PTXOperand::b64; + ins.pq = imm_uint("pq", PTXOperand::b64, 0xf); + ins.a = imm_uint("a", PTXOperand::b64, 0x123456789abcde00ULL); + ins.b.imm_uint = 0x104; + ins.c.imm_uint = 4; + cta->reset(); + cta->eval_Bfi(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0x123456789abcdef0ULL) return false; + + ins.type = PTXOperand::b32; + ins.d.type = PTXOperand::b32; + ins.pq = imm_uint("pq", PTXOperand::b32, 0xf); + ins.a = imm_uint("a", PTXOperand::b32, 0xa5); + ins.b = imm_uint("pos", PTXOperand::u32, 0); + ins.c = imm_uint("len", PTXOperand::u32, 4); + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->reset(); + for (int t = 0; t < threadCount; ++t) { + cta->setRegAsU32(t, 0, 0xdeadbeef); + cta->setRegAsPredicate(t, 3, false); + } + cta->eval_Bfi(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, 0) == 0xdeadbeef; + } + + bool test_Bfind() { + struct Case { PTXOperand::DataType type; PTXU64 value; + bool shift; PTXU32 expected; }; + const Case cases[] = { + { PTXOperand::u32, 0, false, 0xffffffffu }, + { PTXOperand::u32, 1, true, 31 }, + { PTXOperand::u32, 0xffffffffu, false, 31 }, + { PTXOperand::s32, 0xffffffffu, false, 0xffffffffu }, + { PTXOperand::s32, 0xfffffffeu, true, 31 }, + { PTXOperand::s32, 0x80000000u, false, 30 }, + { PTXOperand::u64, 0, false, 0xffffffffu }, + { PTXOperand::u64, 1, true, 63 }, + { PTXOperand::u64, 0xffffffffffffffffULL, false, 63 }, + { PTXOperand::s64, 0xffffffffffffffffULL, false, 0xffffffffu }, + { PTXOperand::s64, 0xfffffffffffffffeULL, true, 63 }, + { PTXOperand::s64, 0x8000000000000000ULL, false, 62 } + }; + PTXInstruction ins; + ins.opcode = PTXInstruction::Bfind; + ins.d = reg("d", PTXOperand::u32, 0); + ins.a = reg("a", PTXOperand::u32, 1); + for (const Case& test : cases) { + ins.type = test.type; + ins.a.type = test.type; + ins.shiftAmount = test.shift; + cta->reset(); + if (test.type == PTXOperand::u32 || test.type == PTXOperand::s32) + cta->setRegAsU32(0, 1, static_cast(test.value)); + else cta->setRegAsU64(0, 1, test.value); + cta->eval_Bfind(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != test.expected) return false; + } + ins.type = PTXOperand::s32; + ins.a.type = PTXOperand::s32; + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->reset(); + cta->setRegAsU32(0, 0, 0xdeadbeefu); + cta->setRegAsU32(0, 1, 0xffffffffu); + cta->setRegAsPredicate(0, 3, false); + cta->eval_Bfind(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, 0) == 0xdeadbeefu; + } + + bool test_Popc() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Popc; + ins.type = PTXOperand::b32; + ins.a = reg("a", PTXOperand::b32, 1); + ins.d = reg("d", PTXOperand::b32, 0); + if (!ins.valid().empty()) return false; + ins.d.type = PTXOperand::u32; + if (!ins.valid().empty()) return false; + ins.d.type = PTXOperand::b32; + ins.modifier = PTXInstruction::rn; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::ftz; + if (ins.valid().empty()) return false; + ins.modifier = PTXInstruction::sat; + if (ins.valid().empty()) return false; + ins.modifier = 0; + ins.carry = PTXInstruction::CC; + if (ins.valid().empty()) return false; + ins.carry = PTXInstruction::None; + ins.type = PTXOperand::u32; + if (ins.valid().empty()) return false; + ins.type = PTXOperand::b32; + ins.a.type = PTXOperand::b64; + if (ins.valid().empty()) return false; + ins.type = PTXOperand::b64; + ins.a.type = PTXOperand::b64; + ins.d.type = PTXOperand::b64; + if (ins.valid().empty()) return false; + ins.d.type = PTXOperand::b32; + if (!ins.valid().empty()) return false; + ins.type = PTXOperand::b32; + ins.a.type = PTXOperand::b32; + cta->reset(); + const PTXU32 values[] = {0, 0xffffffffU, 0x55, 0xaa, 1, 0x80000000U}; + const PTXU32 expected[] = {0, 32, 4, 4, 1, 1}; + for (unsigned i = 0; i < sizeof(values) / sizeof(*values); ++i) { + cta->setRegAsU32(0, 1, values[i]); + cta->eval_Popc(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != expected[i]) return false; + } + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 3; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsU32(0, 0, 0xdeadbeef); + cta->eval_Popc(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0xdeadbeef) return false; + ins.pg.condition = PTXOperand::PT; + ins.type = PTXOperand::b64; + ins.a.type = PTXOperand::b64; + const PTXU64 values64[] = {0, 0xffffffffffffffffULL, 0x55, 0xaa, 1, + 0x8000000000000000ULL}; + const PTXU32 expected64[] = {0, 64, 4, 4, 1, 1}; + for (unsigned i = 0; i < sizeof(values64) / sizeof(*values64); ++i) { + cta->setRegAsU64(0, 1, values64[i]); + cta->eval_Popc(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != expected64[i]) return false; + } + ins.pg.condition = PTXOperand::Pred; + cta->setRegAsPredicate(0, 3, false); + cta->setRegAsU32(0, 0, 0xcafebabe); + cta->eval_Popc(cta->getActiveContext(), ins); + return cta->getRegAsU32(0, 0) == 0xcafebabe; + } + + bool test_Lop3() { + std::stringstream ptx; + ptx << ".version 8.2\n" + << ".target sm_86\n" + << ".address_size 64\n" + << ".visible .entry test_lop3() {\n" + << " .reg .b32 d, a, b, c;\n" + << " .reg .pred p, q;\n" + << " lop3.b32 d, a, b, c, 0x40;\n" + << " lop3.or.b32 d|p, a, b, c, 0x3f, q;\n" + << " lop3.and.b32 _|p, a, b, c, 0x3f, q;\n" + << " ret;\n" + << "}\n"; + Module parsed; + try { + parsed.load(ptx); + } + catch (const hydrazine::Exception& error) { + status << "failed to parse lop3 examples: " << error.what() << "\n"; + return false; + } + + PTXInstruction ins; + ins.opcode = PTXInstruction::Lop3; + ins.type = PTXOperand::b32; + ins.d = reg("d", PTXOperand::b32, 0); + ins.a = reg("a", PTXOperand::b32, 1); + ins.b = reg("b", PTXOperand::b32, 2); + ins.c = reg("c", PTXOperand::b32, 3); + ins.immLut = imm_uint("immLut", PTXOperand::b32, 0x40); + for (int thread = 0; thread < threadCount; ++thread) { + cta->setRegAsU32(thread, 1, 0xffffffffu); + cta->setRegAsU32(thread, 2, 0x0f0f0f0fu); + cta->setRegAsU32(thread, 3, 0x00ff00ffu); + } + cta->eval_Lop3(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU32(thread, 0) != 0x0f000f00u) return false; + } + + ins.booleanOperator = PTXInstruction::BoolOr; + ins.pq = reg("p", PTXOperand::pred, 4); + ins.q = reg("q", PTXOperand::pred, 5); + ins.immLut = imm_uint("immLut", PTXOperand::b32, 0x3f); + for (int thread = 0; thread < threadCount; ++thread) { + cta->setRegAsU32(thread, 1, 0xffffffffu); + cta->setRegAsU32(thread, 2, 0xffffffffu); + cta->setRegAsPredicate(thread, 5, true); + } + cta->eval_Lop3(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU32(thread, 0) != 0 + || !cta->getRegAsPredicate(thread, 4)) return false; + } + + ins.booleanOperator = PTXInstruction::BoolAnd; + ins.d.addressMode = PTXOperand::BitBucket; + for (int thread = 0; thread < threadCount; ++thread) { + cta->setRegAsU32(thread, 1, 0); + cta->setRegAsPredicate(thread, 5, false); + } + cta->eval_Lop3(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsPredicate(thread, 4)) return false; + } + + return true; + } + + // The lop3 predicate source `q` must be walked by assignRegisters like + // every other register operand, or it keeps whatever register id the + // parser left it with instead of the id its defining instruction got. + bool test_Lop3RegisterAllocation() { + std::stringstream ptx; + ptx << ".version 8.2\n" + << ".target sm_86\n" + << ".address_size 64\n" + << ".visible .entry test_lop3_regalloc() {\n" + << " .reg .b32 dd, aa, bb, cc;\n" + << " .reg .pred pp, qq;\n" + << " lop3.b32 dd, aa, bb, cc, 0x40;\n" + << " setp.eq.b32 qq, aa, bb;\n" + << " lop3.or.b32 dd|pp, aa, bb, cc, 0x3f, qq;\n" + << " ret;\n" + << "}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse lop3 register allocation example: " + << error.what() << "\n"; + return false; + } + PTXKernel* kernel = parsed.kernels().begin()->second; + PTXKernel::assignRegisters(*kernel->cfg()); + PTXOperand::RegisterType qqDef = 0, qqUse = 0; + bool sawDef = false, sawUse = false; + for (auto block = kernel->cfg()->begin(); block != kernel->cfg()->end(); ++block) + for (auto instruction = block->instructions.begin(); + instruction != block->instructions.end(); ++instruction) { + const PTXInstruction* instr = + dynamic_cast(*instruction); + if (!instr) continue; + if (instr->opcode == PTXInstruction::SetP) { + qqDef = instr->d.reg; + sawDef = true; + } else if (instr->opcode == PTXInstruction::Lop3 + && instr->q.addressMode == PTXOperand::Register) { + qqUse = instr->q.reg; + sawUse = true; + } + } + if (!sawDef || !sawUse || qqDef != qqUse) { + status << "lop3 q operand register id " << qqUse + << " does not match its definition's id " << qqDef << "\n"; + return false; + } + + return true; + } + /*! Tests several forms of the and instruction */ @@ -2865,7 +7158,7 @@ class TestInstructions: public Test { if (cta->getRegAsB16(t, 0) != expected) { result = false; status << "and.b16 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + << ", got " << cta->getRegAsB16(t, 0) << "\n"; } } } @@ -2887,7 +7180,7 @@ class TestInstructions: public Test { if (cta->getRegAsB32(t, 0) != expected) { result = false; status << "and.b32 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + << ", got " << cta->getRegAsB32(t, 0) << "\n"; } } } @@ -2909,7 +7202,29 @@ class TestInstructions: public Test { if (cta->getRegAsB64(t, 0) != expected) { result = false; status << "and.b64 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + << ", got " << cta->getRegAsB64(t, 0) << "\n"; + } + } + } + + // pred + // + if (result) { + ins.type = PTXOperand::pred; + ins.d = reg("p3", PTXOperand::pred, 0); + ins.a = reg("p1", PTXOperand::pred, 1); + ins.b = reg("p2", PTXOperand::pred, 2); + for (int t = 0; t < threadCount; t++) { + cta->setRegAsPredicate(t, 1, t & 1); + cta->setRegAsPredicate(t, 2, t & 2); + } + cta->eval_And(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; t++) { + const bool expected = (t & 1) && (t & 2); + if (cta->getRegAsPredicate(t, 0) != expected) { + result = false; + status << "and.pred failed\n"; + break; } } } @@ -2944,7 +7259,7 @@ class TestInstructions: public Test { if (cta->getRegAsB16(t, 0) != expected) { result = false; status << "or.b16 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + << ", got " << cta->getRegAsB16(t, 0) << "\n"; } } } @@ -2966,7 +7281,7 @@ class TestInstructions: public Test { if (cta->getRegAsB32(t, 0) != expected) { result = false; status << "or.b32 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + << ", got " << cta->getRegAsB32(t, 0) << "\n"; } } } @@ -2988,7 +7303,29 @@ class TestInstructions: public Test { if (cta->getRegAsB64(t, 0) != expected) { result = false; status << "or.b64 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + << ", got " << cta->getRegAsB64(t, 0) << "\n"; + } + } + } + + // pred + // + if (result) { + ins.type = PTXOperand::pred; + ins.d = reg("p3", PTXOperand::pred, 0); + ins.a = reg("p1", PTXOperand::pred, 1); + ins.b = reg("p2", PTXOperand::pred, 2); + for (int t = 0; t < threadCount; t++) { + cta->setRegAsPredicate(t, 1, t & 1); + cta->setRegAsPredicate(t, 2, t & 2); + } + cta->eval_Or(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; t++) { + const bool expected = (t & 1) || (t & 2); + if (cta->getRegAsPredicate(t, 0) != expected) { + result = false; + status << "or.pred failed\n"; + break; } } } @@ -3023,7 +7360,7 @@ class TestInstructions: public Test { if (cta->getRegAsB16(t, 0) != expected) { result = false; status << "xor.b16 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + << ", got " << cta->getRegAsB16(t, 0) << "\n"; } } } @@ -3045,7 +7382,7 @@ class TestInstructions: public Test { if (cta->getRegAsB32(t, 0) != expected) { result = false; status << "xor.b32 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS32(t, 0) << "\n"; + << ", got " << cta->getRegAsB32(t, 0) << "\n"; } } } @@ -3072,11 +7409,35 @@ class TestInstructions: public Test { } } } + + // pred + // + if (result) { + ins.type = PTXOperand::pred; + ins.d = reg("p3", PTXOperand::pred, 0); + ins.a = reg("p1", PTXOperand::pred, 1); + ins.b = reg("p2", PTXOperand::pred, 2); + for (int t = 0; t < threadCount; t++) { + cta->setRegAsPredicate(t, 1, t & 1); + cta->setRegAsPredicate(t, 2, t & 2); + } + cta->eval_Xor(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; t++) { + const bool a = t & 1; + const bool b = t & 2; + const bool expected = a != b; + if (cta->getRegAsPredicate(t, 0) != expected) { + result = false; + status << "xor.pred failed\n"; + break; + } + } + } return result; } /*! - Tests several forms of the and instruction + Tests several forms of the not instruction */ bool test_Not() { bool result = true; @@ -3101,8 +7462,8 @@ class TestInstructions: public Test { PTXB16 expected = (~t); if (cta->getRegAsB16(t, 0) != expected) { result = false; - status << "xor.b16 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + status << "not.b16 failed (thread " << t << "): expected " << expected + << ", got " << cta->getRegAsB16(t, 0) << "\n"; } } } @@ -3122,8 +7483,8 @@ class TestInstructions: public Test { PTXB32 expected = (~t); if (cta->getRegAsB32(t, 0) != expected) { result = false; - status << "xor.b32 failed (thread " << t << "): expected " << expected - << ", got " << cta->getRegAsS16(t, 0) << "\n"; + status << "not.b32 failed (thread " << t << "): expected " << expected + << ", got " << cta->getRegAsB32(t, 0) << "\n"; } } } @@ -3144,12 +7505,198 @@ class TestInstructions: public Test { PTXB64 got = cta->getRegAsB64(t, 0); if (got != expected) { result = false; - status << "xor.b64 failed (thread " << t << "): expected " << expected + status << "not.b64 failed (thread " << t << "): expected " << expected << ", got " << got << "\n"; } } } - return result; + + // pred + // + if (result) { + ins.type = PTXOperand::pred; + ins.d = reg("p3", PTXOperand::pred, 0); + ins.a = reg("p1", PTXOperand::pred, 1); + for (int t = 0; t < threadCount; t++) { + cta->setRegAsPredicate(t, 1, t & 1); + } + cta->eval_Not(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; t++) { + const bool expected = !(t & 1); + if (cta->getRegAsPredicate(t, 0) != expected) { + result = false; + status << "not.pred failed\n"; + break; + } + } + } + return result; + } + + bool test_Shl() { + bool result = true; + + PTXInstruction ins; + ins.opcode = PTXInstruction::Shl; + ins.type = PTXOperand::b32; + ins.d = reg("r3", PTXOperand::b32, 0); + ins.a = reg("r1", PTXOperand::b32, 1); + ins.b = reg("r2", PTXOperand::u32, 2); + if (!ins.valid().empty()) { + status << "valid shl.b32 rejected: " << ins.valid() << "\n"; + return false; + } + + PTXInstruction invalid = ins; + invalid.d = reg("rd", PTXOperand::b64, 0); + invalid.a = reg("ra", PTXOperand::b64, 1); + if (invalid.valid().empty()) { + status << "shl.b32 accepted 64-bit operands\n"; + return false; + } + invalid = ins; + invalid.b = reg("rf", PTXOperand::f32, 2); + if (invalid.valid().empty()) { + status << "shl.b32 accepted an f32 shift operand\n"; + return false; + } + + cta->reset(); + for (int t = 0; t < threadCount; t++) { + cta->setRegAsB32(t, 1, 1); + cta->setRegAsU32(t, 2, 31 + t % 3); + } + cta->eval_Shl(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; t++) { + const PTXU32 shift = 31 + t % 3; + const PTXB32 expected = shift < 32 ? 1u << shift : 0; + const PTXB32 got = cta->getRegAsB32(t, 0); + if (got != expected) { + result = false; + status << "shl.b32 failed (thread " << t << "): expected " + << expected << ", got " << got << "\n"; + break; + } + } + return result; + } + + bool test_Shf() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Shf; + ins.type = PTXOperand::b32; + ins.d = reg("r4", PTXOperand::b32, 0); + ins.a = reg("r1", PTXOperand::b32, 1); + ins.b = reg("r2", PTXOperand::b32, 2); + ins.c = reg("r3", PTXOperand::u32, 3); + if (!ins.valid().empty()) { + status << "valid shf rejected: " << ins.valid() << "\n"; + return false; + } + PTXInstruction invalid = ins; + invalid.c = reg("f1", PTXOperand::f32, 3); + if (invalid.valid().empty()) { + status << "shf accepted an f32 shift operand\n"; + return false; + } + + const PTXB32 a = 0x12345678u; + const PTXB32 b = 0xabcdef01u; + for (int mode = 0; mode < 2; ++mode) { + ins.shiftMode = mode == 0 ? PTXInstruction::ShiftMode::Wrap + : PTXInstruction::ShiftMode::Clamp; + for (int direction = 0; direction < 2; ++direction) { + ins.shiftDirection = direction == 0 + ? PTXInstruction::ShiftLeft : PTXInstruction::ShiftRight; + cta->reset(); + for (int t = 0; t < threadCount; ++t) { + cta->setRegAsB32(t, 1, a); + cta->setRegAsB32(t, 2, b); + cta->setRegAsU32(t, 3, 31 + t % 4); + } + cta->eval_Shf(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; ++t) { + const PTXU32 c = 31 + t % 4; + const PTXU32 n = mode == 0 ? c & 31 : std::min(c, 32u); + PTXB32 expected; + if (direction == 0) { + expected = n == 0 ? b : n == 32 ? a + : (b << n) | (a >> (32 - n)); + } else { + expected = n == 0 ? a : n == 32 ? b + : (b << (32 - n)) | (a >> n); + } + const PTXB32 got = cta->getRegAsB32(t, 0); + if (got != expected) { + status << "shf failed for count " << c << ": expected " + << expected << ", got " << got << "\n"; + return false; + } + } + } + } + return true; + } + + bool test_Shr() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Shr; + ins.d = reg("r3", PTXOperand::u32, 0); + ins.a = reg("r1", PTXOperand::u32, 1); + ins.b = reg("r2", PTXOperand::u32, 2); + if (!ins.valid().empty()) { + status << "valid shr.u32 rejected: " << ins.valid() << "\n"; + return false; + } + + PTXInstruction invalid = ins; + invalid.d = reg("rd", PTXOperand::u64, 0); + invalid.a = reg("ra", PTXOperand::u64, 1); + if (invalid.valid().empty()) { + status << "shr.u32 accepted 64-bit operands\n"; + return false; + } + invalid = ins; + invalid.b = reg("rf", PTXOperand::f32, 2); + if (invalid.valid().empty()) { + status << "shr.u32 accepted an f32 shift operand\n"; + return false; + } + + cta->reset(); + ins.type = PTXOperand::u32; + for (int t = 0; t < threadCount; t++) { + cta->setRegAsU32(t, 1, 0x80000000u); + cta->setRegAsU32(t, 2, 31 + t % 3); + } + cta->eval_Shr(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; t++) { + const PTXU32 shift = 31 + t % 3; + const PTXU32 expected = shift == 31 ? 1 : 0; + const PTXU32 got = cta->getRegAsU32(t, 0); + if (got != expected) { + status << "shr.u32 failed (thread " << t << "): expected " + << expected << ", got " << got << "\n"; + return false; + } + } + + ins.type = PTXOperand::s32; + ins.d.type = ins.a.type = PTXOperand::s32; + for (int t = 0; t < threadCount; t++) { + cta->setRegAsS32(t, 1, t & 1 ? INT_MIN : INT_MAX); + } + cta->eval_Shr(cta->getActiveContext(), ins); + for (int t = 0; t < threadCount; t++) { + const PTXS32 expected = t & 1 ? -1 : 0; + const PTXS32 got = cta->getRegAsS32(t, 0); + if (got != expected) { + status << "shr.s32 failed (thread " << t << "): expected " + << expected << ", got " << got << "\n"; + return false; + } + } + return true; } ///////////////////////////////////////////////////////////////////////////////////////////////// @@ -3523,283 +8070,1058 @@ class TestInstructions: public Test { cta->setRegAsU32(i, 5, 0); } - cta->eval_Ld(cta->getActiveContext(), ins); + cta->eval_Ld(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + for (int j = 0; j < 4; j++) { + if (cta->getRegAsU32(i, 1+j) != source[j]) { + status << "ld.u32.global.v4 failed\n"; + result = false; + break; + } + } + } + } + + + return result; + } + + bool test_Ld() { + bool result = true; + + // scalar loads + result = (result && test_Ld_global() && test_Ld_shared()); + + // vector loads + result = (result && test_Ld_global_vec()); + + return result; + } + + /*! + Store to global memory + */ + bool test_St_global() { + bool result = true; + + PTXInstruction ins; + ins.opcode = PTXInstruction::St; + + cta->reset(); + + // + // Global memory + // + + ins.addressSpace = PTXInstruction::Global; + + // register indirect + if (result) { + PTXU32 source[64] = { 0 }; + ins.d = reg("ra", PTXOperand::u64, 5); + ins.a = reg("rd", PTXOperand::u32, 0); + ins.d.addressMode = PTXOperand::Indirect; + ins.type = PTXOperand::u32; + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU64(i, 5, (PTXU64)&source[i]); + cta->setRegAsU32(i, 0, i); + } + + cta->eval_St(cta->getActiveContext(), ins); + for (PTXU32 i = 0; i < (PTXU32)threadCount; i++) { + if (source[i] != i) { + result = false; + status << "st.u32.global [reg] failed\n"; + } + } + } + + // register indirect + offset + if (result) { + PTXU32 source[65] = { 0 }; + ins.d = reg("ra", PTXOperand::u64, 5); + ins.a = reg("rd", PTXOperand::u32, 0); + ins.d.addressMode = PTXOperand::Indirect; + ins.d.offset = sizeof(PTXU32); + ins.type = PTXOperand::u32; + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU64(i, 5, (PTXU64)&source[i]); + cta->setRegAsU32(i, 0, i); + } + + cta->eval_St(cta->getActiveContext(), ins); + for (PTXU32 i = 0; i < (PTXU32)threadCount; i++) { + if (source[i+1] != i) { + result = false; + status << "st.u32.global [reg+off] failed. Expected " << (i+1) << ", got " << source[i+1] << "\n"; + } + } + } + + // register indirect + offset + if (result) { + PTXU32 source[65] = { 0 }; + ins.d = reg("ra", PTXOperand::u64, 5); + ins.a = reg("rd", PTXOperand::u32, 0); + ins.d.addressMode = PTXOperand::Immediate; + ins.d.offset = 0; + ins.d.imm_uint = (PTXU64)&source[0]; + ins.type = PTXOperand::u32; + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU32(i, 0, i); + } + + cta->eval_St(cta->getActiveContext(), ins); + + if (source[0] != (PTXU32)threadCount - 1) { + result = false; + status << "st.u32.global [imm] failed\n"; + } + } + + return result; + } + + /*! + Store to global memory + */ + bool test_St_vec() { + bool result = true; + + PTXInstruction ins; + ins.opcode = PTXInstruction::St; + + cta->reset(); + + // + // Global memory + // + + ins.addressSpace = PTXInstruction::Global; + + // register indirect + if (result) { + PTXU32 block[128] __attribute__((aligned(4*sizeof(PTXU32)))) = {0}; + + ins.a = reg("rval", PTXOperand::u32, 1); + ins.a.array.resize( 4 ); + ins.a.array[0] = reg("rval[0]", PTXOperand::u32, 1); + ins.a.array[1] = reg("rval[1]", PTXOperand::u32, 2); + ins.a.array[2] = reg("rval[2]", PTXOperand::u32, 3); + ins.a.array[3] = reg("rval[3]", PTXOperand::u32, 4); + ins.a.vec = PTXOperand::v4; + + ins.d = reg("raddr", PTXOperand::u64, 0); + ins.d.addressMode = PTXOperand::Indirect; + ins.d.offset = 0; + + for (int i = 0; i < threadCount; i++) { + cta->setRegAsU32(i, 0, (PTXU64)&block[i*4]); + cta->setRegAsU32(i, 1, i); + cta->setRegAsU32(i, 2, i*2); + cta->setRegAsU32(i, 3, i*3); + cta->setRegAsU32(i, 4, i*4); + } + + cta->eval_St(cta->getActiveContext(), ins); + for (PTXU32 i = 0; i < (PTXU32)threadCount; i++) { + for (PTXU32 j = 0; j < 4; j++) { + if (block[i*4+j] != i * (j+1)) { + result = false; + status << "st.u32.global.v4 [reg] failed\n"; + } + } + } + } + + return result; + } + + bool test_St() { + bool result = true; + + // scalar stores + result = (result && test_St_global()); + + // vector stores + result = (result && test_St_vec()); + + return result; + } + + ///////////////////////////////////////////////////////////////////////////////////////////////// + // + // mov, cvt + + bool test_Mov() { + bool result = true; + + /* + mov.f32 d,a; + mov.u16 u,v; + mov.f32 k,0.1; + mov.u32 ptr, A; // move address of A into ptr + mov.u32 ptr, A[5]; // move address of A[5] into ptr + mov.b32 addr, myFunc; // get address of myFunc + */ + + PTXInstruction ins; + ins.opcode = PTXInstruction::Mov; + + cta->reset(); + + // from register + + // from special register tidX + if (result) { + ins.d = reg("r6", PTXOperand::u16, 0); + ins.a = sreg(PTXOperand::tid, PTXOperand::ix); + ins.type = PTXOperand::u16; + cta->eval_Mov(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 0) != (PTXU16)i) { + result = false; + status << "mov.u32 r6, tidX failed\n"; + } + } + } + + // from special register tidY + if (result) { + ins.d = reg("r6", PTXOperand::u16, 0); + ins.a = sreg(PTXOperand::tid, PTXOperand::iy); + ins.type = PTXOperand::u16; + cta->eval_Mov(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 0) != 0) { + result = false; + status << "mov.u32 r6, tidY failed\n"; + } + } + } + + // from special register ntidX + if (result) { + ins.d = reg("r6", PTXOperand::u16, 0); + ins.a = sreg(PTXOperand::ntid, PTXOperand::ix); + ins.type = PTXOperand::u16; + cta->eval_Mov(cta->getActiveContext(), ins); + for (int i = 0; i < threadCount; i++) { + if (cta->getRegAsU16(i, 0) != threadCount) { + result = false; + status << "mov.u32 r6, ntidX failed\n"; + } + } + } + + // from special register ntidY + if (result) { + ins.d = reg("r6", PTXOperand::u16, 0); + ins.a = sreg(PTXOperand::ntid, PTXOperand::iy); + ins.type = PTXOperand::u16; + cta->eval_Mov(cta->getActiveContext(), ins); for (int i = 0; i < threadCount; i++) { - for (int j = 0; j < 4; j++) { - if (cta->getRegAsU32(i, 1+j) != source[j]) { - status << "ld.u32.global.v4 failed\n"; - result = false; - break; - } + if (cta->getRegAsU16(i, 0) != 1) { + result = false; + status << "mov.u32 r6, ntidY failed\n"; } } } + // pack a 16-bit immediate and register into a 32-bit destination + if (result) { + ins.d = reg("f10", PTXOperand::f32, 0); + ins.a = PTXOperand(); + ins.a.addressMode = PTXOperand::Register; + ins.a.type = PTXOperand::s16; + ins.a.vec = PTXOperand::v2; + ins.a.array.push_back(imm_uint("0", PTXOperand::s16, 0)); + ins.a.array.push_back(reg("rs1", PTXOperand::s16, 1)); + ins.type = PTXOperand::b32; - return result; - } + for (int i = 0; i < threadCount; ++i) { + cta->setRegAsB16(i, 1, 0x3f80); + } - bool test_Ld() { - bool result = true; - - // scalar loads - result = (result && test_Ld_global() && test_Ld_shared()); + cta->eval_Mov(cta->getActiveContext(), ins); - // vector loads - result = (result && test_Ld_global_vec()); + for (int i = 0; i < threadCount; ++i) { + if (cta->getRegAsB32(i, 0) != 0x3f800000) { + result = false; + status << "mov.b32 f10, {0, rs1} failed\n"; + break; + } + } + } + + // pack two 16-bit immediates into a 32-bit destination + if (result) { + ins.a = PTXOperand(); + ins.a.addressMode = PTXOperand::Register; + ins.a.type = PTXOperand::b16; + ins.a.vec = PTXOperand::v2; + ins.a.array.push_back(imm_uint("5", PTXOperand::b16, 5)); + ins.a.array.push_back(imm_uint("3", PTXOperand::b16, 3)); + + cta->eval_Mov(cta->getActiveContext(), ins); + + for (int i = 0; i < threadCount; ++i) { + if (cta->getRegAsB32(i, 0) != 0x00030005) { + result = false; + status << "mov.b32 f10, {5, 3} failed\n"; + break; + } + } + } + // from label + return result; } - /*! - Store to global memory - */ - bool test_St_global() { + bool test_Cvt() { bool result = true; + std::stringstream cvtPtx; + cvtPtx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_cvt() {\n" + << " .reg .f16x2 d; .reg .f32 a, b;\n" + << " cvt.rn.f16x2.f32 d, a, b; ret; }\n"; + Module parsed; + try { parsed.load(cvtPtx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse .reg .f16x2 CVT: " + << error.what() << "\n"; + return false; + } + std::stringstream packedImmediatePtx; + packedImmediatePtx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry packed_immediate() {\n" + << " .reg .b32 d, td; .reg .b16 bd;\n" + << " cvt.rn.f16x2.f32 d, 0f3f800000, 0f40000000;\n" + << " cvt.rn.bf16x2.f32 d, 0f3f800000, 0f40000000;\n" + << " cvt.rn.bf16.f32 bd, 0f3f800000;\n" + << " cvt.rna.tf32.f32 td, 0f3f800000; ret; }\n"; + try { + Module immediate; + immediate.load(packedImmediatePtx); + Module roundTrip; + std::stringstream serialized(immediate.toString()); + roundTrip.load(serialized); + unsigned int foundA = 0; + unsigned int foundB = 0; + for (Module::StatementVector::const_iterator i = + roundTrip.statements().begin(); i != roundTrip.statements().end(); ++i) { + if (i->directive == PTXStatement::Instr + && i->instruction.opcode == PTXInstruction::Cvt) { + if (i->instruction.a.type == PTXOperand::f32) ++foundA; + if (i->instruction.b.type == PTXOperand::f32) ++foundB; + } + } + if (foundA != 4 || foundB != 2) { + status << "packed cvt immediate lost its f32 source type\n"; + return false; + } + } + catch (const std::exception& error) { + status << "packed cvt immediate parse/round-trip failed: " + << error.what() << "\n"; + return false; + } + std::stringstream invalidBf16Register; + invalidBf16Register << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry invalid_bf16() {\n" + << " .reg .f32 d; .reg .b32 a;\n" + << " cvt.f32.bf16 d, a; ret; }\n"; + try { + Module invalid; + invalid.load(invalidBf16Register); + status << "cvt.f32.bf16 accepted a b32 register\n"; + return false; + } + catch (const std::exception&) { + } + std::stringstream validBf16Register; + validBf16Register << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry valid_bf16() {\n" + << " .reg .f32 d; .reg .b16 a;\n" + << " cvt.f32.bf16 d, a; ret; }\n"; + try { + Module valid; + valid.load(validBf16Register); + } + catch (const std::exception& error) { + status << "cvt.f32.bf16 rejected a b16 register: " + << error.what() << "\n"; + return false; + } + std::stringstream invalidScalarCvtPtx; + invalidScalarCvtPtx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry invalid_cvt() {\n" + << " .reg .u32 d, a, b;\n" + << " cvt.u32.u32 d, a, b; ret; }\n"; + try { + Module invalid; + invalid.load(invalidScalarCvtPtx); + status << "scalar CVT with an extra operand was accepted\n"; + return false; + } + catch (const std::exception&) { + } PTXInstruction ins; - ins.opcode = PTXInstruction::St; + ins.opcode = PTXInstruction::Cvt; + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn | PTXInstruction::ftz; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::s32, 1); + if (!ins.valid().empty()) { + status << "valid cvt.rn.ftz.f32.s32 rejected\n"; + return false; + } + ins.type = PTXOperand::f64; + ins.d = reg("d", PTXOperand::f64, 0); + if (ins.valid().empty()) { + status << "invalid cvt.rn.ftz.f64.s32 accepted\n"; + return false; + } + ins.modifier = 0; + ins.type = PTXOperand::b32; + ins.d = reg("d", PTXOperand::b32, 0); + ins.a = reg("a", PTXOperand::s32, 1); + if (ins.valid().empty()) { + status << "invalid cvt.b32.s32 accepted\n"; + return false; + } + ins.type = PTXOperand::s32; + ins.d = reg("d", PTXOperand::s32, 0); + ins.a = reg("a", PTXOperand::b32, 1); + if (ins.valid().empty()) { + status << "invalid cvt.s32.b32 accepted\n"; + return false; + } + ins.type = PTXOperand::f64; + ins.modifier = PTXInstruction::ftz; + ins.d = reg("d", PTXOperand::f64, 0); + ins.a = reg("a", PTXOperand::b32, 1); + ins.a.relaxedType = PTXOperand::f32; + if (!ins.valid().empty()) { + status << "valid cvt.ftz.f64.f32 with b32 source rejected\n"; + return false; + } + ins.modifier = 0; + ins.type = PTXOperand::f32; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::s32, 1); + if (ins.valid().empty()) { + status << "cvt.f32.s32 accepted without rounding\n"; + return false; + } + ins.type = PTXOperand::s32; + ins.d = reg("d", PTXOperand::s32, 0); + ins.a = reg("a", PTXOperand::f32, 1); + if (ins.valid().empty()) { + status << "cvt.s32.f32 accepted without rounding\n"; + return false; + } + ins.type = PTXOperand::f32; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::f64, 1); + if (ins.valid().empty()) { + status << "cvt.f32.f64 accepted without rounding\n"; + return false; + } + ins.type = PTXOperand::s64; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::s64, 0); + ins.a = reg("a", PTXOperand::s32, 1); + if (ins.valid().empty()) { + status << "cvt.rn.s64.s32 accepted\n"; + return false; + } + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rni; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::f32, 1); + if (!ins.valid().empty()) { + status << "valid cvt.rni.f32.f32 rejected\n"; + return false; + } + ins.type = PTXOperand::s64; + ins.modifier = PTXInstruction::sat; + ins.d = reg("d", PTXOperand::s64, 0); + ins.a = reg("a", PTXOperand::s32, 1); + if (ins.valid().empty()) { + status << "cvt.sat.s64.s32 accepted although saturation is impossible\n"; + return false; + } cta->reset(); + ins.type = PTXOperand::f64; + ins.modifier = PTXInstruction::ftz; + ins.d = reg("d", PTXOperand::f64, 0); + ins.a = reg("a", PTXOperand::f32, 1); + cta->setRegAsU32(0, 1, 0x80000001); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0x8000000000000000ull) { + status << "cvt.ftz.f64.f32 did not flush its f32 input\n"; + return false; + } - // - // Global memory - // + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn | PTXInstruction::ftz; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::f64, 1); + cta->setRegAsF64(0, 1, + -static_cast(std::numeric_limits::denorm_min())); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0x80000000) { + status << "cvt.rn.ftz.f32.f64 did not flush its f32 result\n"; + return false; + } - ins.addressSpace = PTXInstruction::Global; + // f64-to-f32 CVT clears stale upper bits in a wider destination. + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::b64, 0); + ins.a = reg("a", PTXOperand::f64, 1); + cta->setRegAsU64(0, 0, ~0ull); + cta->setRegAsF64(0, 1, 1.0); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0x000000003f800000ull) { + status << "cvt.rn.f32.f64 did not zero-extend its result\n"; + return false; + } + + // cvt.rn.bf16.f32 + ins.type = PTXOperand::bf16; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::b16, 0); + ins.a = reg("a", PTXOperand::f32, 1); + + const PTXU32 input[] = { + 0x3f800000, // exact + 0x3f807fff, // below halfway + 0x3f808000, // halfway, upper even + 0x3f808001, // above halfway + 0x3f818000, // halfway, upper odd + 0x00000000, // +0 + 0x80000000, // -0 + 0x7f800000, // +infinity + 0x7fc00000 // NaN + }; + const PTXU16 expected[] = { + 0x3f80, + 0x3f80, + 0x3f80, + 0x3f81, + 0x3f82, + 0x0000, + 0x8000, + 0x7f80, + 0x7fff + }; + const int cases = sizeof(input) / sizeof(input[0]); + + for (int i = 0; i < threadCount; ++i) { + cta->setRegAsU32(i, 1, input[i % cases]); + cta->setRegAsU16(i, 0, 0); + } + + cta->eval_Cvt(cta->getActiveContext(), ins); + + for (int i = 0; i < threadCount; ++i) { + PTXU16 got = cta->getRegAsU16(i, 0); + if (got != expected[i % cases]) { + status << "cvt.rn.bf16.f32 failed (thread " << i + << "): expected 0x" << hex << expected[i % cases] + << ", got 0x" << got << dec << "\n"; + result = false; + break; + } + } - // register indirect if (result) { - PTXU32 source[64] = { 0 }; - ins.d = reg("ra", PTXOperand::u64, 5); - ins.a = reg("rd", PTXOperand::u32, 0); - ins.d.addressMode = PTXOperand::Indirect; - ins.type = PTXOperand::u32; + ins.modifier = PTXInstruction::rz; + cta->setRegAsU32(0, 1, 0x3f808001); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 0) != 0x3f80) { + status << "cvt.rz.bf16.f32 failed\n"; + result = false; + } + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsU64(i, 5, (PTXU64)&source[i]); - cta->setRegAsU32(i, 0, i); + if (result) { + // cvt.f32.bf16 + ins.type = PTXOperand::f32; + ins.modifier = 0; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::bf16, 1); + + cta->setRegAsU16(0, 1, 0xc020); + cta->eval_Cvt(cta->getActiveContext(), ins); + + if (cta->getRegAsU32(0, 0) != 0xc0200000) { + status << "cvt.f32.bf16 failed\n"; + result = false; } + } - cta->eval_St(cta->getActiveContext(), ins); - for (PTXU32 i = 0; i < (PTXU32)threadCount; i++) { - if (source[i] != i) { - result = false; - status << "st.u32.global [reg] failed\n"; - } + if (result) { + // cvt.rn.f16.f32 + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::b16, 0); + ins.a = reg("a", PTXOperand::f32, 1); + + cta->setRegAsF32(0, 1, 2049.0f); + cta->eval_Cvt(cta->getActiveContext(), ins); + + if (cta->getRegAsU16(0, 0) != 0x6800) { + status << "cvt.rn.f16.f32 failed\n"; + result = false; } } - // register indirect + offset if (result) { - PTXU32 source[65] = { 0 }; - ins.d = reg("ra", PTXOperand::u64, 5); - ins.a = reg("rd", PTXOperand::u32, 0); - ins.d.addressMode = PTXOperand::Indirect; - ins.d.offset = sizeof(PTXU32); - ins.type = PTXOperand::u32; + // cvt.f32.f16 + ins.type = PTXOperand::f32; + ins.modifier = 0; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::f16, 1); - for (int i = 0; i < threadCount; i++) { - cta->setRegAsU64(i, 5, (PTXU64)&source[i]); - cta->setRegAsU32(i, 0, i); + cta->setRegAsU16(0, 1, 0x3c00); + cta->eval_Cvt(cta->getActiveContext(), ins); + + if (cta->getRegAsU32(0, 0) != 0x3f800000) { + status << "cvt.f32.f16 failed\n"; + result = false; } + } - cta->eval_St(cta->getActiveContext(), ins); - for (PTXU32 i = 0; i < (PTXU32)threadCount; i++) { - if (source[i+1] != i) { - result = false; - status << "st.u32.global [reg+off] failed. Expected " << (i+1) << ", got " << source[i+1] << "\n"; - } + if (result) { + // Same-size floating conversions honor integer rounding modifiers. + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rzi; + ins.d = reg("d", PTXOperand::b16, 0); + ins.a = reg("a", PTXOperand::f16, 1); + cta->setRegAsU16(0, 1, 0x3f00); // 1.75 + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 0) != 0x3c00) { + status << "cvt.rzi.f16.f16 failed\n"; + result = false; } } - // register indirect + offset if (result) { - PTXU32 source[65] = { 0 }; - ins.d = reg("ra", PTXOperand::u64, 5); - ins.a = reg("rd", PTXOperand::u32, 0); - ins.d.addressMode = PTXOperand::Immediate; - ins.d.offset = 0; - ins.d.imm_uint = (PTXU64)&source[0]; - ins.type = PTXOperand::u32; + ins.modifier = PTXInstruction::rni; + cta->setRegAsU16(0, 1, 0x4100); // 2.5, ties to even 2 + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 0) != 0x4000) { + status << "cvt.rni.f16.f16 failed\n"; + result = false; + } + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsU32(i, 0, i); + if (result) { + ins.type = PTXOperand::f64; + ins.modifier = PTXInstruction::rzi; + ins.d = reg("d", PTXOperand::f64, 0); + ins.a = reg("a", PTXOperand::f64, 1); + cta->setRegAsF64(0, 1, 1.75); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsF64(0, 0) != 1.0) { + status << "cvt.rzi.f64.f64 failed\n"; + result = false; } + } - cta->eval_St(cta->getActiveContext(), ins); + if (result) { + ins.modifier = PTXInstruction::rni; + cta->setRegAsF64(0, 1, 2.5); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsF64(0, 0) != 2.0) { + status << "cvt.rni.f64.f64 failed\n"; + result = false; + } + } - if (source[0] != (PTXU32)threadCount - 1) { + if (result) { + // .rni is nearest-even regardless of the host mode, then restores it. + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rni; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::f32, 1); + const int previous = hydrazine::fegetround(); + hydrazine::fesetround(FE_UPWARD); + cta->setRegAsF32(0, 1, 2.5f); + cta->eval_Cvt(cta->getActiveContext(), ins); + const bool rounded = cta->getRegAsU32(0, 0) == 0x40000000u; + const bool restored = hydrazine::fegetround() == FE_UPWARD; + hydrazine::fesetround(previous); + if (!rounded || !restored) { + status << "cvt.rni did not force nearest-even or restore host mode\n"; result = false; - status << "st.u32.global [imm] failed\n"; } } - return result; - } + if (result) { + // cvt.rzi.s32.f16 + ins.type = PTXOperand::s32; + ins.modifier = PTXInstruction::rzi; + ins.d = reg("d", PTXOperand::s32, 0); + ins.a = reg("a", PTXOperand::f16, 1); - /*! - Store to global memory - */ - bool test_St_vec() { - bool result = true; + cta->setRegAsU16(0, 1, 0x3e00); + cta->eval_Cvt(cta->getActiveContext(), ins); - PTXInstruction ins; - ins.opcode = PTXInstruction::St; + if (cta->getRegAsS32(0, 0) != 1) { + status << "cvt.rzi.s32.f16 failed\n"; + result = false; + } + } - cta->reset(); + if (result) { + // cvt.rn.f16.s64 + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::b16, 0); + ins.a = reg("a", PTXOperand::s64, 1); - // - // Global memory - // + cta->setRegAsS64(0, 1, 2049); + cta->eval_Cvt(cta->getActiveContext(), ins); - ins.addressSpace = PTXInstruction::Global; + if (cta->getRegAsU16(0, 0) != 0x6800) { + status << "cvt.rn.f16.s64 failed\n"; + result = false; + } + } - // register indirect if (result) { - PTXU32 block[128] __attribute__((aligned(4*sizeof(PTXU32)))) = {0}; + // cvt.rn.f16.f64 + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::b16, 0); + ins.a = reg("a", PTXOperand::f64, 1); - ins.a = reg("rval", PTXOperand::u32, 1); - ins.a.array.resize( 4 ); - ins.a.array[0] = reg("rval[0]", PTXOperand::u32, 1); - ins.a.array[1] = reg("rval[1]", PTXOperand::u32, 2); - ins.a.array[2] = reg("rval[2]", PTXOperand::u32, 3); - ins.a.array[3] = reg("rval[3]", PTXOperand::u32, 4); - ins.a.vec = PTXOperand::v4; + cta->setRegAsF64(0, 1, 1.0); + cta->eval_Cvt(cta->getActiveContext(), ins); - ins.d = reg("raddr", PTXOperand::u64, 0); - ins.d.addressMode = PTXOperand::Indirect; - ins.d.offset = 0; + if (cta->getRegAsU16(0, 0) != 0x3c00) { + status << "cvt.rn.f16.f64 failed\n"; + result = false; + } + } - for (int i = 0; i < threadCount; i++) { - cta->setRegAsU32(i, 0, (PTXU64)&block[i*4]); - cta->setRegAsU32(i, 1, i); - cta->setRegAsU32(i, 2, i*2); - cta->setRegAsU32(i, 3, i*3); - cta->setRegAsU32(i, 4, i*4); + if (result) { + // Integer saturation clamps both ends of the destination range. + ins.type = PTXOperand::u8; + ins.modifier = PTXInstruction::sat; + ins.d = reg("d", PTXOperand::u8, 0); + ins.a = reg("a", PTXOperand::s32, 1); + cta->setRegAsS32(0, 1, 300); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU8(0, 0) != 255) { + status << "cvt.sat.u8.s32 failed\n"; + result = false; } + } - cta->eval_St(cta->getActiveContext(), ins); - for (PTXU32 i = 0; i < (PTXU32)threadCount; i++) { - for (PTXU32 j = 0; j < 4; j++) { - if (block[i*4+j] != i * (j+1)) { - result = false; - status << "st.u32.global.v4 [reg] failed\n"; - } - } + if (result) { + // Floating saturation clamps to [0, 1]. + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn | PTXInstruction::sat; + ins.d = reg("d", PTXOperand::f32, 0); + ins.a = reg("a", PTXOperand::s32, 1); + cta->setRegAsS32(0, 1, -2); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0) { + status << "cvt.rn.sat.f32.s32 failed\n"; + result = false; + } + } + + if (result) { + // Unsigned and floating results zero-extend in wider registers. + ins.type = PTXOperand::u32; + ins.modifier = 0; + ins.d = reg("d", PTXOperand::u64, 0); + ins.a = reg("a", PTXOperand::s32, 1); + cta->setRegAsS32(0, 1, -1); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0x00000000ffffffffull) { + status << "cvt.u32.s32 did not zero-extend its result\n"; + result = false; + } + } + + if (result) { + // Exact exclusive upper bounds must not reach an integer cast. + ins.type = PTXOperand::s32; + ins.modifier = PTXInstruction::rzi; + ins.d = reg("d", PTXOperand::s32, 0); + ins.a = reg("a", PTXOperand::f32, 1); + cta->setRegAsF32(0, 1, std::ldexp(1.0f, 31)); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsS32(0, 0) != INT_MAX) { + status << "cvt.rzi.s32.f32 upper boundary failed\n"; + result = false; } } - return result; - } - - bool test_St() { - bool result = true; - - // scalar stores - result = (result && test_St_global()); - - // vector stores - result = (result && test_St_vec()); + if (result) { + ins.type = PTXOperand::u32; + ins.modifier = PTXInstruction::rzi; + ins.d = reg("d", PTXOperand::u32, 0); + ins.a = reg("a", PTXOperand::f64, 1); + cta->setRegAsF64(0, 1, std::ldexp(1.0, 32)); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != UINT_MAX) { + status << "cvt.rzi.u32.f64 upper boundary failed\n"; + result = false; + } + } - return result; - } + if (result) { + ins.type = PTXOperand::s64; + ins.modifier = PTXInstruction::rzi; + ins.d = reg("d", PTXOperand::s64, 0); + ins.a = reg("a", PTXOperand::f64, 1); + cta->setRegAsF64(0, 1, std::ldexp(1.0, 63)); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsS64(0, 0) != LLONG_MAX) { + status << "cvt.rzi.s64.f64 upper boundary failed\n"; + result = false; + } + } - ///////////////////////////////////////////////////////////////////////////////////////////////// - // - // mov, cvt - - bool test_Mov() { - bool result = true; + if (result) { + ins.type = PTXOperand::f32; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::b64, 0); + ins.a = reg("a", PTXOperand::s32, 1); + if (!ins.valid().empty()) { + status << "cvt.rn.f32.s32 rejected a wider destination register\n"; + result = false; + } else { + cta->setRegAsU64(0, 0, ~0ull); + cta->setRegAsS32(0, 1, 1); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0x000000003f800000ull) { + status << "cvt.rn.f32.s32 did not zero-extend its result\n"; + result = false; + } + } + } - /* - mov.f32 d,a; - mov.u16 u,v; - mov.f32 k,0.1; - mov.u32 ptr, A; // move address of A into ptr - mov.u32 ptr, A[5]; // move address of A[5] into ptr - mov.b32 addr, myFunc; // get address of myFunc - */ + if (result) { + // ReLU and packed conversions supported by sm_86. + ins.type = PTXOperand::f16x2; + ins.modifier = PTXInstruction::rn | PTXInstruction::relu; + ins.d = reg("d", PTXOperand::b32, 0); + ins.a = reg("a", PTXOperand::f32, 1); + ins.b = reg("b", PTXOperand::f32, 2); + cta->setRegAsF32(0, 1, 1.0f); + cta->setRegAsF32(0, 2, -2.0f); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0x3c000000) { + status << "cvt.rn.relu.f16x2.f32 failed\n"; + result = false; + } + } - PTXInstruction ins; - ins.opcode = PTXInstruction::Mov; + if (result) { + ins.type = PTXOperand::bf16x2; + ins.modifier = PTXInstruction::rn; + cta->setRegAsF32(0, 1, 1.0f); + cta->setRegAsF32(0, 2, 2.0f); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0x3f804000) { + status << "cvt.rn.bf16x2.f32 failed\n"; + result = false; + } + } - cta->reset(); + if (result) { + ins.type = PTXOperand::tf32; + ins.modifier = PTXInstruction::rna; + ins.b = PTXOperand(); + cta->setRegAsU32(0, 1, 0x3f801000); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0x3f802000) { + status << "cvt.rna.tf32.f32 failed\n"; + result = false; + } + } - // from register + if (result) { + // CVT sources are read before writing an aliased destination. + ins.type = PTXOperand::u32; + ins.modifier = 0; + ins.d = reg("r0", PTXOperand::b64, 0); + ins.a = reg("r0", PTXOperand::s32, 0); + cta->setRegAsS32(0, 0, -1); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 0xffffffffull) { + status << "aliased scalar cvt lost its source\n"; + result = false; + } + } - // from special register tidX if (result) { - ins.d = reg("r6", PTXOperand::u16, 0); - ins.a = sreg(PTXOperand::tid, PTXOperand::ix); - ins.type = PTXOperand::u16; - cta->eval_Mov(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - if (cta->getRegAsU16(i, 0) != (PTXU16)i) { - result = false; - status << "mov.u32 r6, tidX failed\n"; - } + // Both packed sources may alias the packed destination register. + ins.type = PTXOperand::f16x2; + ins.modifier = PTXInstruction::rn; + ins.d = reg("r0", PTXOperand::b32, 0); + ins.a = reg("r0", PTXOperand::f32, 0); + ins.b = reg("r0", PTXOperand::f32, 0); + cta->setRegAsF32(0, 0, 1.0f); + cta->eval_Cvt(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 0) != 0x3c003c00u) { + status << "aliased packed cvt lost its sources\n"; + result = false; } } - // from special register tidY if (result) { - ins.d = reg("r6", PTXOperand::u16, 0); - ins.a = sreg(PTXOperand::tid, PTXOperand::iy); - ins.type = PTXOperand::u16; - cta->eval_Mov(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - if (cta->getRegAsU16(i, 0) != 0) { - result = false; - status << "mov.u32 r6, tidY failed\n"; - } + // ReLU maps NaN to the canonical FP16 NaN encoding. + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn | PTXInstruction::relu; + ins.d = reg("d", PTXOperand::b16, 0); + ins.a = reg("a", PTXOperand::f32, 1); + ins.b = PTXOperand(); + if (!ins.valid().empty()) { + status << "valid scalar ReLU CVT rejected before evaluation\n"; + result = false; + } + cta->setRegAsU32(0, 1, 0xffc12345u); + if (result) cta->eval_Cvt(cta->getActiveContext(), ins); + if (result && cta->getRegAsU16(0, 0) != 0x7fffu) { + status << "cvt.rn.relu.f16.f32 did not canonicalize NaN\n"; + result = false; } } - // from special register ntidX if (result) { - ins.d = reg("r6", PTXOperand::u16, 0); - ins.a = sreg(PTXOperand::ntid, PTXOperand::ix); - ins.type = PTXOperand::u16; - cta->eval_Mov(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - if (cta->getRegAsU16(i, 0) != threadCount) { - result = false; - status << "mov.u32 r6, ntidX failed\n"; - } + // Packed f16x2 accepts relaxed b32 sources and wider destinations. + ins.type = PTXOperand::f16x2; + ins.modifier = PTXInstruction::rn; + ins.d = reg("d", PTXOperand::b64, 0); + ins.a = reg("a", PTXOperand::b32, 1); + ins.a.relaxedType = PTXOperand::f32; + ins.b = reg("b", PTXOperand::b32, 2); + ins.b.relaxedType = PTXOperand::f32; + if (!ins.valid().empty()) { + status << "packed cvt rejected relaxed sources or destination\n"; + result = false; } } - // from special register ntidY if (result) { - ins.d = reg("r6", PTXOperand::u16, 0); - ins.a = sreg(PTXOperand::ntid, PTXOperand::iy); - ins.type = PTXOperand::u16; - cta->eval_Mov(cta->getActiveContext(), ins); - for (int i = 0; i < threadCount; i++) { - if (cta->getRegAsU16(i, 0) != 1) { + // Scalar CVT rejects invalid ReLU modifiers. + ins.type = PTXOperand::f16; + ins.modifier = PTXInstruction::rn | PTXInstruction::relu + | PTXInstruction::sat; + ins.d = reg("d", PTXOperand::b16, 0); + ins.a = reg("a", PTXOperand::f32, 1); + ins.b = PTXOperand(); + if (ins.valid().empty()) { + status << "cvt.rn.sat.relu.f16.f32 accepted\n"; + result = false; + } + if (result) { + ins.modifier = PTXInstruction::rn | PTXInstruction::relu + | PTXInstruction::ftz; + if (ins.valid().empty()) { + status << "cvt.relu.ftz.f16.f32 accepted\n"; result = false; - status << "mov.u32 r6, ntidY failed\n"; } } } - - // from label return result; } - bool test_Cvt() { - bool result = true; - + bool test_Cvta() { PTXInstruction ins; - ins.opcode = PTXInstruction::Cvt; + ins.opcode = PTXInstruction::Cvta; + ins.addressSpace = PTXInstruction::Param; + ins.type = PTXOperand::u64; + ins.d = reg("d", PTXOperand::u64, 1); + ins.a = reg("a", PTXOperand::u64, 0); + if (!ins.valid().empty()) return false; cta->reset(); + const PTXU64 expected = (PTXU64)kernel->ArgumentMemory + 1; + for (int thread = 0; thread < threadCount; ++thread) { + cta->setRegAsU64(thread, 0, 1); + } + cta->eval_Cvta(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU64(thread, 1) != expected) return false; + cta->setRegAsU64(thread, 0, expected); + } - // - - return result; + ins.toAddrSpace = true; + cta->eval_Cvta(cta->getActiveContext(), ins); + for (int thread = 0; thread < threadCount; ++thread) { + if (cta->getRegAsU64(thread, 1) != 1) return false; + } + cta->reset(); + cta->functionCallStack.pushFrame(0, kernel->registerCount(), 16, 0, 0, 0, 0); + ins.pg.condition = PTXOperand::PT; + ins.addressSpace = PTXInstruction::Local; + for (PTXOperand::DataType type : {PTXOperand::u32, PTXOperand::u64}) { + ins.type = type; + for (PTXOperand::RegisterType source : {0u, 1u}) { + for (PTXU64 offset : {PTXU64(0), PTXU64(15)}) { + ins.toAddrSpace = false; + ins.a = reg("a", type, source); + ins.d = reg("d", type, 2); + for (int thread = 0; thread < 2; ++thread) { + if (type == PTXOperand::u32) + cta->setRegAsU32(thread, source, static_cast(offset)); + else cta->setRegAsU64(thread, source, offset); + } + cta->eval_Cvta(cta->getActiveContext(), ins); + for (int thread = 0; thread < 2; ++thread) { + const PTXU64 expected = reinterpret_cast( + cta->functionCallStack.localMemoryPointer(thread)) + offset; + if ((type == PTXOperand::u32 + ? cta->getRegAsU32(thread, 2) : cta->getRegAsU64(thread, 2)) + != (type == PTXOperand::u32 ? static_cast(expected) : expected)) + return false; + } + ins.toAddrSpace = true; + ins.a = reg("g", type, 2); + ins.d = reg("o", type, 3); + cta->eval_Cvta(cta->getActiveContext(), ins); + for (int thread = 0; thread < 2; ++thread) + if ((type == PTXOperand::u32 + ? cta->getRegAsU32(thread, 3) : cta->getRegAsU64(thread, 3)) + != (type == PTXOperand::u32 ? static_cast(offset) : offset)) + return false; + } + } + } + ins.type = PTXOperand::u64; + ins.toAddrSpace = false; + ins.addressSpace = PTXInstruction::Local; + ins.a = reg("a", PTXOperand::u64, 1); + ins.d = reg("d", PTXOperand::u64, 2); + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 4; + for (int thread = 0; thread < 2; ++thread) { + cta->setRegAsPredicate(thread, 4, false); + cta->setRegAsU64(thread, 2, 0xdeadbeefULL); + } + cta->eval_Cvta(cta->getActiveContext(), ins); + for (int thread = 0; thread < 2; ++thread) + if (cta->getRegAsU64(thread, 2) != 0xdeadbeefULL) return false; + cta->functionCallStack.popFrame(); + return true; } ///////////////////////////////////////////////////////////////////////////////////////////////// @@ -3900,6 +9222,62 @@ class TestInstructions: public Test { } } + if (result) { + // set.eq.f16.f16.and + ins = PTXInstruction(); + ins.opcode = PTXInstruction::Set; + ins.type = PTXOperand::f16; + ins.d = reg("d", PTXOperand::b16, 3); + ins.a = reg("a", PTXOperand::f16, 1); + ins.b = reg("b", PTXOperand::f16, 2); + ins.c = reg("c", PTXOperand::pred, 0); + ins.comparisonOperator = PTXInstruction::Eq; + ins.booleanOperator = PTXInstruction::BoolAnd; + + cta->setRegAsU16(0, 1, 0x3c00); // 1.0 + cta->setRegAsU16(0, 2, 0x3c00); // 1.0 + cta->setRegAsPredicate(0, 0, true); + cta->eval_Set(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 3) != 0x3c00) { + status << "[set.eq.f16.f16.and test] failed\n"; + result = false; + } + + cta->setRegAsPredicate(0, 0, false); + cta->eval_Set(cta->getActiveContext(), ins); + if (cta->getRegAsU16(0, 3) != 0x0000) { + status << "[set.eq.f16.f16.and false predicate test] failed\n"; + result = false; + } + } + + if (result) { + // set.num.u32.f64 is false when either input is NaN. + ins = PTXInstruction(); + ins.opcode = PTXInstruction::Set; + ins.type = PTXOperand::u32; + ins.d = reg("d", PTXOperand::u32, 3); + ins.a = reg("a", PTXOperand::f64, 1); + ins.b = reg("b", PTXOperand::f64, 2); + ins.comparisonOperator = PTXInstruction::Num; + + cta->setRegAsF64(0, 1, + std::numeric_limits::quiet_NaN()); + cta->setRegAsF64(0, 2, 1.0); + cta->eval_Set(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 3) != 0) { + status << "[set.num.u32.f64 NaN test] failed\n"; + result = false; + } + + ins.comparisonOperator = PTXInstruction::Ne; + cta->eval_Set(cta->getActiveContext(), ins); + if (cta->getRegAsU32(0, 3) != 0) { + status << "[set.ne.u32.f64 NaN test] failed\n"; + result = false; + } + } + return result; } @@ -3911,6 +9289,45 @@ class TestInstructions: public Test { cta->reset(); + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_setp_sink() {\n" + << " .reg .pred p;\n .reg .s32 a, b;\n" + << " setp.eq.s32 _|p, a, b;\n" + << " setp.ne.s32 p|_, a, b;\n ret;\n}\n"; + Module parsed; + try { parsed.load(ptx); } + catch (const hydrazine::Exception& error) { + status << "failed to parse setp sink forms: " << error.what() << "\n"; + return false; + } + + PTXOperand sink; + sink.identifier = "_"; + sink.type = PTXOperand::b64; + sink.addressMode = PTXOperand::BitBucket; + sink.reg = 0; + ins.type = PTXOperand::s32; + ins.a = reg("a", PTXOperand::s32, 1); + ins.b = reg("b", PTXOperand::s32, 2); + ins.comparisonOperator = PTXInstruction::Eq; + cta->setRegAsS32(0, 1, 1); + cta->setRegAsS32(0, 2, 2); + cta->setRegAsPredicate(0, 0, true); + + ins.d = sink; + ins.pq = reg("q", PTXOperand::pred, 3); + cta->eval_SetP(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0) + || !cta->getRegAsPredicate(0, 3)) return false; + + ins.d = reg("p", PTXOperand::pred, 3); + ins.pq = sink; + cta->setRegAsS32(0, 2, 1); + cta->eval_SetP(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0) + || !cta->getRegAsPredicate(0, 3)) return false; + if (result) { // setp.s32.lt p|q, a, b; // p = (a < b); q = !(a < b); // @@ -4119,6 +9536,65 @@ class TestInstructions: public Test { } } + if (result) { + // Unordered floating-point comparisons are true when an input is NaN. + ins = PTXInstruction(); + ins.opcode = PTXInstruction::SetP; + ins.type = PTXOperand::f64; + ins.d = reg("p", PTXOperand::pred, 3); + ins.pq = reg("q", PTXOperand::pred, 4); + ins.a = reg("a", PTXOperand::f64, 1); + ins.b = reg("b", PTXOperand::f64, 2); + ins.comparisonOperator = PTXInstruction::Equ; + + cta->setRegAsF64(0, 1, + std::numeric_limits::quiet_NaN()); + cta->setRegAsF64(0, 2, 1.0); + cta->eval_SetP(cta->getActiveContext(), ins); + + if (!cta->getRegAsPredicate(0, 3) || + cta->getRegAsPredicate(0, 4)) { + status << "[f64 Equ NaN test] " << ins.toString() + << " failed\n"; + result = false; + } + } + + if (result) { + // Half inputs are widened exactly, with FTZ applied before widening. + ins = PTXInstruction(); + ins.opcode = PTXInstruction::SetP; + ins.type = PTXOperand::f16; + ins.d = reg("p", PTXOperand::pred, 3); + ins.pq = reg("q", PTXOperand::pred, 4); + ins.a = reg("a", PTXOperand::b16, 1); + ins.b = reg("b", PTXOperand::b16, 2); + ins.comparisonOperator = PTXInstruction::Lt; + + cta->setRegAsU16(0, 1, 0x3c00); // 1.0 + cta->setRegAsU16(0, 2, 0x4000); // 2.0 + cta->eval_SetP(cta->getActiveContext(), ins); + const bool normal = cta->getRegAsPredicate(0, 3) && + !cta->getRegAsPredicate(0, 4); + + ins.comparisonOperator = PTXInstruction::Eq; + cta->setRegAsU16(0, 1, 0x0001); // minimum half subnormal + cta->setRegAsU16(0, 2, 0x0000); + cta->eval_SetP(cta->getActiveContext(), ins); + const bool preserved = !cta->getRegAsPredicate(0, 3) && + cta->getRegAsPredicate(0, 4); + + ins.modifier = PTXInstruction::ftz; + cta->eval_SetP(cta->getActiveContext(), ins); + const bool flushed = cta->getRegAsPredicate(0, 3) && + !cta->getRegAsPredicate(0, 4); + + if (!normal || !preserved || !flushed) { + status << "[f16 widening/FTZ test] failed\n"; + result = false; + } + } + return result; } @@ -4220,6 +9696,40 @@ class TestInstructions: public Test { cta->reset(); + // .ftz is controlled by the f32 comparison type, not the result type. + ins.type = PTXOperand::u64; + ins.modifier = PTXInstruction::ftz; + ins.d = reg("d", PTXOperand::u64, 0); + ins.a = reg("a", PTXOperand::u64, 1); + ins.b = reg("b", PTXOperand::u64, 2); + ins.c = reg("c", PTXOperand::f32, 3); + if (!ins.valid().empty()) { + status << "slct.ftz.u64.f32 should be valid\n"; + result = false; + } + PTXInstruction invalid = ins; + invalid.type = PTXOperand::f32; + invalid.d = reg("d", PTXOperand::f32, 0); + invalid.a = reg("a", PTXOperand::f32, 1); + invalid.b = reg("b", PTXOperand::f32, 2); + invalid.c = reg("c", PTXOperand::s32, 3); + if (invalid.valid().empty()) { + status << "slct.ftz.f32.s32 should be invalid\n"; + result = false; + } + + if (result) { + cta->setRegAsU64(0, 1, 11); + cta->setRegAsU64(0, 2, 22); + cta->setRegAsF32(0, 3, + -std::numeric_limits::denorm_min()); + cta->eval_SlCt(cta->getActiveContext(), ins); + if (cta->getRegAsU64(0, 0) != 11) { + status << "slct.ftz.u64.f32 did not select a\n"; + result = false; + } + } + if (result) { // slct.f32.f32 r, a, b, c // @@ -4260,67 +9770,303 @@ class TestInstructions: public Test { } bool test_TestP() { - bool result = false; - /* PTXInstruction ins; ins.opcode = PTXInstruction::TestP; - + ins.d = reg("p", PTXOperand::pred, 0); cta->reset(); + for (int typeIndex = 0; typeIndex < 2; ++typeIndex) { + ins.type = typeIndex == 0 ? PTXOperand::f32 : PTXOperand::f64; + ins.a = reg("a", ins.type, 1); + if (ins.type == PTXOperand::f32) { + cta->setRegAsF32(0, 1, 0.0f); + cta->setRegAsF32(1, 1, -0.0f); + cta->setRegAsF32(2, 1, std::numeric_limits::denorm_min()); + cta->setRegAsF32(3, 1, 1.0f); + } else { + cta->setRegAsF64(0, 1, 0.0); + cta->setRegAsF64(1, 1, -0.0); + cta->setRegAsF64(2, 1, std::numeric_limits::denorm_min()); + cta->setRegAsF64(3, 1, 1.0); + } + for (int mode = 0; mode < 2; ++mode) { + const bool normal = mode == 0; + ins.floatingPointMode = normal + ? PTXInstruction::Normal : PTXInstruction::SubNormal; + cta->eval_TestP(cta->getActiveContext(), ins); + if (cta->getRegAsPredicate(0, 0) + || cta->getRegAsPredicate(1, 0) + || cta->getRegAsPredicate(2, 0) != !normal + || cta->getRegAsPredicate(3, 0) != normal) { + status << "testp normal/subnormal classification failed\n"; + return false; + } + } + } + return true; + } - // f32 - // - if (result) { - // testp.op.type p, a - // - // op: .finite, .infinite, .number, .notanumber, .normal, .subnormal - // type: .f32, .f64 - // - ins.type = PTXOperand::f32; - ins.d = reg("p", PTXOperand::pred, 0); - ins.a = reg("a", PTXOperand::f32, 1); - - ir::PTXInstruction::FloatingPointMode floatModes[] = { - ir::PTXInstruction::Finite, - ir::PTXInstruction::Infinite, - ir::PTXInstruction::Number, - ir::PTXInstruction::NotANumber, - ir::PTXInstruction::Normal, - ir::PTXInstruction::SubNormal, - ir::PTXInstruction::FloatingPointMode_Invalid - }; - - PTXF32 floatValues[] = { - -1, 0, 1, FLT_EPSILON, -FLT_EPSILON, 0 - }; - - for (int mode = 0; floatModes[mode] != ir::PTXInstruction::FloatingPointMode_Invalid; mode++) { - ins.opcode = PTXInstruction::TestP; - ins.floatingPointMode = floatModes[mode]; - ins.d = reg("p", PTXOperand::pred, 0); - ins.a = reg("a", PTXOperand::f32, 1); - - - - } - + bool test_Isspacep() { + std::stringstream ptx; + ptx << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry test_isspacep() {\n" + << " .reg .pred %p; .reg .u32 %r; .reg .u64 %rd;\n" + << " isspacep.shared %p, %r;\n" + << " isspacep.shared::cta %p, %rd; ret; }\n"; + Module parsed; + try { parsed.load(ptx); } + catch (...) { + status << "isspacep.shared::cta failed to parse\n"; + return false; } - - // f64 - // - if (result) { - // testp.op.type p, a - // - // op: .finite, .infinite, .number, .notanumber, .normal, .subnormal - // type: .f32, .f64 - // - ins.type = PTXOperand::f32; - ins.d = reg("p", PTXOperand::pred, 0); - ins.a = reg("a", PTXOperand::f32, 1); - - + PTXInstruction parsed32, parsed64; + unsigned int parsedIsspacep = 0; + for (Module::StatementVector::const_iterator i = parsed.statements().begin(); + i != parsed.statements().end(); ++i) { + if (i->directive != PTXStatement::Instr + || i->instruction.opcode != PTXInstruction::Isspacep) continue; + if (parsedIsspacep++ == 0) parsed32 = i->instruction; + else parsed64 = i->instruction; + } + if (parsedIsspacep != 2 || parsed32.a.type != PTXOperand::u32 + || parsed64.a.type != PTXOperand::u64 + || parsed32.addressSpace != PTXInstruction::Shared + || parsed64.addressSpace != PTXInstruction::Shared) return false; + std::stringstream serialized(parsed.toString()); + try { Module reparsed; reparsed.load(serialized); } + catch (...) { return false; } + if (parsed.toString().find("shared::cta") != std::string::npos) return false; + for (const char* space : {"global::cta", "local::cta", "shared.cta", + "shared::cluster"}) { + std::stringstream invalid; + invalid << ".version 8.0\n.target sm_86\n.address_size 64\n" + << ".visible .entry invalid() { .reg .pred %p; .reg .u64 %rd;\n" + << "isspacep." << space << " %p, %rd; ret; }\n"; + try { Module rejected; rejected.load(invalid); return false; } + catch (...) {} } - */ - return result; + PTXInstruction ins; + ins.opcode = PTXInstruction::Isspacep; + ins.addressSpace = PTXInstruction::Shared; + ins.d = reg("p", PTXOperand::pred, 0); + ins.a = reg("a", PTXOperand::u64, 1); + if (!ins.valid().empty()) return false; + ins.d.type = PTXOperand::u32; + if (ins.valid().empty()) return false; + ins.d.type = PTXOperand::pred; + for (PTXOperand::DataType type : {PTXOperand::s32, PTXOperand::f32, + PTXOperand::u16, PTXOperand::b64}) { + ins.a.type = type; + if (ins.valid().empty()) return false; + } + ins.a.type = PTXOperand::u64; + ins.a.vec = PTXOperand::v2; + if (ins.valid().empty()) return false; + ins.a.vec = PTXOperand::v1; + + cta->reset(); + cta->functionCallStack.pushFrame(0, kernel->registerCount(), 16, 64, + 0, 0, 0); + PTXU64 shared64, local64; + hydrazine::bit_cast(shared64, cta->functionCallStack.sharedMemoryPointer()); + hydrazine::bit_cast(local64, cta->functionCallStack.localMemoryPointer(0)); + for (const PTXInstruction& qualified : {parsed32, parsed64}) { + const PTXU64 base = shared64; + const PTXU64 size = cta->functionCallStack.sharedMemorySize(); + for (PTXU64 offset : {PTXU64(0), size - 1, size}) { + if (qualified.a.type == PTXOperand::u32) + cta->setRegAsU32(0, qualified.a.reg, static_cast(base + offset)); + else cta->setRegAsU64(0, qualified.a.reg, base + offset); + cta->eval_Isspacep(cta->getActiveContext(), qualified); + if (cta->getRegAsPredicate(0, qualified.d.reg) != (offset < size)) return false; + } + } + ins.a.type = PTXOperand::u64; + ins.addressSpace = PTXInstruction::Shared; + cta->setRegAsU64(0, 1, shared64); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + cta->setRegAsU64(0, 1, local64); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (cta->getRegAsPredicate(0, 0)) return false; + ins.addressSpace = PTXInstruction::Local; + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + ins.addressSpace = PTXInstruction::Global; + cta->setRegAsU64(0, 1, shared64); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (cta->getRegAsPredicate(0, 0)) return false; + cta->setRegAsU64(0, 1, 0x123456789ULL); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + PTXU32 shared32; + hydrazine::bit_cast(shared32, cta->functionCallStack.sharedMemoryPointer()); + ins.a.type = PTXOperand::u32; + ins.addressSpace = PTXInstruction::Shared; + cta->setRegAsU32(0, 1, shared32); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + PTXU32 local32; + hydrazine::bit_cast(local32, cta->functionCallStack.localMemoryPointer(0)); + ins.addressSpace = PTXInstruction::Local; + cta->setRegAsU32(0, 1, local32); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + ins.addressSpace = PTXInstruction::Global; + cta->setRegAsU32(0, 1, local32); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (cta->getRegAsPredicate(0, 0)) return false; + cta->setRegAsPredicate(0, 0, true); + cta->setRegAsPredicate(0, 2, false); + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 2; + ins.addressSpace = PTXInstruction::Shared; + cta->setRegAsU32(0, 1, local32); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + cta->setRegAsPredicate(0, 0, false); + cta->setRegAsU32(0, 1, shared32); + cta->eval_Isspacep(cta->getActiveContext(), ins); + bool predicatedOff = !cta->getRegAsPredicate(0, 0); + + *reinterpret_cast(local64) = 0x12345678; + cta->functionCallStack.pushFrame(0, kernel->registerCount(), 16, 64, + 0, 0, 0); + PTXInstruction load; + load.opcode = PTXInstruction::Ld; + load.addressSpace = PTXInstruction::Generic; + load.type = PTXOperand::u32; + load.d = reg("d", PTXOperand::u32, 3); + load.a = reg("a", PTXOperand::u64, 1); + load.a.addressMode = PTXOperand::Indirect; + load.pg.condition = PTXOperand::Pred; + load.pg.reg = 4; + for (int thread = 0; thread < threadCount; ++thread) + cta->setRegAsPredicate(thread, 4, thread == 0); + cta->setRegAsU64(0, 1, local64); + cta->eval_Ld(cta->getActiveContext(), load); + bool callerLocal = cta->getRegAsU32(0, 3) == 0x12345678; + ins.a.type = PTXOperand::u64; + ins.pg.condition = PTXOperand::PT; + ins.addressSpace = PTXInstruction::Local; + cta->eval_Isspacep(cta->getActiveContext(), ins); + callerLocal &= cta->getRegAsPredicate(0, 0); + ins.addressSpace = PTXInstruction::Global; + cta->eval_Isspacep(cta->getActiveContext(), ins); + callerLocal &= !cta->getRegAsPredicate(0, 0); + ins.a.type = PTXOperand::u32; + cta->setRegAsU32(0, 1, static_cast(local64)); + ins.addressSpace = PTXInstruction::Local; + cta->eval_Isspacep(cta->getActiveContext(), ins); + callerLocal &= cta->getRegAsPredicate(0, 0); + ins.addressSpace = PTXInstruction::Global; + cta->eval_Isspacep(cta->getActiveContext(), ins); + callerLocal &= !cta->getRegAsPredicate(0, 0); + cta->functionCallStack.popFrame(); + cta->functionCallStack.popFrame(); + return predicatedOff && callerLocal; + } + + bool test_IsspacepConst() { + ConstMemoryTestKernel constKernel( + module.getKernel("_Z17k_simple_sequencePi")); + constKernel.initialize(); + constKernel.setKernelShape(1, 1, 1); + constKernel.setConstMemory(16); + CooperativeThreadArray constCta(&constKernel, ir::Dim3(), false); + constCta.functionCallStack.pushFrame(0, constKernel.registerCount(), + 16, 64, 0, 0, 0); + PTXInstruction ins; + ins.opcode = PTXInstruction::Isspacep; + ins.addressSpace = PTXInstruction::Const; + ins.d = reg("p", PTXOperand::pred, 0); + ins.a = reg("a", PTXOperand::u64, 1); + if (!ins.valid().empty()) return false; + PTXU64 base; + hydrazine::bit_cast(base, constKernel.ConstMemory); + constCta.setRegAsU64(0, 1, base + 15); + constCta.eval_Isspacep(constCta.getActiveContext(), ins); + if (!constCta.getRegAsPredicate(0, 0)) return false; + constCta.setRegAsU64(0, 1, base + 16); + constCta.eval_Isspacep(constCta.getActiveContext(), ins); + if (constCta.getRegAsPredicate(0, 0)) return false; + constCta.setRegAsU64(0, 1, base - 1); + constCta.eval_Isspacep(constCta.getActiveContext(), ins); + if (constCta.getRegAsPredicate(0, 0)) return false; + constCta.setRegAsPredicate(0, 2, false); + constCta.setRegAsPredicate(0, 0, true); + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 2; + constCta.setRegAsU64(0, 1, base + 16); + constCta.eval_Isspacep(constCta.getActiveContext(), ins); + if (!constCta.getRegAsPredicate(0, 0)) return false; + ins.addressSpace = PTXInstruction::Global; + ins.pg.condition = PTXOperand::PT; + const PTXU64 addresses[] = {base - 1, base, base + 15, base + 16}; + for (PTXOperand::DataType type : {PTXOperand::u64, PTXOperand::u32}) { + ins.a.type = type; + for (int i = 0; i < 4; ++i) { + if (type == PTXOperand::u32) + constCta.setRegAsU32(0, 1, static_cast(addresses[i])); + else constCta.setRegAsU64(0, 1, addresses[i]); + constCta.eval_Isspacep(constCta.getActiveContext(), ins); + if (constCta.getRegAsPredicate(0, 0) != (i == 0 || i == 3)) return false; + } + } + constCta.setRegAsPredicate(0, 0, true); + constCta.setRegAsPredicate(0, 2, false); + ins.a.type = PTXOperand::u64; + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 2; + constCta.setRegAsU64(0, 1, base); + constCta.eval_Isspacep(constCta.getActiveContext(), ins); + if (!constCta.getRegAsPredicate(0, 0)) return false; + constCta.setRegAsPredicate(0, 2, true); + ins.addressSpace = PTXInstruction::Const; + constCta.setRegAsU32(0, 1, static_cast(base)); + ins.a.type = PTXOperand::u32; + constCta.eval_Isspacep(constCta.getActiveContext(), ins); + return constCta.getRegAsPredicate(0, 0); + } + + bool test_IsspacepParam() { + PTXInstruction ins; + ins.opcode = PTXInstruction::Isspacep; + ins.addressSpace = PTXInstruction::Param; + ins.d = reg("p", PTXOperand::pred, 0); + ins.a = reg("a", PTXOperand::u64, 1); + if (!ins.valid().empty() || kernel->argumentMemorySize() == 0) + return false; + PTXU64 base = reinterpret_cast(kernel->ArgumentMemory); + PTXU64 size = kernel->argumentMemorySize(); + cta->reset(); + for (PTXU64 address : {base, base + size - 1, base - 1, base + size}) { + cta->setRegAsU64(0, 1, address); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (cta->getRegAsPredicate(0, 0) != + (address >= base && address - base < size)) return false; + } + ins.addressSpace = PTXInstruction::Global; + cta->setRegAsU64(0, 1, base); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + ins.addressSpace = PTXInstruction::Param; + ins.a.type = PTXOperand::u32; + const PTXU32 base32 = static_cast(base); + cta->setRegAsU32(0, 1, base32 + static_cast(size - 1)); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + ins.addressSpace = PTXInstruction::Global; + cta->setRegAsU32(0, 1, base32); + cta->eval_Isspacep(cta->getActiveContext(), ins); + if (!cta->getRegAsPredicate(0, 0)) return false; + ins.addressSpace = PTXInstruction::Param; + cta->setRegAsPredicate(0, 0, true); + cta->setRegAsPredicate(0, 2, false); + ins.pg.condition = PTXOperand::Pred; + ins.pg.reg = 2; + cta->setRegAsU32(0, 1, base32 + static_cast(size)); + cta->eval_Isspacep(cta->getActiveContext(), ins); + return cta->getRegAsPredicate(0, 0); } ///////////////////////////////////////////////////////////////////////////////////////////////// @@ -4443,14 +10189,32 @@ class TestInstructions: public Test { result = (result && test_Mov()); // cvt instruction - + result = (result && test_Cvt()); + result = (result && test_Cvta()); + result = (result && test_Isspacep()); + result = (result && test_IsspacepConst()); + result = (result && test_IsspacepParam()); + result = (result && test_LdLu()); + result = (result && test_SuldSustCacheOperator()); + result = (result && test_StCacheOperator()); + result = (result && test_GrammarSharedActionFix()); + result = (result && test_PrefetchOpcode()); + result = (result && test_Red()); + result = (result && test_Fence()); + result = (result && test_AtomicRedSemantics()); + result = (result && test_LdStOrdering()); + result = (result && test_LdStMmio()); + result = (result && test_AtomicRedTypeTable()); + // arithmetic instructions result = (result && test_Abs()); result = (result && test_Add()); result = (result && test_Sub()); + result = (result && test_AddSubRounding()); result = (result && test_Div()); result = (result && test_Neg()); result = (result && test_Rem()); + result = (result && test_Sad()); result = (result && test_Min()); result = (result && test_Max()); if (prolix && result) { @@ -4460,8 +10224,10 @@ class TestInstructions: public Test { // difficult arithmetic instructions result = (result && test_Mad()); result = (result && test_Mul()); + result = (result && test_MulRounding()); result = (result && test_AddC()); result = (result && test_SubC()); + result = (result && test_Dp()); if (prolix && result) { status << "pass: exotic arithmetic instructions\n"; } @@ -4469,20 +10235,43 @@ class TestInstructions: public Test { // floating-point instructions result = (result && test_Cos()); result = (result && test_Sin()); + result = (result && test_Tanh()); + result = (result && test_CopySign()); result = (result && test_Ex2()); + result = (result && test_Fma()); + result = (result && test_F16Fma()); + result = (result && test_Bf16Fma()); + result = (result && test_MmaWarpParticipation()); + result = (result && test_MmaFloatShapes()); + result = (result && test_Mma()); + result = (result && test_MmaInt8()); + result = (result && test_MmaInt8WideShapes()); result = (result && test_Lg2()); result = (result && test_Sqrt()); result = (result && test_Rsqrt()); result = (result && test_Rcp()); + result = (result && test_TestP()); if (prolix && result) { status << "pass: floating-point instructions\n"; } // logical and shift instructions + result = (result && test_Fns()); + result = (result && test_Szext()); + result = (result && test_Bfe()); + result = (result && test_Bfi()); + result = (result && test_Bfind()); + result = (result && test_Popc()); + result = (result && test_Bmsk()); + result = (result && test_Lop3()); + result = (result && test_Lop3RegisterAllocation()); result = (result && test_And()); result = (result && test_Or()); result = (result && test_Xor()); result = (result && test_Not()); + result = (result && test_Shf()); + result = (result && test_Shl()); + result = (result && test_Shr()); if (prolix && result) { status << "pass: logical instructions\n"; } @@ -4540,4 +10329,3 @@ int main(int argc, char **argv) { return test.passed(); } - diff --git a/ocelot/src/ir/PTXInstruction.cpp b/ocelot/src/ir/PTXInstruction.cpp index 099a0493f..c6204db75 100644 --- a/ocelot/src/ir/PTXInstruction.cpp +++ b/ocelot/src/ir/PTXInstruction.cpp @@ -19,6 +19,18 @@ std::string ir::PTXInstruction::toString( Level l ) { return ""; } +std::string ir::PTXInstruction::toString( Semantics s ) { + switch( s ) { + case Sc: return "sc"; break; + case AcqRel: return "acq_rel"; break; + case Acquire: return "acquire"; break; + case Release: return "release"; break; + case Relaxed: return "relaxed"; break; + default: break; + } + return ""; +} + std::string ir::PTXInstruction::toString(CacheLevel cache) { switch( cache ) { case L1: return "L1"; @@ -35,6 +47,7 @@ std::string ir::PTXInstruction::toStringLoad(CacheOperation operation) { case Cs: return "cs"; case Cv: return "cv"; case Nc: return "nc"; + case Lu: return "lu"; default: break; } return ""; @@ -201,6 +214,9 @@ std::string ir::PTXInstruction::modifierString( unsigned int modifier, else if( modifier & rn ) { result += "rn."; } + else if( modifier & rna ) { + result += "rna."; + } else if( modifier & rz ) { result += "rz."; } @@ -214,9 +230,21 @@ std::string ir::PTXInstruction::modifierString( unsigned int modifier, if( modifier & ftz ) { result += "ftz."; } + if( modifier & nan ) { + result += "NaN."; + } + if( modifier & xorsign ) { + result += "xorsign."; + } + if( modifier & abs ) { + result += "abs."; + } if( modifier & sat ) { result += "sat."; } + if( modifier & relu ) { + result += "relu."; + } if( carry == CC ) { result += "cc."; } @@ -230,11 +258,16 @@ std::string ir::PTXInstruction::toString( Modifier modifier ) { case wide: return "wide"; break; case sat: return "sat"; break; case rn: return "rn"; break; + case rna: return "rna"; break; case rz: return "rz"; break; case rm: return "rm"; break; case rp: return "rp"; break; case approx: return "approx"; break; case ftz: return "ftz"; break; + case nan: return "NaN"; break; + case xorsign:return "xorsign";break; + case abs: return "abs"; break; + case relu: return "relu"; break; default: break; } return ""; @@ -362,6 +395,7 @@ std::string ir::PTXInstruction::toString( Opcode opcode ) { case Bfe: return "bfe"; break; case Bfi: return "bfi"; break; case Bfind: return "bfind"; break; + case Bmsk: return "bmsk"; break; case Bra: return "bra"; break; case Brev: return "brev"; break; case Brkpt: return "brkpt"; break; @@ -373,18 +407,24 @@ std::string ir::PTXInstruction::toString( Opcode opcode ) { case Cvt: return "cvt"; break; case Cvta: return "cvta"; break; case Div: return "div"; break; + case Dp2a: return "dp2a"; break; + case Dp4a: return "dp4a"; break; case Ex2: return "ex2"; break; case Exit: return "exit"; break; case Fma: return "fma"; break; + case Fns: return "fns"; break; case Isspacep: return "isspacep"; break; case Ld: return "ld"; break; case Ldu: return "ldu"; break; case Lg2: return "lg2"; break; + case Lop3: return "lop3"; break; case Mad24: return "mad24"; break; case Mad: return "mad"; break; + case Mma: return "mma"; break; case MadC: return "madc"; break; case Max: return "max"; break; case Membar: return "membar"; break; + case Fence: return "fence"; break; case Min: return "min"; break; case Mov: return "mov"; break; case Mul24: return "mul24"; break; @@ -420,6 +460,8 @@ std::string ir::PTXInstruction::toString( Opcode opcode ) { case Sured: return "sured"; break; case Sust: return "sust"; break; case Suq: return "suq"; break; + case Szext: return "szext"; break; + case Tanh: return "tanh"; break; case TestP: return "testp"; break; case Tex: return "tex"; break; case Tld4: return "tld4"; break; @@ -454,10 +496,14 @@ ir::PTXInstruction::PTXInstruction( Opcode op, const PTXOperand& _d, : opcode(op), d(_d), a(_a), b(_b), c(_c) { ISA = Instruction::PTX; type = PTXOperand::s32; + bType = PTXOperand::TypeSpecifier_invalid; modifier = 0; reconvergeInstruction = 0; branchTargetInstruction = 0; vec = PTXOperand::v1; + mmaShape = MmaShape_Invalid; + mmaAColumnMajor = false; + mmaBColumnMajor = true; pg.condition = PTXOperand::PT; pg.type = PTXOperand::pred; barrierOperation = BarSync; @@ -467,6 +513,11 @@ ir::PTXInstruction::PTXInstruction( Opcode op, const PTXOperand& _d, cc = 0; addressSpace = AddressSpace_Invalid; tailCall = false; + mmio = false; + cacheOperation = Ca; + booleanOperator = BoolAnd; + semantics = AcqRel; + scope = Level_Invalid; } ir::PTXInstruction::~PTXInstruction() { @@ -478,11 +529,29 @@ bool ir::PTXInstruction::operator==( const PTXInstruction& i ) const { } std::string ir::PTXInstruction::valid() const { + if (opcode != Mma && (d.vec == PTXOperand::v8 || a.vec == PTXOperand::v8 || + b.vec == PTXOperand::v8 || c.vec == PTXOperand::v8)) return "eight-register fragments require mma"; + if( opcode == Min || opcode == Max ) { + const unsigned int fp32Modifiers = nan | xorsign | abs; + if( (modifier & fp32Modifiers) && type != PTXOperand::f32 + && type != PTXOperand::f16 && type != PTXOperand::f16x2 + && type != PTXOperand::bf16 && type != PTXOperand::bf16x2 ) { + return "NaN and xorsign.abs modifiers require a floating type"; + } + if ((modifier & ftz) && (type == PTXOperand::bf16 || type == PTXOperand::bf16x2)) + return "ftz is invalid for bf16 min/max"; + if( (modifier & (xorsign | abs)) != 0 + && (modifier & (xorsign | abs)) != (xorsign | abs) ) { + return "xorsign and abs modifiers must be specified together"; + } + } switch (opcode) { case Abs: { if ( !( type == PTXOperand::s16 || type == PTXOperand::s32 || type == PTXOperand::s64 || type == PTXOperand::f32 || - type == PTXOperand::f64 ) ) { + type == PTXOperand::f64 || type == PTXOperand::f16 || + type == PTXOperand::bf16 || type == PTXOperand::f16x2 || + type == PTXOperand::bf16x2 ) ) { return "invalid instruction type " + PTXOperand::toString( type ); } @@ -504,11 +573,15 @@ std::string ir::PTXInstruction::valid() const { } case Add: { if ( !( type != PTXOperand::s8 && type != PTXOperand::u8 && - type != PTXOperand::b8 && type != PTXOperand::f16 + type != PTXOperand::b8 && type != PTXOperand::pred ) ) { return "invalid instruction type " + PTXOperand::toString( type ); } + if ((type == PTXOperand::f16 || type == PTXOperand::f16x2) + && (modifier & ~(rn | ftz | sat))) { + return "add.f16 only supports .rn, .ftz, and .sat"; + } if( carry == CC ) { if( ( modifier & sat ) ) { return "saturate not supported with carry out"; @@ -578,34 +651,33 @@ std::string ir::PTXInstruction::valid() const { break; } case Atom: { - if( !PTXOperand::valid( PTXOperand::b32, type ) - && !PTXOperand::valid( PTXOperand::b64, type ) - && ( atomicOperation == AtomicAnd || atomicOperation == AtomicOr - || atomicOperation == AtomicXor || atomicOperation == AtomicCas - || atomicOperation == AtomicExch ) ) { - return "invalid instruction type " - + PTXOperand::toString( type ) + " for atomic " + if( ( atomicOperation == AtomicAnd || atomicOperation == AtomicOr + || atomicOperation == AtomicXor || atomicOperation == AtomicCas + || atomicOperation == AtomicExch ) + && type != PTXOperand::b32 && type != PTXOperand::b64 ) { + return "invalid instruction type " + + PTXOperand::toString( type ) + " for atomic " + toString( atomicOperation ); } - - if( !PTXOperand::valid( PTXOperand::u32, type ) - && !PTXOperand::valid( PTXOperand::u64, type ) - && !PTXOperand::valid( PTXOperand::s32, type ) - && ( atomicOperation == AtomicInc - || atomicOperation == AtomicDec ) ) { - return "invalid instruction type " - + PTXOperand::toString( type ) + " for atomic " + if( ( atomicOperation == AtomicInc || atomicOperation == AtomicDec ) + && type != PTXOperand::u32 ) { + return "invalid instruction type " + + PTXOperand::toString( type ) + " for atomic " + toString( atomicOperation ); } - if( !PTXOperand::valid( PTXOperand::f32, type ) - && !PTXOperand::valid( PTXOperand::u32, type ) - && !PTXOperand::valid( PTXOperand::u64, type ) - && !PTXOperand::valid( PTXOperand::s32, type ) - && ( atomicOperation == AtomicAdd - || atomicOperation == AtomicMin - || atomicOperation == AtomicMax ) ) { - return "invalid instruction type " - + PTXOperand::toString( type ) + " for atomic " + if( atomicOperation == AtomicAdd + && type != PTXOperand::u32 && type != PTXOperand::s32 + && type != PTXOperand::u64 && type != PTXOperand::f32 + && type != PTXOperand::f64 ) { + return "invalid instruction type " + + PTXOperand::toString( type ) + " for atomic " + + toString( atomicOperation ); + } + if( ( atomicOperation == AtomicMin || atomicOperation == AtomicMax ) + && type != PTXOperand::u32 && type != PTXOperand::s32 + && type != PTXOperand::u64 && type != PTXOperand::s64 ) { + return "invalid instruction type " + + PTXOperand::toString( type ) + " for atomic " + toString( atomicOperation ); } if( !( addressSpace == Shared || addressSpace == Global ) ) { @@ -637,12 +709,12 @@ std::string ir::PTXInstruction::valid() const { + " cannot be assigned to " + PTXOperand::toString( type ); } if( !PTXOperand::valid( PTXOperand::u32, b.type ) - && a.addressMode != PTXOperand::Immediate ) { + && b.addressMode != PTXOperand::Immediate ) { return "operand 1 type " + PTXOperand::toString( b.type ) + " cannot be assigned to " + PTXOperand::toString( PTXOperand::u32 ); } if( !PTXOperand::valid( PTXOperand::u32, c.type ) - && a.addressMode != PTXOperand::Immediate ) { + && c.addressMode != PTXOperand::Immediate ) { return "operand 1 type " + PTXOperand::toString( c.type ) + " cannot be assigned to " + PTXOperand::toString( PTXOperand::u32 ); } @@ -672,7 +744,7 @@ std::string ir::PTXInstruction::valid() const { + " cannot be assigned to " + PTXOperand::toString( PTXOperand::u32 ); } - if( !PTXOperand::valid( PTXOperand::u32, b.type ) ) { + if( !PTXOperand::valid( PTXOperand::u32, c.type ) ) { return "operand 4 type " + PTXOperand::toString( c.type ) + " cannot be assigned to " + PTXOperand::toString( PTXOperand::u32 ); @@ -696,6 +768,46 @@ std::string ir::PTXInstruction::valid() const { } break; } + case Bmsk: { + if( type != PTXOperand::b32 ) { + return "bmsk instruction requires .b32 type"; + } + if( shiftMode != ShiftMode::Clamp && shiftMode != ShiftMode::Wrap ) { + return "bmsk instruction requires .clamp or .wrap mode"; + } + if( !PTXOperand::valid( type, d.type ) + || !PTXOperand::valid( type, a.type ) + || !PTXOperand::valid( type, b.type ) ) { + return "bmsk operands must have 32-bit integer types"; + } + break; + } + case Lop3: { + if( type != PTXOperand::b32 ) { + return "lop3 instruction requires .b32 type"; + } + const bool predicateResult = pq.addressMode != PTXOperand::Invalid; + if( (d.addressMode == PTXOperand::BitBucket && !predicateResult) + || (d.addressMode != PTXOperand::BitBucket + && !PTXOperand::valid( type, d.type )) + || !PTXOperand::valid( type, a.type ) + || !PTXOperand::valid( type, b.type ) ) { + return "lop3 data operands must have 32-bit integer types"; + } + if( !PTXOperand::valid( type, c.type ) ) { + return "lop3 data operands must have 32-bit integer types"; + } + if( immLut.addressMode != PTXOperand::Immediate + || immLut.imm_uint > 0xff ) { + return "lop3 immLut must be an integer constant from 0 to 255"; + } + if( predicateResult && (pq.type != PTXOperand::pred + || q.type != PTXOperand::pred + || (booleanOperator != BoolAnd && booleanOperator != BoolOr)) ) { + return "lop3 predicate form requires .and or .or and predicates p, q"; + } + break; + } case Bra: { if( !( d.addressMode == PTXOperand::Label || d.addressMode == PTXOperand::Register ) ) { @@ -825,6 +937,9 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString( type ); } + if( modifier != approx && modifier != ( approx | ftz ) ) { + return "cos requires .approx with optional .ftz"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -833,27 +948,151 @@ std::string ir::PTXInstruction::valid() const { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } - if( modifier & ftz ) { - if( PTXOperand::isInt( type ) ) { - return toString( ftz ) - + " only valid for float point instructions."; - } - } break; } case Cvt: { - if( type == PTXOperand::pred ) { + PTXOperand::DataType sourceType = a.relaxedType == + PTXOperand::TypeSpecifier_invalid ? a.type : a.relaxedType; + const int cvtModifiers = rni | rzi | rmi | rpi | rn | rna | rz + | rm | rp | ftz | sat | relu; + if( modifier & ~cvtModifiers ) return "invalid cvt modifier"; + const bool packed = type == PTXOperand::f16x2 + || type == PTXOperand::bf16x2; + if (!packed && b.addressMode != PTXOperand::Invalid) { + return "scalar cvt accepts only two operands"; + } + if( packed ) { + const PTXOperand::DataType bSourceType = b.relaxedType == + PTXOperand::TypeSpecifier_invalid ? b.type : b.relaxedType; + if( sourceType != PTXOperand::f32 + || bSourceType != PTXOperand::f32 ) { + return "packed cvt requires two f32 source operands"; + } + const int rounding = modifier & (rn | rz); + if( rounding != rn && rounding != rz ) { + return "packed cvt requires .rn or .rz"; + } + if( modifier & ~(rn | rz | relu) ) { + return "invalid packed cvt modifier"; + } + if( type == PTXOperand::f16x2 + ? !PTXOperand::relaxedValid(type, d.type) + : d.type != PTXOperand::b32 ) { + return "packed cvt requires a b32 destination"; + } + break; + } + if( type == PTXOperand::tf32 ) { + if( sourceType != PTXOperand::f32 || modifier != rna ) { + return "sm_86 tf32 cvt requires cvt.rna.tf32.f32"; + } + if( d.type != PTXOperand::b32 ) { + return "tf32 cvt requires a b32 destination"; + } + break; + } + if( type == PTXOperand::bf16 ) { + const int rounding = modifier & (rn | rz); + if( sourceType != PTXOperand::f32 + || (modifier & ~(rn | rz | relu)) + || (rounding != rn && rounding != rz) ) { + return "sm_86 bf16 cvt requires cvt.{rn,rz}{.relu}.bf16.f32"; + } + if( d.type != PTXOperand::b16 ) { + return "bf16 cvt requires a b16 destination"; + } + break; + } + if( sourceType == PTXOperand::bf16 ) { + if( type != PTXOperand::f32 || modifier ) { + return "sm_86 supports only cvt.f32.bf16"; + } + if( !PTXOperand::relaxedValid(type, d.type) ) { + return "cvt.f32.bf16 requires an f32-compatible destination"; + } + break; + } + if( sourceType == PTXOperand::f16x2 + || sourceType == PTXOperand::bf16x2 + || sourceType == PTXOperand::tf32 ) { + return "cvt source type is not supported on sm_86"; + } + if( !PTXOperand::isInt(type) && !PTXOperand::isFloat(type) + ) { return "invalid instruction type " + PTXOperand::toString( type ); - } + } + if( !PTXOperand::isInt(sourceType) + && !PTXOperand::isFloat(sourceType) ) { + return "invalid source type " + PTXOperand::toString(sourceType); + } + const bool destinationFloat = PTXOperand::isFloat(type); + const bool sourceFloat = PTXOperand::isFloat(sourceType); + const bool destinationInt = PTXOperand::isInt(type); + const bool sourceInt = PTXOperand::isInt(sourceType); + const int integerRounding = modifier & (rni | rzi | rmi | rpi); + const int floatRounding = modifier & (rn | rz | rm | rp); + const int rounding = integerRounding | floatRounding; + if( rounding && (rounding & (rounding - 1)) ) { + return "cvt accepts only one rounding modifier"; + } + const bool narrowerFloat = sourceFloat && destinationFloat + && PTXOperand::bytes(type) < PTXOperand::bytes(sourceType); + const bool sameSizeFloat = sourceFloat && destinationFloat + && PTXOperand::bytes(type) == PTXOperand::bytes(sourceType); + if( integerRounding + && !(sourceFloat && (destinationInt || sameSizeFloat)) ) { + return "integer rounding is invalid for this cvt conversion"; + } + if( floatRounding + && !(destinationFloat && (sourceInt || narrowerFloat)) ) { + return "floating-point rounding is invalid for this cvt conversion"; + } + if( sourceFloat && destinationInt && !integerRounding ) { + return "float-to-integer cvt requires integer rounding"; + } + if( sourceInt && destinationFloat && !floatRounding ) { + return "integer-to-float cvt requires floating-point rounding"; + } + if( narrowerFloat && !floatRounding ) { + return "narrowing float cvt requires floating-point rounding"; + } + if( modifier & rna ) { + return ".rna is valid only for cvt.rna.tf32.f32 on sm_86"; + } + if( modifier & relu ) { + const int rounding = modifier & (rn | rz); + if( type != PTXOperand::f16 || sourceType != PTXOperand::f32 + || (rounding != rn && rounding != rz) + || (modifier & ~(rn | rz | relu)) ) { + return ".relu requires cvt.{rn,rz}.f16.f32"; + } + } + if( modifier & sat ) { + if( !destinationInt && type != PTXOperand::f16 + && type != PTXOperand::f32 && type != PTXOperand::f64 ) { + return ".sat is invalid for this cvt destination type"; + } + if( destinationInt && sourceInt ) { + const unsigned int destinationBits = PTXOperand::bytes(type) * 8; + const unsigned int sourceBits = PTXOperand::bytes(sourceType) * 8; + const bool destinationSigned = PTXOperand::isSigned(type); + const bool sourceSigned = PTXOperand::isSigned(sourceType); + const bool destinationSuperset = destinationSigned == sourceSigned + ? destinationBits >= sourceBits + : destinationSigned && destinationBits > sourceBits; + if( destinationSuperset ) { + return ".sat is invalid when the destination range contains the source range"; + } + } + } if( d.bytes() < PTXOperand::bytes( type ) ) { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned from " + PTXOperand::toString( type ); } if( modifier & ftz ) { - if( !(PTXOperand::isFloat( type ) || PTXOperand::isFloat(a.type))) { - return toString( ftz ) - + " only valid for float point instructions."; + if( type != PTXOperand::f32 && sourceType != PTXOperand::f32 ) { + return ".ftz requires an f32 source or destination"; } } if( vec == PTXOperand::v1 @@ -870,7 +1109,8 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString(type); } if (!(addressSpace == Global || addressSpace == Local - || addressSpace == Shared || addressSpace == Const)) { + || addressSpace == Shared || addressSpace == Const + || addressSpace == Param)) { return "invalid address space " + toString(addressSpace); } break; @@ -879,20 +1119,30 @@ std::string ir::PTXInstruction::valid() const { if( ( modifier & sat ) ) { return "no support for saturating divide."; } - if( ( modifier & approx ) ) { - if( type != PTXOperand::f32 ) { - return "only f32 supported for approximate"; + if( modifier & ~(rn | rz | rm | rp | approx | full | ftz) ) { + return "invalid divide modifier"; + } + const int rounding = modifier & (rn | rz | rm | rp); + const bool oneRounding = rounding == rn || rounding == rz + || rounding == rm || rounding == rp; + if( type == PTXOperand::f32 ) { + const int modes = ((modifier & approx) ? 1 : 0) + + ((modifier & full) ? 1 : 0) + (rounding ? 1 : 0); + if( modes != 1 || (rounding && !oneRounding) ) { + return "div.f32 requires exactly one of approx, full, or rounding"; } } if( type == PTXOperand::f64 ) { - if( !( modifier & rn ) && !( modifier & rz ) - && !( modifier & rm ) && !( modifier & rp ) ) { - return "requires a rounding modifier"; + if( modifier & (approx | full | ftz) ) { + return "approx, full, and ftz are invalid for div.f64"; } - if( !( modifier & rn ) ) { - return "only nearest rounding supported"; + if( !oneRounding ) { + return "div.f64 requires exactly one rounding modifier"; } } + if( PTXOperand::isInt( type ) && modifier ) { + return "integer divide does not accept modifiers"; + } if( !( type == PTXOperand::u16 || type == PTXOperand::u32 || type == PTXOperand::u64 || type == PTXOperand::s16 || type == PTXOperand::s32 || type == PTXOperand::s64 @@ -920,11 +1170,34 @@ std::string ir::PTXInstruction::valid() const { } break; } + case Dp2a: + case Dp4a: { + if( (type != PTXOperand::u32 && type != PTXOperand::s32) + || (bType != PTXOperand::u32 && bType != PTXOperand::s32) ) { + return "dp2a/dp4a type qualifiers must be .u32 or .s32"; + } + if( opcode == Dp2a && modifier != lo && modifier != hi ) { + return "dp2a instruction requires .lo or .hi mode"; + } + const PTXOperand::DataType resultType = type == PTXOperand::u32 + && bType == PTXOperand::u32 ? PTXOperand::u32 : PTXOperand::s32; + if( !PTXOperand::valid(type, a.type) + || !PTXOperand::valid(bType, b.type) + || !PTXOperand::valid(resultType, c.type) + || !PTXOperand::valid(resultType, d.type) ) { + return "dp2a/dp4a operands must have compatible 32-bit integer types"; + } + break; + } case Ex2: { - if( !( type == PTXOperand::f32 ) ) { + if( !( type == PTXOperand::f32 || type == PTXOperand::f16 ) ) { return "invalid instruction type " + PTXOperand::toString( type ); } + if( type == PTXOperand::f32 && modifier != approx + && modifier != ( approx | ftz ) ) { + return "ex2.f32 requires .approx with optional .ftz"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -933,35 +1206,90 @@ std::string ir::PTXInstruction::valid() const { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } - if( modifier & ftz ) { - if( PTXOperand::isInt( type ) ) { - return toString( ftz ) - + " only valid for float point instructions."; - } - } break; } case Exit: { break; } case Fma: { - if (!(type == ir::PTXOperand::f32 || type == ir::PTXOperand::f64)) { + if (!(type == ir::PTXOperand::f16 || type == ir::PTXOperand::f16x2 + || type == ir::PTXOperand::f32 + || type == ir::PTXOperand::f64 || type == ir::PTXOperand::bf16 + || type == ir::PTXOperand::bf16x2)) { return "invalid instruction type " + PTXOperand::toString( type ); } + const int rounding = modifier & (rn | rz | rm | rp); + if (rounding != rn && rounding != rz + && rounding != rm && rounding != rp) { + return "fma requires exactly one rounding modifier"; + } + if (type == PTXOperand::f64 && (modifier & (ftz | sat))) { + return "ftz and sat modifiers are invalid for fma.f64"; + } + if ((type == PTXOperand::f16 || type == PTXOperand::f16x2) + && rounding != rn) { + return "fma.f16 requires .rn"; + } + if ((type == PTXOperand::f16 || type == PTXOperand::f16x2) + && (modifier & sat) && (modifier & relu)) { + return "fma.f16 sat and relu are mutually exclusive"; + } + if ((type == PTXOperand::bf16 || type == PTXOperand::bf16x2) + && (rounding != rn || (modifier & (ftz | sat)))) { + return "fma.bf16 requires .rn without ftz or sat"; + } + if( !PTXOperand::valid( type, a.type ) ) { + return "operand A type " + PTXOperand::toString( a.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + if( !PTXOperand::valid( type, b.type ) ) { + return "operand B type " + PTXOperand::toString( b.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + if( !PTXOperand::valid( type, c.type ) ) { + return "operand C type " + PTXOperand::toString( c.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } if( !PTXOperand::valid( type, d.type ) ) { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } break; } + case Fns: { + if( type != PTXOperand::b32 ) { + return "fns instruction only supports .b32 type"; + } + if( d.type != PTXOperand::b32 ) { + return "fns destination operand must have .b32 type"; + } + if( a.bytes() != 4 ) { + return "fns mask operand must be 32-bit"; + } + if( b.type != PTXOperand::b32 && b.type != PTXOperand::u32 + && b.type != PTXOperand::s32 ) { + return "fns base operand must have .b32, .u32, or .s32 type"; + } + if( c.type != PTXOperand::s32 ) { + return "fns offset operand must have .s32 type"; + } + if( b.addressMode == PTXOperand::Immediate && b.imm_uint > 31 ) { + return "fns base operand must be between 0 and 31"; + } + break; + } case Isspacep: { - if (!(addressSpace == PTXInstruction::Global + if (!(addressSpace == PTXInstruction::Const + || addressSpace == PTXInstruction::Global + || addressSpace == PTXInstruction::Param || addressSpace == PTXInstruction::Shared || addressSpace == PTXInstruction::Local)) { return "invalid address space " + toString(addressSpace); } if (!(d.addressMode == PTXOperand::Register - && a.addressMode == PTXOperand::Register)) { + && d.type == PTXOperand::pred && d.vec == PTXOperand::v1 + && a.addressMode == PTXOperand::Register && a.vec == PTXOperand::v1 + && (a.type == PTXOperand::u32 || a.type == PTXOperand::u64))) { return "invalid address mode for operands"; } break; @@ -981,9 +1309,12 @@ std::string ir::PTXInstruction::valid() const { } if( addressSpace != Global && addressSpace != Shared && volatility == Volatile && addressSpace != Generic ) { - return "only shared and global address spaces supported " + return "only shared and global address spaces supported " "for volatile loads"; } + if( cacheOperation == Lu && volatility == Volatile ) { + return "ld.lu cannot be combined with .volatile"; + } if( d.addressMode != PTXOperand::Register ) { return "operand D must be a register not a " + PTXOperand::toString( d.addressMode ); @@ -1004,6 +1335,9 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString( type ); } + if( modifier != approx && modifier != ( approx | ftz ) ) { + return "lg2 requires .approx with optional .ftz"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1012,12 +1346,6 @@ std::string ir::PTXInstruction::valid() const { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } - if( modifier & ftz ) { - if( PTXOperand::isInt( type ) ) { - return toString( ftz ) - + " only valid for float point instructions."; - } - } break; } case Mad24: { @@ -1044,6 +1372,33 @@ std::string ir::PTXInstruction::valid() const { break; } case Mad: { + if( type == PTXOperand::f32 || type == PTXOperand::f64 ) { + const int rounding = modifier & (rn | rz | rm | rp); + if( rounding != rn && rounding != rz + && rounding != rm && rounding != rp ) { + return "mad requires exactly one rounding modifier"; + } + if( type == PTXOperand::f64 && (modifier & (ftz | sat)) ) { + return "ftz and sat modifiers are invalid for mad.f64"; + } + if( !PTXOperand::valid( type, a.type ) ) { + return "operand A type " + PTXOperand::toString( a.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + if( !PTXOperand::valid( type, b.type ) ) { + return "operand B type " + PTXOperand::toString( b.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + if( !PTXOperand::valid( type, c.type ) ) { + return "operand C type " + PTXOperand::toString( c.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + if( !PTXOperand::valid( type, d.type ) ) { + return "operand D type " + PTXOperand::toString( d.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + break; + } if( !( type != PTXOperand::s8 && type != PTXOperand::u8 && type != PTXOperand::b8 && type != PTXOperand::f16 && type != PTXOperand::pred ) ) { @@ -1094,6 +1449,246 @@ std::string ir::PTXInstruction::valid() const { } break; } + case Mma: { + for (const PTXOperand* operand : {&a, &b, &c, &d}) { + if (operand->addressMode != PTXOperand::Register) { + return "mma fragments must be register operands"; + } + for (const auto& element : operand->array) { + if (element.addressMode != PTXOperand::Register) { + return "mma fragment elements must be registers"; + } + if (element.vec != PTXOperand::v1 || !element.array.empty()) { + return "mma fragment elements must be scalar registers"; + } + } + } + if ((mmaAColumnMajor || !mmaBColumnMajor) && + (mmaShape != MmaM8N8K4 || a.type != PTXOperand::f16)) { + return "this mma form requires row.col layouts"; + } + if (a.type == PTXOperand::b1) { + const bool small = mmaShape == MmaM8N8K128; + const bool k256 = mmaShape == MmaM16N8K256; + if ((!small && !k256 && mmaShape != MmaM16N8K128) || + type != PTXOperand::s32 || b.type != a.type || c.type != type || d.type != type || modifier || + (booleanOperator != BoolAnd && booleanOperator != BoolXor)) return "invalid binary mma shape/type/modifier"; + const unsigned aCount = small ? 1u : k256 ? 4u : 2u; + const unsigned bCount = k256 ? 2u : 1u; + const unsigned cdCount = small ? 2u : 4u; + for (const PTXOperand* operand : {&a, &b, &c, &d}) { + const bool input = operand == &a || operand == &b; + const unsigned count = operand == &a ? aCount : operand == &b ? bCount : cdCount; + if (operand->array.size() != count || operand->vec != static_cast(count)) return "invalid binary mma fragment size"; + for (const auto& element : operand->array) + if (input ? element.type != PTXOperand::b32 : !PTXOperand::relaxedValid(type, element.type)) return "invalid binary mma register type"; + } + break; + } + if (a.type == PTXOperand::f64) { + if (mmaShape != MmaM8N8K4 || type != PTXOperand::f64 || + b.type != type || c.type != type || d.type != type || + (modifier & ~(rn | rz | rm | rp))) return "invalid f64 mma shape/type/modifier"; + if (modifier && (modifier & (modifier - 1))) { + return "f64 mma accepts only one rounding modifier"; + } + for (const PTXOperand* operand : {&a, &b, &c, &d}) { + const int count = (operand == &a || operand == &b) ? 1 : 2; + if (operand->array.size() != count || operand->vec != (count == 1 ? PTXOperand::v1 : PTXOperand::v2)) + return "invalid f64 mma fragment size"; + for (const auto& element : operand->array) + if (!PTXOperand::relaxedValid(type, element.type)) return "invalid f64 mma register type"; + } + break; + } + if (mmaShape == MmaM8N8K4) { + if ((type != PTXOperand::f16 && type != PTXOperand::f32) || a.type != PTXOperand::f16 || + b.type != a.type || (c.type != PTXOperand::f16 && c.type != PTXOperand::f32) || + d.type != type || (c.type == PTXOperand::f32 && type != PTXOperand::f32) || modifier) + return "unsupported m8n8k4 type/modifier combination"; + for (const PTXOperand* operand : {&a, &b, &c, &d}) { + const bool input = operand == &a || operand == &b; + const unsigned count = input ? 2 : operand->type == PTXOperand::f32 ? 8 : 4; + if (operand->vec != static_cast(count) || operand->array.size() != count) + return "invalid m8n8k4 f16 fragment size"; + for (const auto& element : operand->array) + if (input || operand->type == PTXOperand::f16 ? + (element.type != PTXOperand::b32 && element.type != PTXOperand::f16x2) : + !PTXOperand::relaxedValid(PTXOperand::f32, element.type)) return "invalid mma fragment register type"; + } + break; + } + if (a.type == PTXOperand::s4 || a.type == PTXOperand::u4) { + const bool small = mmaShape == MmaM8N8K32; + const bool k64 = mmaShape == MmaM16N8K64; + if ((!small && !k64 && mmaShape != MmaM16N8K32) || (b.type != PTXOperand::s4 && b.type != PTXOperand::u4) || + type != PTXOperand::s32 || c.type != type || d.type != type || (modifier & ~satfinite)) + return "unsupported sub-byte mma shape/type/modifier combination"; + const unsigned aCount = small ? 1u : k64 ? 4u : 2u; + const unsigned bCount = k64 ? 2u : 1u; + const unsigned cdCount = small ? 2u : 4u; + for (const PTXOperand* operand : {&a, &b, &c, &d}) { + const bool input = operand == &a || operand == &b; + const unsigned count = operand == &a ? aCount : operand == &b ? bCount : cdCount; + if (operand->vec != static_cast(count) || operand->array.size() != count) + return "invalid sub-byte mma fragment size"; + for (const auto& element : operand->array) + if (input ? element.type != PTXOperand::b32 : !PTXOperand::relaxedValid(type, element.type)) return "invalid sub-byte mma register type"; + } + break; + } + const bool intInput = a.type == PTXOperand::s8 || + a.type == PTXOperand::u8; + if (intInput) { + if (modifier & ~satfinite) { + return "integer mma accepts only satfinite"; + } + if (mmaShape != MmaM16N8K16 && mmaShape != MmaM8N8K16 && + mmaShape != MmaM16N8K32) { + return "integer mma requires m8n8k16, m16n8k16, or m16n8k32"; + } + if (type != PTXOperand::s32) { + return "integer mma requires s32 accumulators"; + } + if (b.type != PTXOperand::s8 && b.type != PTXOperand::u8) { + return "integer mma B type must be s8 or u8"; + } + if (d.type != PTXOperand::s32 || c.type != PTXOperand::s32) { + return "integer mma C and D types must be s32"; + } + const PTXOperand::Vec aVec = mmaShape == MmaM16N8K32 + ? PTXOperand::v4 : mmaShape == MmaM8N8K16 + ? PTXOperand::v1 : PTXOperand::v2; + const PTXOperand::Vec bVec = mmaShape == MmaM8N8K16 + ? PTXOperand::v1 : mmaShape == MmaM16N8K32 + ? PTXOperand::v2 : PTXOperand::v1; + const PTXOperand::Vec cdVec = mmaShape == MmaM8N8K16 + ? PTXOperand::v2 : PTXOperand::v4; + const unsigned int aCount = mmaShape == MmaM16N8K32 + ? 4u : mmaShape == MmaM8N8K16 ? 1u : 2u; + const unsigned int bCount = mmaShape == MmaM8N8K16 + ? 1u : mmaShape == MmaM16N8K32 ? 2u : 1u; + const unsigned int cdCount = mmaShape == MmaM8N8K16 ? 2u : 4u; + if (d.vec != cdVec || c.vec != cdVec || + a.vec != aVec || b.vec != bVec) { + return "integer mma has invalid fragment sizes"; + } + if (d.array.size() != cdCount || c.array.size() != cdCount || + a.array.size() != aCount || b.array.size() != bCount) { + return "integer mma has invalid fragment register counts"; + } + for (PTXOperand::Array::const_iterator element = a.array.begin(); + element != a.array.end(); ++element) { + if (element->type != PTXOperand::b32) { + return "mma A fragment registers must be 32-bit " + "packed values"; + } + } + for (PTXOperand::Array::const_iterator element = b.array.begin(); + element != b.array.end(); ++element) { + if (element->type != PTXOperand::b32) { + return "mma B fragment registers must be 32-bit " + "packed values"; + } + } + for (PTXOperand::Array::const_iterator element = c.array.begin(); + element != c.array.end(); ++element) { + if (!PTXOperand::relaxedValid(PTXOperand::s32, element->type)) { + return "mma C fragment registers must be s32 or b32"; + } + } + for (PTXOperand::Array::const_iterator element = d.array.begin(); + element != d.array.end(); ++element) { + if (!PTXOperand::relaxedValid(PTXOperand::s32, element->type)) { + return "mma D fragment registers must be s32 or b32"; + } + } + break; + } + if (modifier) { + return "f16/bf16/tf32 mma does not accept modifiers"; + } + const bool m16n8k8 = mmaShape == MmaM16N8K8; + const bool m16n8k4 = mmaShape == MmaM16N8K4; + const bool tf32Input = a.type == PTXOperand::tf32; + const bool halfAccumulator = type == PTXOperand::f16; + if ((m16n8k4 && !tf32Input) || + (!m16n8k4 && !m16n8k8 && mmaShape != MmaM16N8K16)) { + return "unsupported floating-point mma shape/type combination"; + } + if (!halfAccumulator && type != PTXOperand::f32) { + return "mma requires f16 or f32 accumulators"; + } + if (a.type != PTXOperand::f16 && a.type != PTXOperand::bf16 && + a.type != PTXOperand::tf32) { + return "mma A type must be f16, bf16, or tf32"; + } + if (tf32Input && ((!m16n8k8 && !m16n8k4) || type != PTXOperand::f32 || + b.type != PTXOperand::tf32)) { + return "tf32 mma requires m16n8k4 or m16n8k8 with f32 accumulators"; + } + if (halfAccumulator && a.type != PTXOperand::f16) { + return "f16 mma accumulators require f16 inputs"; + } + if (b.type != a.type) { + return "mma A and B types must match"; + } + if (d.type != type || c.type != type) { + return halfAccumulator ? "mma C and D types must be f16" + : "mma C and D types must be f32"; + } + const PTXOperand::Vec accumulatorVec = halfAccumulator + ? PTXOperand::v2 : PTXOperand::v4; + const unsigned int accumulatorRegisters = halfAccumulator ? 2 : 4; + const bool compactInputFragment = m16n8k4 || (m16n8k8 && !tf32Input); + const PTXOperand::Vec inputAVec = compactInputFragment + ? PTXOperand::v2 : PTXOperand::v4; + const PTXOperand::Vec inputBVec = compactInputFragment + ? PTXOperand::v1 : PTXOperand::v2; + if (d.vec != accumulatorVec || c.vec != accumulatorVec || + a.vec != inputAVec || b.vec != inputBVec) { + return "mma has invalid fragment sizes"; + } + if (d.array.size() != accumulatorRegisters || + c.array.size() != accumulatorRegisters || + a.array.size() != (compactInputFragment ? 2u : 4u) || + b.array.size() != (compactInputFragment ? 1u : 2u)) { + return "mma has invalid fragment register counts"; + } + for (PTXOperand::Array::const_iterator element = a.array.begin(); + element != a.array.end(); ++element) { + if (element->type != PTXOperand::b32 && + (tf32Input || element->type != PTXOperand::f16x2)) { + return "mma A fragment registers must be 32-bit packed values"; + } + } + for (PTXOperand::Array::const_iterator element = b.array.begin(); + element != b.array.end(); ++element) { + if (element->type != PTXOperand::b32 && + (tf32Input || element->type != PTXOperand::f16x2)) { + return "mma B fragment registers must be 32-bit packed values"; + } + } + for (PTXOperand::Array::const_iterator element = c.array.begin(); + element != c.array.end(); ++element) { + if (halfAccumulator ? (element->type != PTXOperand::b32 && + element->type != PTXOperand::f16x2) + : !PTXOperand::relaxedValid(PTXOperand::f32, element->type)) { + return halfAccumulator ? "mma C fragment registers must be b32 or f16x2" + : "mma C fragment registers must be f32 or b32"; + } + } + for (PTXOperand::Array::const_iterator element = d.array.begin(); + element != d.array.end(); ++element) { + if (halfAccumulator ? (element->type != PTXOperand::b32 && + element->type != PTXOperand::f16x2) + : !PTXOperand::relaxedValid(PTXOperand::f32, element->type)) { + return halfAccumulator ? "mma D fragment registers must be b32 or f16x2" + : "mma D fragment registers must be f32 or b32"; + } + } + break; + } case MadC: { if( !( type == PTXOperand::u32 || type == PTXOperand::s32 ) ) { return "invalid instruction type " @@ -1122,11 +1717,21 @@ std::string ir::PTXInstruction::valid() const { } case Max: { if( !( type != PTXOperand::s8 && type != PTXOperand::u8 && - type != PTXOperand::b8 && type != PTXOperand::f16 + type != PTXOperand::b8 && type != PTXOperand::pred ) ) { return "invalid instruction type " + PTXOperand::toString( type ); } + if (type == PTXOperand::b16 || type == PTXOperand::b32 || + type == PTXOperand::b64) { + return "invalid instruction type " + PTXOperand::toString(type); + } + if ((type == PTXOperand::u16 || type == PTXOperand::u32 || + type == PTXOperand::u64 || type == PTXOperand::s16 || + type == PTXOperand::s32 || type == PTXOperand::s64) && + (modifier || carry != None)) { + return "integer max does not support modifiers or carry out"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1150,13 +1755,25 @@ std::string ir::PTXInstruction::valid() const { case Membar: { break; } + case Fence: { + break; + } case Min: { if( !( type != PTXOperand::s8 && type != PTXOperand::u8 && - type != PTXOperand::b8 && type != PTXOperand::f16 - && type != PTXOperand::pred ) ) { + type != PTXOperand::b8 && type != PTXOperand::pred ) ) { return "invalid instruction type " + PTXOperand::toString( type ); } + if (type == PTXOperand::b16 || type == PTXOperand::b32 || + type == PTXOperand::b64) { + return "invalid instruction type " + PTXOperand::toString(type); + } + if ((type == PTXOperand::u16 || type == PTXOperand::u32 || + type == PTXOperand::u64 || type == PTXOperand::s16 || + type == PTXOperand::s32 || type == PTXOperand::s64) && + (modifier || carry != None)) { + return "integer min does not support modifiers or carry out"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1178,14 +1795,16 @@ std::string ir::PTXInstruction::valid() const { break; } case Mov: { - if ( ( a.type == PTXOperand::f16 ) && + if ( type != PTXOperand::b16 && + a.type == PTXOperand::f16 && + a.array.empty() && a.addressMode != PTXOperand::Address && a.addressMode != PTXOperand::Immediate ) { - return "invalid type for operand A " + return "invalid type for operand A " + PTXOperand::toString( a.type ); } if ( !( d.type != PTXOperand::s8 && d.type != PTXOperand::u8 - && d.type != PTXOperand::b8 && d.type != PTXOperand::f16 ) ) { + && d.type != PTXOperand::b8 ) ) { return "invalid type for operand D " + PTXOperand::toString( d.type ); } @@ -1222,11 +1841,15 @@ std::string ir::PTXInstruction::valid() const { } case Mul: { if( type == PTXOperand::s8 || type == PTXOperand::u8 - || type == PTXOperand::b8 || type == PTXOperand::f16 + || type == PTXOperand::b8 || type == PTXOperand::pred ) { return "invalid instruction type " + PTXOperand::toString( type ); } + if ((type == PTXOperand::f16 || type == PTXOperand::f16x2) + && (modifier & ~(rn | ftz | sat))) { + return "mul.f16 only supports .rn, .ftz, and .sat"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1264,10 +1887,16 @@ std::string ir::PTXInstruction::valid() const { case Neg: { if( type != PTXOperand::s16 && type != PTXOperand::s32 && type != PTXOperand::s64 && type != PTXOperand::f32 && - type != PTXOperand::f64 ) { + type != PTXOperand::f64 && type != PTXOperand::f16 && + type != PTXOperand::f16x2 && type != PTXOperand::bf16x2 && + type != PTXOperand::bf16 ) { return "invalid instruction type " + PTXOperand::toString( type ); } + if ((type == PTXOperand::s16 || type == PTXOperand::s32 || + type == PTXOperand::s64) && (modifier || carry != None)) { + return "integer neg does not support modifiers or carry out"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1315,6 +1944,9 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString( type ); } + if (modifier || carry != None) { + return "popc does not support modifiers or carry out"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1402,17 +2034,24 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString( type ); } - if( type == PTXOperand::f64 ) { - if( modifier & ftz ) { - if( !( modifier & approx ) ) { - return "requires .approx.ftz for f64"; - } + const int rounding = modifier & (rn | rz | rm | rp); + if( modifier & ~(approx | ftz | rn | rz | rm | rp) ) { + return "invalid modifier for rcp"; + } + if( type == PTXOperand::f64 && (modifier & approx) ) { + if( modifier != (approx | ftz) ) { + return "rcp.approx.f64 requires .ftz and no rounding modifier"; } - else if ( !( modifier & rn ) && !( modifier & rz ) - && !( modifier & rm ) && !( modifier & rp ) ) { - return "rounding mode required"; + } else { + if( modifier & approx ) { + if( rounding ) return "rcp.approx.f32 cannot specify rounding"; + } else if( rounding != rn && rounding != rz + && rounding != rm && rounding != rp ) { + return "rcp requires exactly one rounding modifier"; + } + if( type == PTXOperand::f64 && (modifier & ftz) ) { + return "rcp.rnd.f64 does not support .ftz"; } - } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) @@ -1422,43 +2061,38 @@ std::string ir::PTXInstruction::valid() const { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } - if( modifier & ftz ) { - if( PTXOperand::isInt( type ) ) { - return toString( ftz ) - + " only valid for float point instructions."; - } - } break; } case Red: { - if( ( reductionOperation == ReductionAnd + if( ( reductionOperation == ReductionAnd || reductionOperation == ReductionOr - || reductionOperation == ReductionXor ) - && type != PTXOperand::b32 ) { - return "invalid instruction type " - + PTXOperand::toString( type ) + " for reduction " + || reductionOperation == ReductionXor ) + && type != PTXOperand::b32 && type != PTXOperand::b64 ) { + return "invalid instruction type " + + PTXOperand::toString( type ) + " for reduction " + toString( reductionOperation ); } if( reductionOperation == ReductionAdd - && ( type != PTXOperand::u32 && type != PTXOperand::s32 - && type != PTXOperand::f32 && type != PTXOperand::u64 ) ) { - return "invalid instruction type " - + PTXOperand::toString( type ) + " for reduction " + && ( type != PTXOperand::u32 && type != PTXOperand::s32 + && type != PTXOperand::u64 && type != PTXOperand::f32 + && type != PTXOperand::f64 ) ) { + return "invalid instruction type " + + PTXOperand::toString( type ) + " for reduction " + toString( reductionOperation ); } if( ( reductionOperation == ReductionInc || reductionOperation == ReductionDec ) && type != PTXOperand::u32 ) { - return "invalid instruction type " - + PTXOperand::toString( type ) + " for reduction " + return "invalid instruction type " + + PTXOperand::toString( type ) + " for reduction " + toString( reductionOperation ); } if( ( reductionOperation == ReductionMin || reductionOperation == ReductionMax ) - && ( type != PTXOperand::u32 && type != PTXOperand::s32 - && type != PTXOperand::f32 ) ) { - return "invalid instruction type " - + PTXOperand::toString( type ) + " for reduction " + && ( type != PTXOperand::u32 && type != PTXOperand::s32 + && type != PTXOperand::u64 && type != PTXOperand::s64 ) ) { + return "invalid instruction type " + + PTXOperand::toString( type ) + " for reduction " + toString( reductionOperation ); } if( a.addressMode != PTXOperand::Address @@ -1501,6 +2135,9 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString( type ); } + if( modifier != approx && modifier != (approx | ftz) ) { + return "rsqrt requires .approx with optional .ftz"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1537,7 +2174,7 @@ std::string ir::PTXInstruction::valid() const { + " cannot be assigned to " + PTXOperand::toString( type ); } if( !PTXOperand::valid( type, c.type ) ) { - return "operand C type " + PTXOperand::toString( b.type ) + return "operand C type " + PTXOperand::toString( c.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } break; @@ -1576,14 +2213,15 @@ std::string ir::PTXInstruction::valid() const { && type != PTXOperand::u32 && type != PTXOperand::u64 && type != PTXOperand::b16 && type != PTXOperand::b32 && type != PTXOperand::b64 && type != PTXOperand::f32 - && type != PTXOperand::f64 ) { + && type != PTXOperand::f64 && type != PTXOperand::f16 ) { return "invalid instruction type " + PTXOperand::toString( type ); } - if( d.type != PTXOperand::s32 && d.type != PTXOperand::f32 - && d.type != PTXOperand::u32 ) { + if( d.type != PTXOperand::s32 && d.type != PTXOperand::f32 + && d.type != PTXOperand::u32 && d.type != PTXOperand::b32 + && d.type != PTXOperand::b16 && d.type != PTXOperand::f16 ) { return "operand D type " + PTXOperand::toString( d.type ) - + " invalid (must be u32, s32, or f32)"; + + " invalid (must be b16, f16, b32, u32, s32, or f32)"; } if( c.type != PTXOperand::pred && c.addressMode != PTXOperand::Invalid ) { @@ -1598,13 +2236,19 @@ std::string ir::PTXInstruction::valid() const { break; } case SetP: { - if( d.type != PTXOperand::pred ) { + if( d.type != PTXOperand::pred + && d.addressMode != PTXOperand::BitBucket ) { return "destination must be a predicate"; } if( pq.type != PTXOperand::pred - && pq.addressMode != PTXOperand::Invalid ) { + && pq.addressMode != PTXOperand::Invalid + && pq.addressMode != PTXOperand::BitBucket ) { return "Pq must be a predicate"; } + if( d.addressMode == PTXOperand::BitBucket + && pq.addressMode == PTXOperand::BitBucket ) { + return "only one setp destination may be a sink"; + } if( c.type != PTXOperand::pred && c.addressMode != PTXOperand::Invalid ) { return "operand C type " + PTXOperand::toString( c.type ) @@ -1615,7 +2259,7 @@ std::string ir::PTXInstruction::valid() const { && type != PTXOperand::u32 && type != PTXOperand::u64 && type != PTXOperand::b16 && type != PTXOperand::b32 && type != PTXOperand::b64 && type != PTXOperand::f32 - && type != PTXOperand::f64 ) { + && type != PTXOperand::f64 && type != PTXOperand::f16 ) { return "invalid instruction type " + PTXOperand::toString( type ); } @@ -1680,10 +2324,9 @@ std::string ir::PTXInstruction::valid() const { stream << "second source operand must be 32-bit, got " << b.bytes() << " bytes"; return stream.str(); } - if ( c.bytes() != 4 && c.addressMode != PTXOperand::Immediate ) { - std::stringstream stream; - stream << "shift amount must be 32-bit, got " << c.bytes() << " bytes"; - return stream.str(); + if ( c.addressMode != PTXOperand::Immediate + && !PTXOperand::valid( PTXOperand::u32, c.type ) ) { + return "shift amount must have a 32-bit integer type"; } break; } @@ -1693,18 +2336,19 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString( type ); } - if( d.bytes() != a.bytes() - && a.addressMode != PTXOperand::Immediate ) { - std::stringstream stream; - stream << "size of operand A " << a.bytes() - << " does not match size of operand D " << d.bytes(); - return stream.str(); + if( !PTXOperand::valid( type, d.type ) ) { + return "operand D type " + PTXOperand::toString( d.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); } - if( b.bytes() != 4 && b.addressMode != PTXOperand::Immediate ) { - std::stringstream stream; - stream << "size of operand B " << a.bytes() - << " must be 4 bytes"; - return stream.str(); + if( a.addressMode != PTXOperand::Immediate + && !PTXOperand::valid( type, a.type ) ) { + return "operand A type " + PTXOperand::toString( a.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + if( b.addressMode != PTXOperand::Immediate + && !PTXOperand::valid( PTXOperand::u32, b.type ) ) { + return "operand B type " + PTXOperand::toString( b.type ) + + " cannot be assigned to u32"; } break; } @@ -1717,17 +2361,19 @@ std::string ir::PTXInstruction::valid() const { return "invalid instruction type " + PTXOperand::toString( type ); } - if( d.bytes() != a.bytes() ) { - std::stringstream stream; - stream << "size of operand A " << a.bytes() - << " does not match size of operand D " << d.bytes(); - return stream.str(); + if( !PTXOperand::valid( type, d.type ) ) { + return "operand D type " + PTXOperand::toString( d.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); } - if( b.bytes() != 4 && b.addressMode != PTXOperand::Immediate ) { - std::stringstream stream; - stream << "size of operand B " << a.bytes() - << " must be 4 bytes"; - return stream.str(); + if( a.addressMode != PTXOperand::Immediate + && !PTXOperand::valid( type, a.type ) ) { + return "operand A type " + PTXOperand::toString( a.type ) + + " cannot be assigned to " + PTXOperand::toString( type ); + } + if( b.addressMode != PTXOperand::Immediate + && !PTXOperand::valid( PTXOperand::u32, b.type ) ) { + return "operand B type " + PTXOperand::toString( b.type ) + + " cannot be assigned to u32"; } break; } @@ -1735,7 +2381,10 @@ std::string ir::PTXInstruction::valid() const { if( !( type == PTXOperand::f32 ) ) { return "invalid instruction type " + PTXOperand::toString( type ); - } + } + if( modifier != approx && modifier != ( approx | ftz ) ) { + return "sin requires .approx with optional .ftz"; + } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) + " cannot be assigned to " + PTXOperand::toString( type ); @@ -1744,12 +2393,6 @@ std::string ir::PTXInstruction::valid() const { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } - if( modifier & ftz ) { - if( PTXOperand::isInt( type ) ) { - return toString( ftz ) - + " only valid for float point instructions."; - } - } break; } case SlCt: { @@ -1786,32 +2429,34 @@ std::string ir::PTXInstruction::valid() const { return "operand C must be either s32 or f32 assignable"; } if( modifier & ftz ) { - if( PTXOperand::isInt( type ) ) { + if( c.type != PTXOperand::f32 ) { return toString( ftz ) - + " only valid for float point instructions."; + + " only valid for an f32 comparison."; } } break; } case Sqrt: { - if( ( modifier & approx ) ) { - if( type != PTXOperand::f32 ) { - return "only f32 supported for approximate"; - } - } if( type != PTXOperand::f32 && type != PTXOperand::f64 ) { return "invalid instruction type " + PTXOperand::toString( type ); } - if( type == PTXOperand::f64 ) { - if( !( modifier & rn ) && !( modifier & rz ) - && !( modifier & rm ) && !( modifier & rp ) ) { - return "requires a rounding modifier"; - } - if( !( modifier & rn ) ) { - return "only nearest rounding supported"; + const int rounding = modifier & (rn | rz | rm | rp); + if( modifier & ~(approx | ftz | rn | rz | rm | rp) ) { + return "invalid modifier for sqrt"; + } + if( modifier & approx ) { + if( type != PTXOperand::f32 ) { + return "sqrt.approx only supports f32"; } + if( rounding ) return "sqrt.approx cannot specify rounding"; + } else if( rounding != rn && rounding != rz + && rounding != rm && rounding != rp ) { + return "sqrt requires exactly one rounding modifier"; + } + if( type == PTXOperand::f64 && (modifier & ftz) ) { + return "sqrt.rnd.f64 does not support .ftz"; } if( !PTXOperand::valid( type, a.type ) ) { return "operand A type " + PTXOperand::toString( a.type ) @@ -1821,12 +2466,6 @@ std::string ir::PTXInstruction::valid() const { return "operand D type " + PTXOperand::toString( d.type ) + " cannot be assigned to " + PTXOperand::toString( type ); } - if( modifier & ftz ) { - if( PTXOperand::isInt( type ) ) { - return toString( ftz ) - + " only valid for float point instructions."; - } - } break; } case St: { @@ -1862,11 +2501,15 @@ std::string ir::PTXInstruction::valid() const { } case Sub: { if ( !( type != PTXOperand::s8 && type != PTXOperand::u8 && - type != PTXOperand::b8 && type != PTXOperand::f16 + type != PTXOperand::b8 && type != PTXOperand::pred ) ) { return "invalid instruction type " + PTXOperand::toString( type ); } + if ((type == PTXOperand::f16 || type == PTXOperand::f16x2) + && (modifier & ~(rn | ftz | sat))) { + return "sub.f16 only supports .rn, .ftz, and .sat"; + } if( carry == CC ) { if( ( modifier & sat ) ) { return "saturate not supported with carry out"; @@ -1915,6 +2558,17 @@ std::string ir::PTXInstruction::valid() const { } break; } + case Tanh: { + if( !(type == PTXOperand::f32 || type == PTXOperand::f16 + || type == PTXOperand::f16x2) || modifier != approx ) { + return "tanh requires .approx with f32, f16, or f16x2"; + } + if( !PTXOperand::valid(type, a.type) + || !PTXOperand::valid(type, d.type) ) { + return "tanh operands must have compatible types"; + } + break; + } case TestP: { if( !( type == PTXOperand::f32 || type == PTXOperand::f64 ) ) { return "invalid instruction type " @@ -1980,17 +2634,20 @@ std::string ir::PTXInstruction::valid() const { } case Suld: { if (formatMode == Formatted && !(type == ir::PTXOperand::b32 - || type == ir::PTXOperand::u32 + || type == ir::PTXOperand::u32 || type == ir::PTXOperand::s32 || type == ir::PTXOperand::f32)) { return "sust.p - data type must be .b32, .u32, .s32, or .f32"; } else if (formatMode == Unformatted && !(type == ir::PTXOperand::b8 - || type == ir::PTXOperand::b16 + || type == ir::PTXOperand::b16 || type == ir::PTXOperand::b32 || type == ir::PTXOperand::b64)) { return "sust.b - data type must be .b8, .b16, .b32, or .b64"; } + if (cacheOperation == Nc || cacheOperation == Lu) { + return "suld cache operator must be .ca, .cg, .cs, or .cv"; + } break; } case Suq: { @@ -2003,6 +2660,24 @@ std::string ir::PTXInstruction::valid() const { } break; } + case Szext: { + if( type != PTXOperand::u32 && type != PTXOperand::s32 ) { + return "szext instruction requires .u32 or .s32 type"; + } + if( shiftMode != ShiftMode::Clamp && shiftMode != ShiftMode::Wrap ) { + return "szext instruction requires .clamp or .wrap mode"; + } + if( !PTXOperand::valid( type, d.type ) ) { + return "invalid szext destination type " + PTXOperand::toString(d.type); + } + if( !PTXOperand::valid( type, a.type ) ) { + return "invalid szext source type " + PTXOperand::toString(a.type); + } + if( b.type != PTXOperand::u32 ) { + return "szext operand B must have .u32 type"; + } + break; + } case Sured: { if (!(reductionOperation == ReductionAdd || reductionOperation == ReductionMin @@ -2019,17 +2694,20 @@ std::string ir::PTXInstruction::valid() const { } case Sust: { if (formatMode == Formatted && !(type == ir::PTXOperand::b32 - || type == ir::PTXOperand::u32 + || type == ir::PTXOperand::u32 || type == ir::PTXOperand::s32 || type == ir::PTXOperand::f32)) { return "sust.p - data type must be .b32, .u32, .s32, or .f32"; } else if (formatMode == Unformatted && !(type == ir::PTXOperand::b8 - || type == ir::PTXOperand::b16 + || type == ir::PTXOperand::b16 || type == ir::PTXOperand::b32 || type == ir::PTXOperand::b64)) { return "sust.b - data type must be .b8, .b16, .b32, or .b64"; } + if (cacheOperation == Nc || cacheOperation == Lu) { + return "sust cache operator must be .ca, .cg, .cs, or .cv"; + } break; } case Trap: { @@ -2112,10 +2790,19 @@ std::string ir::PTXInstruction::toString() const { + b.toString(); } case Atom: { - std::string result = guard() + "atom." + toString( addressSpace ) - + "." + toString( atomicOperation ) + "." + std::string result = guard() + "atom."; + if( semantics != Relaxed ) { + result += toString( semantics ) + "."; + } + if( scope != Level_Invalid ) { + // same ".gpu" spelling quirk as fence's printer + result += ( ( scope == GlobalLevel ) ? "gpu" : toString( scope ) ) + + std::string( "." ); + } + result += toString( addressSpace ) + + "." + toString( atomicOperation ) + "." + PTXOperand::toString( type ) + " " - + d.toString() + ", [" + a.toString() + "], " + + d.toString() + ", [" + a.toString() + "], " + b.toString(); if( c.addressMode != PTXOperand::Invalid ) { result += ", " + c.toString(); @@ -2164,6 +2851,22 @@ std::string ir::PTXInstruction::toString() const { + d.toString() + ", " + a.toString(); return result; } + case Bmsk: { + return guard() + "bmsk." + toString(shiftMode) + ".b32 " + + d.toString() + ", " + a.toString() + ", " + b.toString(); + } + case Lop3: { + std::string result = guard() + "lop3"; + if( pq.addressMode != PTXOperand::Invalid ) { + result += "." + toString(booleanOperator); + } + result += ".b32 " + d.toString(); + if( pq.addressMode != PTXOperand::Invalid ) result += "|" + pq.toString(); + result += ", " + a.toString() + ", " + b.toString() + ", " + + c.toString() + ", " + immLut.toString(); + if( pq.addressMode != PTXOperand::Invalid ) result += ", " + q.toString(); + return result; + } case Bra: { std::string result = guard() + "bra"; if( uni ) { @@ -2218,38 +2921,11 @@ std::string ir::PTXInstruction::toString() const { } case Cvt: { std::string result = guard() + "cvt."; - if( PTXOperand::isFloat( d.type )) { - if ((d.type == PTXOperand::f32 && a.type == PTXOperand::f64) - || PTXOperand::isInt(a.type)) { - result += modifierString( modifier, carry ); - } - } - else { - if( modifier & rn ) { - result += "rn."; - } else if( modifier & rz ) { - result += "rz."; - } else if( modifier & rm ) { - result += "rm."; - } else if( modifier & rp ) { - result += "rp."; - } else if( modifier & rni ) { - result += "rni."; - } else if( modifier & rzi ) { - result += "rzi."; - } else if( modifier & rmi ) { - result += "rmi."; - } else if( modifier & rpi ) { - result += "rpi."; - } - - if( modifier & ftz ) { - result += "ftz."; - } - if( modifier & sat ) { - result += "sat."; - } - } + if( modifier & rni ) result += "rni."; + else if( modifier & rzi ) result += "rzi."; + else if( modifier & rmi ) result += "rmi."; + else if( modifier & rpi ) result += "rpi."; + result += modifierString(modifier, carry); PTXOperand::DataType sourceType = a.type; @@ -2260,6 +2936,11 @@ std::string ir::PTXInstruction::toString() const { result += PTXOperand::toString( type ) + "." + PTXOperand::toString( sourceType ) + " " + d.toString() + ", " + a.toString(); + if( type == PTXOperand::f16x2 || type == PTXOperand::bf16x2 ) { + result += ", " + b.toString(); + } else if (b.addressMode != PTXOperand::Invalid) { + result += ", " + b.toString(); + } return result; } case Cvta: { @@ -2276,7 +2957,7 @@ std::string ir::PTXInstruction::toString() const { } case Div: { std::string result = guard() + "div."; - if( divideFull ) { + if( modifier & full ) { result += "full."; } result += modifierString( modifier, carry ); @@ -2284,6 +2965,14 @@ std::string ir::PTXInstruction::toString() const { + a.toString() + ", " + b.toString(); return result; } + case Dp2a: + case Dp4a: { + std::string result = guard() + toString(opcode) + "."; + if( opcode == Dp2a ) result += modifierString(modifier); + return result + PTXOperand::toString(type) + "." + + PTXOperand::toString(bType) + " " + d.toString() + ", " + + a.toString() + ", " + b.toString() + ", " + c.toString(); + } case Ex2: { std::string result = guard() + "ex2."; result += modifierString( modifier, carry ); @@ -2301,6 +2990,10 @@ std::string ir::PTXInstruction::toString() const { + a.toString() + ", " + b.toString() + ", " + c.toString(); return result; } + case Fns: { + return guard() + "fns.b32 " + d.toString() + ", " + a.toString() + + ", " + b.toString() + ", " + c.toString(); + } case Isspacep: { std::string result = guard() + "isspacep." + toString(addressSpace) + " " + d.toString() + ", " + a.toString(); @@ -2310,6 +3003,15 @@ std::string ir::PTXInstruction::toString() const { std::string result = guard() + "ld."; if( volatility == Volatile ) { result += "volatile."; + } else if( semantics != Weak ) { + if( mmio ) { + result += "mmio."; + } + result += toString( semantics ) + "."; + if( scope != Level_Invalid ) { + result += ( ( scope == GlobalLevel ) ? "gpu" : toString( scope ) ) + + std::string( "." ); + } } if( cacheOperation != Ca ) { result += toStringLoad(cacheOperation) + "."; @@ -2363,6 +3065,37 @@ std::string ir::PTXInstruction::toString() const { + a.toString() + ", " + b.toString() + ", " + c.toString(); return result; } + case Mma: { + std::string shapeName; + switch (mmaShape) { + case MmaM8N8K32: shapeName = "8n8k32"; break; + case MmaM16N8K64: shapeName = "16n8k64"; break; + case MmaM8N8K128: shapeName = "8n8k128"; break; + case MmaM16N8K128: shapeName = "16n8k128"; break; + case MmaM16N8K256: shapeName = "16n8k256"; break; + case MmaM8N8K4: shapeName = "8n8k4"; break; + case MmaM16N8K4: shapeName = "16n8k4"; break; + case MmaM16N8K8: shapeName = "16n8k8"; break; + case MmaM8N8K16: shapeName = "8n8k16"; break; + case MmaM16N8K32: shapeName = "16n8k32"; break; + default: shapeName = "16n8k16"; break; + } + std::string result = guard() + + "mma.sync.aligned.m" + + shapeName + + (mmaAColumnMajor ? ".col" : ".row") + + (mmaBColumnMajor ? ".col." : ".row.") + + ((modifier & satfinite) ? "satfinite." : "") + + PTXOperand::toString(type) + "." + + PTXOperand::toString(a.type) + "." + + PTXOperand::toString(b.type) + "." + + PTXOperand::toString(c.type) + + (a.type == PTXOperand::b1 ? "." + toString(booleanOperator) + ".popc" : "") + + (type == PTXOperand::f64 && modifier ? "." + toString((Modifier)modifier) : "") + " " + + d.toString() + ", " + a.toString() + ", " + + b.toString() + ", " + c.toString(); + return result; + } case MadC: { std::string result = guard() + "madc."; result += modifierString( modifier, carry ); @@ -2378,6 +3111,12 @@ std::string ir::PTXInstruction::toString() const { case Membar: { return guard() + "membar." + toString( level ); } + case Fence: { + // fence spells the device-wide scope ".gpu", membar spells + // the same GlobalLevel value ".gl" -- toString(Level) can't be reused as-is + std::string scope = ( level == GlobalLevel ) ? "gpu" : toString( level ); + return guard() + "fence." + toString( semantics ) + "." + scope; + } case Min: { return guard() + "min." + modifierString(modifier, carry) + PTXOperand::toString( type ) + " " @@ -2447,10 +3186,19 @@ std::string ir::PTXInstruction::toString() const { return result; } case Red: { - return guard() + "red." + toString( addressSpace ) + "." - + toString( reductionOperation ) + "." - + PTXOperand::toString( type ) + " " + d.toString() + ", " + std::string result = guard() + "red."; + if( semantics != Relaxed ) { + result += toString( semantics ) + "."; + } + if( scope != Level_Invalid ) { + result += ( ( scope == GlobalLevel ) ? "gpu" : toString( scope ) ) + + std::string( "." ); + } + result += toString( addressSpace ) + "." + + toString( reductionOperation ) + "." + + PTXOperand::toString( type ) + " " + d.toString() + ", " + a.toString(); + return result; } case Rem: { return guard() + "rem." + PTXOperand::toString( type ) + " " @@ -2562,6 +3310,15 @@ std::string ir::PTXInstruction::toString() const { std::string result = guard() + "st."; if( volatility == Volatile ) { result += "volatile."; + } else if( semantics != Weak ) { + if( mmio ) { + result += "mmio."; + } + result += toString( semantics ) + "."; + if( scope != Level_Invalid ) { + result += ( ( scope == GlobalLevel ) ? "gpu" : toString( scope ) ) + + std::string( "." ); + } } if( cacheOperation != Wb ) { result += toStringStore(cacheOperation) + "."; @@ -2602,6 +3359,11 @@ std::string ir::PTXInstruction::toString() const { + "." + PTXOperand::toString(type) + " " + d.toString() + ", [" + a.toString() + "]"; } + case Szext: { + return guard() + "szext." + toString(shiftMode) + "." + + PTXOperand::toString(type) + " " + d.toString() + ", " + + a.toString() + ", " + b.toString(); + } case Sured: { return guard() + "sured" + toString(formatMode) + "." + toString(reductionOperation) + "." + @@ -2616,6 +3378,11 @@ std::string ir::PTXInstruction::toString() const { + toString(clamp) + " [" + d.toString() +", " + a.toString() + "], " + b.toString(); } + case Tanh: { + return guard() + "tanh." + modifierString(modifier) + + PTXOperand::toString(type) + " " + d.toString() + ", " + + a.toString(); + } case TestP: { return guard() + "testp." + toString( floatingPointMode ) + "." + PTXOperand::toString( type ) + " " + d.toString() + ", " diff --git a/ocelot/src/ir/PTXKernel.cpp b/ocelot/src/ir/PTXKernel.cpp index 12e474175..01f831553 100644 --- a/ocelot/src/ir/PTXKernel.cpp +++ b/ocelot/src/ir/PTXKernel.cpp @@ -450,14 +450,14 @@ PTXKernel::RegisterMap PTXKernel::assignRegisters( ControlFlowGraph& cfg ) instruction != block->instructions.end(); ++instruction) { PTXInstruction& instr = *static_cast( *instruction); - PTXOperand PTXInstruction:: * operands[] = - { &PTXInstruction::a, &PTXInstruction::b, &PTXInstruction::c, - &PTXInstruction::d, &PTXInstruction::pg, - &PTXInstruction::pq }; - + PTXOperand PTXInstruction:: * operands[] = + { &PTXInstruction::a, &PTXInstruction::b, &PTXInstruction::c, + &PTXInstruction::d, &PTXInstruction::pg, + &PTXInstruction::pq, &PTXInstruction::q }; + report( " For instruction '" << instr.toString() << "'" ); - - for (int i = 0; i < 6; i++) { + + for (int i = 0; i < 7; i++) { if ((instr.*operands[i]).addressMode == PTXOperand::Invalid) { continue; diff --git a/ocelot/src/ir/PTXOperand.cpp b/ocelot/src/ir/PTXOperand.cpp index 7ea6c4c93..ba32319b6 100644 --- a/ocelot/src/ir/PTXOperand.cpp +++ b/ocelot/src/ir/PTXOperand.cpp @@ -30,12 +30,16 @@ std::string ir::PTXOperand::toString(Vec index) { case v1: return "v1"; break; case v2: return "v2"; break; case v4: return "v4"; break; + case v8: return "v8"; break; } return ""; } std::string ir::PTXOperand::toString( DataType type ) { switch( type ) { + case s4: return "s4"; break; + case b1: return "b1"; break; + case u4: return "u4"; break; case s8: return "s8"; break; case s16: return "s16"; break; case s32: return "s32"; break; @@ -49,7 +53,11 @@ std::string ir::PTXOperand::toString( DataType type ) { case b32: return "b32"; break; case b64: return "b64"; break; case f16: return "f16"; break; + case f16x2:return "f16x2";break; case f32: return "f32"; break; + case bf16: return "bf16"; break; + case bf16x2:return "bf16x2";break; + case tf32: return "tf32"; break; case f64: return "f64"; break; case pred: return "pred"; break; default: break; @@ -149,7 +157,11 @@ bool ir::PTXOperand::isFloat( DataType type ) { bool result = false; switch( type ) { case f16: /* fall through */ + case f16x2: /* fall through */ case f32: /* fall through */ + case bf16:/* fall through */ + case bf16x2:/* fall through */ + case tf32:/* fall through */ case f64: result = true; default: break; } @@ -194,10 +206,14 @@ unsigned int ir::PTXOperand::bytes( DataType type ) { case u16: /* fall through */ case f16: /* fall through */ case b16: /* fall through */ + case bf16: /* fall through */ case s16: return 2; break; case u32: /* fall through */ case b32: /* fall through */ + case f16x2: /* fall through */ + case bf16x2: /* fall through */ case f32: /* fall through */ + case tf32: /* fall through */ case s32: return 4; break; case f64: /* fall through */ case u64: /* fall through */ @@ -235,6 +251,7 @@ bool ir::PTXOperand::valid( DataType destination, DataType source ) { case s16: /* fall through */ case u16: /* fall through */ case f16: /* fall through */ + case bf16: /* fall through */ case b16: return true; break; default: break; } @@ -345,6 +362,15 @@ bool ir::PTXOperand::valid( DataType destination, DataType source ) { } break; } + case f16x2: { + return source == b32 || source == f16x2; + } + case bf16: { + return source == b16; + break; + } + case bf16x2: return source == b32; + case tf32: return source == b32; case pred: { return source == pred; break; @@ -539,6 +565,8 @@ bool ir::PTXOperand::relaxedValid( DataType instructionType, } case f32: { switch( operand ) { + case b64: /* fall through */ + case f64: /* fall through */ case b32: /* fall through */ case f32: return true; break; default: break; @@ -547,12 +575,33 @@ bool ir::PTXOperand::relaxedValid( DataType instructionType, } case f16: { switch( operand ) { + case b64: /* fall through */ + case f64: /* fall through */ + case b32: /* fall through */ + case f32: /* fall through */ case b16: /* fall through */ case f16: return true; break; default: break; } break; } + case f16x2: { + return operand == b32 || operand == f16x2 + || operand == b64 || operand == f64; + } + case bf16: { + switch( operand ) { + case b16: return true; break; + default: break; + } + break; + } + case bf16x2: { + return operand == b32; + } + case tf32: { + return operand == b32; + } case pred: { return operand == pred; break; @@ -774,7 +823,8 @@ std::string ir::PTXOperand::toString() const { else if( vec != v1 ) { if( !array.empty() ) { assert( ( vec == v2 && array.size() == 2 ) - || ( vec == v4 && array.size() == 4 ) ); + || ( vec == v4 && array.size() == 4 ) + || ( vec == v8 && array.size() == 8 ) ); std::string result = "{"; for( Array::const_iterator fi = array.begin(); fi != array.end(); ++fi ) { @@ -864,5 +914,3 @@ bool ir::PTXOperand::isRegister() const { bool ir::PTXOperand::isVector() const { return isRegister() && vec != v1; } - - diff --git a/ocelot/src/parser/PTXLexer.cpp b/ocelot/src/parser/PTXLexer.cpp index 8d9c3c01d..6cc7ef466 100644 --- a/ocelot/src/parser/PTXLexer.cpp +++ b/ocelot/src/parser/PTXLexer.cpp @@ -59,12 +59,16 @@ namespace parser CASE(OPCODE_SUB) CASE(OPCODE_EX2) CASE(OPCODE_LG2) + CASE(OPCODE_LOP3) CASE(OPCODE_RCP) CASE(OPCODE_SIN) CASE(OPCODE_REM) CASE(OPCODE_MUL24) CASE(OPCODE_MAD24) + CASE(OPCODE_MMA) CASE(OPCODE_DIV) + CASE(OPCODE_DP2A) + CASE(OPCODE_DP4A) CASE(OPCODE_ABS) CASE(OPCODE_NEG) CASE(OPCODE_MIN) @@ -99,6 +103,8 @@ namespace parser CASE(OPCODE_RED) CASE(OPCODE_NOT) CASE(OPCODE_CNOT) + CASE(OPCODE_FNS) + CASE(OPCODE_BMSK) CASE(OPCODE_VOTE) CASE(OPCODE_SHF) CASE(OPCODE_SHR) @@ -112,6 +118,8 @@ namespace parser CASE(OPCODE_SUST) CASE(OPCODE_SURED) CASE(OPCODE_SUQ) + CASE(OPCODE_SZEXT) + CASE(OPCODE_TANH) CASE(OPCODE_TXQ) CASE(PREPROCESSOR_INCLUDE) CASE(PREPROCESSOR_DEFINE) @@ -151,6 +159,7 @@ namespace parser CASE(TOKEN_PRAGMA) CASE(TOKEN_REG) CASE(TOKEN_SHARED) + CASE(TOKEN_SHARED_CTA) CASE(TOKEN_SAMPLERREF) CASE(TOKEN_SURFREF) CASE(TOKEN_TEXREF) @@ -172,6 +181,7 @@ namespace parser CASE(TOKEN_F16) CASE(TOKEN_F64) CASE(TOKEN_F32) + CASE(TOKEN_BF16) CASE(TOKEN_PRED) CASE(TOKEN_EQ) CASE(TOKEN_NE) @@ -216,6 +226,9 @@ namespace parser CASE(TOKEN_FTZ) CASE(TOKEN_APPROX) CASE(TOKEN_FULL) + CASE(TOKEN_NAN_MODIFIER) + CASE(TOKEN_XORSIGN) + CASE(TOKEN_ABS_MODIFIER) CASE(TOKEN_V2) CASE(TOKEN_V4) CASE(TOKEN_X) @@ -260,6 +273,11 @@ namespace parser CASE(TOKEN_ARRIVE) CASE(TOKEN_RED) CASE(TOKEN_SYNC) + CASE(TOKEN_ALIGNED) + CASE(TOKEN_M16N8K8) + CASE(TOKEN_M16N8K16) + CASE(TOKEN_ROW) + CASE(TOKEN_COL) CASE(TOKEN_POPC) CASE(TOKEN_BALLOT) CASE(TOKEN_F4E) @@ -367,4 +385,3 @@ namespace parser } #endif - diff --git a/ocelot/src/parser/PTXParser.cpp b/ocelot/src/parser/PTXParser.cpp index 636237bed..acc505aed 100644 --- a/ocelot/src/parser/PTXParser.cpp +++ b/ocelot/src/parser/PTXParser.cpp @@ -152,6 +152,66 @@ namespace parser operand.type = instruction.type; } } + //https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#integer-arithmetic-instructions-bfi + //https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#integer-arithmetic-instructions-bfe + if( instruction.opcode == ir::PTXInstruction::Bfi + || instruction.opcode == ir::PTXInstruction::Bfe ) + { + if ( instruction.b.addressMode == ir::PTXOperand::AddressMode::Immediate) + instruction.b.type = ir::PTXOperand::u32; + if ( instruction.c.addressMode == ir::PTXOperand::AddressMode::Immediate) + instruction.c.type = ir::PTXOperand::u32; + } + else if( instruction.opcode == ir::PTXInstruction::Fns ) + { + if( instruction.c.addressMode == ir::PTXOperand::Immediate ) + instruction.c.type = ir::PTXOperand::s32; + } + else if( instruction.opcode == ir::PTXInstruction::Szext + && instruction.b.addressMode == ir::PTXOperand::Immediate ) + { + instruction.b.type = ir::PTXOperand::u32; + } + } + + void PTXParser::State::_setMovVectorImmediateTypes() + { + ir::PTXInstruction& instruction = statement.instruction; + ir::PTXOperand& source = instruction.a; + + if( instruction.opcode != ir::PTXInstruction::Mov || + source.array.empty() ) return; + + unsigned int totalBytes = ir::PTXOperand::bytes( instruction.type ); + if( totalBytes % source.array.size() != 0 ) + { + throw_exception( "Invalid mov vector element size.", + InvalidInstruction ); + } + + unsigned int elementBytes = totalBytes / source.array.size(); + if( source.type == ir::PTXOperand::TypeSpecifier_invalid ) + { + switch( elementBytes ) + { + case 1: source.type = ir::PTXOperand::b8; break; + case 2: source.type = ir::PTXOperand::b16; break; + case 4: source.type = ir::PTXOperand::b32; break; + case 8: source.type = ir::PTXOperand::b64; break; + default: + throw_exception( "Invalid mov vector element size.", + InvalidInstruction ); + } + } + + for( ir::PTXOperand::Array::iterator element = source.array.begin(); + element != source.array.end(); ++element ) + { + if( element->addressMode == ir::PTXOperand::Immediate ) + { + element->type = source.type; + } + } } static std::string strip(const std::string& name) @@ -606,6 +666,8 @@ namespace parser else if( token == TOKEN_SM21 ) statement.targets.push_back( "sm_21" ); else if( token == TOKEN_SM30 ) statement.targets.push_back( "sm_30" ); else if( token == TOKEN_SM35 ) statement.targets.push_back( "sm_35" ); + else if( token == TOKEN_SM50 ) statement.targets.push_back( "sm_50" ); + else if( token == TOKEN_SM86 ) statement.targets.push_back( "sm_86" ); else if( token == TOKEN_MAP_F64_TO_F32 ) { statement.targets.push_back( "map_f64_to_f32" ); @@ -724,7 +786,7 @@ namespace parser statement.column = location.first_column; report( " At (" << statement.line << "," << statement.column - << ") : Parsed statement " << statements.size() + << ") : Parsed statement " << statements.size() << " \"" << statement.toString() << "\"" ); statements.push_back( statement ); @@ -1523,7 +1585,34 @@ namespace parser operandVector.push_back( OperandWrapper( operand, mode->space ) ); } - + void PTXParser::State::vectorOperand( unsigned int elements ) + { + assert( elements == 2 || elements == 4 ); + assert( operandVector.size() >= elements ); + + OperandVector::iterator begin = operandVector.end() - elements; + + ir::PTXOperand vector; + vector.addressMode = ir::PTXOperand::Register; + vector.type = ir::PTXOperand::TypeSpecifier_invalid; + vector.vec = elements == 2 ? ir::PTXOperand::v2 : ir::PTXOperand::v4; + + for( OperandVector::iterator element = begin; + element != operandVector.end(); ++element ) + { + if( vector.type == ir::PTXOperand::TypeSpecifier_invalid && + element->operand.addressMode != ir::PTXOperand::Immediate ) + { + vector.type = element->operand.type; + } + + vector.array.push_back( element->operand ); + } + + operandVector.erase( begin, operandVector.end() ); + operandVector.push_back( vector ); + } + void PTXParser::State::addressableOperand( const std::string& name, long long int value, YYLTYPE& location, bool invert ) { @@ -1570,13 +1659,27 @@ namespace parser OperandWrapper* mode = _getOperand( identifiers.front() ); - if( identifiers.size() > 4 ) + if (identifiers.size() == 8) { + operand.addressMode = ir::PTXOperand::Register; + operand.vec = ir::PTXOperand::v8; + operand.array.clear(); + for (const auto& name : identifiers) { + auto* element = _getOperand(name); + if (!element) throw_exception(toString(location, *this) << "Operand " << name << " not declared.", NoDeclaration); + operand.array.push_back(element->operand); + } + operand.type = operand.array.front().type; + operandVector.push_back(operand); + operand.array.clear(); + return; + } + if( identifiers.size() == 3 || identifiers.size() > 4 ) { throw_exception( toString( location, *this ) << "Array operand \"" << hydrazine::toString( identifiers.begin(), identifiers.end(), "," ) - << "\" has more than 4 elements.", InvalidArray ); + << "\" must have 1, 2, 4, or 8 elements.", InvalidArray ); } if( mode == 0 ) @@ -1721,7 +1824,7 @@ namespace parser void PTXParser::State::full() { - statement.instruction.divideFull = true; + statement.instruction.modifier |= ir::PTXInstruction::full; } void PTXParser::State::modifier( int token ) @@ -1781,7 +1884,30 @@ namespace parser { statement.instruction.level = tokenToLevel( token ); } - + + void PTXParser::State::semantics( int token ) + { + statement.instruction.semantics = tokenToSemantics( token ); + } + + void PTXParser::State::scope( int token ) + { + statement.instruction.scope = tokenToLevel( token ); + } + + void PTXParser::State::mmio( bool condition ) + { + statement.instruction.mmio = condition; + } + + void PTXParser::State::finalizeMmioAddressSpace() + { + if( statement.instruction.mmio ) + { + statement.instruction.addressSpace = ir::PTXInstruction::Global; + } + } + void PTXParser::State::permute( int token ) { statement.instruction.permuteMode = tokenToPermuteMode( token ); @@ -1822,7 +1948,10 @@ namespace parser if( operandVector.size() > index ) { - if( ( operandVector[ index ].operand.type == ir::PTXOperand::pred + if( ( ( operandVector[ index ].operand.type == ir::PTXOperand::pred + || ( statement.instruction.opcode == ir::PTXInstruction::SetP + && operandVector[index].operand.addressMode + == ir::PTXOperand::BitBucket ) ) && operandVector.size() > 4 ) || operandVector.size() == 6 ) { statement.instruction.pq = operandVector[index++].operand; @@ -1843,6 +1972,7 @@ namespace parser } _setImmediateTypes(); + _setMovVectorImmediateTypes(); } void PTXParser::State::instruction( const std::string& opcode ) @@ -1850,6 +1980,94 @@ namespace parser instruction( opcode, TOKEN_B64 ); } + void PTXParser::State::dotType( int token ) + { + ir::PTXInstruction& instruction = statement.instruction; + instruction.bType = tokenToDataType( token ); + if( instruction.b.addressMode == ir::PTXOperand::Immediate ) { + instruction.b.type = instruction.bType; + } + const ir::PTXOperand::DataType resultType = + instruction.type == ir::PTXOperand::u32 + && instruction.bType == ir::PTXOperand::u32 + ? ir::PTXOperand::u32 : ir::PTXOperand::s32; + if( instruction.c.addressMode == ir::PTXOperand::Immediate ) { + instruction.c.type = resultType; + } + } + + void PTXParser::State::lop3() + { + assert( operandVector.size() == 6 || operandVector.size() == 8 ); + statement.directive = ir::PTXStatement::Instr; + statement.instruction.opcode = ir::PTXInstruction::Lop3; + statement.instruction.type = ir::PTXOperand::b32; + statement.instruction.pg = operandVector[0].operand; + + unsigned int index = 1; + statement.instruction.d = operandVector[index++].operand; + if( operandVector.size() == 8 ) { + statement.instruction.pq = operandVector[index++].operand; + } + statement.instruction.a = operandVector[index++].operand; + statement.instruction.b = operandVector[index++].operand; + statement.instruction.c = operandVector[index++].operand; + statement.instruction.immLut = operandVector[index++].operand; + if( operandVector.size() == 8 ) { + statement.instruction.q = operandVector[index].operand; + } + _setImmediateTypes(); + } + + void PTXParser::State::mma( int shapeToken, int accumulatorToken, + int aToken, int bToken, int cToken, bool aColumnMajor, bool bColumnMajor ) + { + assert( operandVector.size() == 5 ); + ir::PTXOperand::DataType accumulatorType = tokenToDataType( accumulatorToken ); + ir::PTXOperand::DataType aType = tokenToDataType( aToken ); + ir::PTXOperand::DataType bType = tokenToDataType( bToken ); + ir::PTXOperand::DataType cType = tokenToDataType( cToken ); + + statement.directive = ir::PTXStatement::Instr; + statement.instruction.opcode = ir::PTXInstruction::Mma; + statement.instruction.mmaShape = shapeToken == TOKEN_M8N8K32 + ? ir::PTXInstruction::MmaM8N8K32 + : shapeToken == TOKEN_M8N8K128 + ? ir::PTXInstruction::MmaM8N8K128 + : shapeToken == TOKEN_M16N8K128 + ? ir::PTXInstruction::MmaM16N8K128 + : shapeToken == TOKEN_M16N8K256 + ? ir::PTXInstruction::MmaM16N8K256 + : shapeToken == TOKEN_M16N8K64 + ? ir::PTXInstruction::MmaM16N8K64 + : shapeToken == TOKEN_M16N8K8 + ? ir::PTXInstruction::MmaM16N8K8 + : shapeToken == TOKEN_M8N8K4 + ? ir::PTXInstruction::MmaM8N8K4 + : shapeToken == TOKEN_M16N8K4 + ? ir::PTXInstruction::MmaM16N8K4 + : shapeToken == TOKEN_M8N8K16 + ? ir::PTXInstruction::MmaM8N8K16 + : shapeToken == TOKEN_M16N8K32 + ? ir::PTXInstruction::MmaM16N8K32 + : ir::PTXInstruction::MmaM16N8K16; + statement.instruction.mmaAColumnMajor = aColumnMajor; + statement.instruction.mmaBColumnMajor = bColumnMajor; + statement.instruction.type = accumulatorType; + statement.instruction.pg = operandVector[0].operand; + statement.instruction.d = operandVector[1].operand; + statement.instruction.a = operandVector[2].operand; + statement.instruction.b = operandVector[3].operand; + statement.instruction.c = operandVector[4].operand; + + statement.instruction.d.type = accumulatorType; + statement.instruction.c.type = cType; + statement.instruction.a.type = aType; + statement.instruction.b.type = bType; + + _setImmediateTypes(); + } + void PTXParser::State::tex( int dataType ) { report( " Rule: instruction : tex" ); @@ -2044,7 +2262,17 @@ namespace parser << " using relaxed conversion rules.", InvalidDataType ); } + if (statement.instruction.a.addressMode == ir::PTXOperand::Immediate) { + statement.instruction.a.type = tokenToDataType(token); + } statement.instruction.a.relaxedType = tokenToDataType( token ); + if (statement.instruction.b.addressMode == ir::PTXOperand::Immediate) { + statement.instruction.b.type = tokenToDataType(token); + } else if (statement.instruction.b.addressMode == ir::PTXOperand::Register + && ir::PTXOperand::relaxedValid(tokenToDataType(token), + statement.instruction.b.type)) { + statement.instruction.b.relaxedType = tokenToDataType(token); + } } void PTXParser::State::cvtaTo() @@ -2440,6 +2668,9 @@ namespace parser { switch( token ) { + case TOKEN_S4: return ir::PTXOperand::s4; + case TOKEN_B1: return ir::PTXOperand::b1; + case TOKEN_U4: return ir::PTXOperand::u4; case TOKEN_U8: return ir::PTXOperand::u8; break; case TOKEN_U16: return ir::PTXOperand::u16; break; case TOKEN_U32: return ir::PTXOperand::u32; break; @@ -2454,7 +2685,11 @@ namespace parser case TOKEN_B64: return ir::PTXOperand::b64; break; case TOKEN_PRED: return ir::PTXOperand::pred; break; case TOKEN_F16: return ir::PTXOperand::f16; break; + case TOKEN_F16X2:return ir::PTXOperand::f16x2; break; case TOKEN_F32: return ir::PTXOperand::f32; break; + case TOKEN_BF16: return ir::PTXOperand::bf16; break; + case TOKEN_BF16X2:return ir::PTXOperand::bf16x2; break; + case TOKEN_TF32: return ir::PTXOperand::tf32; break; case TOKEN_F64: return ir::PTXOperand::f64; break; default: { @@ -2512,6 +2747,7 @@ namespace parser if( string == "bfe" ) return ir::PTXInstruction::Bfe; if( string == "bfi" ) return ir::PTXInstruction::Bfi; if( string == "bfind" ) return ir::PTXInstruction::Bfind; + if( string == "bmsk" ) return ir::PTXInstruction::Bmsk; if( string == "bra" ) return ir::PTXInstruction::Bra; if( string == "brev" ) return ir::PTXInstruction::Brev; if( string == "brkpt" ) return ir::PTXInstruction::Brkpt; @@ -2523,18 +2759,24 @@ namespace parser if( string == "cvt" ) return ir::PTXInstruction::Cvt; if( string == "cvta" ) return ir::PTXInstruction::Cvta; if( string == "div" ) return ir::PTXInstruction::Div; + if( string == "dp2a" ) return ir::PTXInstruction::Dp2a; + if( string == "dp4a" ) return ir::PTXInstruction::Dp4a; if( string == "ex2" ) return ir::PTXInstruction::Ex2; if( string == "exit" ) return ir::PTXInstruction::Exit; if( string == "fma" ) return ir::PTXInstruction::Fma; + if( string == "fns" ) return ir::PTXInstruction::Fns; if( string == "isspacep" ) return ir::PTXInstruction::Isspacep; if( string == "ld" ) return ir::PTXInstruction::Ld; if( string == "ldu" ) return ir::PTXInstruction::Ldu; if( string == "lg2" ) return ir::PTXInstruction::Lg2; + if( string == "lop3" ) return ir::PTXInstruction::Lop3; if( string == "mad24" ) return ir::PTXInstruction::Mad24; if( string == "mad" ) return ir::PTXInstruction::Mad; if( string == "madc" ) return ir::PTXInstruction::MadC; + if( string == "mma" ) return ir::PTXInstruction::Mma; if( string == "max" ) return ir::PTXInstruction::Max; if( string == "membar" ) return ir::PTXInstruction::Membar; + if( string == "fence" ) return ir::PTXInstruction::Fence; if( string == "min" ) return ir::PTXInstruction::Min; if( string == "mov" ) return ir::PTXInstruction::Mov; if( string == "mul24" ) return ir::PTXInstruction::Mul24; @@ -2543,6 +2785,8 @@ namespace parser if( string == "not" ) return ir::PTXInstruction::Not; if( string == "pmevent" ) return ir::PTXInstruction::Pmevent; if( string == "popc" ) return ir::PTXInstruction::Popc; + if( string == "prefetch" ) return ir::PTXInstruction::Prefetch; + if( string == "prefetchu" ) return ir::PTXInstruction::Prefetchu; if( string == "prmt" ) return ir::PTXInstruction::Prmt; if( string == "or" ) return ir::PTXInstruction::Or; if( string == "rcp" ) return ir::PTXInstruction::Rcp; @@ -2568,6 +2812,8 @@ namespace parser if( string == "sust" ) return ir::PTXInstruction::Sust; if( string == "sured" ) return ir::PTXInstruction::Sured; if( string == "suq" ) return ir::PTXInstruction::Suq; + if( string == "szext" ) return ir::PTXInstruction::Szext; + if( string == "tanh" ) return ir::PTXInstruction::Tanh; if( string == "tex" ) return ir::PTXInstruction::Tex; if( string == "testp" ) return ir::PTXInstruction::TestP; if( string == "tld4" ) return ir::PTXInstruction::Tld4; @@ -2588,8 +2834,10 @@ namespace parser case TOKEN_LO: return ir::PTXInstruction::lo; break; case TOKEN_WIDE: return ir::PTXInstruction::wide; break; case TOKEN_SAT: return ir::PTXInstruction::sat; break; + case TOKEN_SATFINITE: return ir::PTXInstruction::satfinite; break; case TOKEN_RNI: return ir::PTXInstruction::rni; break; case TOKEN_RN: return ir::PTXInstruction::rn; break; + case TOKEN_RNA: return ir::PTXInstruction::rna; break; case TOKEN_RZI: return ir::PTXInstruction::rzi; break; case TOKEN_RZ: return ir::PTXInstruction::rz; break; case TOKEN_RMI: return ir::PTXInstruction::rmi; break; @@ -2598,6 +2846,10 @@ namespace parser case TOKEN_RP: return ir::PTXInstruction::rp; break; case TOKEN_FTZ: return ir::PTXInstruction::ftz; break; case TOKEN_APPROX: return ir::PTXInstruction::approx; break; + case TOKEN_NAN_MODIFIER: return ir::PTXInstruction::nan; break; + case TOKEN_XORSIGN: return ir::PTXInstruction::xorsign; break; + case TOKEN_ABS_MODIFIER: return ir::PTXInstruction::abs; break; + case TOKEN_RELU: return ir::PTXInstruction::relu; break; default: break; } @@ -2714,6 +2966,7 @@ namespace parser case TOKEN_CV: return ir::PTXInstruction::Cv; case TOKEN_WT: return ir::PTXInstruction::Wt; case TOKEN_NC: return ir::PTXInstruction::Nc; + case TOKEN_LU: return ir::PTXInstruction::Lu; default: break; } return ir::PTXInstruction::CacheOperation_Invalid; @@ -2862,11 +3115,28 @@ namespace parser { case TOKEN_CTA: return ir::PTXInstruction::CtaLevel; break; case TOKEN_GL: return ir::PTXInstruction::GlobalLevel; break; + case TOKEN_GPU: return ir::PTXInstruction::GlobalLevel; break; case TOKEN_SYS: return ir::PTXInstruction::SystemLevel; break; default: break; } - - return ir::PTXInstruction::Level_Invalid; + + return ir::PTXInstruction::Level_Invalid; + } + + ir::PTXInstruction::Semantics PTXParser::tokenToSemantics( int token ) + { + switch( token ) + { + case TOKEN_SC: return ir::PTXInstruction::Sc; break; + case TOKEN_ACQ_REL: return ir::PTXInstruction::AcqRel; break; + case TOKEN_ACQUIRE: return ir::PTXInstruction::Acquire; break; + case TOKEN_RELEASE: return ir::PTXInstruction::Release; break; + case TOKEN_RELAXED: return ir::PTXInstruction::Relaxed; break; + case TOKEN_WEAK: return ir::PTXInstruction::Weak; break; + default: break; + } + + return ir::PTXInstruction::Semantics_Invalid; } ir::PTXInstruction::PermuteMode PTXParser::tokenToPermuteMode( int token ) @@ -3000,4 +3270,3 @@ namespace parser } #endif - diff --git a/ocelot/src/parser/ptx.ll b/ocelot/src/parser/ptx.ll index e3d747ef8..fa6ac2527 100644 --- a/ocelot/src/parser/ptx.ll +++ b/ocelot/src/parser/ptx.ll @@ -120,6 +120,8 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return OPCODE_BFE; } "bfind" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_BFIND; } +"bmsk" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_BMSK; } "bra" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_BRA; } "brev" { sstrcpy( yylval->text, yytext, 1024 ); \ @@ -142,12 +144,20 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return OPCODE_CVTA; } "div" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_DIV; } +"dp2a" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_DP2A; } +"dp4a" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_DP4A; } "ex2" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_EX2; } "exit" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_EXIT; } +"fence" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_FENCE; } "fma" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_FMA; } +"fns" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_FNS; } "isspacep" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_ISSPACEP; } "ld" { sstrcpy( yylval->text, yytext, 1024 ); \ @@ -156,6 +166,8 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return OPCODE_LDU; } "lg2" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_LG2; } +"lop3" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_LOP3; } "membar" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_MEMBAR; } "min" { sstrcpy( yylval->text, yytext, 1024 ); \ @@ -166,6 +178,8 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return OPCODE_MADC; } "mad24" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_MAD24; } +"mma" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_MMA; } "max" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_MAX; } "mov" { sstrcpy( yylval->text, yytext, 1024 ); \ @@ -236,6 +250,10 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return OPCODE_SURED; } "suq" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_SUQ; } +"szext" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_SZEXT; } +"tanh" { sstrcpy( yylval->text, yytext, 1024 ); \ + return OPCODE_TANH; } "testp" { sstrcpy( yylval->text, yytext, 1024 ); \ return OPCODE_TESTP; } "tex" { sstrcpy( yylval->text, yytext, 1024 ); \ @@ -296,6 +314,8 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return TOKEN_SAMPLERREF; } ".section" { yylval->value = TOKEN_SECTION; \ return TOKEN_SECTION; } +".shared::cta" { yylval->value = TOKEN_SHARED; \ + return TOKEN_SHARED_CTA; } ".shared" { yylval->value = TOKEN_SHARED; \ return TOKEN_SHARED;} ".shiftamt" { yylval->value = TOKEN_SHIFT_AMOUNT; \ @@ -316,6 +336,13 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") ".cta" { yylval->value = TOKEN_CTA; return TOKEN_CTA; } ".gl" { yylval->value = TOKEN_GL; return TOKEN_GL; } ".sys" { yylval->value = TOKEN_SYS; return TOKEN_SYS; } +".gpu" { yylval->value = TOKEN_GPU; return TOKEN_GPU; } +".sc" { yylval->value = TOKEN_SC; return TOKEN_SC; } +".acq_rel" { yylval->value = TOKEN_ACQ_REL; return TOKEN_ACQ_REL; } +".acquire" { yylval->value = TOKEN_ACQUIRE; return TOKEN_ACQUIRE; } +".release" { yylval->value = TOKEN_RELEASE; return TOKEN_RELEASE; } +".relaxed" { yylval->value = TOKEN_RELAXED; return TOKEN_RELAXED; } +".mmio" { yylval->value = TOKEN_MMIO; return TOKEN_MMIO; } "sm_10" { yylval->value = TOKEN_SM10; return TOKEN_SM10; } @@ -333,6 +360,10 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return TOKEN_SM30; } "sm_35" { yylval->value = TOKEN_SM35; return TOKEN_SM35; } +"sm_50" { yylval->value = TOKEN_SM50; + return TOKEN_SM50; } +"sm_86" { yylval->value = TOKEN_SM86; + return TOKEN_SM86; } "map_f64_to_f32" { yylval->value = TOKEN_MAP_F64_TO_F32; return TOKEN_MAP_F64_TO_F32; } "texmode_independent" { yylval->value = TOKEN_TEXMODE_INDEPENDENT; @@ -342,6 +373,14 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") ".u32" { yylval->value = TOKEN_U32; return TOKEN_U32; } ".s32" { yylval->value = TOKEN_S32; return TOKEN_S32; } +".s4" { yylval->value = TOKEN_S4; return TOKEN_S4; } +".u4" { yylval->value = TOKEN_U4; return TOKEN_U4; } +".m8n8k32" { yylval->value = TOKEN_M8N8K32; return TOKEN_M8N8K32; } +".m16n8k64" { yylval->value = TOKEN_M16N8K64; return TOKEN_M16N8K64; } +".m8n8k128" { yylval->value = TOKEN_M8N8K128; return TOKEN_M8N8K128; } +".m16n8k128" { yylval->value = TOKEN_M16N8K128; return TOKEN_M16N8K128; } +".m16n8k256" { yylval->value = TOKEN_M16N8K256; return TOKEN_M16N8K256; } +".b1" { yylval->value = TOKEN_B1; return TOKEN_B1; } ".s8" { yylval->value = TOKEN_S8; return TOKEN_S8; } ".s16" { yylval->value = TOKEN_S16; return TOKEN_S16; } ".s64" { yylval->value = TOKEN_S64; return TOKEN_S64; } @@ -353,8 +392,12 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") ".b32" { yylval->value = TOKEN_B32; return TOKEN_B32; } ".b64" { yylval->value = TOKEN_B64; return TOKEN_B64; } ".f16" { yylval->value = TOKEN_F16; return TOKEN_F16; } +".f16x2" { yylval->value = TOKEN_F16X2; return TOKEN_F16X2; } ".f64" { yylval->value = TOKEN_F64; return TOKEN_F64; } ".f32" { yylval->value = TOKEN_F32; return TOKEN_F32; } +".bf16" { yylval->value = TOKEN_BF16; return TOKEN_BF16; } +".bf16x2" { yylval->value = TOKEN_BF16X2; return TOKEN_BF16X2; } +".tf32" { yylval->value = TOKEN_TF32; return TOKEN_TF32; } ".pred" { yylval->value = TOKEN_PRED; \ return TOKEN_PRED; } @@ -374,6 +417,9 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") ".geu" { yylval->value = TOKEN_GEU; return TOKEN_GEU; } ".num" { yylval->value = TOKEN_NUM; return TOKEN_NUM; } ".nan" { yylval->value = TOKEN_NAN; return TOKEN_NAN; } +".NaN" { yylval->value = TOKEN_NAN_MODIFIER; return TOKEN_NAN_MODIFIER; } +".xorsign" { yylval->value = TOKEN_XORSIGN; return TOKEN_XORSIGN; } +".abs" { yylval->value = TOKEN_ABS_MODIFIER; return TOKEN_ABS_MODIFIER; } ".and" { yylval->value = TOKEN_AND; return TOKEN_AND; } ".or" { yylval->value = TOKEN_OR; return TOKEN_OR; } @@ -382,6 +428,7 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") ".hi" { yylval->value = TOKEN_HI; return TOKEN_HI; } ".lo" { yylval->value = TOKEN_LO; return TOKEN_LO; } ".rn" { yylval->value = TOKEN_RN; return TOKEN_RN; } +".rna" { yylval->value = TOKEN_RNA; return TOKEN_RNA; } ".rm" { yylval->value = TOKEN_RM; return TOKEN_RM; } ".rz" { yylval->value = TOKEN_RZ; return TOKEN_RZ; } ".rp" { yylval->value = TOKEN_RP; return TOKEN_RP; } @@ -389,8 +436,10 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") ".rmi" { yylval->value = TOKEN_RMI; return TOKEN_RMI; } ".rzi" { yylval->value = TOKEN_RZI; return TOKEN_RZI; } ".rpi" { yylval->value = TOKEN_RPI; return TOKEN_RPI; } +".satfinite" { yylval->value = TOKEN_SATFINITE; return TOKEN_SATFINITE; } ".sat" { yylval->value = TOKEN_SAT; return TOKEN_SAT; } ".ftz" { yylval->value = TOKEN_FTZ; return TOKEN_FTZ; } +".relu" { yylval->value = TOKEN_RELU; return TOKEN_RELU; } ".approx" { yylval->value = TOKEN_APPROX; \ return TOKEN_APPROX; } @@ -494,6 +543,21 @@ LABEL ({IDENTIFIER}{WHITESPACE}":") return TOKEN_RED; } ".sync" { yylval->value = TOKEN_SYNC; \ return TOKEN_SYNC; } +".aligned" { yylval->value = TOKEN_ALIGNED; \ + return TOKEN_ALIGNED; } +".m8n8k4" { yylval->value = TOKEN_M8N8K4; return TOKEN_M8N8K4; } +".m16n8k4" { yylval->value = TOKEN_M16N8K4; \ + return TOKEN_M16N8K4; } +".m16n8k8" { yylval->value = TOKEN_M16N8K8; \ + return TOKEN_M16N8K8; } +".m16n8k16" { yylval->value = TOKEN_M16N8K16; \ + return TOKEN_M16N8K16; } +".m8n8k16" { yylval->value = TOKEN_M8N8K16; \ + return TOKEN_M8N8K16; } +".m16n8k32" { yylval->value = TOKEN_M16N8K32; \ + return TOKEN_M16N8K32; } +".row" { yylval->value = TOKEN_ROW; return TOKEN_ROW; } +".col" { yylval->value = TOKEN_COL; return TOKEN_COL; } ".popc" { yylval->value = TOKEN_POPC; \ return TOKEN_POPC; } @@ -648,4 +712,3 @@ void sstrcpy( char* destination, const char* source, unsigned int max ) #endif /******************************************************************************/ - diff --git a/ocelot/src/parser/ptxgrammar.yy b/ocelot/src/parser/ptxgrammar.yy index 32ab8e38a..8e563686f 100644 --- a/ocelot/src/parser/ptxgrammar.yy +++ b/ocelot/src/parser/ptxgrammar.yy @@ -54,18 +54,20 @@ %token OPCODE_COPYSIGN OPCODE_COS OPCODE_SQRT OPCODE_ADD OPCODE_RSQRT %token OPCODE_MUL OPCODE_SAD OPCODE_SUB OPCODE_EX2 OPCODE_LG2 OPCODE_ADDC %token OPCODE_RCP OPCODE_SIN OPCODE_REM OPCODE_MUL24 OPCODE_MAD24 -%token OPCODE_DIV OPCODE_ABS OPCODE_NEG OPCODE_MIN OPCODE_MAX +%token OPCODE_DIV OPCODE_DP2A OPCODE_DP4A OPCODE_ABS OPCODE_NEG OPCODE_MIN OPCODE_MAX %token OPCODE_MAD OPCODE_MADC OPCODE_SET OPCODE_SETP OPCODE_SELP %token OPCODE_SLCT OPCODE_MOV OPCODE_ST OPCODE_CVT OPCODE_AND OPCODE_XOR -%token OPCODE_OR OPCODE_CVTA OPCODE_ISSPACEP OPCODE_LDU -%token OPCODE_SULD OPCODE_TXQ OPCODE_SUST OPCODE_SURED OPCODE_SUQ +%token OPCODE_OR OPCODE_LOP3 OPCODE_CVTA OPCODE_ISSPACEP OPCODE_LDU +%token OPCODE_SULD OPCODE_TXQ OPCODE_SUST OPCODE_SURED OPCODE_SUQ OPCODE_SZEXT %token OPCODE_BRA OPCODE_CALL OPCODE_RET OPCODE_EXIT OPCODE_TRAP %token OPCODE_BRKPT OPCODE_SUBC OPCODE_TEX OPCODE_LD OPCODE_BARSYNC %token OPCODE_ATOM OPCODE_RED OPCODE_NOT OPCODE_CNOT OPCODE_VOTE %token OPCODE_SHR OPCODE_SHL OPCODE_FMA OPCODE_MEMBAR OPCODE_PMEVENT %token OPCODE_POPC OPCODE_PRMT OPCODE_CLZ OPCODE_BFIND OPCODE_BREV -%token OPCODE_BFI OPCODE_BFE OPCODE_TESTP OPCODE_TLD4 OPCODE_BAR +%token OPCODE_BFI OPCODE_BFE OPCODE_BMSK OPCODE_FNS OPCODE_TANH OPCODE_TESTP OPCODE_TLD4 OPCODE_BAR %token OPCODE_PREFETCH OPCODE_PREFETCHU OPCODE_SHFL OPCODE_SHF +%token OPCODE_MMA +%token OPCODE_FENCE %token PREPROCESSOR_INCLUDE PREPROCESSOR_DEFINE PREPROCESSOR_IF %token PREPROCESSOR_IFDEF PREPROCESSOR_ELSE PREPROCESSOR_ENDIF @@ -77,26 +79,32 @@ %token TOKEN_MAXNREG TOKEN_MAXNTID TOKEN_MAXNCTAPERSM TOKEN_MINNCTAPERSM %token TOKEN_SM11 TOKEN_SM12 TOKEN_SM13 TOKEN_SM20 TOKEN_MAP_F64_TO_F32 -%token TOKEN_SM21 TOKEN_SM10 TOKEN_SM30 TOKEN_SM35 +%token TOKEN_SM21 TOKEN_SM10 TOKEN_SM30 TOKEN_SM35 TOKEN_SM50 TOKEN_SM86 %token TOKEN_TEXMODE_INDEPENDENT TOKEN_TEXMODE_UNIFIED %token TOKEN_CONST TOKEN_GLOBAL TOKEN_LOCAL TOKEN_PARAM TOKEN_PRAGMA TOKEN_PTR -%token TOKEN_REG TOKEN_SHARED TOKEN_TEXREF TOKEN_CTA TOKEN_SURFREF +%token TOKEN_REG TOKEN_SHARED TOKEN_SHARED_CTA TOKEN_TEXREF TOKEN_CTA TOKEN_SURFREF %token TOKEN_GL TOKEN_SYS TOKEN_SAMPLERREF +%token TOKEN_GPU TOKEN_SC TOKEN_ACQ_REL TOKEN_ACQUIRE TOKEN_RELEASE TOKEN_RELAXED +%token TOKEN_MMIO +%token TOKEN_S4 TOKEN_U4 TOKEN_M8N8K32 TOKEN_M16N8K64 +%token TOKEN_B1 TOKEN_M8N8K128 TOKEN_M16N8K128 TOKEN_M16N8K256 %token TOKEN_U32 TOKEN_S32 TOKEN_S8 TOKEN_S16 TOKEN_S64 TOKEN_U8 %token TOKEN_U16 TOKEN_U64 TOKEN_B8 TOKEN_B16 TOKEN_B32 TOKEN_B64 -%token TOKEN_F16 TOKEN_F64 TOKEN_F32 TOKEN_PRED +%token TOKEN_F16 TOKEN_F16X2 TOKEN_F64 TOKEN_F32 TOKEN_BF16 TOKEN_BF16X2 +%token TOKEN_TF32 TOKEN_PRED %token TOKEN_EQ TOKEN_NE TOKEN_LT TOKEN_LE TOKEN_GT TOKEN_GE %token TOKEN_LS TOKEN_HS TOKEN_EQU TOKEN_NEU TOKEN_LTU TOKEN_LEU %token TOKEN_GTU TOKEN_GEU TOKEN_NUM TOKEN_NAN %token TOKEN_HI TOKEN_LO TOKEN_AND TOKEN_OR TOKEN_XOR -%token TOKEN_RN TOKEN_RM TOKEN_RZ TOKEN_RP TOKEN_SAT TOKEN_VOLATILE +%token TOKEN_RN TOKEN_RNA TOKEN_RM TOKEN_RZ TOKEN_RP TOKEN_SAT TOKEN_SATFINITE TOKEN_VOLATILE %token TOKEN_TAIL TOKEN_UNI TOKEN_ALIGN TOKEN_BYTE TOKEN_WIDE TOKEN_CARRY %token TOKEN_RNI TOKEN_RMI TOKEN_RZI TOKEN_RPI %token TOKEN_FTZ TOKEN_APPROX TOKEN_FULL TOKEN_SHIFT_AMOUNT +%token TOKEN_NAN_MODIFIER TOKEN_XORSIGN TOKEN_ABS_MODIFIER TOKEN_RELU %token TOKEN_R TOKEN_G TOKEN_B TOKEN_A TOKEN_L %token TOKEN_TO @@ -127,7 +135,8 @@ %token TOKEN_TRAP TOKEN_CLAMP TOKEN_ZERO TOKEN_WRAP -%token TOKEN_ARRIVE TOKEN_RED TOKEN_POPC TOKEN_SYNC +%token TOKEN_ARRIVE TOKEN_RED TOKEN_POPC TOKEN_SYNC TOKEN_ALIGNED +%token TOKEN_M8N8K4 TOKEN_M16N8K4 TOKEN_M16N8K8 TOKEN_M16N8K16 TOKEN_M8N8K16 TOKEN_M16N8K32 TOKEN_ROW TOKEN_COL %token TOKEN_BALLOT @@ -136,6 +145,8 @@ %token TOKEN_FINITE TOKEN_INFINITE TOKEN_NUMBER TOKEN_NOT_A_NUMBER %token TOKEN_NORMAL TOKEN_SUBNORMAL +%type mmaLayout mmaShape mmaAccumulatorTypeId mmaInputTypeId mmaIntTypeId + %token TOKEN_DECIMAL_CONSTANT %token TOKEN_UNSIGNED_DECIMAL_CONSTANT @@ -260,7 +271,7 @@ singleInitializer : singleList | '{' singleList '}' | '{' singleListSingle '}' | singleListSingle; shaderModel : TOKEN_SM10 | TOKEN_SM11 | TOKEN_SM12 | TOKEN_SM13 | TOKEN_SM20 - | TOKEN_SM21 | TOKEN_SM30 | TOKEN_SM35; + | TOKEN_SM21 | TOKEN_SM30 | TOKEN_SM35 | TOKEN_SM50 | TOKEN_SM86; floatingPointOption : TOKEN_MAP_F64_TO_F32; textureOption: TOKEN_TEXMODE_INDEPENDENT | TOKEN_TEXMODE_UNIFIED; @@ -302,7 +313,8 @@ pointerDataTypeId: TOKEN_U64 | TOKEN_U32; dataTypeId : TOKEN_U8 | TOKEN_U16 | TOKEN_U32 | TOKEN_U64 | TOKEN_S8 | TOKEN_S16 | TOKEN_S32 | TOKEN_S64 | TOKEN_B8 | TOKEN_B16 | TOKEN_B32 - | TOKEN_B64 | TOKEN_F16 | TOKEN_F32 | TOKEN_F64 | TOKEN_PRED; + | TOKEN_B64 | TOKEN_F16 | TOKEN_F32 | TOKEN_F64 + | TOKEN_BF16 | TOKEN_F16X2 | TOKEN_PRED; dataType : dataTypeId { @@ -671,20 +683,67 @@ initializable : externOrVisible initializableAddress opcode : OPCODE_COS | OPCODE_SQRT | OPCODE_ADD | OPCODE_RSQRT | OPCODE_ADDC | OPCODE_MUL | OPCODE_SAD | OPCODE_SUB | OPCODE_EX2 | OPCODE_LG2 | OPCODE_RCP | OPCODE_SIN | OPCODE_REM | OPCODE_MUL24 | OPCODE_MAD24 - | OPCODE_DIV | OPCODE_ABS | OPCODE_NEG | OPCODE_MIN | OPCODE_MAX + | OPCODE_DIV | OPCODE_DP2A | OPCODE_DP4A | OPCODE_ABS | OPCODE_NEG | OPCODE_MIN | OPCODE_MAX | OPCODE_MAD | OPCODE_MADC | OPCODE_SET | OPCODE_SETP | OPCODE_SELP | OPCODE_SLCT | OPCODE_MOV | OPCODE_ST | OPCODE_COPYSIGN | OPCODE_SHFL | OPCODE_SHF | OPCODE_CVT | OPCODE_CVTA | OPCODE_ISSPACEP - | OPCODE_AND | OPCODE_XOR | OPCODE_OR + | OPCODE_AND | OPCODE_XOR | OPCODE_OR | OPCODE_LOP3 | OPCODE_BRA | OPCODE_CALL | OPCODE_RET | OPCODE_EXIT | OPCODE_TRAP | OPCODE_BRKPT | OPCODE_SUBC | OPCODE_TEX | OPCODE_LD | OPCODE_LDU | OPCODE_BARSYNC | OPCODE_SULD | OPCODE_TXQ | OPCODE_SUST | OPCODE_SURED - | OPCODE_SUQ | OPCODE_ATOM | OPCODE_RED | OPCODE_NOT | OPCODE_CNOT - | OPCODE_VOTE | OPCODE_SHR | OPCODE_SHL | OPCODE_MEMBAR | OPCODE_FMA + | OPCODE_SUQ | OPCODE_SZEXT | OPCODE_ATOM | OPCODE_RED | OPCODE_NOT | OPCODE_CNOT + | OPCODE_VOTE | OPCODE_SHR | OPCODE_SHL | OPCODE_MEMBAR | OPCODE_FENCE | OPCODE_FMA | OPCODE_PMEVENT | OPCODE_POPC | OPCODE_CLZ | OPCODE_BFIND | OPCODE_BREV - | OPCODE_BFI | OPCODE_TESTP | OPCODE_TLD4 + | OPCODE_BFI | OPCODE_BMSK | OPCODE_FNS | OPCODE_TANH | OPCODE_TESTP | OPCODE_TLD4 | OPCODE_PREFETCH | OPCODE_PREFETCHU; +mma : OPCODE_MMA TOKEN_SYNC TOKEN_ALIGNED mmaShape mmaLayout mmaLayout + mmaAccumulatorTypeId mmaInputTypeId mmaInputTypeId mmaAccumulatorTypeId + arrayOperand ',' arrayOperand ',' arrayOperand ',' arrayOperand ';' +{ + state.mma( $4, $7, $8, $9, $10, + $5 == TOKEN_COL, $6 == TOKEN_COL ); +}; + +mma : OPCODE_MMA TOKEN_SYNC TOKEN_ALIGNED mmaShape mmaLayout mmaLayout + optionalSatfinite TOKEN_S32 mmaIntTypeId mmaIntTypeId TOKEN_S32 + arrayOperand ',' arrayOperand ',' arrayOperand ',' arrayOperand ';' +{ + state.mma( $4, $8, $9, $10, $11, + $5 == TOKEN_COL, $6 == TOKEN_COL ); +}; + +mma : OPCODE_MMA TOKEN_SYNC TOKEN_ALIGNED mmaShape mmaLayout mmaLayout + TOKEN_F64 TOKEN_F64 TOKEN_F64 TOKEN_F64 optionalFloatRounding + arrayOperand ',' arrayOperand ',' arrayOperand ',' arrayOperand ';' +{ state.mma($4, $7, $8, $9, $10, + $5 == TOKEN_COL, $6 == TOKEN_COL); }; +mmaLayout : TOKEN_ROW | TOKEN_COL; + +mma : OPCODE_MMA TOKEN_SYNC TOKEN_ALIGNED mmaShape mmaLayout mmaLayout + optionalSatfinite TOKEN_S32 TOKEN_B1 TOKEN_B1 TOKEN_S32 mmaBitOp TOKEN_POPC + arrayOperand ',' arrayOperand ',' arrayOperand ',' arrayOperand ';' +{ + state.mma($4, $8, $9, $10, $11, + $5 == TOKEN_COL, $6 == TOKEN_COL); + state.boolean($12); +}; +mmaBitOp : TOKEN_XOR | TOKEN_AND; + +mmaShape : TOKEN_M8N8K128 | TOKEN_M16N8K128 | TOKEN_M16N8K256 | TOKEN_M16N8K64 | TOKEN_M8N8K32 | TOKEN_M8N8K4 | TOKEN_M16N8K4 | TOKEN_M16N8K8 | TOKEN_M16N8K16 | TOKEN_M8N8K16 | TOKEN_M16N8K32; + +mmaAccumulatorTypeId : TOKEN_F16 | TOKEN_F32; + +mmaInputTypeId : TOKEN_F16 | TOKEN_BF16 | TOKEN_TF32; + +mmaIntTypeId : TOKEN_S8 | TOKEN_U8 | TOKEN_S4 | TOKEN_U4; + +optionalSatfinite : TOKEN_SATFINITE +{ + state.modifier( $1 ); +} +| /* empty string */; + uninitializableDeclaration : uninitializable addressableVariablePrefix identifier arrayDimensions ';' { @@ -851,11 +910,11 @@ intRounding : intRoundingToken optionalFloatRounding : floatRounding | /* empty string */; instruction : ftzInstruction2 | ftzInstruction3 | approxInstruction2 - | basicInstruction3 | bfe | bfi | bfind | brev | branch | addOrSub - | addCOrSubC | atom | bar | brkpt | clz | cvt | cvta | isspacep | div | exit - | ld | ldu | mad | mad24 | madc | membar | mov | mul24 | mul | notInstruction + | basicInstruction3 | bfe | bfi | bfind | bmsk | brev | branch | addOrSub + | addCOrSubC | atom | bar | brkpt | clz | cvt | cvta | isspacep | div | dp2a | dp4a | exit + | fence | fns | ld | ldu | lop3 | mad | mad24 | madc | mma | membar | mov | mul24 | mul | notInstruction | pmevent | popc | prefetch | prefetchu | prmt | rcpSqrtInstruction | red - | ret | sad | selp | set | setp | slct | st | suld | suq | sured | sust + | ret | sad | selp | set | setp | slct | st | suld | suq | sured | sust | szext | testp | tex | tld4 | trap | txq | vote | shfl | shf; basicInstruction3Opcode : OPCODE_AND | OPCODE_OR | OPCODE_SHF @@ -867,8 +926,21 @@ basicInstruction3 : basicInstruction3Opcode dataType operand ',' operand ',' state.instruction( $1, $2 ); }; +dp4a : OPCODE_DP4A dataType dataType operand ',' operand ',' operand ',' operand ';' +{ + state.instruction( $1, $2 ); + state.dotType( $3 ); +}; + +dp2a : OPCODE_DP2A hiOrLo dataType dataType operand ',' operand ',' operand ',' operand ';' +{ + state.instruction( $1, $3 ); + state.modifier( $2 ); + state.dotType( $4 ); +}; + approxInstruction2Opcode : OPCODE_RSQRT | OPCODE_SIN | OPCODE_COS | OPCODE_LG2 - | OPCODE_EX2; + | OPCODE_EX2 | OPCODE_TANH; approximate : TOKEN_APPROX { @@ -903,12 +975,38 @@ ftzInstruction2 : ftzInstruction2Opcode optionalFtz dataType operand ',' state.instruction( $1, $3 ); }; +ftzInstruction2 : ftzInstruction2Opcode TOKEN_BF16X2 operand ',' + operand ';' +{ + state.instruction( $1, $2 ); +}; + ftzInstruction3Opcode : OPCODE_MAX | OPCODE_MIN; -ftzInstruction3 : ftzInstruction3Opcode optionalFtz dataType operand ',' - operand ',' operand ';' +optionalNanModifier : TOKEN_NAN_MODIFIER { - state.instruction( $1, $3 ); + state.modifier( $1 ); +} +| /* empty string */; + +optionalXorsignAbs : TOKEN_XORSIGN TOKEN_ABS_MODIFIER +{ + state.modifier( $1 ); + state.modifier( $2 ); +} +| /* empty string */; + +minMaxDataTypeId : dataTypeId | TOKEN_BF16X2; + +minMaxDataType : minMaxDataTypeId +{ + state.dataType( $1 ); +}; + +ftzInstruction3 : ftzInstruction3Opcode optionalFtz optionalNanModifier + optionalXorsignAbs minMaxDataType operand ',' operand ',' operand ';' +{ + state.instruction( $1, $5 ); }; optionalUni : /* empty string */ @@ -1028,16 +1126,32 @@ atomModifier: /* empty string */ state.addressSpace(TOKEN_GLOBAL); } -atom : OPCODE_ATOM atomModifier atomicOperation dataType operand ',' '[' +atomicSemantics : TOKEN_RELAXED { state.semantics( $1 ); } + | TOKEN_ACQUIRE { state.semantics( $1 ); } + | TOKEN_RELEASE { state.semantics( $1 ); } + | TOKEN_ACQ_REL { state.semantics( $1 ); } + ; + +optionalAtomicSemantics : atomicSemantics + | /* empty */ { state.semantics( TOKEN_RELAXED ); } + ; + +atomicScope : fenceScopeType { state.scope( $1 ); }; + +optionalAtomicScope : atomicScope | /* empty */; + +atom : OPCODE_ATOM optionalAtomicSemantics optionalAtomicScope atomModifier + atomicOperation dataType operand ',' '[' memoryOperand ']' ',' operand ';' { - state.instruction( $1, $4 ); + state.instruction( $1, $6 ); }; -atom : OPCODE_ATOM atomModifier atomicOperation dataType operand ',' '[' +atom : OPCODE_ATOM optionalAtomicSemantics optionalAtomicScope atomModifier + atomicOperation dataType operand ',' '[' memoryOperand ']' ',' operand ',' operand ';' { - state.instruction( $1, $4 ); + state.instruction( $1, $6 ); }; shiftAmount : TOKEN_SHIFT_AMOUNT @@ -1062,16 +1176,32 @@ bfi : OPCODE_BFI dataType operand ',' operand ',' operand state.instruction( $1, $2 ); }; -bfind : OPCODE_BFIND shiftAmount dataType operand ',' operand ';' +lop3 : OPCODE_LOP3 TOKEN_B32 operand ',' operand ',' operand ',' operand + ',' operand ';' { - state.instruction( $1, $3 ); + state.lop3(); }; -barrierOperation : TOKEN_ARRIVE | TOKEN_RED | TOKEN_SYNC +lop3BoolOperator : TOKEN_AND { state.boolean( $1 ); } + | TOKEN_OR { state.boolean( $1 ); } + ; + +lop3 : OPCODE_LOP3 lop3BoolOperator TOKEN_B32 operand '|' operand ',' operand + ',' operand ',' operand ',' operand ',' operand ';' { - state.barrierOperation( $1, @1 ); + state.lop3(); }; +bfind : OPCODE_BFIND shiftAmount dataType operand ',' operand ';' +{ + state.instruction( $1, $3 ); +}; + +barrierOperation : TOKEN_ARRIVE { state.barrierOperation( $1, @1 ); } + | TOKEN_RED { state.barrierOperation( $1, @1 ); } + | TOKEN_SYNC { state.barrierOperation( $1, @1 ); } + ; + optionalBarrierOperator : reductionOperation dataType | /* or nothing */ ; operandSequence: operand operandSequence | /* empty */ ; @@ -1096,6 +1226,11 @@ clz : OPCODE_CLZ dataType operand ',' operand ';' state.instruction( $1, $2 ); }; +fns : OPCODE_FNS TOKEN_B32 operand ',' operand ',' operand ',' operand ';' +{ + state.instruction( $1, $2 ); +}; + floatRoundingModifier : floatRounding { state.modifier( $1 ); @@ -1108,12 +1243,33 @@ intRoundingModifier : intRounding cvtRoundingModifier : intRoundingModifier | floatRoundingModifier; -cvtModifier : cvtRoundingModifier optionalFtz sat; -cvtModifier : cvtRoundingModifier optionalFtz; -cvtModifier : optionalFtz sat; -cvtModifier : optionalFtz; +cvtRoundingModifier : TOKEN_RNA +{ + state.modifier( $1 ); +}; + +optionalCvtRounding : cvtRoundingModifier | /* empty string */; +optionalRelu : TOKEN_RELU +{ + state.modifier( $1 ); +}; +optionalRelu : /* empty string */; + +cvtModifier : optionalCvtRounding optionalFtz optionalSaturate optionalRelu; + +cvtDataTypeId : dataTypeId | TOKEN_BF16X2 | TOKEN_TF32; +cvtDataType : cvtDataTypeId +{ + state.dataType( $1 ); +}; + +cvt : OPCODE_CVT cvtModifier cvtDataType cvtDataType operand ',' operand ';' +{ + state.instruction( $1, $3 ); + state.relaxedConvert( $4, @1 ); +}; -cvt : OPCODE_CVT cvtModifier dataType dataType operand ',' operand ';' +cvt : OPCODE_CVT cvtModifier cvtDataType cvtDataType operand ',' operand ',' operand ';' { state.instruction( $1, $3 ); state.relaxedConvert( $4, @1 ); @@ -1147,15 +1303,12 @@ divApproxModifier : TOKEN_APPROX optionalFtz state.modifier($1); }; -divRnModifier : TOKEN_RN optionalFtz -{ - state.modifier($1); -}; +divRoundingModifier : floatRounding optionalFtz; -divRnModifier : /* empty string */; +divRoundingModifier : /* empty string */; divModifier : divFullModifier | divApproxModifier - | divRnModifier; + | divRoundingModifier; div : OPCODE_DIV divModifier dataType operand ',' operand ',' operand ';' { @@ -1167,7 +1320,10 @@ exit : OPCODE_EXIT ';' state.instruction( $1 ); }; -isspacep : OPCODE_ISSPACEP addressSpace operand ',' operand ';' +isspacepAddressSpace : addressSpace + | TOKEN_SHARED_CTA { state.addressSpace( TOKEN_SHARED ); }; + +isspacep : OPCODE_ISSPACEP isspacepAddressSpace operand ',' operand ';' { state.instruction( $1, TOKEN_U32 ); } @@ -1184,8 +1340,41 @@ optionalVolatile : /* empty string */ state.volatileFlag( false ); }; -ldModifier : optionalVolatile optionalAddressSpace optionalCacheOperation - optionalInstructionVectorType; +mmioLdSemantics : TOKEN_ACQUIRE { state.semantics( $1 ); } + | TOKEN_RELAXED { state.semantics( $1 ); } + ; + +mmioStSemantics : TOKEN_RELAXED { state.semantics( $1 ); } + | TOKEN_RELEASE { state.semantics( $1 ); } + ; + +ldOrdering : volatileModifier { state.semantics( TOKEN_WEAK ); } + | TOKEN_WEAK { state.semantics( $1 ); state.volatileFlag( false ); } + | TOKEN_RELAXED fenceScopeType { state.semantics( $1 ); + state.scope( $2 ); state.volatileFlag( false ); } + | TOKEN_ACQUIRE fenceScopeType { state.semantics( $1 ); + state.scope( $2 ); state.volatileFlag( false ); } + | TOKEN_MMIO mmioLdSemantics TOKEN_SYS { state.mmio( true ); + state.scope( $3 ); state.volatileFlag( false ); } + | /* empty */ { state.semantics( TOKEN_WEAK ); state.volatileFlag( false ); } + ; + +stOrdering : volatileModifier { state.semantics( TOKEN_WEAK ); } + | TOKEN_WEAK { state.semantics( $1 ); state.volatileFlag( false ); } + | TOKEN_RELAXED fenceScopeType { state.semantics( $1 ); + state.scope( $2 ); state.volatileFlag( false ); } + | TOKEN_RELEASE fenceScopeType { state.semantics( $1 ); + state.scope( $2 ); state.volatileFlag( false ); } + | TOKEN_MMIO mmioStSemantics TOKEN_SYS { state.mmio( true ); + state.scope( $3 ); state.volatileFlag( false ); } + | /* empty */ { state.semantics( TOKEN_WEAK ); state.volatileFlag( false ); } + ; + +ldModifier : ldOrdering optionalAddressSpace optionalCacheOperation + optionalInstructionVectorType { state.finalizeMmioAddressSpace(); }; + +stModifier : stOrdering optionalAddressSpace optionalStoreCacheOperation + optionalInstructionVectorType { state.finalizeMmioAddressSpace(); }; ld : OPCODE_LD ldModifier dataType arrayOperand ',' '[' memoryOperand ']' ';' { @@ -1261,12 +1450,43 @@ membar : OPCODE_MEMBAR membarSpace ';' state.instruction( $1 ); }; +fenceSemanticsType : TOKEN_SC | TOKEN_ACQ_REL | TOKEN_ACQUIRE | TOKEN_RELEASE; + +fenceSemantics : fenceSemanticsType +{ + state.semantics( $1 ); +}; + +optionalFenceSemantics : fenceSemantics | /* empty */; + +fenceScopeType : TOKEN_CTA | TOKEN_GPU | TOKEN_SYS; + +fenceScope : fenceScopeType +{ + state.level( $1 ); +}; + +fence : OPCODE_FENCE optionalFenceSemantics fenceScope ';' +{ + state.instruction( $1 ); +}; + movIndexedOperand : identifier '[' TOKEN_DECIMAL_CONSTANT ']' { state.indexedOperand( $1, @1, $3 ); }; -movSourceOperand : arrayOperand | offsetAddressableOperand | movIndexedOperand; +movVectorOperand : '{' operand ',' operand '}' +{ + state.vectorOperand(2); +}; + +movVectorOperand : '{' operand ',' operand ',' operand ',' operand '}' +{ + state.vectorOperand(4); +}; + +movSourceOperand : operand | offsetAddressableOperand | movIndexedOperand | movVectorOperand; mov : OPCODE_MOV dataType arrayOperand ',' movSourceOperand ';' { @@ -1320,10 +1540,9 @@ permuteMode : /* empty string */ state.defaultPermute(); }; -cacheLevel : TOKEN_L1 | TOKEN_L2 -{ - state.cacheLevel( $1 ); -}; +cacheLevel : TOKEN_L1 { state.cacheLevel( $1 ); } + | TOKEN_L2 { state.cacheLevel( $1 ); } + ; prefetch : OPCODE_PREFETCH addressSpace cacheLevel '[' memoryOperand ']' ';' { @@ -1347,7 +1566,7 @@ rcpSqrtModifier : TOKEN_APPROX optionalFtz }; rcpSqrtModifier : /* empty string */; -rcpSqrtModifier : TOKEN_RN optionalFtz +rcpSqrtModifier : floatRoundingToken optionalFtz { state.modifier( $1 ); }; @@ -1368,10 +1587,19 @@ reductionOperation : reductionOperationId state.reduction( $1 ); }; -red : OPCODE_RED addressSpace reductionOperation dataType operand ',' +reductionSemantics : TOKEN_RELAXED { state.semantics( $1 ); } + | TOKEN_RELEASE { state.semantics( $1 ); } + ; + +optionalReductionSemantics : reductionSemantics + | /* empty */ { state.semantics( TOKEN_RELAXED ); } + ; + +red : OPCODE_RED optionalReductionSemantics optionalAtomicScope addressSpace + reductionOperation dataType operand ',' operand ';' { - state.instruction( $1, $4 ); + state.instruction( $1, $6 ); }; ret : OPCODE_RET optionalUni ';' @@ -1465,6 +1693,18 @@ shf : OPCODE_SHF shfDirection shfMode TOKEN_B32 operand ',' operand ',' operand state.shiftMode( $3 ); }; +szext : OPCODE_SZEXT shfMode dataType operand ',' operand ',' operand ';' +{ + state.instruction( $1, $3 ); + state.shiftMode( $2 ); +}; + +bmsk : OPCODE_BMSK shfMode TOKEN_B32 operand ',' operand ',' operand ';' +{ + state.instruction( $1, $3 ); + state.shiftMode( $2 ); +}; + shuffleModifierId : TOKEN_UP | TOKEN_DOWN | TOKEN_BFLY | TOKEN_IDX; shuffleModifier : shuffleModifierId @@ -1485,7 +1725,7 @@ slct : OPCODE_SLCT optionalFtz dataType dataType operand ',' operand ',' state.convertC( $4, @1 ); }; -st : OPCODE_ST ldModifier dataType '[' memoryOperand ']' ',' arrayOperand ';' +st : OPCODE_ST stModifier dataType '[' memoryOperand ']' ',' arrayOperand ';' { state.instruction( $1, $3 ); }; @@ -1535,13 +1775,17 @@ tld4 : OPCODE_TLD4 colorComponent TOKEN_2D TOKEN_V4 dataType dataType // Surface sampling // -surfaceQuery : TOKEN_WIDTH | TOKEN_HEIGHT | TOKEN_DEPTH - | TOKEN_CHANNEL_DATA_TYPE | TOKEN_CHANNEL_ORDER | TOKEN_NORMALIZED_COORDS - | TOKEN_FILTER_MODE | TOKEN_ADDR_MODE_0 | TOKEN_ADDR_MODE_1 - | TOKEN_ADDR_MODE_2 -{ - state.surfaceQuery( $1 ); -}; +surfaceQuery : TOKEN_WIDTH { state.surfaceQuery( $1 ); } + | TOKEN_HEIGHT { state.surfaceQuery( $1 ); } + | TOKEN_DEPTH { state.surfaceQuery( $1 ); } + | TOKEN_CHANNEL_DATA_TYPE { state.surfaceQuery( $1 ); } + | TOKEN_CHANNEL_ORDER { state.surfaceQuery( $1 ); } + | TOKEN_NORMALIZED_COORDS { state.surfaceQuery( $1 ); } + | TOKEN_FILTER_MODE { state.surfaceQuery( $1 ); } + | TOKEN_ADDR_MODE_0 { state.surfaceQuery( $1 ); } + | TOKEN_ADDR_MODE_1 { state.surfaceQuery( $1 ); } + | TOKEN_ADDR_MODE_2 { state.surfaceQuery( $1 ); } + ; txq : OPCODE_TXQ surfaceQuery dataType operand ',' '[' operand ']' ';' { @@ -1555,22 +1799,32 @@ suq : OPCODE_SUQ surfaceQuery dataType operand ',' '[' operand ']' ';' state.surfaceQuery( $2 ); }; -cacheOperation : TOKEN_CA | TOKEN_CG | TOKEN_CS | TOKEN_CV | TOKEN_NC -{ - state.cacheOperation( $1 ); -}; +cacheOperation : TOKEN_CA { state.cacheOperation( $1 ); } + | TOKEN_CG { state.cacheOperation( $1 ); } + | TOKEN_CS { state.cacheOperation( $1 ); } + | TOKEN_CV { state.cacheOperation( $1 ); } + | TOKEN_NC { state.cacheOperation( $1 ); } + | TOKEN_LU { state.cacheOperation( $1 ); } + ; optionalCacheOperation : cacheOperation | /* empty */; -clampOperation : TOKEN_CLAMP | TOKEN_ZERO | TOKEN_TRAP -{ - state.clampOperation( $1 ); -}; +storeCacheOperation : TOKEN_WB { state.cacheOperation( $1 ); } + | TOKEN_CG { state.cacheOperation( $1 ); } + | TOKEN_CS { state.cacheOperation( $1 ); } + | TOKEN_WT { state.cacheOperation( $1 ); } + ; -formatMode : TOKEN_B | TOKEN_P -{ - state.formatMode( $1 ); -}; +optionalStoreCacheOperation : storeCacheOperation | /* empty */; + +clampOperation : TOKEN_CLAMP { state.clampOperation( $1 ); } + | TOKEN_ZERO { state.clampOperation( $1 ); } + | TOKEN_TRAP { state.clampOperation( $1 ); } + ; + +formatMode : TOKEN_B { state.formatMode( $1 ); } + | TOKEN_P { state.formatMode( $1 ); } + ; suld : OPCODE_SULD formatMode geometry optionalCacheOperation instructionVectorType dataType clampOperation arrayOperand ','