【问题标题】:Is it possible to express a field type dependent on another parametric type?是否可以表示依赖于另一种参数类型的字段类型?
【发布时间】:2023-03-30 10:05:01
【问题描述】:

作为一个人工示例,假设我有一个参数结构,其中T <: AbstractFloat

mutable struct Summary{T<:AbstractFloat}

    count
    sum::T

end

我想在T === Float16 时将count 字段输入为UInt16,在T === Float32 时输入为UInt32,在所有其他情况下输入为UInt64

我目前的方法是为 count 字段使用联合类型 Union{UInt16, UInt32, UInt64}

module SummaryStats

export Summary, avg

const CounterType = Union{UInt16, UInt32, UInt64}

mutable struct Summary{T<:AbstractFloat}

    count::CounterType
    sum::T
    # explicitly typed no-arg constructor
    Summary{T}() where {T<:AbstractFloat} = new(_counter(T), zero(T))
end

# untyped no-arg constructor defaults to Float64
Summary() = Summary{Float64}()

function avg(summary::Summary{T})::T where {T <: AbstractFloat}
    if summary.count > zero(_counter(typeof(T)))
        summary.sum / summary.count
    else
        zero(T)
    end
end

# internal helper functions, not exported
Base.@pure _counter(::Type{Float16})::UInt16 = UInt16(0)
Base.@pure _counter(::Type{Float32})::UInt32 = UInt32(0)
Base.@pure _counter(::DataType)::UInt64 = UInt64(0)

end # module

这似乎可行,但显然@code_warntypecount 字段的联合类型不满意。

我想知道是否有可能根据上面列出的规则以某种方式计算出正确的具体类型?

【问题讨论】:

    标签: julia


    【解决方案1】:

    "outer-only" constructors 主要用于这些用例:

    julia> const CounterType = Union{UInt16, UInt32, UInt64}
    Union{UInt16, UInt32, UInt64}
    
    julia> mutable struct Summary{T<:AbstractFloat, S<:CounterType}
               count::S
               sum::T
               function Summary{T}() where {T<:AbstractFloat}
                   S = T === Float16 ? UInt16 : 
                       T === Float32 ? UInt32 :
                       T === Float64 ? UInt64 : throw(ArgumentError("unexpected type: $(T)!"))
                   new{T,S}(zero(S), zero(T))
               end
           end
    
    julia> Summary() = Summary{Float64}()
    Summary
    
    julia> function avg(summary::Summary{T})::T where {T <: AbstractFloat}
           if summary.count > zero(summary.count)
               summary.sum / summary.count
           else
               zero(T)
           end
       end
    avg (generic function with 1 method)
    
    julia> avg(Summary())
    0.0
    
    julia> @code_warntype avg(Summary())
    Body::Float64
    1 ─ %1 = (Base.getfield)(summary, :count)::UInt64
    │   %2 = (Base.ult_int)(0x0000000000000000, %1)::Bool
    └──      goto #3 if not %2
    2 ─ %4 = (Base.getfield)(summary, :sum)::Float64
    │   %5 = (Base.getfield)(summary, :count)::UInt64
    │   %6 = (Base.uitofp)(Float64, %5)::Float64
    │   %7 = (Base.div_float)(%4, %6)::Float64
    └──      return %7
    3 ─      return 0.0
    
    julia> @code_warntype avg(Summary{Float32}())
    Body::Float32
    1 ─ %1 = (Base.getfield)(summary, :count)::UInt32
    │   %2 = (Base.ult_int)(0x00000000, %1)::Bool
    └──      goto #3 if not %2
    2 ─ %4 = (Base.getfield)(summary, :sum)::Float32
    │   %5 = (Base.getfield)(summary, :count)::UInt32
    │   %6 = (Base.uitofp)(Float32, %5)::Float32
    │   %7 = (Base.div_float)(%4, %6)::Float32
    └──      return %7
    3 ─      return 0.0f0
    

    【讨论】:

    • +1 整洁! IIUC,Float32 avg() 调用需要在 UInt32 和 UInt64 之间转换,以便与 0 ((Core.zext_int)(Core.UInt64, %1)::UInt64) 进行比较。出于好奇,这可以避免吗?
    • 是的。 summary.count &gt; zero(summary.count)
    猜你喜欢
    • 1970-01-01
    • 2023-03-21
    • 2019-02-08
    • 2014-06-21
    • 1970-01-01
    • 2020-08-21
    • 2020-11-20
    • 2022-11-24
    • 1970-01-01
    相关资源
    最近更新 更多