summaryrefslogtreecommitdiff
path: root/xmake/rules/cuda/gencodes/xmake.lua
blob: bf2a52d2725e56f15495f65bfd95488c95b59430 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
--!A cross-platform build utility based on Lua
--
-- Licensed under the Apache License, Version 2.0 (the "License");
-- you may not use this file except in compliance with the License.
-- You may obtain a copy of the License at
--
--     http://www.apache.org/licenses/LICENSE-2.0
--
-- Unless required by applicable law or agreed to in writing, software
-- distributed under the License is distributed on an "AS IS" BASIS,
-- WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
-- See the License for the specific language governing permissions and
-- limitations under the License.
--
-- Copyright (C) 2015-present, Xmake Open Source Community.
--
-- @author      ruki
-- @file        xmake.lua
--

-- define rule: gencodes
rule("cuda.gencodes")

    -- add cuda `-gencode` flags to target
    --
    -- the gpu arch format syntax
    -- - compute_xx                   --> `-gencode arch=compute_xx,code=compute_xx`
    -- - sm_xx                        --> `-gencode arch=compute_xx,code=sm_xx`
    -- - sm_xx,sm_yy                  --> `-gencode arch=compute_xx,code=[sm_xx,sm_yy]`
    -- - compute_xx,sm_yy             --> `-gencode arch=compute_xx,code=sm_yy`
    -- - compute_xx,sm_yy,sm_zz       --> `-gencode arch=compute_xx,code=[sm_yy,sm_zz]`
    -- - native                       --> match the fastest cuda device on current host,
    --                                    eg. for a Tesla P100, `-gencode arch=compute_60,code=sm_60` will be added,
    --                                    if no available device is found, no `-gencode` flags will be added
    --                                    @seealso xmake/modules/lib/detect/find_cudadevices
    --
    on_config(function (target)

        -- imports
        import("core.platform.platform")
        import("lib.detect.find_cudadevices")
        import("core.base.hashset")

        -- sm_20 and compute_20 is supported until CUDA 8
        -- sm_30 and compute_30 is supported until CUDA 10
        -- sm_37 and compute_37 is supported until CUDA 11
        -- sm_72 and compute_72 is supported until CUDA 12
        local known_v_archs = hashset.of(20, 30, 32, 35, 37, 50, 52, 53, 60, 61, 62, 70, 72, 75, 80, 86, 87, 89, 90, 100, 103, 110, 120, 121)
        local known_r_archs = hashset.of(20, 30, 32, 35, 37, 50, 52, 53, 60, 61, 62, 70, 72, 75, 80, 86, 87, 89, 90, 100, 103, 110, 120, 121)

        local function nf_cugencode(archs)
            if type(archs) ~= "string" then
                return nil
            end
            archs = archs:trim():lower()
            if archs == "native" then
                local cuda_envs
                for _, toolchain_inst in ipairs(target:toolchains()) do
                    if toolchain_inst:name() == "cuda" then
                        cuda_envs = toolchain_inst:runenvs()
                        break
                    end
                end
                local device = find_cudadevices({skip_compute_mode_prohibited = true, order_by_flops = true, envs = cuda_envs, plat = target:plat(), arch = target:arch()})[1]
                if device then
                    return nf_cugencode("sm_" .. device.major .. device.minor)
                end
                return nil
            end

            local v_arch = nil
            local r_archs = {}

            -- full legal value list could be found in nvcc docs
            -- https://docs.nvidia.com/cuda/cuda-compiler-driver-nvcc/index.html#gpu-name-gpuname-arch
            -- examples: sm_75, compute_75, sm_90a, compute_100, sm_100f, compute_100a, etc.
            -- For robustness, xmake support a unoffical format: sm75, compute75, etc. New version still support it.
            -- examples: sm75, compute75, sm90a, compute100, sm100f, compute100a, etc.
            local function parse_arch(value, prefix, know_list)
                if not value:startswith(prefix) then
                    return nil
                end
                local arch_str = value:sub(#prefix + 1)
                if arch_str:startswith("_") then
                    arch_str = arch_str:sub(2)
                end

                -- a legal arch_str should be like: 75, 90a, 100f, etc.
                local arch_ver, suffix = arch_str:match("^(%d+)([af]?)$")
                local arch = tonumber(arch_ver)
                if arch == nil then
                    raise("unknown architecture: " .. value)
                end
                if suffix == 'a' and arch < 90 then
                    raise("unknown architecture: " .. prefix .. "_" .. arch .. suffix)
                end
                if suffix == 'f' and arch < 100 then
                    raise("unknown architecture: " .. prefix .. "_" .. arch .. suffix)
                end
                if not know_list:has(arch) then
                    if arch <= table.maxn(know_list:data()) then
                        raise("unknown architecture: " .. prefix .. "_" .. arch)
                    else
                        utils.warning("unknown architecture: " .. prefix .. "_" .. arch)
                    end
                end
                if suffix and #suffix > 0 then
                    return arch .. suffix
                end
                return arch
            end

            for _, v in ipairs(archs:split(',')) do
                local arch = v:trim()
                local temp_r_arch = parse_arch(arch, "sm", known_r_archs)
                if temp_r_arch then
                    table.insert(r_archs, temp_r_arch)
                end

                local temp_v_arch = parse_arch(arch, "compute", known_v_archs)
                if temp_v_arch then
                    if v_arch ~= nil then
                        raise("more than one virtual architecture is defined in one gpu gencode option: compute_" .. v_arch .. " and compute_" .. temp_v_arch)
                    end
                    v_arch = temp_v_arch
                end
                if not (temp_r_arch or temp_v_arch) then
                    raise("unknown architecture: " .. arch)
                end
            end

            if v_arch == nil and #r_archs == 0 then
                return nil
            end

            if #r_archs == 0 then
                return {
                    clang = "--cuda-gpu-arch=sm_" .. v_arch,
                    nvcc = "-gencode arch=compute_" .. v_arch .. ",code=compute_" .. v_arch
                }
            end

            if v_arch then
                table.insert(r_archs, v_arch)
            else
                v_arch = r_archs[1]
                local v_arch_ver = type(v_arch) == "string" and tonumber(v_arch:match("^(%d+)")) or v_arch
                for i = 2, #r_archs do
                    local r_arch = r_archs[i]
                    local r_arch_ver = type(r_arch) == "string" and tonumber(r_arch:match("^(%d+)")) or r_arch
                    if r_arch_ver < v_arch_ver then
                        v_arch = r_arch
                        v_arch_ver = r_arch_ver
                    elseif r_arch_ver == v_arch_ver and type(r_arch) == "number" and type(v_arch) == "string" then
                        v_arch = r_arch
                    end
                end
            end
            r_archs = table.unique(r_archs)

            local clang_flags = {}
            for _, r_arch in ipairs(r_archs) do
                table.insert(clang_flags, "--cuda-gpu-arch=sm_" .. r_arch)
            end

            local nvcc_flags = nil
            if #r_archs == 1 then
                nvcc_flags = "-gencode arch=compute_" .. v_arch .. ",code=sm_" .. r_archs[1]
            else
                nvcc_flags = "-gencode arch=compute_" .. v_arch .. ",code=[sm_" .. table.concat(r_archs, ",sm_") .. "]"
            end

            return { clang = clang_flags, nvcc = nvcc_flags }
        end

        local cugencodes = table.wrap(target:get("cugencodes"))
        for _, opt in ipairs(target:orderopts()) do
            table.join2(cugencodes, opt:get("cugencodes"))
        end
        for _, v in ipairs(cugencodes) do
            local flag = nf_cugencode(v)
            if flag then
                if target:has_tool("cu", "nvcc") then
                    target:add("cuflags", flag.nvcc)
                else
                    target:add("cuflags", flag.clang)
                end
                target:add("culdflags", flag.nvcc)
            end
        end
    end)
rule_end()