diff --git a/src/schema.jl b/src/schema.jl index d05b0ad7..c8b7e453 100644 --- a/src/schema.jl +++ b/src/schema.jl @@ -56,8 +56,8 @@ Base.keys(schema::Schema) = keys(schema.schema) Base.haskey(schema::Schema, key) = haskey(schema.schema, key) """ - schema([terms::AbstractVector{<:AbstractTerm}, ]data, hints::Dict{Symbol}) - schema(term::AbstractTerm, data, hints::Dict{Symbol}) + schema([terms::AbstractVector{<:AbstractTerm}, ]data, hints::Dict{Symbol}; statistics=true) + schema(term::AbstractTerm, data, hints::Dict{Symbol}; statistics=true) Compute all the invariants necessary to fit a model with `terms`. A schema is a dict that maps `Term`s to their concrete instantiations (either `CategoricalTerm`s or @@ -67,6 +67,15 @@ the appropriate term type will be guessed based on the data type from the data c numeric data is assumed to be continuous, and any non-numeric data is assumed to be categorical. +By default, creating a `ContinuousTerm` computes the summary statistics (mean, variance, +and extrema) of the corresponding column, which requires several passes over the data. +Pass `statistics=false` to skip this computation and fill those fields with `NaN` +placeholders instead. Use this only when nothing downstream reads the summary statistics +(they are needed, e.g., by packages that use them as default centering or scaling values). +The keyword only affects the terms that StatsModels itself concretizes as continuous +(numeric columns without a hint, and columns with a `ContinuousTerm` hint); terms with +other hints are created exactly as they would be otherwise. + Returns a [`StatsModels.Schema`](@ref), which is a wrapper around a `Dict` mapping `Term`s to their concrete instantiations (`ContinuousTerm` or `CategoricalTerm`). @@ -114,21 +123,39 @@ julia> sch[term(:y)] y(continuous) ``` """ -schema(data, hints=Dict{Symbol,Any}()) = schema(columntable(data), hints) -schema(dt::D, hints=Dict{Symbol,Any}()) where {D<:ColumnTable} = - schema(Term.(collect(fieldnames(D))), dt, hints) -schema(ts::AbstractVector{<:AbstractTerm}, data, hints::Dict{Symbol}) = - schema(ts, columntable(data), hints) +schema(data, hints=Dict{Symbol,Any}(); kwargs...) = schema(columntable(data), hints; kwargs...) +schema(dt::D, hints=Dict{Symbol,Any}(); kwargs...) where {D<:ColumnTable} = + schema(Term.(collect(fieldnames(D))), dt, hints; kwargs...) +schema(ts::AbstractVector{<:AbstractTerm}, data, hints::Dict{Symbol}; kwargs...) = + schema(ts, columntable(data), hints; kwargs...) # handle hints: -schema(ts::AbstractVector{<:AbstractTerm}, dt::ColumnTable, - hints::Dict{Symbol}=Dict{Symbol,Any}()) = - sch = Schema(t=>concrete_term(t, dt, hints) for t in ts) +function schema(ts::AbstractVector{<:AbstractTerm}, dt::ColumnTable, + hints::Dict{Symbol}=Dict{Symbol,Any}(); statistics::Bool=true) + sch = Schema() + for t in ts + # route `statistics=false` only into the paths where the statistics are + # computed by StatsModels itself, so that `concrete_term` methods that + # other packages define for their own hint types never see the keyword + if !statistics && t isa Term + msg = checkcol(dt, t.sym) + msg != "" && throw(ArgumentError(msg)) + col = getproperty(dt, t.sym) + hint = get(hints, t.sym, nothing) + if hint === ContinuousTerm || (hint === nothing && col isa AbstractVector{<:Number}) + sch.schema[t] = concrete_term(t, col, ContinuousTerm; statistics=false) + continue + end + end + sch.schema[t] = concrete_term(t, dt, hints) + end + return sch +end -schema(f::TermOrTerms, data, hints::Dict{Symbol}) = - schema(filter(needs_schema, terms(f)), data, hints) +schema(f::TermOrTerms, data, hints::Dict{Symbol}; kwargs...) = + schema(filter(needs_schema, terms(f)), data, hints; kwargs...) -schema(f::TermOrTerms, data) = schema(f, data, Dict{Symbol,Any}()) +schema(f::TermOrTerms, data; kwargs...) = schema(f, data, Dict{Symbol,Any}(); kwargs...) """ concrete_term(t::Term, data[, hint]) @@ -146,6 +173,11 @@ If no hint is provided (or `hint==nothing`), the `eltype` of the data is used: `Number`s are assumed to be continuous, and all others are assumed to be categorical. +The `ContinuousTerm`-hint method additionally accepts a `statistics::Bool=true` +keyword: `concrete_term(t, xs, ContinuousTerm, statistics=false)` skips computing +the summary statistics (mean, variance, and extrema) stored in the term, filling +those fields with `NaN` placeholders instead (see [`schema`](@ref)). + # Example ```jldoctest @@ -198,10 +230,18 @@ concrete_term(t::Term, x, hint::AbstractTerm) = hint concrete_term(t, d, hint) = t concrete_term(t::Term, xs::AbstractVector{<:Number}, ::Nothing) = concrete_term(t, xs, ContinuousTerm) -function concrete_term(t::Term, xs::AbstractVector, ::Type{ContinuousTerm}) - μ, σ2 = StatsBase.mean_and_var(xs) - min, max = extrema(xs) - ContinuousTerm(t.sym, promote(μ, σ2, min, max)...) +function concrete_term(t::Term, xs::AbstractVector, ::Type{ContinuousTerm}; + statistics::Bool=true) + if statistics + μ, σ2 = StatsBase.mean_and_var(xs) + min, max = extrema(xs) + return ContinuousTerm(t.sym, promote(μ, σ2, min, max)...) + else + # keep the field type the statistics would have had, without the O(n) passes + E = eltype(xs) + nan = E <: Number ? convert(float(E), NaN) : NaN + return ContinuousTerm(t.sym, nan, nan, nan, nan) + end end # default contrasts: dummy coding concrete_term(t::Term, xs::AbstractVector, ::Nothing) = concrete_term(t, xs, CategoricalTerm) diff --git a/src/terms.jl b/src/terms.jl index 3617e28b..f6b99ffd 100644 --- a/src/terms.jl +++ b/src/terms.jl @@ -196,6 +196,9 @@ Represents a continuous variable, with a name and summary statistics. * `var::T`: Variance * `min::T`: Minimum value * `max::T`: Maximum value + +The summary statistics are `NaN` placeholders when the term was created with +`statistics=false` (see [`schema`](@ref)). """ struct ContinuousTerm{T} <: AbstractTerm sym::Symbol diff --git a/test/schema.jl b/test/schema.jl index 9f6cf7f3..d5cfff99 100644 --- a/test/schema.jl +++ b/test/schema.jl @@ -70,6 +70,41 @@ end + @testset "statistics=false" begin + f = @formula(y ~ 1 + a + b + c) + d = (y = rand(10), a = rand(10), b = Float32.(1:10), + c = repeat(["u","v"], 5)) + + sch = schema(f, d, statistics=false) + t = sch[term(:a)] + @test t isa ContinuousTerm + @test isnan(t.mean) && isnan(t.var) && isnan(t.min) && isnan(t.max) + # the statistics keep the type they would have had + @test sch[term(:b)] isa ContinuousTerm{Float32} + # categorical terms are unaffected + c1, c2 = sch[term(:c)], schema(f, d)[term(:c)] + @test c1 isa CategoricalTerm && c1.sym == c2.sym && c1.contrasts == c2.contrasts + + # skipping statistics changes nothing else downstream + ff = apply_schema(f, sch) + @test modelmatrix(ff.rhs, d) == modelmatrix(apply_schema(f, schema(f, d)).rhs, d) + + # hints still work, and never see the keyword + sch1 = schema(f, d, Dict(:a => CategoricalTerm), statistics=false) + @test sch1[term(:a)] isa CategoricalTerm{DummyCoding} + sch2 = schema(f, d, Dict(:a => EffectsCoding()), statistics=false) + @test sch2[term(:a)] isa CategoricalTerm{EffectsCoding} + # a ContinuousTerm hint skips the statistics too + sch3 = schema(f, d, Dict(:a => ContinuousTerm), statistics=false) + @test isnan(sch3[term(:a)].mean) + # an AbstractTerm hint is included as is + hint = schema(f, d)[term(:a)] + @test schema(f, d, Dict(:a => hint), statistics=false)[term(:a)] === hint + + @test isnan(concrete_term(term(:a), [1, 2, 3], ContinuousTerm, statistics=false).mean) + @test concrete_term(term(:a), [1, 2, 3], ContinuousTerm, statistics=true).mean == 2.0 + end + @testset "nice errors" begin d = (yyy = rand(10), aaa = rand(10), bbb = repeat([:a, :b], 5)) @test_throws ArgumentError("There isn't a variable called 'aa' in your data; the nearest names appear to be: aaa") concrete_term(Term(:aa), d, nothing)